https://github.com/RiverDave created 
https://github.com/llvm/llvm-project/pull/226649

Opened to address a portion of 
https://github.com/llvm/llvm-project/issues/226629


In CUDA, `__shared__ int sh` has type `int` but lives in AS 3. Classic codegen 
casts the address to the declared type's AS where it's formed, so users just 
see a generic pointer. We weren't doing that, so things like `return &sh;` 
bitcast the slot instead, and NVPTX never got a `cvta.shared`. This patch does 
the same cast in `getAddrOfGlobalVar` and wherever static locals are fetched.

This also drops the comment claiming lowering would emit the cast for us. 
That's only true for OpenCL, where the declared type already carries the AS. 
LowerToLLVM never inserts casts on its own.

Assisted-by: Claude / Opus 5.5

>From 7515b8d2fec7c97e9ef3b6cd096f420082aaeb46 Mon Sep 17 00:00:00 2001
From: David Rivera <[email protected]>
Date: Sat, 26 Sep 2026 00:18:39 -0400
Subject: [PATCH] [CIR] Cast global addresses to their declared address space

A global can live in a different address space than its declared type,
e.g. a CUDA __device__ or __shared__ variable. Classic CodeGen casts the
address once where it is formed (GetOrCreateLLVMGlobal and
getOrCreateStaticVarDecl), so every user sees a pointer in the declared
(generic) address space. CIR kept the global's address space on the value
and relied on each user to cast it. Users that did not, such as returning
or storing the address, bitcast the destination slot instead,
reinterpreting the pointer:

  __shared__ int sh;
  __device__ int *f() { return &sh; }

stored the raw shared-window address as a generic pointer, with no
cvta.shared on NVPTX. After #226455 the same applies to function-local
__shared__ variables.

Do the same check as classic CodeGen where CIR first has the address as a
value: in getAddrOfGlobalVar, and where static locals are fetched with
get_global. Unlike an LLVM constant cast, a cir.cast is emitted per
function, so it cannot be cached with the static local. getAddrOfGlobal
now returns the GlobalOp directly instead of the defining op of the
(possibly cast) address.
---
 clang/lib/CIR/CodeGen/CIRGenDecl.cpp          | 12 +--
 clang/lib/CIR/CodeGen/CIRGenExpr.cpp          |  3 +-
 clang/lib/CIR/CodeGen/CIRGenModule.cpp        | 21 +++-
 clang/lib/CIR/CodeGen/CIRGenModule.h          |  4 +
 .../CIR/CodeGen/amdgpu-array-addrspace.cpp    | 37 +++++---
 clang/test/CIR/CodeGenCUDA/address-spaces.cu  |  5 +-
 .../CIR/CodeGenCUDA/global-addrspace-cast.cu  | 95 +++++++++++++++++++
 7 files changed, 151 insertions(+), 26 deletions(-)
 create mode 100644 clang/test/CIR/CodeGenCUDA/global-addrspace-cast.cu

