[llvm] [InferAddressSpaces] Add a TTI hook for address space joins (PR #225161)

Alex MacLean via llvm-commits llvm-commits at lists.llvm.org
Mon Sep 21 11:54:54 PDT 2026


https://github.com/AlexMaclean updated https://github.com/llvm/llvm-project/pull/225161

>From c8d7f9ef2a85b4dd5bea03392107fc3b3870b20e Mon Sep 17 00:00:00 2001
From: Alex Maclean <amaclean at nvidia.com>
Date: Mon, 21 Sep 2026 11:15:34 -0700
Subject: [PATCH 1/2] [InferAddressSpaces] Add a TTI hook for address space
 joins

---
 .../llvm/Analysis/TargetTransformInfo.h       |   5 +
 .../llvm/Analysis/TargetTransformInfoImpl.h   |   4 +
 llvm/lib/Analysis/TargetTransformInfo.cpp     |   5 +
 .../Target/NVPTX/NVPTXTargetTransformInfo.h   |  11 +
 .../Transforms/Scalar/InferAddressSpaces.cpp  |  46 ++--
 .../CodeGen/NVPTX/infer-address-space-join.ll |  29 +++
 .../NVPTX/address-space-join.ll               | 223 ++++++++++++++++++
 7 files changed, 299 insertions(+), 24 deletions(-)
 create mode 100644 llvm/test/CodeGen/NVPTX/infer-address-space-join.ll
 create mode 100644 llvm/test/Transforms/InferAddressSpaces/NVPTX/address-space-join.ll

diff --git a/llvm/include/llvm/Analysis/TargetTransformInfo.h b/llvm/include/llvm/Analysis/TargetTransformInfo.h
index f1731bc364555..c585c250062b3 100644
--- a/llvm/include/llvm/Analysis/TargetTransformInfo.h
+++ b/llvm/include/llvm/Analysis/TargetTransformInfo.h
@@ -565,6 +565,11 @@ class TargetTransformInfo {
   /// optimize away.
   LLVM_ABI unsigned getFlatAddressSpace() const;
 
+  /// Return the most specific common address space containing AS1 and AS2.
+  /// Pointers from either input space must be convertible to the result with
+  /// addrspacecast.
+  LLVM_ABI unsigned getAddressSpaceJoin(unsigned AS1, unsigned AS2) const;
+
   /// Return any intrinsic address operand indexes which may be rewritten if
   /// they use a flat address space pointer.
   ///
diff --git a/llvm/include/llvm/Analysis/TargetTransformInfoImpl.h b/llvm/include/llvm/Analysis/TargetTransformInfoImpl.h
index 84cb3a6e664b9..8f0d9abc195f9 100644
--- a/llvm/include/llvm/Analysis/TargetTransformInfoImpl.h
+++ b/llvm/include/llvm/Analysis/TargetTransformInfoImpl.h
@@ -146,6 +146,10 @@ class LLVM_ABI TargetTransformInfoImplBase {
 
   virtual unsigned getFlatAddressSpace() const { return -1; }
 
+  virtual unsigned getAddressSpaceJoin(unsigned AS1, unsigned AS2) const {
+    return AS1 == AS2 ? AS1 : getFlatAddressSpace();
+  }
+
   virtual bool collectFlatAddressOperands(SmallVectorImpl<int> &OpIndexes,
                                           Intrinsic::ID IID) const {
     return false;
diff --git a/llvm/lib/Analysis/TargetTransformInfo.cpp b/llvm/lib/Analysis/TargetTransformInfo.cpp
index a9f76a735a48c..83f2fe46b5765 100644
--- a/llvm/lib/Analysis/TargetTransformInfo.cpp
+++ b/llvm/lib/Analysis/TargetTransformInfo.cpp
@@ -316,6 +316,11 @@ unsigned TargetTransformInfo::getFlatAddressSpace() const {
   return TTIImpl->getFlatAddressSpace();
 }
 
+unsigned TargetTransformInfo::getAddressSpaceJoin(unsigned AS1,
+                                                  unsigned AS2) const {
+  return TTIImpl->getAddressSpaceJoin(AS1, AS2);
+}
+
 bool TargetTransformInfo::collectFlatAddressOperands(
     SmallVectorImpl<int> &OpIndexes, Intrinsic::ID IID) const {
   return TTIImpl->collectFlatAddressOperands(OpIndexes, IID);
diff --git a/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h b/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h
index 8bdafd6b905f1..0c5d8c588d3d4 100644
--- a/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h
+++ b/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h
@@ -54,6 +54,17 @@ class NVPTXTTIImpl final : public BasicTTIImplBase<NVPTXTTIImpl> {
     return AddressSpace::ADDRESS_SPACE_GENERIC;
   }
 
+  unsigned getAddressSpaceJoin(unsigned AS1, unsigned AS2) const override {
+    if (AS1 == AS2)
+      return AS1;
+    if ((AS1 == AddressSpace::ADDRESS_SPACE_SHARED &&
+         AS2 == AddressSpace::ADDRESS_SPACE_SHARED_CLUSTER) ||
+        (AS2 == AddressSpace::ADDRESS_SPACE_SHARED &&
+         AS1 == AddressSpace::ADDRESS_SPACE_SHARED_CLUSTER))
+      return AddressSpace::ADDRESS_SPACE_SHARED_CLUSTER;
+    return AddressSpace::ADDRESS_SPACE_GENERIC;
+  }
+
   bool
   canHaveNonUndefGlobalInitializerInAddressSpace(unsigned AS) const override {
     return AS != AddressSpace::ADDRESS_SPACE_SHARED &&
diff --git a/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp b/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp
index 364dea2c8e3a5..d2596bc7c2f75 100644
--- a/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp
+++ b/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp
@@ -70,7 +70,7 @@
 // The monotone transfer function moves the address space of a pointer down a
 // lattice path from uninitialized to specific and then to generic. A join
 // operation of two different specific address spaces pushes the expression down
-// to the generic address space. The analysis completes once it reaches a fixed
+// to their common address space. The analysis completes once it reaches a fixed
 // point.
 //
 // Second, IR rewriting in Step 2 also needs to be circular. For example,
@@ -791,27 +791,25 @@ static Value *operandWithNewAddressSpaceOrCreatePoison(
   if (Constant *C = dyn_cast<Constant>(Operand))
     return ConstantExpr::getAddrSpaceCast(C, NewPtrTy);
 
-  if (Value *NewOperand = ValueWithNewAddrSpace.lookup(Operand))
-    return NewOperand;
-
   Instruction *Inst = cast<Instruction>(OperandUse.getUser());
-  auto I = PredicatedAS.find(std::make_pair(Inst, Operand));
-  if (I != PredicatedAS.end()) {
-    // Insert an addrspacecast on that operand before the user.
-    unsigned NewAS = I->second;
-    Type *NewPtrTy = getPtrOrVecOfPtrsWithNewAS(Operand->getType(), NewAS);
-    auto *NewI = new AddrSpaceCastInst(Operand, NewPtrTy);
-
-    if (LLVM_UNLIKELY(Inst->getOpcode() == Instruction::PHI))
-      return phiNodeOperandWithNewAddressSpace(NewI, Operand);
-
-    NewI->insertBefore(Inst->getIterator());
-    NewI->setDebugLoc(Inst->getDebugLoc());
-    return NewI;
+  if (Value *NewOperand = ValueWithNewAddrSpace.lookup(Operand)) {
+    Operand = NewOperand;
+  } else if (!PredicatedAS.contains(std::make_pair(Inst, Operand))) {
+    assert(PoisonUsesToFix && "missing inferred operand replacement");
+    PoisonUsesToFix->push_back(&OperandUse);
+    return PoisonValue::get(NewPtrTy);
   }
 
-  PoisonUsesToFix->push_back(&OperandUse);
-  return PoisonValue::get(NewPtrTy);
+  if (Operand->getType() == NewPtrTy)
+    return Operand;
+
+  auto *NewI = new AddrSpaceCastInst(Operand, NewPtrTy);
+  if (LLVM_UNLIKELY(Inst->getOpcode() == Instruction::PHI))
+    return phiNodeOperandWithNewAddressSpace(NewI, OperandUse.get());
+
+  NewI->insertBefore(Inst->getIterator());
+  NewI->setDebugLoc(Inst->getDebugLoc());
+  return NewI;
 }
 
 // A helper function for cloneInstructionWithNewAddressSpace. Handles the
@@ -1109,8 +1107,7 @@ unsigned InferAddressSpacesImpl::joinAddressSpaces(unsigned AS1,
   if (AS2 == UninitializedAddressSpace)
     return AS1;
 
-  // The join of two different specific address spaces is flat.
-  return (AS1 == AS2) ? AS1 : FlatAddrSpace;
+  return TTI->getAddressSpaceJoin(AS1, AS2);
 }
 
 bool InferAddressSpacesImpl::run(Function &CurFn) {
@@ -1589,9 +1586,10 @@ bool InferAddressSpacesImpl::rewriteWithNewAddressSpaces(
 
     unsigned OperandNo = PoisonUse->getOperandNo();
     assert(isa<PoisonValue>(NewV->getOperand(OperandNo)));
-    WeakTrackingVH NewOp = ValueWithNewAddrSpace.lookup(PoisonUse->get());
-    assert(NewOp &&
-           "poison replacements in ValueWithNewAddrSpace shouldn't be null");
+    unsigned NewAS =
+        NewV->getOperand(OperandNo)->getType()->getPointerAddressSpace();
+    Value *NewOp = operandWithNewAddressSpaceOrCreatePoison(
+        *PoisonUse, NewAS, ValueWithNewAddrSpace, PredicatedAS, nullptr);
     NewV->setOperand(OperandNo, NewOp);
   }
 
diff --git a/llvm/test/CodeGen/NVPTX/infer-address-space-join.ll b/llvm/test/CodeGen/NVPTX/infer-address-space-join.ll
new file mode 100644
index 0000000000000..b0cea461f30a7
--- /dev/null
+++ b/llvm/test/CodeGen/NVPTX/infer-address-space-join.ll
@@ -0,0 +1,29 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 6
+; RUN: llc -mtriple=nvptx64-nvidia-cuda -mcpu=sm_90 -mattr=+ptx78 %s -o - | FileCheck %s
+
+define i32 @mixed_select(ptr addrspace(3) %local, ptr addrspace(7) %cluster, i1 %cond) {
+; CHECK-LABEL: mixed_select(
+; CHECK:       {
+; CHECK-NEXT:    .reg .pred %p<2>;
+; CHECK-NEXT:    .reg .b16 %rs<3>;
+; CHECK-NEXT:    .reg .b32 %r<2>;
+; CHECK-NEXT:    .reg .b64 %rd<6>;
+; CHECK-EMPTY:
+; CHECK-NEXT:  // %bb.0:
+; CHECK-NEXT:    ld.param.b8 %rs1, [mixed_select_param_2];
+; CHECK-NEXT:    and.b16 %rs2, %rs1, 1;
+; CHECK-NEXT:    setp.ne.b16 %p1, %rs2, 0;
+; CHECK-NEXT:    ld.param.b64 %rd1, [mixed_select_param_0];
+; CHECK-NEXT:    cvta.shared.u64 %rd2, %rd1;
+; CHECK-NEXT:    cvta.to.shared::cluster.u64 %rd3, %rd2;
+; CHECK-NEXT:    ld.param.b64 %rd4, [mixed_select_param_1];
+; CHECK-NEXT:    selp.b64 %rd5, %rd3, %rd4, %p1;
+; CHECK-NEXT:    ld.shared::cluster.b32 %r1, [%rd5];
+; CHECK-NEXT:    st.param.b32 [func_retval0], %r1;
+; CHECK-NEXT:    ret;
+  %lg = addrspacecast ptr addrspace(3) %local to ptr
+  %cg = addrspacecast ptr addrspace(7) %cluster to ptr
+  %p = select i1 %cond, ptr %lg, ptr %cg
+  %value = load i32, ptr %p, align 4
+  ret i32 %value
+}
diff --git a/llvm/test/Transforms/InferAddressSpaces/NVPTX/address-space-join.ll b/llvm/test/Transforms/InferAddressSpaces/NVPTX/address-space-join.ll
new file mode 100644
index 0000000000000..8a1f79cc29b3e
--- /dev/null
+++ b/llvm/test/Transforms/InferAddressSpaces/NVPTX/address-space-join.ll
@@ -0,0 +1,223 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --version 6
+; RUN: opt -S -mtriple=nvptx64-nvidia-cuda -passes=infer-address-spaces -verify-each %s | FileCheck %s
+; RUN: opt -S -mtriple=nvptx64-nvidia-cuda -passes='infer-address-spaces,infer-address-spaces' -verify-each %s | FileCheck %s
+
+; Generic pointers with AS3 and AS7 origins have the common address space AS7.
+define i32 @mixed_select(ptr addrspace(3) %local, ptr addrspace(7) %cluster, i1 %cond) {
+; CHECK-LABEL: define i32 @mixed_select(
+; CHECK-SAME: ptr addrspace(3) [[LOCAL:%.*]], ptr addrspace(7) [[CLUSTER:%.*]], i1 [[COND:%.*]]) {
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr addrspace(3) [[LOCAL]] to ptr addrspace(7)
+; CHECK-NEXT:    [[P:%.*]] = select i1 [[COND]], ptr addrspace(7) [[TMP1]], ptr addrspace(7) [[CLUSTER]]
+; CHECK-NEXT:    [[VALUE:%.*]] = load i32, ptr addrspace(7) [[P]], align 4
+; CHECK-NEXT:    ret i32 [[VALUE]]
+;
+  %lg = addrspacecast ptr addrspace(3) %local to ptr
+  %cg = addrspacecast ptr addrspace(7) %cluster to ptr
+  %p = select i1 %cond, ptr %lg, ptr %cg
+  %value = load i32, ptr %p, align 4
+  ret i32 %value
+}
+
+define i32 @mixed_phi(ptr addrspace(3) %local, ptr addrspace(7) %cluster, i1 %cond) {
+; CHECK-LABEL: define i32 @mixed_phi(
+; CHECK-SAME: ptr addrspace(3) [[LOCAL:%.*]], ptr addrspace(7) [[CLUSTER:%.*]], i1 [[COND:%.*]]) {
+; CHECK-NEXT:  [[ENTRY:.*:]]
+; CHECK-NEXT:    br i1 [[COND]], label %[[LEFT:.*]], label %[[RIGHT:.*]]
+; CHECK:       [[LEFT]]:
+; CHECK-NEXT:    [[TMP0:%.*]] = addrspacecast ptr addrspace(3) [[LOCAL]] to ptr addrspace(7)
+; CHECK-NEXT:    br label %[[EXIT:.*]]
+; CHECK:       [[RIGHT]]:
+; CHECK-NEXT:    br label %[[EXIT]]
+; CHECK:       [[EXIT]]:
+; CHECK-NEXT:    [[P:%.*]] = phi ptr addrspace(7) [ [[TMP0]], %[[LEFT]] ], [ [[CLUSTER]], %[[RIGHT]] ]
+; CHECK-NEXT:    [[VALUE:%.*]] = load i32, ptr addrspace(7) [[P]], align 4
+; CHECK-NEXT:    ret i32 [[VALUE]]
+;
+entry:
+  br i1 %cond, label %left, label %right
+left:
+  %lg = addrspacecast ptr addrspace(3) %local to ptr
+  br label %exit
+right:
+  %cg = addrspacecast ptr addrspace(7) %cluster to ptr
+  br label %exit
+exit:
+  %p = phi ptr [ %lg, %left ], [ %cg, %right ]
+  %value = load i32, ptr %p, align 4
+  ret i32 %value
+}
+
+; Inference must converge as the loop's initial AS3 result joins with AS7.
+define i32 @mixed_loop(ptr addrspace(3) %local, ptr addrspace(7) %cluster, i1 %choose, i1 %again) {
+; CHECK-LABEL: define i32 @mixed_loop(
+; CHECK-SAME: ptr addrspace(3) [[LOCAL:%.*]], ptr addrspace(7) [[CLUSTER:%.*]], i1 [[CHOOSE:%.*]], i1 [[AGAIN:%.*]]) {
+; CHECK-NEXT:  [[ENTRY:.*]]:
+; CHECK-NEXT:    [[TMP0:%.*]] = addrspacecast ptr addrspace(3) [[LOCAL]] to ptr addrspace(7)
+; CHECK-NEXT:    br label %[[LOOP:.*]]
+; CHECK:       [[LOOP]]:
+; CHECK-NEXT:    [[P:%.*]] = phi ptr addrspace(7) [ [[TMP0]], %[[ENTRY]] ], [ [[NEXT:%.*]], %[[LOOP]] ]
+; CHECK-NEXT:    [[NEXT]] = select i1 [[CHOOSE]], ptr addrspace(7) [[P]], ptr addrspace(7) [[CLUSTER]]
+; CHECK-NEXT:    [[VALUE:%.*]] = load i32, ptr addrspace(7) [[P]], align 4
+; CHECK-NEXT:    br i1 [[AGAIN]], label %[[LOOP]], label %[[EXIT:.*]]
+; CHECK:       [[EXIT]]:
+; CHECK-NEXT:    ret i32 [[VALUE]]
+;
+entry:
+  %lg = addrspacecast ptr addrspace(3) %local to ptr
+  %cg = addrspacecast ptr addrspace(7) %cluster to ptr
+  br label %loop
+loop:
+  %p = phi ptr [ %lg, %entry ], [ %next, %loop ]
+  %next = select i1 %choose, ptr %p, ptr %cg
+  %value = load i32, ptr %p, align 4
+  br i1 %again, label %loop, label %exit
+exit:
+  ret i32 %value
+}
+
+; A deferred AS3 operand needs a cast when fixing the AS7 PHI's backedge.
+define i32 @mixed_backedge(ptr addrspace(7) %cluster, ptr addrspace(3) %local, i1 %again) {
+; CHECK-LABEL: define i32 @mixed_backedge(
+; CHECK-SAME: ptr addrspace(7) [[CLUSTER:%.*]], ptr addrspace(3) [[LOCAL:%.*]], i1 [[AGAIN:%.*]]) {
+; CHECK-NEXT:  [[ENTRY:.*]]:
+; CHECK-NEXT:    br label %[[LOOP:.*]]
+; CHECK:       [[LOOP]]:
+; CHECK-NEXT:    [[P:%.*]] = phi ptr addrspace(7) [ [[CLUSTER]], %[[ENTRY]] ], [ [[TMP0:%.*]], %[[LOOP]] ]
+; CHECK-NEXT:    [[BACK:%.*]] = getelementptr inbounds i32, ptr addrspace(3) [[LOCAL]], i64 1
+; CHECK-NEXT:    [[TMP0]] = addrspacecast ptr addrspace(3) [[BACK]] to ptr addrspace(7)
+; CHECK-NEXT:    [[BACK_VALUE:%.*]] = load i32, ptr addrspace(3) [[BACK]], align 4
+; CHECK-NEXT:    [[VALUE:%.*]] = load i32, ptr addrspace(7) [[P]], align 4
+; CHECK-NEXT:    [[SUM:%.*]] = add i32 [[BACK_VALUE]], [[VALUE]]
+; CHECK-NEXT:    br i1 [[AGAIN]], label %[[LOOP]], label %[[EXIT:.*]]
+; CHECK:       [[EXIT]]:
+; CHECK-NEXT:    ret i32 [[SUM]]
+;
+entry:
+  %cg = addrspacecast ptr addrspace(7) %cluster to ptr
+  %lg = addrspacecast ptr addrspace(3) %local to ptr
+  br label %loop
+loop:
+  %p = phi ptr [ %cg, %entry ], [ %back, %loop ]
+  %back = getelementptr inbounds i32, ptr %lg, i64 1
+  %back_value = load i32, ptr %back, align 4
+  %value = load i32, ptr %p, align 4
+  %sum = add i32 %back_value, %value
+  br i1 %again, label %loop, label %exit
+exit:
+  ret i32 %sum
+}
+
+; The widening cast must be on the normal path of the invoke.
+declare ptr addrspace(3) @local_source()
+declare i32 @personality(...)
+define i32 @invoke_source_phi(ptr addrspace(7) %cluster, i1 %cond) personality ptr @personality {
+; CHECK-LABEL: define i32 @invoke_source_phi(
+; CHECK-SAME: ptr addrspace(7) [[CLUSTER:%.*]], i1 [[COND:%.*]]) personality ptr @personality {
+; CHECK-NEXT:  [[ENTRY:.*:]]
+; CHECK-NEXT:    [[LOCAL:%.*]] = invoke ptr addrspace(3) @local_source()
+; CHECK-NEXT:            to label %[[LEFT:.*]] unwind label %[[EH:.*]]
+; CHECK:       [[LEFT]]:
+; CHECK-NEXT:    [[TMP0:%.*]] = addrspacecast ptr addrspace(3) [[LOCAL]] to ptr addrspace(7)
+; CHECK-NEXT:    br i1 [[COND]], label %[[EXIT:.*]], label %[[RIGHT:.*]]
+; CHECK:       [[RIGHT]]:
+; CHECK-NEXT:    br label %[[EXIT]]
+; CHECK:       [[EXIT]]:
+; CHECK-NEXT:    [[P:%.*]] = phi ptr addrspace(7) [ [[TMP0]], %[[LEFT]] ], [ [[CLUSTER]], %[[RIGHT]] ]
+; CHECK-NEXT:    [[VALUE:%.*]] = load i32, ptr addrspace(7) [[P]], align 4
+; CHECK-NEXT:    ret i32 [[VALUE]]
+; CHECK:       [[EH]]:
+; CHECK-NEXT:    [[PAD:%.*]] = landingpad { ptr, i32 }
+; CHECK-NEXT:            cleanup
+; CHECK-NEXT:    resume { ptr, i32 } [[PAD]]
+;
+entry:
+  %local = invoke ptr addrspace(3) @local_source() to label %left unwind label %eh
+left:
+  %lg = addrspacecast ptr addrspace(3) %local to ptr
+  br i1 %cond, label %exit, label %right
+right:
+  %cg = addrspacecast ptr addrspace(7) %cluster to ptr
+  br label %exit
+exit:
+  %p = phi ptr [ %lg, %left ], [ %cg, %right ]
+  %value = load i32, ptr %p, align 4
+  ret i32 %value
+eh:
+  %pad = landingpad { ptr, i32 } cleanup
+  resume { ptr, i32 } %pad
+}
+
+; Adding a global origin forces the result back to generic.
+define i32 @mixed_global(ptr addrspace(3) %local, ptr addrspace(7) %cluster, ptr addrspace(1) %global, i1 %cond, i1 %cond2) {
+; CHECK-LABEL: define i32 @mixed_global(
+; CHECK-SAME: ptr addrspace(3) [[LOCAL:%.*]], ptr addrspace(7) [[CLUSTER:%.*]], ptr addrspace(1) [[GLOBAL:%.*]], i1 [[COND:%.*]], i1 [[COND2:%.*]]) {
+; CHECK-NEXT:    [[GG:%.*]] = addrspacecast ptr addrspace(1) [[GLOBAL]] to ptr
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr addrspace(3) [[LOCAL]] to ptr addrspace(7)
+; CHECK-NEXT:    [[SHARED:%.*]] = select i1 [[COND]], ptr addrspace(7) [[TMP1]], ptr addrspace(7) [[CLUSTER]]
+; CHECK-NEXT:    [[TMP2:%.*]] = addrspacecast ptr addrspace(7) [[SHARED]] to ptr
+; CHECK-NEXT:    [[P:%.*]] = select i1 [[COND2]], ptr [[TMP2]], ptr [[GG]]
+; CHECK-NEXT:    [[VALUE:%.*]] = load i32, ptr [[P]], align 4
+; CHECK-NEXT:    ret i32 [[VALUE]]
+;
+  %lg = addrspacecast ptr addrspace(3) %local to ptr
+  %cg = addrspacecast ptr addrspace(7) %cluster to ptr
+  %gg = addrspacecast ptr addrspace(1) %global to ptr
+  %shared = select i1 %cond, ptr %lg, ptr %cg
+  %p = select i1 %cond2, ptr %shared, ptr %gg
+  %value = load i32, ptr %p, align 4
+  ret i32 %value
+}
+
+declare i1 @llvm.nvvm.isspacep.shared(ptr)
+declare void @llvm.assume(i1)
+
+; The assumed operand space can be more specific than the user's join.
+define i32 @predicated_select(ptr %local, ptr addrspace(7) %cluster, i1 %cond) {
+; CHECK-LABEL: define i32 @predicated_select(
+; CHECK-SAME: ptr [[LOCAL:%.*]], ptr addrspace(7) [[CLUSTER:%.*]], i1 [[COND:%.*]]) {
+; CHECK-NEXT:    [[IS_SHARED:%.*]] = call i1 @llvm.nvvm.isspacep.shared(ptr [[LOCAL]])
+; CHECK-NEXT:    call void @llvm.assume(i1 [[IS_SHARED]])
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr [[LOCAL]] to ptr addrspace(7)
+; CHECK-NEXT:    [[P:%.*]] = select i1 [[COND]], ptr addrspace(7) [[CLUSTER]], ptr addrspace(7) [[TMP1]]
+; CHECK-NEXT:    [[VALUE:%.*]] = load i32, ptr addrspace(7) [[P]], align 4
+; CHECK-NEXT:    ret i32 [[VALUE]]
+;
+  %is_shared = call i1 @llvm.nvvm.isspacep.shared(ptr %local)
+  call void @llvm.assume(i1 %is_shared)
+  %cg = addrspacecast ptr addrspace(7) %cluster to ptr
+  %p = select i1 %cond, ptr %cg, ptr %local
+  %value = load i32, ptr %p, align 4
+  ret i32 %value
+}
+
+; Constant operands must use the inferred common space too.
+define ptr @mixed_null(ptr addrspace(3) %local, ptr addrspace(7) %cluster, i1 %cond, i1 %cond2) {
+; CHECK-LABEL: define ptr @mixed_null(
+; CHECK-SAME: ptr addrspace(3) [[LOCAL:%.*]], ptr addrspace(7) [[CLUSTER:%.*]], i1 [[COND:%.*]], i1 [[COND2:%.*]]) {
+; CHECK-NEXT:    [[TMP1:%.*]] = addrspacecast ptr addrspace(3) [[LOCAL]] to ptr addrspace(7)
+; CHECK-NEXT:    [[SHARED:%.*]] = select i1 [[COND]], ptr addrspace(7) [[TMP1]], ptr addrspace(7) [[CLUSTER]]
+; CHECK-NEXT:    [[P:%.*]] = select i1 [[COND2]], ptr addrspace(7) [[SHARED]], ptr addrspace(7) addrspacecast (ptr null to ptr addrspace(7))
+; CHECK-NEXT:    [[TMP2:%.*]] = addrspacecast ptr addrspace(7) [[P]] to ptr
+; CHECK-NEXT:    ret ptr [[TMP2]]
+;
+  %lg = addrspacecast ptr addrspace(3) %local to ptr
+  %cg = addrspacecast ptr addrspace(7) %cluster to ptr
+  %shared = select i1 %cond, ptr %lg, ptr %cg
+  %p = select i1 %cond2, ptr %shared, ptr null
+  ret ptr %p
+}
+
+; Only generic expressions are analyzed; existing AS7 expressions retain AS7.
+define i32 @cluster_expression(ptr addrspace(3) %local) {
+; CHECK-LABEL: define i32 @cluster_expression(
+; CHECK-SAME: ptr addrspace(3) [[LOCAL:%.*]]) {
+; CHECK-NEXT:    [[CLUSTER:%.*]] = addrspacecast ptr addrspace(3) [[LOCAL]] to ptr addrspace(7)
+; CHECK-NEXT:    [[P:%.*]] = getelementptr inbounds i32, ptr addrspace(7) [[CLUSTER]], i64 1
+; CHECK-NEXT:    [[VALUE:%.*]] = load i32, ptr addrspace(7) [[P]], align 4
+; CHECK-NEXT:    ret i32 [[VALUE]]
+;
+  %cluster = addrspacecast ptr addrspace(3) %local to ptr addrspace(7)
+  %p = getelementptr inbounds i32, ptr addrspace(7) %cluster, i64 1
+  %value = load i32, ptr addrspace(7) %p, align 4
+  ret i32 %value
+}

>From c070a6dbbb14449e755079cfb037ce5460efa83e Mon Sep 17 00:00:00 2001
From: Alex Maclean <amaclean at nvidia.com>
Date: Mon, 21 Sep 2026 11:51:52 -0700
Subject: [PATCH 2/2] [InferAddressSpaces] Add a TTI hook for address space
 joins

---
 llvm/include/llvm/Analysis/TargetTransformInfo.h             | 5 ++++-
 llvm/include/llvm/Analysis/TargetTransformInfoImpl.h         | 2 +-
 llvm/lib/Analysis/TargetTransformInfo.cpp                    | 1 +
 llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h             | 2 --
 llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp            | 3 +++
 .../InferAddressSpaces/NVPTX/address-space-join.ll           | 3 +--
 6 files changed, 10 insertions(+), 6 deletions(-)

diff --git a/llvm/include/llvm/Analysis/TargetTransformInfo.h b/llvm/include/llvm/Analysis/TargetTransformInfo.h
index c585c250062b3..86daffa4fba9e 100644
--- a/llvm/include/llvm/Analysis/TargetTransformInfo.h
+++ b/llvm/include/llvm/Analysis/TargetTransformInfo.h
@@ -566,8 +566,11 @@ class TargetTransformInfo {
   LLVM_ABI unsigned getFlatAddressSpace() const;
 
   /// Return the most specific common address space containing AS1 and AS2.
+  /// AS1 and AS2 must be distinct, and pointers from both spaces must be
+  /// convertible to the target's flat address space with addrspacecast.
   /// Pointers from either input space must be convertible to the result with
-  /// addrspacecast.
+  /// addrspacecast. Return getFlatAddressSpace() if no more specific common
+  /// address space is available.
   LLVM_ABI unsigned getAddressSpaceJoin(unsigned AS1, unsigned AS2) const;
 
   /// Return any intrinsic address operand indexes which may be rewritten if
diff --git a/llvm/include/llvm/Analysis/TargetTransformInfoImpl.h b/llvm/include/llvm/Analysis/TargetTransformInfoImpl.h
index 8f0d9abc195f9..593be031218c4 100644
--- a/llvm/include/llvm/Analysis/TargetTransformInfoImpl.h
+++ b/llvm/include/llvm/Analysis/TargetTransformInfoImpl.h
@@ -147,7 +147,7 @@ class LLVM_ABI TargetTransformInfoImplBase {
   virtual unsigned getFlatAddressSpace() const { return -1; }
 
   virtual unsigned getAddressSpaceJoin(unsigned AS1, unsigned AS2) const {
-    return AS1 == AS2 ? AS1 : getFlatAddressSpace();
+    return getFlatAddressSpace();
   }
 
   virtual bool collectFlatAddressOperands(SmallVectorImpl<int> &OpIndexes,
diff --git a/llvm/lib/Analysis/TargetTransformInfo.cpp b/llvm/lib/Analysis/TargetTransformInfo.cpp
index 83f2fe46b5765..269dea28b3dbe 100644
--- a/llvm/lib/Analysis/TargetTransformInfo.cpp
+++ b/llvm/lib/Analysis/TargetTransformInfo.cpp
@@ -318,6 +318,7 @@ unsigned TargetTransformInfo::getFlatAddressSpace() const {
 
 unsigned TargetTransformInfo::getAddressSpaceJoin(unsigned AS1,
                                                   unsigned AS2) const {
+  assert(AS1 != AS2 && "Expected distinct address spaces");
   return TTIImpl->getAddressSpaceJoin(AS1, AS2);
 }
 
diff --git a/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h b/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h
index 0c5d8c588d3d4..c6313a5a0fb51 100644
--- a/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h
+++ b/llvm/lib/Target/NVPTX/NVPTXTargetTransformInfo.h
@@ -55,8 +55,6 @@ class NVPTXTTIImpl final : public BasicTTIImplBase<NVPTXTTIImpl> {
   }
 
   unsigned getAddressSpaceJoin(unsigned AS1, unsigned AS2) const override {
-    if (AS1 == AS2)
-      return AS1;
     if ((AS1 == AddressSpace::ADDRESS_SPACE_SHARED &&
          AS2 == AddressSpace::ADDRESS_SPACE_SHARED_CLUSTER) ||
         (AS2 == AddressSpace::ADDRESS_SPACE_SHARED &&
diff --git a/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp b/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp
index d2596bc7c2f75..6000e6d602f25 100644
--- a/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp
+++ b/llvm/lib/Transforms/Scalar/InferAddressSpaces.cpp
@@ -1099,6 +1099,9 @@ Value *InferAddressSpacesImpl::cloneValueWithNewAddressSpace(
 // comments).
 unsigned InferAddressSpacesImpl::joinAddressSpaces(unsigned AS1,
                                                    unsigned AS2) const {
+  if (AS1 == AS2)
+    return AS1;
+
   if (AS1 == FlatAddrSpace || AS2 == FlatAddrSpace)
     return FlatAddrSpace;
 
diff --git a/llvm/test/Transforms/InferAddressSpaces/NVPTX/address-space-join.ll b/llvm/test/Transforms/InferAddressSpaces/NVPTX/address-space-join.ll
index 8a1f79cc29b3e..512b6867a2c84 100644
--- a/llvm/test/Transforms/InferAddressSpaces/NVPTX/address-space-join.ll
+++ b/llvm/test/Transforms/InferAddressSpaces/NVPTX/address-space-join.ll
@@ -1,6 +1,5 @@
 ; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --version 6
-; RUN: opt -S -mtriple=nvptx64-nvidia-cuda -passes=infer-address-spaces -verify-each %s | FileCheck %s
-; RUN: opt -S -mtriple=nvptx64-nvidia-cuda -passes='infer-address-spaces,infer-address-spaces' -verify-each %s | FileCheck %s
+; RUN: opt -S -mtriple=nvptx64-nvidia-cuda -passes=infer-address-spaces %s | FileCheck %s
 
 ; Generic pointers with AS3 and AS7 origins have the common address space AS7.
 define i32 @mixed_select(ptr addrspace(3) %local, ptr addrspace(7) %cluster, i1 %cond) {



More information about the llvm-commits mailing list