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

>From 4ad5c70f08d3d0c406c3017171ffc02e01e31384 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

---
 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