diff --git a/clang/lib/CIR/CodeGen/CIRGenDecl.cpp 
b/clang/lib/CIR/CodeGen/CIRGenDecl.cpp
index 451f6f8af7fd11..6d6627c12d8d36 100644
--- a/clang/lib/CIR/CodeGen/CIRGenDecl.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenDecl.cpp
@@ -558,15 +558,8 @@ CIRGenModule::getOrCreateStaticVarDecl(const VarDecl &d,
 
   setGVProperties(gv, &d);
 
-  // OG checks if the expected address space, denoted by the type, is the
-  // same as the actual address space indicated by attributes. If they aren't
-  // the same, an addrspacecast is emitted when this variable is accessed.
-  // In CIR however, cir.get_global already carries that information in
-  // !cir.ptr type - if this global is in OpenCL local address space, then its
-  // type would be !cir.ptr<..., addrspace(offload_local)>. Therefore we don't
-  // need an explicit address space cast in CIR: they will get emitted when
-  // lowering to LLVM IR.
-
+  // The global may live in a different address space than the declared type.
+  // Users of the address cast it through castGlobalToDeclAddrSpace.
   setStaticLocalDeclAddress(&d, gv);
 
   // Ensure that the static local gets initialized by making sure the parent
@@ -807,6 +800,7 @@ void CIRGenFunction::emitStaticVarDecl(const VarDecl &d,
   // RAUW's the GV uses of this constant will be invalid.
   mlir::Value castedAddr =
       builder.createBitcast(getAddrOp.getAddr(), expectedType);
+  castedAddr = cgm.castGlobalToDeclAddrSpace(castedAddr, d);
   localDeclMap.find(&d)->second = Address(castedAddr, elemTy, alignment);
   cgm.setStaticLocalDeclAddress(&d, var);
 
diff --git a/clang/lib/CIR/CodeGen/CIRGenExpr.cpp 
b/clang/lib/CIR/CodeGen/CIRGenExpr.cpp
index 7688fcc3cc337c..bf526237e8928a 100644
--- a/clang/lib/CIR/CodeGen/CIRGenExpr.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenExpr.cpp
@@ -1138,7 +1138,8 @@ LValue CIRGenFunction::emitDeclRefLValue(const 
DeclRefExpr *e) {
       auto getGlob = getGlobVal.getDefiningOp<cir::GetGlobalOp>();
       getGlob.setStaticLocal(var.getStaticLocalGuard().has_value());
       getGlob.setTls(vd->getTLSKind() != VarDecl::TLS_None);
-      addr = Address(getGlob, convertTypeForMem(vd->getType()),
+      addr = Address(cgm.castGlobalToDeclAddrSpace(getGlob, *vd),
+                     convertTypeForMem(vd->getType()),
                      getContext().getDeclAlign(vd));
     } else {
       llvm_unreachable("DeclRefExpr for Decl not entered in localDeclMap?");
diff --git a/clang/lib/CIR/CodeGen/CIRGenModule.cpp 
b/clang/lib/CIR/CodeGen/CIRGenModule.cpp
index adffa7dfe29694..0debc67ccaa502 100644
--- a/clang/lib/CIR/CodeGen/CIRGenModule.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenModule.cpp
@@ -421,8 +421,8 @@ CIRGenModule::getAddrOfGlobal(GlobalDecl gd, 
ForDefinition_t isForDefinition) {
                              isForDefinition);
   }
 
-  return getAddrOfGlobalVar(cast<VarDecl>(d), /*ty=*/nullptr, isForDefinition)
-      .getDefiningOp();
+  return getOrCreateCIRGlobal(cast<VarDecl>(d), /*ty=*/nullptr,
+                              isForDefinition);
 }
 
 void CIRGenModule::emitGlobalDecl(const clang::GlobalDecl &d) {
@@ -1444,10 +1444,25 @@ mlir::Value CIRGenModule::getAddrOfGlobalVar(const 
VarDecl *d, mlir::Type ty,
   bool tlsAccess = d->getTLSKind() != VarDecl::TLS_None;
   cir::GlobalOp g = getOrCreateCIRGlobal(d, ty, isForDefinition);
   mlir::Type ptrTy = builder.getPointerTo(g.getSymType(), 
g.getAddrSpaceAttr());
-  return cir::GetGlobalOp::create(
+  mlir::Value addr = cir::GetGlobalOp::create(
       builder, getLoc(d->getSourceRange()), ptrTy, g.getSymNameAttr(),
       tlsAccess,
       /*static_local=*/g.getStaticLocalGuard().has_value());
+  return castGlobalToDeclAddrSpace(addr, *d);
+}
+
+mlir::Value CIRGenModule::castGlobalToDeclAddrSpace(mlir::Value addr,
+                                                    const VarDecl &vd) {
+  // A global may live in a different address space than its declared type,
+  // e.g. a CUDA __shared__ variable. Like classic CodeGen, cast once where
+  // the address is formed so every user sees the declared type.
+  auto ptrTy = mlir::cast<cir::PointerType>(addr.getType());
+  mlir::ptr::MemorySpaceAttrInterface declAS =
+      getTypes().getPointerAddressSpace(vd.getType());
+  if (ptrTy.getAddrSpace() == declAS)
+    return addr;
+  return builder.createAddrSpaceCast(
+      addr, builder.getPointerTo(ptrTy.getPointee(), declAS));
 }
 
 cir::GlobalViewAttr CIRGenModule::getAddrOfGlobalVarAttr(const VarDecl *d) {
diff --git a/clang/lib/CIR/CodeGen/CIRGenModule.h 
b/clang/lib/CIR/CodeGen/CIRGenModule.h
index 3fb95f346536dc..83ef80090faefa 100644
--- a/clang/lib/CIR/CodeGen/CIRGenModule.h
+++ b/clang/lib/CIR/CodeGen/CIRGenModule.h
@@ -339,6 +339,10 @@ class CIRGenModule : public CIRGenTypeCache {
   getAddrOfGlobalVar(const VarDecl *d, mlir::Type ty = {},
                      ForDefinition_t isForDefinition = NotForDefinition);
 
+  /// Cast \p addr, the address of the global \p vd, to the address space of
+  /// the declared type of \p vd if they differ.
+  mlir::Value castGlobalToDeclAddrSpace(mlir::Value addr, const VarDecl &vd);
+
   /// Get or create a thunk function with the given name and type.
   cir::FuncOp getAddrOfThunk(StringRef name, mlir::Type fnTy, GlobalDecl gd);
 
diff --git a/clang/test/CIR/CodeGen/amdgpu-array-addrspace.cpp 
b/clang/test/CIR/CodeGen/amdgpu-array-addrspace.cpp
index 2ff665edf2dc10..47d050d6e0a5d8 100644
--- a/clang/test/CIR/CodeGen/amdgpu-array-addrspace.cpp
+++ b/clang/test/CIR/CodeGen/amdgpu-array-addrspace.cpp
@@ -10,16 +10,30 @@
 
 int globalArr[10] = {0};
 
+// A dynamic initializer stores through the flat address of the global, as in
+// classic CodeGen.
+
+int f();
+int dyn = f();
+
+// CIR:       cir.func {{.*}}@__cxx_global_var_init
+// CIR:         %[[DYN:.*]] = cir.get_global @dyn : !cir.ptr<!s32i, 
target_address_space(1)>
+// CIR-NEXT:    %[[FLAT:.*]] = cir.cast address_space %[[DYN]] : 
!cir.ptr<!s32i, target_address_space(1)> -> !cir.ptr<!s32i>
+// CIR:         cir.store align(4) %{{.*}}, %[[FLAT]] : !s32i, !cir.ptr<!s32i>
+
+// LLVM:        store i32 %{{.*}}, ptr addrspacecast (ptr addrspace(1) @dyn to 
ptr), align 4
+// OGCG:        store i32 %{{.*}}, ptr addrspacecast (ptr addrspace(1) @dyn to 
ptr), align 4
+
 void takes_ptr(int *p);
 
-// The array_to_ptrdecay cast must preserve the address space of the base
-// pointer, followed by an address_space cast.
+// The address of the global is cast to the declared (flat) address space
+// before the array decays.
 
 // CIR-LABEL: cir.func{{.*}} @_Z17pass_global_arrayv()
 // CIR:         %[[ARR:.*]] = cir.get_global @globalArr : 
!cir.ptr<!cir.array<!s32i x 10>, target_address_space(1)>
-// CIR-NEXT:    %[[DECAY:.*]] = cir.cast array_to_ptrdecay %[[ARR]] : 
!cir.ptr<!cir.array<!s32i x 10>, target_address_space(1)> -> !cir.ptr<!s32i, 
target_address_space(1)>
-// CIR-NEXT:    %[[FLAT:.*]] = cir.cast address_space %[[DECAY]] : 
!cir.ptr<!s32i, target_address_space(1)> -> !cir.ptr<!s32i>
-// CIR-NEXT:    cir.call @_Z9takes_ptrPi(%[[FLAT]])
+// CIR-NEXT:    %[[FLAT:.*]] = cir.cast address_space %[[ARR]] : 
!cir.ptr<!cir.array<!s32i x 10>, target_address_space(1)> -> 
!cir.ptr<!cir.array<!s32i x 10>>
+// CIR-NEXT:    %[[DECAY:.*]] = cir.cast array_to_ptrdecay %[[FLAT]] : 
!cir.ptr<!cir.array<!s32i x 10>> -> !cir.ptr<!s32i>
+// CIR-NEXT:    cir.call @_Z9takes_ptrPi(%[[DECAY]])
 
 // LLVM-LABEL: define{{.*}} void @_Z17pass_global_arrayv()
 // LLVM:         call void @_Z9takes_ptrPi(ptr noundef addrspacecast (ptr 
addrspace(1) @globalArr to ptr))
@@ -30,17 +44,17 @@ void pass_global_array() {
   takes_ptr(globalArr);
 }
 
-// The get_element op must preserve the address space of the base pointer
-// so that the subsequent load uses the correct address space.
+// Indexing goes through the flat address, as in classic CodeGen.
 
 // CIR-LABEL: cir.func{{.*}} @_Z18index_global_arrayi
 // CIR:         %[[ARR:.*]] = cir.get_global @globalArr : 
!cir.ptr<!cir.array<!s32i x 10>, target_address_space(1)>
-// CIR-NEXT:    %[[ELEM:.*]] = cir.get_element %[[ARR]][%{{.*}} : !s64i] : 
!cir.ptr<!cir.array<!s32i x 10>, target_address_space(1)> -> !cir.ptr<!s32i, 
target_address_space(1)>
-// CIR-NEXT:    %{{.*}} = cir.load align(4) %[[ELEM]] : !cir.ptr<!s32i, 
target_address_space(1)>, !s32i
+// CIR-NEXT:    %[[FLAT:.*]] = cir.cast address_space %[[ARR]] : 
!cir.ptr<!cir.array<!s32i x 10>, target_address_space(1)> -> 
!cir.ptr<!cir.array<!s32i x 10>>
+// CIR-NEXT:    %[[ELEM:.*]] = cir.get_element %[[FLAT]][%{{.*}} : !s64i] : 
!cir.ptr<!cir.array<!s32i x 10>> -> !cir.ptr<!s32i>
+// CIR-NEXT:    %{{.*}} = cir.load align(4) %[[ELEM]] : !cir.ptr<!s32i>, !s32i
 
 // LLVM-LABEL: define{{.*}} i32 @_Z18index_global_arrayi
-// LLVM:         %[[GEP:.*]] = getelementptr [10 x i32], ptr addrspace(1) 
@globalArr, i32 0, i64 %{{.*}}
-// LLVM-NEXT:    %{{.*}} = load i32, ptr addrspace(1) %[[GEP]], align 4
+// LLVM:         %[[GEP:.*]] = getelementptr [10 x i32], ptr addrspacecast 
(ptr addrspace(1) @globalArr to ptr), i32 0, i64 %{{.*}}
+// LLVM-NEXT:    %{{.*}} = load i32, ptr %[[GEP]], align 4
 
 // OGCG-LABEL: define{{.*}} i32 @_Z18index_global_arrayi
 // OGCG:         getelementptr inbounds [10 x i32], ptr addrspacecast (ptr 
addrspace(1) @globalArr to ptr)
@@ -48,3 +62,4 @@ void pass_global_array() {
 int index_global_array(int i) {
   return globalArr[i];
 }
+
diff --git a/clang/test/CIR/CodeGenCUDA/address-spaces.cu 
b/clang/test/CIR/CodeGenCUDA/address-spaces.cu
index 6637100fd76c90..c6d7943c79a60b 100644
--- a/clang/test/CIR/CodeGenCUDA/address-spaces.cu
+++ b/clang/test/CIR/CodeGenCUDA/address-spaces.cu
@@ -162,15 +162,16 @@ __global__ void fn() {
 // CIR-DEVICE:   %[[ZERO:.*]] = cir.const #cir.int<0> : !s32i
 // CIR-DEVICE:   cir.store {{.*}}%[[ZERO]], %[[ALLOCA]] : !s32i, 
!cir.ptr<!s32i>
 // CIR-DEVICE:   %[[J:.*]] = cir.get_global @_ZZ2fnvE1j : !cir.ptr<!s32i, 
target_address_space(3)>
+// CIR-DEVICE:   %[[J_CAST:.*]] = cir.cast address_space %[[J]] : 
!cir.ptr<!s32i, target_address_space(3)> -> !cir.ptr<!s32i>
 // CIR-DEVICE:   %[[VAL:.*]] = cir.load {{.*}}%[[ALLOCA]] : !cir.ptr<!s32i>, 
!s32i
-// CIR-DEVICE:   cir.store {{.*}}%[[VAL]], %[[J]] : !s32i, !cir.ptr<!s32i, 
target_address_space(3)>
+// CIR-DEVICE:   cir.store {{.*}}%[[VAL]], %[[J_CAST]] : !s32i, !cir.ptr<!s32i>
 // CIR-DEVICE:   cir.return
 
 // LLVM-DEVICE: define dso_local ptx_kernel void @_Z2fnv()
 // LLVM-DEVICE:   %[[ALLOCA:.*]] = alloca i32, align 4
 // LLVM-DEVICE:   store i32 0, ptr %[[ALLOCA]], align 4
 // LLVM-DEVICE:   %[[VAL:.*]] = load i32, ptr %[[ALLOCA]], align 4
-// LLVM-DEVICE:   store i32 %[[VAL]], ptr addrspace(3) @_ZZ2fnvE1j, align 4
+// LLVM-DEVICE:   store i32 %[[VAL]], ptr addrspacecast (ptr addrspace(3) 
@_ZZ2fnvE1j to ptr), align 4
 // LLVM-DEVICE:   ret void
 
 // OGCG-DEVICE: define dso_local ptx_kernel void @_Z2fnv()
diff --git a/clang/test/CIR/CodeGenCUDA/global-addrspace-cast.cu 
b/clang/test/CIR/CodeGenCUDA/global-addrspace-cast.cu
new file mode 100644
index 00000000000000..7796a73869ef1d
--- /dev/null
+++ b/clang/test/CIR/CodeGenCUDA/global-addrspace-cast.cu
@@ -0,0 +1,95 @@
+#include "Inputs/cuda.h"
+
+// RUN: %clang_cc1 -triple nvptx64-nvidia-cuda -x cuda -fcuda-is-device \
+// RUN:   -fclangir -emit-cir %s -o %t.cir
+// RUN: FileCheck --check-prefix=CIR --input-file=%t.cir %s
+// RUN: %clang_cc1 -triple nvptx64-nvidia-cuda -x cuda -fcuda-is-device \
+// RUN:   -fclangir -emit-llvm %s -o %t-cir.ll
+// RUN: FileCheck --check-prefix=LLVM --input-file=%t-cir.ll %s
+// RUN: %clang_cc1 -triple nvptx64-nvidia-cuda -x cuda -fcuda-is-device \
+// RUN:   -emit-llvm %s -o %t.ll
+// RUN: FileCheck --check-prefix=OGCG --input-file=%t.ll %s
+
+// RUN: %clang_cc1 -triple amdgcn-amd-amdhsa -x hip -fcuda-is-device \
+// RUN:   -fclangir -emit-llvm %s -o %t-cir-amdgcn.ll
+// RUN: FileCheck --check-prefix=LLVM --input-file=%t-cir-amdgcn.ll %s
+// RUN: %clang_cc1 -triple amdgcn-amd-amdhsa -x hip -fcuda-is-device \
+// RUN:   -emit-llvm %s -o %t-amdgcn.ll
+// RUN: FileCheck --check-prefix=OGCG --input-file=%t-amdgcn.ll %s
+
+// The address of a global whose address space differs from its declared type
+// is cast to the declared (generic) address space where it is formed.
+
+__device__ int g;
+__device__ int arr[4];
+__shared__ int sh;
+extern __shared__ int dyn[];
+
+__device__ int *addr_of_global() { return &g; }
+
+// CIR-LABEL: cir.func {{.*}}@_Z14addr_of_globalv
+// CIR: %[[G:.*]] = cir.get_global @g : !cir.ptr<{{.*}}, 
target_address_space(1)>
+// CIR: cir.cast address_space %[[G]] : !cir.ptr<{{.*}}, 
target_address_space(1)> -> !cir.ptr<{{.*}}>
+
+// LLVM-LABEL: @_Z14addr_of_globalv
+// LLVM: store ptr addrspacecast (ptr addrspace(1) @g to ptr)
+// OGCG-LABEL: @_Z14addr_of_globalv
+// OGCG: ret ptr addrspacecast (ptr addrspace(1) @g to ptr)
+
+__device__ int *array_decay() { return arr; }
+
+// CIR-LABEL: cir.func {{.*}}@_Z11array_decayv
+// CIR: %[[G:.*]] = cir.get_global @arr : !cir.ptr<{{.*}}, 
target_address_space(1)>
+// CIR: cir.cast address_space %[[G]] : !cir.ptr<{{.*}}, 
target_address_space(1)> -> !cir.ptr<{{.*}}>
+
+// LLVM-LABEL: @_Z11array_decayv
+// LLVM: store ptr addrspacecast (ptr addrspace(1) @arr to ptr)
+// OGCG-LABEL: @_Z11array_decayv
+// OGCG: ret ptr addrspacecast (ptr addrspace(1) @arr to ptr)
+
+__device__ int &bind_ref() { return g; }
+
+// CIR-LABEL: cir.func {{.*}}@_Z8bind_refv
+// CIR: %[[G:.*]] = cir.get_global @g : !cir.ptr<{{.*}}, 
target_address_space(1)>
+// CIR: cir.cast address_space %[[G]] : !cir.ptr<{{.*}}, 
target_address_space(1)> -> !cir.ptr<{{.*}}>
+
+// LLVM-LABEL: @_Z8bind_refv
+// LLVM: store ptr addrspacecast (ptr addrspace(1) @g to ptr)
+// OGCG-LABEL: @_Z8bind_refv
+// OGCG: ret ptr addrspacecast (ptr addrspace(1) @g to ptr)
+
+__device__ int *addr_of_shared() { return &sh; }
+
+// CIR-LABEL: cir.func {{.*}}@_Z14addr_of_sharedv
+// CIR: %[[G:.*]] = cir.get_global @sh : !cir.ptr<{{.*}}, 
target_address_space(3)>
+// CIR: cir.cast address_space %[[G]] : !cir.ptr<{{.*}}, 
target_address_space(3)> -> !cir.ptr<{{.*}}>
+
+// LLVM-LABEL: @_Z14addr_of_sharedv
+// LLVM: store ptr addrspacecast (ptr addrspace(3) @sh to ptr)
+// OGCG-LABEL: @_Z14addr_of_sharedv
+// OGCG: ret ptr addrspacecast (ptr addrspace(3) @sh to ptr)
+
+__device__ int *dynamic_shared() { return dyn; }
+
+// CIR-LABEL: cir.func {{.*}}@_Z14dynamic_sharedv
+// CIR: %[[G:.*]] = cir.get_global @dyn : !cir.ptr<{{.*}}, 
target_address_space(3)>
+// CIR: cir.cast address_space %[[G]] : !cir.ptr<{{.*}}, 
target_address_space(3)> -> !cir.ptr<{{.*}}>
+
+// LLVM-LABEL: @_Z14dynamic_sharedv
+// LLVM: store ptr addrspacecast (ptr addrspace(3) @dyn to ptr)
+// OGCG-LABEL: @_Z14dynamic_sharedv
+// OGCG: ret ptr addrspacecast (ptr addrspace(3) @dyn to ptr)
+
+__device__ int *addr_of_static_shared() {
+  __shared__ int s;
+  return &s;
+}
+
+// CIR-LABEL: cir.func {{.*}}@_Z21addr_of_static_sharedv
+// CIR: %[[G:.*]] = cir.get_global @_ZZ21addr_of_static_sharedvE1s : 
!cir.ptr<{{.*}}, target_address_space(3)>
+// CIR: cir.cast address_space %[[G]] : !cir.ptr<{{.*}}, 
target_address_space(3)> -> !cir.ptr<{{.*}}>
+
+// LLVM-LABEL: @_Z21addr_of_static_sharedv
+// LLVM: store ptr addrspacecast (ptr addrspace(3) 
@_ZZ21addr_of_static_sharedvE1s to ptr)
+// OGCG-LABEL: @_Z21addr_of_static_sharedv
+// OGCG: ret ptr addrspacecast (ptr addrspace(3) 
@_ZZ21addr_of_static_sharedvE1s to ptr)

_______________________________________________
cfe-commits mailing list
[email protected]
https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits

Reply via email to