[llvm] [InferAddressSpaces] Do not commute ptrmask with cast when it may not preserve null (PR #219472)

Arseniy Obolenskiy via llvm-commits llvm-commits at lists.llvm.org
Wed Sep 23 06:21:14 PDT 2026


https://github.com/aobolensk updated https://github.com/llvm/llvm-project/pull/219472

>From 2397e7cca03f0790378b089fefb70ecfe061267a Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Fri, 28 Aug 2026 15:36:20 +0200
Subject: [PATCH 1/2] [InferAddressSpaces] Do not commute ptrmask with cast
 when it may not preserve null

A non-noop addrspacecast can map null to a different value, so folding the mask into the cast is only safe when the pointer is known non-null
---
 .../Transforms/Scalar/InferAddressSpaces.cpp  |  46 ++++---
 .../InferAddressSpaces/AMDGPU/ptrmask.ll      | 116 +++++++++++++++---
 2 files changed, 122 insertions(+), 40 deletions(-)

diff --git a/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp b/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp
index 3820f3e1e45ad4..f9ad24107466eb 100644
--- a/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp
+++ b/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp
@@ -827,17 +827,16 @@ Value *InferAddressSpacesImpl::clonePtrMaskWithNewAddressSpace(
 
   KnownBits OldPtrBits{DL->getPointerSizeInBits(OldAddrSpace)};
   KnownBits NewPtrBits{DL->getPointerSizeInBits(NewAddrSpace)};
-  if (!TTI->isNoopAddrSpaceCast(OldAddrSpace, NewAddrSpace)) {
+  bool IsNoopCast = TTI->isNoopAddrSpaceCast(OldAddrSpace, NewAddrSpace);
+  if (!IsNoopCast) {
     std::tie(OldPtrBits, NewPtrBits) =
         TTI->computeKnownBitsAddrSpaceCast(NewAddrSpace, *PtrOpUse.get());
   }
 
-  // If the pointers in both addrspaces have a bitwise representation and if the
-  // representation of the new pointer is smaller (fewer bits) than the old one,
-  // check if the mask is applicable to the ptr in the new addrspace. Any
-  // masking only clearing the low bits will also apply in the new addrspace
-  // Note: checking if the mask clears high bits is not sufficient as those
-  // might have already been 0 in the old ptr.
+  bool KnownNonZero = OldPtrBits.isNonZero();
+
+  // If narrower, check the mask only clears low bits that fit in it.
+  bool MaskFits = true;
   if (OldPtrBits.getBitWidth() > NewPtrBits.getBitWidth()) {
     KnownBits MaskBits =
         computeKnownBits(MaskOp, *DL, /*AssumptionCache=*/nullptr, I);
@@ -846,19 +845,26 @@ Value *InferAddressSpacesImpl::clonePtrMaskWithNewAddressSpace(
     OldPtrBits.One |= ~OldPtrBits.Zero;
     // Check which bits are cleared by the mask in the old ptr.
     KnownBits ClearedBits = KnownBits::sub(OldPtrBits, OldPtrBits & MaskBits);
-
-    // If the mask isn't applicable to the new ptr, leave the ptrmask as-is and
-    // insert an addrspacecast after it.
-    if (ClearedBits.countMaxActiveBits() > NewPtrBits.countMaxActiveBits()) {
-      std::optional<BasicBlock::iterator> InsertPoint =
-          I->getInsertionPointAfterDef();
-      assert(InsertPoint && "insertion after ptrmask should be possible");
-      Type *NewPtrType = getPtrOrVecOfPtrsWithNewAS(I->getType(), NewAddrSpace);
-      Instruction *AddrSpaceCast =
-          new AddrSpaceCastInst(I, NewPtrType, "", *InsertPoint);
-      AddrSpaceCast->setDebugLoc(I->getDebugLoc());
-      return AddrSpaceCast;
-    }
+    MaskFits =
+        ClearedBits.countMaxActiveBits() <= NewPtrBits.countMaxActiveBits();
+  }
+
+  // Commuting past a non-noop cast is unsound for a maybe-null pointer,
+  // since the cast may not preserve null.
+  if (!IsNoopCast && MaskFits && !KnownNonZero)
+    KnownNonZero =
+        isKnownNonZero(PtrOpUse.get(), SimplifyQuery(*DL, DT, &AC, I));
+  bool CanCommuteWithCast = IsNoopCast || (MaskFits && KnownNonZero);
+
+  if (!CanCommuteWithCast) {
+    std::optional<BasicBlock::iterator> InsertPoint =
+        I->getInsertionPointAfterDef();
+    assert(InsertPoint && "insertion after ptrmask should be possible");
+    Type *NewPtrType = getPtrOrVecOfPtrsWithNewAS(I->getType(), NewAddrSpace);
+    Instruction *AddrSpaceCast =
+        new AddrSpaceCastInst(I, NewPtrType, "", *InsertPoint);
+    AddrSpaceCast->setDebugLoc(I->getDebugLoc());
+    return AddrSpaceCast;
   }
 
   IRBuilder<> B(I);
diff --git a/llvm/test/Transforms/InferAddressSpaces/AMDGPU/ptrmask.ll b/llvm/test/Transforms/InferAddressSpaces/AMDGPU/ptrmask.ll
index 28d089b920e584..40c2784e9b8f11 100644
--- a/llvm/test/Transforms/InferAddressSpaces/AMDGPU/ptrmask.ll
+++ b/llvm/test/Transforms/InferAddressSpaces/AMDGPU/ptrmask.ll
@@ -128,9 +128,7 @@ define <3 x ptr addrspace(3)> @ptrmask_vector_cast_flat_to_local(<3 x ptr> %src.
   ret <3 x ptr addrspace(3)> %cast
 }
 
-; Casting null *does* result in null again if addrspace 0 is casted to a
-; smaller addrspace (by default we assume that casting to a smaller addrspace =
-; truncating)
+; Flat null does not cast to local null, so the mask must stay in flat.
 define i8 @ptrmask_cast_flat_null_to_local(i64 %mask) {
 ; CHECK-LABEL: @ptrmask_cast_flat_null_to_local(
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) addrspacecast (ptr null to ptr addrspace(3)), align 1
@@ -338,7 +336,9 @@ define i8 @ptrmask_cast_local_null_to_flat_const_mask_7fffffffffffffff() {
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffff00000000(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_ffffffff00000000(
-; CHECK-NEXT:    [[TMP1:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 0)
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
+; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -4294967296)
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -350,7 +350,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffff00000000(ptr addrspace(
 
 define <3 x ptr addrspace(3)> @ptrmask_vector_cast_local_to_flat_const_mask_ffffffff00000000(<3 x ptr addrspace(3)> %src.ptr) {
 ; CHECK-LABEL: @ptrmask_vector_cast_local_to_flat_const_mask_ffffffff00000000(
-; CHECK-NEXT:    [[TMP1:%.*]] = call <3 x ptr addrspace(3)> @llvm.ptrmask.v3p3.v3i32(<3 x ptr addrspace(3)> [[SRC_PTR:%.*]], <3 x i32> zeroinitializer)
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast <3 x ptr addrspace(3)> [[SRC_PTR:%.*]] to <3 x ptr>
+; CHECK-NEXT:    [[MASKED:%.*]] = call <3 x ptr> @llvm.ptrmask.v3p0.v3i64(<3 x ptr> [[CAST]], <3 x i64> splat (i64 -4294967296))
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast <3 x ptr> [[MASKED]] to <3 x ptr addrspace(3)>
 ; CHECK-NEXT:    ret <3 x ptr addrspace(3)> [[TMP1]]
 ;
   %cast = addrspacecast <3 x ptr addrspace(3)> %src.ptr to <3 x ptr>
@@ -361,7 +363,9 @@ define <3 x ptr addrspace(3)> @ptrmask_vector_cast_local_to_flat_const_mask_ffff
 
 define i8 @ptrmask_cast_local_null_to_flat_const_mask_ffffffff00000000() {
 ; CHECK-LABEL: @ptrmask_cast_local_null_to_flat_const_mask_ffffffff00000000(
-; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) null, align 1
+; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr addrspacecast (ptr addrspace(3) null to ptr), i64 -4294967296)
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
+; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
   %cast = addrspacecast ptr addrspace(3) zeroinitializer to ptr
@@ -372,7 +376,9 @@ define i8 @ptrmask_cast_local_null_to_flat_const_mask_ffffffff00000000() {
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffff80000000(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_ffffffff80000000(
-; CHECK-NEXT:    [[TMP1:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -2147483648)
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
+; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -2147483648)
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -382,10 +388,12 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffff80000000(ptr addrspace(
   ret i8 %load
 }
 
-; Test some align-down patterns. These only touch the low bits, which are preserved through the cast.
+; Align-down patterns, but a flat null pointer keeps the mask in flat.
 define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffff0000(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_ffffffffffff0000(
-; CHECK-NEXT:    [[TMP1:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -65536)
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
+; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -65536)
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -397,7 +405,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffff0000(ptr addrspace(
 
 define <3 x ptr addrspace(3)> @ptrmask_vector_cast_local_to_flat_const_mask_ffffffffffff0000(<3 x ptr addrspace(3)> %src.ptr) {
 ; CHECK-LABEL: @ptrmask_vector_cast_local_to_flat_const_mask_ffffffffffff0000(
-; CHECK-NEXT:    [[TMP1:%.*]] = call <3 x ptr addrspace(3)> @llvm.ptrmask.v3p3.v3i32(<3 x ptr addrspace(3)> [[SRC_PTR:%.*]], <3 x i32> splat (i32 -65536))
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast <3 x ptr addrspace(3)> [[SRC_PTR:%.*]] to <3 x ptr>
+; CHECK-NEXT:    [[MASKED:%.*]] = call <3 x ptr> @llvm.ptrmask.v3p0.v3i64(<3 x ptr> [[CAST]], <3 x i64> splat (i64 -65536))
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast <3 x ptr> [[MASKED]] to <3 x ptr addrspace(3)>
 ; CHECK-NEXT:    ret <3 x ptr addrspace(3)> [[TMP1]]
 ;
   %cast = addrspacecast <3 x ptr addrspace(3)> %src.ptr to <3 x ptr>
@@ -408,7 +418,9 @@ define <3 x ptr addrspace(3)> @ptrmask_vector_cast_local_to_flat_const_mask_ffff
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffffff00(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_ffffffffffffff00(
-; CHECK-NEXT:    [[TMP1:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -256)
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
+; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -256)
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -420,7 +432,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffffff00(ptr addrspace(
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffffffe0(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_ffffffffffffffe0(
-; CHECK-NEXT:    [[TMP1:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -32)
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
+; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -32)
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -432,7 +446,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffffffe0(ptr addrspace(
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffff0(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_fffffffffffffff0(
-; CHECK-NEXT:    [[TMP1:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -16)
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
+; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -16)
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -444,7 +460,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffff0(ptr addrspace(
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffff8(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_fffffffffffffff8(
-; CHECK-NEXT:    [[TMP1:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -8)
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
+; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -8)
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -456,7 +474,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffff8(ptr addrspace(
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffffc(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_fffffffffffffffc(
-; CHECK-NEXT:    [[TMP1:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -4)
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
+; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -4)
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -468,7 +488,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffffc(ptr addrspace(
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffffe(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_fffffffffffffffe(
-; CHECK-NEXT:    [[TMP1:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -2)
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
+; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -2)
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -489,12 +511,45 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffffffff(ptr addrspace(
   ret i8 %load
 }
 
+; -1 & -4096 != -1, so the align-down must not move to the local pointer.
+define i1 @ptrmask_cast_local_to_flat_const_mask_icmp_null(ptr addrspace(3) %src.ptr) {
+; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_icmp_null(
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
+; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -4096)
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
+; CHECK-NEXT:    [[CMP:%.*]] = icmp eq ptr addrspace(3) [[TMP1]], addrspacecast (ptr null to ptr addrspace(3))
+; CHECK-NEXT:    ret i1 [[CMP]]
+;
+  %cast = addrspacecast ptr addrspace(3) %src.ptr to ptr
+  %masked = call ptr @llvm.ptrmask.p0.i64(ptr %cast, i64 -4096)
+  %cmp = icmp eq ptr %masked, null
+  ret i1 %cmp
+}
+
+; A known non-null flat pointer allows the mask on the local pointer.
+define i8 @ptrmask_cast_local_to_flat_const_mask_nonnull(ptr addrspace(3) %src.ptr) {
+; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_nonnull(
+; CHECK-NEXT:    [[NONNULL:%.*]] = icmp ne ptr addrspace(3) [[SRC_PTR:%.*]], addrspacecast (ptr null to ptr addrspace(3))
+; CHECK-NEXT:    call void @llvm.assume(i1 [[NONNULL]])
+; CHECK-NEXT:    [[TMP1:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR]], i32 -4096)
+; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
+; CHECK-NEXT:    ret i8 [[LOAD]]
+;
+  %cast = addrspacecast ptr addrspace(3) %src.ptr to ptr
+  %nonnull = icmp ne ptr %cast, null
+  call void @llvm.assume(i1 %nonnull)
+  %masked = call ptr @llvm.ptrmask.p0.i64(ptr %cast, i64 -4096)
+  %load = load i8, ptr %masked
+  ret i8 %load
+}
+
 ; Make sure non-constant masks can also be handled.
 define i8 @ptrmask_cast_local_to_flat_load_range_mask(ptr addrspace(3) %src.ptr, ptr addrspace(1) %mask.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_load_range_mask(
 ; CHECK-NEXT:    [[LOAD_MASK:%.*]] = load i64, ptr addrspace(1) [[MASK_PTR:%.*]], align 8, !range [[RNG0:![0-9]+]]
-; CHECK-NEXT:    [[TMP1:%.*]] = trunc i64 [[LOAD_MASK]] to i32
-; CHECK-NEXT:    [[TMP2:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 [[TMP1]])
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
+; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 [[LOAD_MASK]])
+; CHECK-NEXT:    [[TMP2:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP2]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -508,8 +563,9 @@ define i8 @ptrmask_cast_local_to_flat_load_range_mask(ptr addrspace(3) %src.ptr,
 define <2 x ptr addrspace(3)> @ptrmask_vector_cast_local_to_flat_load_range_mask(<2 x ptr addrspace(3)> %src.ptr, ptr addrspace(1) %mask.ptr) {
 ; CHECK-LABEL: @ptrmask_vector_cast_local_to_flat_load_range_mask(
 ; CHECK-NEXT:    [[LOAD_MASK:%.*]] = load <2 x i64>, ptr addrspace(1) [[MASK_PTR:%.*]], align 16, !range [[RNG0]]
-; CHECK-NEXT:    [[TMP1:%.*]] = trunc <2 x i64> [[LOAD_MASK]] to <2 x i32>
-; CHECK-NEXT:    [[TMP2:%.*]] = call <2 x ptr addrspace(3)> @llvm.ptrmask.v2p3.v2i32(<2 x ptr addrspace(3)> [[SRC_PTR:%.*]], <2 x i32> [[TMP1]])
+; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast <2 x ptr addrspace(3)> [[SRC_PTR:%.*]] to <2 x ptr>
+; CHECK-NEXT:    [[MASKED:%.*]] = call <2 x ptr> @llvm.ptrmask.v2p0.v2i64(<2 x ptr> [[CAST]], <2 x i64> [[LOAD_MASK]])
+; CHECK-NEXT:    [[TMP2:%.*]] = addrspacecast <2 x ptr> [[MASKED]] to <2 x ptr addrspace(3)>
 ; CHECK-NEXT:    ret <2 x ptr addrspace(3)> [[TMP2]]
 ;
   %load.mask = load <2 x i64>, ptr addrspace(1) %mask.ptr, align 16, !range !0
@@ -519,6 +575,25 @@ define <2 x ptr addrspace(3)> @ptrmask_vector_cast_local_to_flat_load_range_mask
   ret <2 x ptr addrspace(3)> %cast2
 }
 
+define i8 @ptrmask_cast_local_to_flat_load_range_mask_nonnull(ptr addrspace(3) %src.ptr, ptr addrspace(1) %mask.ptr) {
+; CHECK-LABEL: @ptrmask_cast_local_to_flat_load_range_mask_nonnull(
+; CHECK-NEXT:    [[LOAD_MASK:%.*]] = load i64, ptr addrspace(1) [[MASK_PTR:%.*]], align 8, !range [[RNG0]]
+; CHECK-NEXT:    [[NONNULL:%.*]] = icmp ne ptr addrspace(3) [[SRC_PTR:%.*]], addrspacecast (ptr null to ptr addrspace(3))
+; CHECK-NEXT:    call void @llvm.assume(i1 [[NONNULL]])
+; CHECK-NEXT:    [[TMP1:%.*]] = trunc i64 [[LOAD_MASK]] to i32
+; CHECK-NEXT:    [[TMP2:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR]], i32 [[TMP1]])
+; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP2]], align 1
+; CHECK-NEXT:    ret i8 [[LOAD]]
+;
+  %load.mask = load i64, ptr addrspace(1) %mask.ptr, align 8, !range !0
+  %cast = addrspacecast ptr addrspace(3) %src.ptr to ptr
+  %nonnull = icmp ne ptr %cast, null
+  call void @llvm.assume(i1 %nonnull)
+  %masked = call ptr @llvm.ptrmask.p0.i64(ptr %cast, i64 %load.mask)
+  %load = load i8, ptr %masked
+  ret i8 %load
+}
+
 ; Non-const masks with no known range should not prevent other ptr-manipulating
 ; instructions (such as gep) from being converted.
 define i8 @ptrmask_cast_local_to_flat_unknown_mask(ptr addrspace(3) %src.ptr, i64 %mask, i64 %idx) {
@@ -574,6 +649,7 @@ declare <3 x ptr> @llvm.ptrmask.v3p0.v3i64(<3 x ptr>, <3 x i64>) #0
 declare ptr addrspace(5) @llvm.ptrmask.p5.i32(ptr addrspace(5), i32) #0
 declare ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3), i32) #0
 declare ptr addrspace(1) @llvm.ptrmask.p1.i64(ptr addrspace(1), i64) #0
+declare void @llvm.assume(i1)
 
 attributes #0 = { nounwind readnone speculatable willreturn }
 

>From 9b9ccb2687d808cecfaf0cfb040a9cc0fd3afa14 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Mon, 14 Sep 2026 18:04:47 +0200
Subject: [PATCH 2/2] Get rid of 3 -> 0 -> 3 roundtrip

---
 .../llvm/Analysis/TargetTransformInfo.h       |   5 +
 .../llvm/Analysis/TargetTransformInfoImpl.h   |   6 +
 llvm/lib/Analysis/TargetTransformInfo.cpp     |   5 +
 .../Target/AMDGPU/AMDGPUTargetTransformInfo.h |   5 +
 .../Transforms/Scalar/InferAddressSpaces.cpp  | 136 ++++++++++++++----
 .../InferAddressSpaces/AMDGPU/ptrmask.ll      |  94 ++++++------
 6 files changed, 175 insertions(+), 76 deletions(-)

diff --git a/llvm/include/llvm/Analysis/TargetTransformInfo.h b/llvm/include/llvm/Analysis/TargetTransformInfo.h
index 5e22b890ec9751..67bf1ce107606b 100644
--- a/llvm/include/llvm/Analysis/TargetTransformInfo.h
+++ b/llvm/include/llvm/Analysis/TargetTransformInfo.h
@@ -574,6 +574,11 @@ class TargetTransformInfo {
 
   LLVM_ABI bool isNoopAddrSpaceCast(unsigned FromAS, unsigned ToAS) const;
 
+  /// Return the bit pattern of the null pointer in address space \p AS, or
+  /// std::nullopt if the target does not define one. The result has the pointer
+  /// size of \p AS and need not be zero.
+  LLVM_ABI std::optional<APInt> getNullPointerValue(unsigned AS) const;
+
   // Given an address space cast of the given pointer value, calculate the known
   // bits of the source pointer in the source addrspace and the destination
   // pointer in the destination addrspace.
diff --git a/llvm/include/llvm/Analysis/TargetTransformInfoImpl.h b/llvm/include/llvm/Analysis/TargetTransformInfoImpl.h
index 4240eb090e9526..fdc5ddc06da7ac 100644
--- a/llvm/include/llvm/Analysis/TargetTransformInfoImpl.h
+++ b/llvm/include/llvm/Analysis/TargetTransformInfoImpl.h
@@ -153,6 +153,12 @@ class LLVM_ABI TargetTransformInfoImplBase {
 
   virtual bool isNoopAddrSpaceCast(unsigned, unsigned) const { return false; }
 
+  virtual std::optional<APInt> getNullPointerValue(unsigned AS) const {
+    if (DL.isNonIntegralAddressSpace(AS))
+      return std::nullopt;
+    return APInt::getZero(DL.getPointerSizeInBits(AS));
+  }
+
   virtual std::pair<KnownBits, KnownBits>
   computeKnownBitsAddrSpaceCast(unsigned ToAS, const Value &PtrOp) const {
     const Type *PtrTy = PtrOp.getType();
diff --git a/llvm/lib/Analysis/TargetTransformInfo.cpp b/llvm/lib/Analysis/TargetTransformInfo.cpp
index 67dee54028c0e6..e85df0fc8fb201 100644
--- a/llvm/lib/Analysis/TargetTransformInfo.cpp
+++ b/llvm/lib/Analysis/TargetTransformInfo.cpp
@@ -326,6 +326,11 @@ bool TargetTransformInfo::isNoopAddrSpaceCast(unsigned FromAS,
   return TTIImpl->isNoopAddrSpaceCast(FromAS, ToAS);
 }
 
+std::optional<APInt>
+TargetTransformInfo::getNullPointerValue(unsigned AS) const {
+  return TTIImpl->getNullPointerValue(AS);
+}
+
 std::pair<KnownBits, KnownBits>
 TargetTransformInfo::computeKnownBitsAddrSpaceCast(unsigned ToAS,
                                                    const Value &PtrOp) const {
diff --git a/llvm/lib/Target/AMDGPU/AMDGPUTargetTransformInfo.h b/llvm/lib/Target/AMDGPU/AMDGPUTargetTransformInfo.h
index 3f9aacaa036bd7..4c26c56a6815f2 100644
--- a/llvm/lib/Target/AMDGPU/AMDGPUTargetTransformInfo.h
+++ b/llvm/lib/Target/AMDGPU/AMDGPUTargetTransformInfo.h
@@ -201,6 +201,11 @@ class GCNTTIImpl final : public BasicTTIImplBase<GCNTTIImpl> {
     return AMDGPU::addrspacesMayAlias(AS0, AS1);
   }
 
+  std::optional<APInt> getNullPointerValue(unsigned AS) const override {
+    return APInt(getDataLayout().getPointerSizeInBits(AS),
+                 AMDGPU::getNullPointerValue(AS), /*isSigned=*/true);
+  }
+
   unsigned getFlatAddressSpace() const override {
     // Don't bother running InferAddressSpaces pass on graphics shaders which
     // don't use flat addressing.
diff --git a/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp b/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp
index f9ad24107466eb..8ef0b5f814dcb9 100644
--- a/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp
+++ b/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp
@@ -246,6 +246,11 @@ class InferAddressSpacesImpl {
 
   bool isSafeToCastConstAddrSpace(Constant *C, unsigned NewAS) const;
 
+  bool isKnownNotNullPointer(const Value &PtrOp, const APInt &Null,
+                             const SimplifyQuery &SQ) const;
+
+  static bool maskPreservesNull(const APInt &Null, const KnownBits &MaskBits);
+
   Value *clonePtrMaskWithNewAddressSpace(
       IntrinsicInst *I, unsigned NewAddrSpace,
       const ValueToValueMapTy &ValueWithNewAddrSpace,
@@ -777,7 +782,8 @@ static Value *phiNodeOperandWithNewAddressSpace(AddrSpaceCastInst *NewI,
 
 // A helper function for cloneInstructionWithNewAddressSpace. Returns the clone
 // of OperandUse.get() in the new address space. If the clone is not ready yet,
-// returns poison in the new address space as a placeholder.
+// returns poison in the new address space as a placeholder, or null when
+// PoisonUsesToFix is null.
 static Value *operandWithNewAddressSpaceOrCreatePoison(
     const Use &OperandUse, unsigned NewAddrSpace,
     const ValueToValueMapTy &ValueWithNewAddrSpace,
@@ -809,10 +815,49 @@ static Value *operandWithNewAddressSpaceOrCreatePoison(
     return NewI;
   }
 
+  if (!PoisonUsesToFix)
+    return nullptr;
+
   PoisonUsesToFix->push_back(&OperandUse);
   return PoisonValue::get(NewPtrTy);
 }
 
+// Returns the operand in NewAddrSpace, or null if only a poison placeholder is
+// available. A placeholder is patched up at a single operand position, so a
+// caller using the operand more than once needs the real value.
+static Value *getAvailableNewAddressSpaceOperand(
+    const Use &OperandUse, unsigned NewAddrSpace,
+    const ValueToValueMapTy &ValueWithNewAddrSpace,
+    const PredicatedAddrSpaceMapTy &PredicatedAS) {
+  // Postorder puts a ptrmask before the addrspacecast it feeds on, so the cast
+  // is not in ValueWithNewAddrSpace yet.
+  if (AddrSpaceCastInst *ASC = dyn_cast<AddrSpaceCastInst>(OperandUse.get()))
+    if (ASC->getSrcAddressSpace() == NewAddrSpace)
+      return ASC->getPointerOperand();
+
+  return operandWithNewAddressSpaceOrCreatePoison(
+      OperandUse, NewAddrSpace, ValueWithNewAddrSpace, PredicatedAS,
+      /*PoisonUsesToFix=*/nullptr);
+}
+
+// Returns true if PtrOp is known not to be the null pointer of its address
+// space, whose bit pattern is Null.
+bool InferAddressSpacesImpl::isKnownNotNullPointer(
+    const Value &PtrOp, const APInt &Null, const SimplifyQuery &SQ) const {
+  if (Null.isZero())
+    return isKnownNonZero(&PtrOp, SQ);
+  return KnownBits::ne(computeKnownBits(&PtrOp, SQ),
+                       KnownBits::makeConstant(Null))
+      .value_or(false);
+}
+
+// Returns true if masking the null pointer with bit pattern Null by a mask with
+// the given known bits leaves it unchanged.
+bool InferAddressSpacesImpl::maskPreservesNull(const APInt &Null,
+                                               const KnownBits &MaskBits) {
+  return Null.isSubsetOf(MaskBits.One.zextOrTrunc(Null.getBitWidth()));
+}
+
 // A helper function for cloneInstructionWithNewAddressSpace. Handles the
 // conversion of a ptrmask intrinsic instruction.
 Value *InferAddressSpacesImpl::clonePtrMaskWithNewAddressSpace(
@@ -833,30 +878,53 @@ Value *InferAddressSpacesImpl::clonePtrMaskWithNewAddressSpace(
         TTI->computeKnownBitsAddrSpaceCast(NewAddrSpace, *PtrOpUse.get());
   }
 
-  bool KnownNonZero = OldPtrBits.isNonZero();
-
-  // If narrower, check the mask only clears low bits that fit in it.
-  bool MaskFits = true;
-  if (OldPtrBits.getBitWidth() > NewPtrBits.getBitWidth()) {
-    KnownBits MaskBits =
-        computeKnownBits(MaskOp, *DL, /*AssumptionCache=*/nullptr, I);
-    // Set all unknown bits of the old ptr to 1, so that we are conservative in
-    // checking which bits are cleared by the mask.
-    OldPtrBits.One |= ~OldPtrBits.Zero;
-    // Check which bits are cleared by the mask in the old ptr.
-    KnownBits ClearedBits = KnownBits::sub(OldPtrBits, OldPtrBits & MaskBits);
-    MaskFits =
-        ClearedBits.countMaxActiveBits() <= NewPtrBits.countMaxActiveBits();
-  }
-
-  // Commuting past a non-noop cast is unsound for a maybe-null pointer,
-  // since the cast may not preserve null.
-  if (!IsNoopCast && MaskFits && !KnownNonZero)
-    KnownNonZero =
-        isKnownNonZero(PtrOpUse.get(), SimplifyQuery(*DL, DT, &AC, I));
-  bool CanCommuteWithCast = IsNoopCast || (MaskFits && KnownNonZero);
+  bool CanCommuteWithCast = IsNoopCast;
+  // Set when the fold needs an explicit guard for the null pointer, which the
+  // cast maps to a different bit pattern.
+  Value *GuardedNewPtr = nullptr;
+  std::optional<APInt> NewNull;
+  if (!IsNoopCast) {
+    const SimplifyQuery SQ(*DL, DT, &AC, I);
+    KnownBits MaskBits = computeKnownBits(MaskOp, SQ);
+
+    // If narrower, check the mask only clears low bits that fit in it.
+    bool MaskFits = true;
+    if (OldPtrBits.getBitWidth() > NewPtrBits.getBitWidth()) {
+      // Set all unknown bits of the old ptr to 1, so that we are conservative
+      // in checking which bits are cleared by the mask.
+      OldPtrBits.One |= ~OldPtrBits.Zero;
+      // Check which bits are cleared by the mask in the old ptr.
+      KnownBits ClearedBits = KnownBits::sub(OldPtrBits, OldPtrBits & MaskBits);
+      MaskFits =
+          ClearedBits.countMaxActiveBits() <= NewPtrBits.countMaxActiveBits();
+    }
 
-  if (!CanCommuteWithCast) {
+    // A non-noop addrspacecast maps the null pointer of one address space to
+    // the null pointer of the other, and those need not share a bit pattern,
+    // so a mask applied on one side of the cast is not generally applicable on
+    // the other.
+    std::optional<APInt> OldNull = TTI->getNullPointerValue(OldAddrSpace);
+    NewNull = TTI->getNullPointerValue(NewAddrSpace);
+    // Masking a null pointer leaves it unchanged, so a null input still
+    // produces null on that side of the cast.
+    bool MaskKeepsOldNull = OldNull && maskPreservesNull(*OldNull, MaskBits);
+    bool MaskKeepsNewNull = NewNull && maskPreservesNull(*NewNull, MaskBits);
+    bool KnownNonNull =
+        OldNull && isKnownNotNullPointer(*PtrOpUse.get(), *OldNull, SQ);
+
+    // The masked pointer never is the old null pointer, or masking null gives
+    // null on either side. Either way the mask applies as is.
+    CanCommuteWithCast =
+        MaskFits && (KnownNonNull || (MaskKeepsOldNull && MaskKeepsNewNull));
+
+    // Otherwise the null case still folds if it is guarded explicitly, which
+    // needs the new address space operand up front.
+    if (!CanCommuteWithCast && MaskFits && MaskKeepsOldNull && NewNull)
+      GuardedNewPtr = getAvailableNewAddressSpaceOperand(
+          PtrOpUse, NewAddrSpace, ValueWithNewAddrSpace, PredicatedAS);
+  }
+
+  if (!CanCommuteWithCast && !GuardedNewPtr) {
     std::optional<BasicBlock::iterator> InsertPoint =
         I->getInsertionPointAfterDef();
     assert(InsertPoint && "insertion after ptrmask should be possible");
@@ -872,11 +940,21 @@ Value *InferAddressSpacesImpl::clonePtrMaskWithNewAddressSpace(
     MaskTy = MaskTy->getWithNewBitWidth(NewPtrBits.getBitWidth());
     MaskOp = B.CreateTrunc(MaskOp, MaskTy);
   }
-  Value *NewPtr = operandWithNewAddressSpaceOrCreatePoison(
-      PtrOpUse, NewAddrSpace, ValueWithNewAddrSpace, PredicatedAS,
-      PoisonUsesToFix);
-  return B.CreateIntrinsic(Intrinsic::ptrmask, {NewPtr->getType(), MaskTy},
-                           {NewPtr, MaskOp});
+  Value *NewPtr = GuardedNewPtr
+                      ? GuardedNewPtr
+                      : operandWithNewAddressSpaceOrCreatePoison(
+                            PtrOpUse, NewAddrSpace, ValueWithNewAddrSpace,
+                            PredicatedAS, PoisonUsesToFix);
+  Value *Masked = B.CreateIntrinsic(
+      Intrinsic::ptrmask, {NewPtr->getType(), MaskTy}, {NewPtr, MaskOp});
+  if (!GuardedNewPtr)
+    return Masked;
+
+  // The cast is injective and maps the new null pointer to the old one, so the
+  // null test is equivalent here. Spell null out as its bit pattern, which a
+  // target may define to be nonzero.
+  Constant *NewNullPtr = Constant::getIntegerValue(NewPtr->getType(), *NewNull);
+  return B.CreateSelect(B.CreateICmpEQ(NewPtr, NewNullPtr), NewNullPtr, Masked);
 }
 
 // Returns a clone of `I` with its operands converted to those specified in
diff --git a/llvm/test/Transforms/InferAddressSpaces/AMDGPU/ptrmask.ll b/llvm/test/Transforms/InferAddressSpaces/AMDGPU/ptrmask.ll
index 40c2784e9b8f11..a5dd120156a7b2 100644
--- a/llvm/test/Transforms/InferAddressSpaces/AMDGPU/ptrmask.ll
+++ b/llvm/test/Transforms/InferAddressSpaces/AMDGPU/ptrmask.ll
@@ -336,9 +336,9 @@ define i8 @ptrmask_cast_local_null_to_flat_const_mask_7fffffffffffffff() {
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffff00000000(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_ffffffff00000000(
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
-; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -4294967296)
-; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
+; CHECK-NEXT:    [[TMP3:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 0)
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq ptr addrspace(3) [[SRC_PTR]], inttoptr (i32 -1 to ptr addrspace(3))
+; CHECK-NEXT:    [[TMP1:%.*]] = select i1 [[TMP2]], ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) [[TMP3]]
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -350,9 +350,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffff00000000(ptr addrspace(
 
 define <3 x ptr addrspace(3)> @ptrmask_vector_cast_local_to_flat_const_mask_ffffffff00000000(<3 x ptr addrspace(3)> %src.ptr) {
 ; CHECK-LABEL: @ptrmask_vector_cast_local_to_flat_const_mask_ffffffff00000000(
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast <3 x ptr addrspace(3)> [[SRC_PTR:%.*]] to <3 x ptr>
-; CHECK-NEXT:    [[MASKED:%.*]] = call <3 x ptr> @llvm.ptrmask.v3p0.v3i64(<3 x ptr> [[CAST]], <3 x i64> splat (i64 -4294967296))
-; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast <3 x ptr> [[MASKED]] to <3 x ptr addrspace(3)>
+; CHECK-NEXT:    [[TMP3:%.*]] = call <3 x ptr addrspace(3)> @llvm.ptrmask.v3p3.v3i32(<3 x ptr addrspace(3)> [[SRC_PTR:%.*]], <3 x i32> zeroinitializer)
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq <3 x ptr addrspace(3)> [[SRC_PTR]], <ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3))>
+; CHECK-NEXT:    [[TMP1:%.*]] = select <3 x i1> [[TMP2]], <3 x ptr addrspace(3)> <ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3))>, <3 x ptr addrspace(3)> [[TMP3]]
 ; CHECK-NEXT:    ret <3 x ptr addrspace(3)> [[TMP1]]
 ;
   %cast = addrspacecast <3 x ptr addrspace(3)> %src.ptr to <3 x ptr>
@@ -363,9 +363,7 @@ define <3 x ptr addrspace(3)> @ptrmask_vector_cast_local_to_flat_const_mask_ffff
 
 define i8 @ptrmask_cast_local_null_to_flat_const_mask_ffffffff00000000() {
 ; CHECK-LABEL: @ptrmask_cast_local_null_to_flat_const_mask_ffffffff00000000(
-; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr addrspacecast (ptr addrspace(3) null to ptr), i64 -4294967296)
-; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
-; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
+; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) null, align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
   %cast = addrspacecast ptr addrspace(3) zeroinitializer to ptr
@@ -376,9 +374,9 @@ define i8 @ptrmask_cast_local_null_to_flat_const_mask_ffffffff00000000() {
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffff80000000(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_ffffffff80000000(
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
-; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -2147483648)
-; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
+; CHECK-NEXT:    [[TMP3:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -2147483648)
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq ptr addrspace(3) [[SRC_PTR]], inttoptr (i32 -1 to ptr addrspace(3))
+; CHECK-NEXT:    [[TMP1:%.*]] = select i1 [[TMP2]], ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) [[TMP3]]
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -388,12 +386,12 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffff80000000(ptr addrspace(
   ret i8 %load
 }
 
-; Align-down patterns, but a flat null pointer keeps the mask in flat.
+; Align-down patterns, guarded because the mask does not preserve the local null.
 define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffff0000(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_ffffffffffff0000(
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
-; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -65536)
-; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
+; CHECK-NEXT:    [[TMP3:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -65536)
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq ptr addrspace(3) [[SRC_PTR]], inttoptr (i32 -1 to ptr addrspace(3))
+; CHECK-NEXT:    [[TMP1:%.*]] = select i1 [[TMP2]], ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) [[TMP3]]
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -405,9 +403,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffff0000(ptr addrspace(
 
 define <3 x ptr addrspace(3)> @ptrmask_vector_cast_local_to_flat_const_mask_ffffffffffff0000(<3 x ptr addrspace(3)> %src.ptr) {
 ; CHECK-LABEL: @ptrmask_vector_cast_local_to_flat_const_mask_ffffffffffff0000(
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast <3 x ptr addrspace(3)> [[SRC_PTR:%.*]] to <3 x ptr>
-; CHECK-NEXT:    [[MASKED:%.*]] = call <3 x ptr> @llvm.ptrmask.v3p0.v3i64(<3 x ptr> [[CAST]], <3 x i64> splat (i64 -65536))
-; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast <3 x ptr> [[MASKED]] to <3 x ptr addrspace(3)>
+; CHECK-NEXT:    [[TMP3:%.*]] = call <3 x ptr addrspace(3)> @llvm.ptrmask.v3p3.v3i32(<3 x ptr addrspace(3)> [[SRC_PTR:%.*]], <3 x i32> splat (i32 -65536))
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq <3 x ptr addrspace(3)> [[SRC_PTR]], <ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3))>
+; CHECK-NEXT:    [[TMP1:%.*]] = select <3 x i1> [[TMP2]], <3 x ptr addrspace(3)> <ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3))>, <3 x ptr addrspace(3)> [[TMP3]]
 ; CHECK-NEXT:    ret <3 x ptr addrspace(3)> [[TMP1]]
 ;
   %cast = addrspacecast <3 x ptr addrspace(3)> %src.ptr to <3 x ptr>
@@ -418,9 +416,9 @@ define <3 x ptr addrspace(3)> @ptrmask_vector_cast_local_to_flat_const_mask_ffff
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffffff00(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_ffffffffffffff00(
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
-; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -256)
-; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
+; CHECK-NEXT:    [[TMP3:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -256)
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq ptr addrspace(3) [[SRC_PTR]], inttoptr (i32 -1 to ptr addrspace(3))
+; CHECK-NEXT:    [[TMP1:%.*]] = select i1 [[TMP2]], ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) [[TMP3]]
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -432,9 +430,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffffff00(ptr addrspace(
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffffffe0(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_ffffffffffffffe0(
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
-; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -32)
-; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
+; CHECK-NEXT:    [[TMP3:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -32)
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq ptr addrspace(3) [[SRC_PTR]], inttoptr (i32 -1 to ptr addrspace(3))
+; CHECK-NEXT:    [[TMP1:%.*]] = select i1 [[TMP2]], ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) [[TMP3]]
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -446,9 +444,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffffffe0(ptr addrspace(
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffff0(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_fffffffffffffff0(
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
-; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -16)
-; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
+; CHECK-NEXT:    [[TMP3:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -16)
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq ptr addrspace(3) [[SRC_PTR]], inttoptr (i32 -1 to ptr addrspace(3))
+; CHECK-NEXT:    [[TMP1:%.*]] = select i1 [[TMP2]], ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) [[TMP3]]
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -460,9 +458,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffff0(ptr addrspace(
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffff8(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_fffffffffffffff8(
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
-; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -8)
-; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
+; CHECK-NEXT:    [[TMP3:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -8)
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq ptr addrspace(3) [[SRC_PTR]], inttoptr (i32 -1 to ptr addrspace(3))
+; CHECK-NEXT:    [[TMP1:%.*]] = select i1 [[TMP2]], ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) [[TMP3]]
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -474,9 +472,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffff8(ptr addrspace(
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffffc(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_fffffffffffffffc(
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
-; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -4)
-; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
+; CHECK-NEXT:    [[TMP3:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -4)
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq ptr addrspace(3) [[SRC_PTR]], inttoptr (i32 -1 to ptr addrspace(3))
+; CHECK-NEXT:    [[TMP1:%.*]] = select i1 [[TMP2]], ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) [[TMP3]]
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -488,9 +486,9 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffffc(ptr addrspace(
 
 define i8 @ptrmask_cast_local_to_flat_const_mask_fffffffffffffffe(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_fffffffffffffffe(
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
-; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -2)
-; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
+; CHECK-NEXT:    [[TMP3:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -2)
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq ptr addrspace(3) [[SRC_PTR]], inttoptr (i32 -1 to ptr addrspace(3))
+; CHECK-NEXT:    [[TMP1:%.*]] = select i1 [[TMP2]], ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) [[TMP3]]
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP1]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -511,12 +509,12 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_ffffffffffffffff(ptr addrspace(
   ret i8 %load
 }
 
-; -1 & -4096 != -1, so the align-down must not move to the local pointer.
+; -1 & -4096 != -1, so the align-down needs a null guard on the local pointer.
 define i1 @ptrmask_cast_local_to_flat_const_mask_icmp_null(ptr addrspace(3) %src.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_const_mask_icmp_null(
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
-; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 -4096)
-; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
+; CHECK-NEXT:    [[TMP3:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 -4096)
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq ptr addrspace(3) [[SRC_PTR]], inttoptr (i32 -1 to ptr addrspace(3))
+; CHECK-NEXT:    [[TMP1:%.*]] = select i1 [[TMP2]], ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) [[TMP3]]
 ; CHECK-NEXT:    [[CMP:%.*]] = icmp eq ptr addrspace(3) [[TMP1]], addrspacecast (ptr null to ptr addrspace(3))
 ; CHECK-NEXT:    ret i1 [[CMP]]
 ;
@@ -547,9 +545,10 @@ define i8 @ptrmask_cast_local_to_flat_const_mask_nonnull(ptr addrspace(3) %src.p
 define i8 @ptrmask_cast_local_to_flat_load_range_mask(ptr addrspace(3) %src.ptr, ptr addrspace(1) %mask.ptr) {
 ; CHECK-LABEL: @ptrmask_cast_local_to_flat_load_range_mask(
 ; CHECK-NEXT:    [[LOAD_MASK:%.*]] = load i64, ptr addrspace(1) [[MASK_PTR:%.*]], align 8, !range [[RNG0:![0-9]+]]
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast ptr addrspace(3) [[SRC_PTR:%.*]] to ptr
-; CHECK-NEXT:    [[MASKED:%.*]] = call ptr @llvm.ptrmask.p0.i64(ptr [[CAST]], i64 [[LOAD_MASK]])
-; CHECK-NEXT:    [[TMP2:%.*]] = addrspacecast ptr [[MASKED]] to ptr addrspace(3)
+; CHECK-NEXT:    [[TMP1:%.*]] = trunc i64 [[LOAD_MASK]] to i32
+; CHECK-NEXT:    [[TMP4:%.*]] = call ptr addrspace(3) @llvm.ptrmask.p3.i32(ptr addrspace(3) [[SRC_PTR:%.*]], i32 [[TMP1]])
+; CHECK-NEXT:    [[TMP3:%.*]] = icmp eq ptr addrspace(3) [[SRC_PTR]], inttoptr (i32 -1 to ptr addrspace(3))
+; CHECK-NEXT:    [[TMP2:%.*]] = select i1 [[TMP3]], ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) [[TMP4]]
 ; CHECK-NEXT:    [[LOAD:%.*]] = load i8, ptr addrspace(3) [[TMP2]], align 1
 ; CHECK-NEXT:    ret i8 [[LOAD]]
 ;
@@ -563,9 +562,10 @@ define i8 @ptrmask_cast_local_to_flat_load_range_mask(ptr addrspace(3) %src.ptr,
 define <2 x ptr addrspace(3)> @ptrmask_vector_cast_local_to_flat_load_range_mask(<2 x ptr addrspace(3)> %src.ptr, ptr addrspace(1) %mask.ptr) {
 ; CHECK-LABEL: @ptrmask_vector_cast_local_to_flat_load_range_mask(
 ; CHECK-NEXT:    [[LOAD_MASK:%.*]] = load <2 x i64>, ptr addrspace(1) [[MASK_PTR:%.*]], align 16, !range [[RNG0]]
-; CHECK-NEXT:    [[CAST:%.*]] = addrspacecast <2 x ptr addrspace(3)> [[SRC_PTR:%.*]] to <2 x ptr>
-; CHECK-NEXT:    [[MASKED:%.*]] = call <2 x ptr> @llvm.ptrmask.v2p0.v2i64(<2 x ptr> [[CAST]], <2 x i64> [[LOAD_MASK]])
-; CHECK-NEXT:    [[TMP2:%.*]] = addrspacecast <2 x ptr> [[MASKED]] to <2 x ptr addrspace(3)>
+; CHECK-NEXT:    [[TMP1:%.*]] = trunc <2 x i64> [[LOAD_MASK]] to <2 x i32>
+; CHECK-NEXT:    [[TMP4:%.*]] = call <2 x ptr addrspace(3)> @llvm.ptrmask.v2p3.v2i32(<2 x ptr addrspace(3)> [[SRC_PTR:%.*]], <2 x i32> [[TMP1]])
+; CHECK-NEXT:    [[TMP3:%.*]] = icmp eq <2 x ptr addrspace(3)> [[SRC_PTR]], <ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3))>
+; CHECK-NEXT:    [[TMP2:%.*]] = select <2 x i1> [[TMP3]], <2 x ptr addrspace(3)> <ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3)), ptr addrspace(3) inttoptr (i32 -1 to ptr addrspace(3))>, <2 x ptr addrspace(3)> [[TMP4]]
 ; CHECK-NEXT:    ret <2 x ptr addrspace(3)> [[TMP2]]
 ;
   %load.mask = load <2 x i64>, ptr addrspace(1) %mask.ptr, align 16, !range !0



More information about the llvm-commits mailing list