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

via llvm-commits llvm-commits at lists.llvm.org
Fri Aug 28 06:38:39 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-llvm-transforms

Author: Arseniy Obolenskiy (aobolensk)

<details>
<summary>Changes</summary>

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

---
Full diff: https://github.com/llvm/llvm-project/pull/219472.diff


2 Files Affected:

- (modified) llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp (+26-20) 
- (modified) llvm/test/Transforms/InferAddressSpaces/AMDGPU/ptrmask.ll (+96-20) 


``````````diff
diff --git a/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp b/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp
index 3820f3e1e45ad..f9ad24107466e 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 28d089b920e58..40c2784e9b8f1 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 }
 

``````````

</details>


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


More information about the llvm-commits mailing list