[llvm] 45250b2 - [SPIRV][NFCI] Refactor selection of atomic operations with pointer operands (#208294)

via llvm-commits llvm-commits at lists.llvm.org
Thu Jul 23 09:00:12 PDT 2026


Author: Nick Sarnie
Date: 2026-07-23T16:00:07Z
New Revision: 45250b2db177add3433035ce09069b1c219499d4

URL: https://github.com/llvm/llvm-project/commit/45250b2db177add3433035ce09069b1c219499d4
DIFF: https://github.com/llvm/llvm-project/commit/45250b2db177add3433035ce09069b1c219499d4.diff

LOG: [SPIRV][NFCI] Refactor selection of atomic operations with pointer operands (#208294)

We do similar things for the three different cases (load, store,
exchange), so share some logic.

Suggested
[here](https://github.com/llvm/llvm-project/pull/207830#issuecomment-4912670227).

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply at anthropic.com>

Co-authored-by: Claude Opus 4.8 (1M context) <noreply at anthropic.com>

Added: 
    

Modified: 
    llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp

Removed: 
    


################################################################################
diff  --git a/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp b/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
index aa78e924f46e5..c7c1b727776b6 100644
--- a/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
@@ -22,6 +22,7 @@
 #include "SPIRVTypeInst.h"
 #include "SPIRVUtils.h"
 #include "llvm/ADT/APFloat.h"
+#include "llvm/ADT/STLFunctionalExtras.h"
 #include "llvm/ADT/StringExtras.h"
 #include "llvm/CodeGen/GlobalISel/GIMatchTableExecutorImpl.h"
 #include "llvm/CodeGen/GlobalISel/GenericMachineInstrs.h"
@@ -177,6 +178,21 @@ class SPIRVInstructionSelector : public InstructionSelector {
   bool selectAtomicRMW(Register ResVReg, SPIRVTypeInst ResType, MachineInstr &I,
                        unsigned NewOpcode, unsigned NegateOpcode = 0) const;
 
+  // Creates an integer-typed register with bitwidth equal to pointer size.
+  Register createPtrSizedIntReg(MachineIRBuilder &MIRBuilder) const;
+  // Emit an OpConvertPtrToU that converts the pointer value in \p PtrVal into
+  // an integer of equal bitwidth, returning the register holding the result.
+  Register convertPtrToInt(Register PtrVal, MachineIRBuilder &MIRBuilder) const;
+  // Emit an OpBitcast that reinterprets the pointer \p Ptr as a pointer to an
+  // integer of pointer size in storage class \p SC, returning the result.
+  Register castPtrToPtrToInt(Register Ptr, SPIRV::StorageClass::StorageClass SC,
+                             MachineIRBuilder &MIRBuilder) const;
+  // Handle atomic loads, stores and exchanges of pointer types by casting
+  // to/from integer types as needed.
+  bool selectAtomicPtrValue(
+      Register ResVReg, SPIRVTypeInst ResType, MachineIRBuilder &MIRBuilder,
+      function_ref<Register(SPIRVTypeInst IntType)> EmitAtomic) const;
+
   bool selectInterlockedOp(Register ResVReg, SPIRVTypeInst ResType,
                            MachineInstr &I, unsigned Opcode) const;
 
@@ -1997,6 +2013,70 @@ bool SPIRVInstructionSelector::selectLoad(Register ResVReg,
   return true;
 }
 
+Register SPIRVInstructionSelector::createPtrSizedIntReg(
+    MachineIRBuilder &MIRBuilder) const {
+  SPIRVTypeInst IntType =
+      GR.getOrCreateSPIRVIntegerType(GR.getPointerSize(), MIRBuilder);
+  Register Reg =
+      MRI->createGenericVirtualRegister(LLT::scalar(GR.getPointerSize()));
+  MRI->setRegClass(Reg, GR.getRegClass(IntType));
+  GR.assignSPIRVTypeToVReg(IntType, Reg, MIRBuilder.getMF());
+  return Reg;
+}
+
+Register
+SPIRVInstructionSelector::convertPtrToInt(Register PtrVal,
+                                          MachineIRBuilder &MIRBuilder) const {
+  SPIRVTypeInst IntType =
+      GR.getOrCreateSPIRVIntegerType(GR.getPointerSize(), MIRBuilder);
+  Register IntReg = createPtrSizedIntReg(MIRBuilder);
+  MIRBuilder.buildInstr(SPIRV::OpConvertPtrToU)
+      .addDef(IntReg)
+      .addUse(GR.getSPIRVTypeID(IntType)) // Result type
+      .addUse(PtrVal)                     // Pointer operand
+      .constrainAllUses(TII, TRI, RBI);
+  return IntReg;
+}
+
+Register SPIRVInstructionSelector::castPtrToPtrToInt(
+    Register Ptr, SPIRV::StorageClass::StorageClass SC,
+    MachineIRBuilder &MIRBuilder) const {
+  SPIRVTypeInst IntType =
+      GR.getOrCreateSPIRVIntegerType(GR.getPointerSize(), MIRBuilder);
+  SPIRVTypeInst PtrType =
+      GR.getOrCreateSPIRVPointerType(IntType, MIRBuilder, SC);
+  Register CastedPtr =
+      MRI->createGenericVirtualRegister(LLT::scalar(GR.getPointerSize()));
+  MRI->setRegClass(CastedPtr, GR.getRegClass(PtrType));
+  GR.assignSPIRVTypeToVReg(PtrType, CastedPtr, MIRBuilder.getMF());
+  MIRBuilder.buildInstr(SPIRV::OpBitcast)
+      .addDef(CastedPtr)
+      .addUse(GR.getSPIRVTypeID(PtrType))
+      .addUse(Ptr)
+      .constrainAllUses(TII, TRI, RBI);
+  return CastedPtr;
+}
+
+bool SPIRVInstructionSelector::selectAtomicPtrValue(
+    Register ResVReg, SPIRVTypeInst ResType, MachineIRBuilder &MIRBuilder,
+    function_ref<Register(SPIRVTypeInst IntType)> EmitAtomic) const {
+  // Pointer-typed atomics are lowered by bitcasting the Ptr operand to a
+  // pointer to an integer of the same size as the pointer, so that the actual
+  // atomic instruction operates on integers as required by the spec. Value
+  // operands and results are converted with OpConvertPtrToU/OpConvertUToPtr.
+  unsigned PtrSize = GR.getPointerSize();
+  SPIRVTypeInst IntType = GR.getOrCreateSPIRVIntegerType(PtrSize, MIRBuilder);
+
+  Register IntResult = EmitAtomic(IntType);
+  if (IntResult.isValid())
+    MIRBuilder.buildInstr(SPIRV::OpConvertUToPtr)
+        .addDef(ResVReg)
+        .addUse(GR.getSPIRVTypeID(ResType))
+        .addUse(IntResult)
+        .constrainAllUses(TII, TRI, RBI);
+  return true;
+}
+
 bool SPIRVInstructionSelector::selectAtomicLoad(Register ResVReg,
                                                 SPIRVTypeInst ResType,
                                                 MachineInstr &I) const {
@@ -2035,45 +2115,23 @@ bool SPIRVInstructionSelector::selectAtomicLoad(Register ResVReg,
              "allowed for pointer types for physical addressing model");
     // If data to load is a pointer type we bitcast the Ptr parameter to pointer
     // to an integer type of the same size as the pointer size and then generate
-    // OpAtomicLoad the return value of that OpAtomicLoad is an integet that is
+    // OpAtomicLoad the return value of that OpAtomicLoad is an integer that is
     // converted back to a pointer type using OpConvertUToPtr.
-
-    unsigned PtrSize = GR.getPointerSize();
-    SPIRVTypeInst PtrAsIntSpirvType =
-        GR.getOrCreateSPIRVIntegerType(PtrSize, MIRBuilder);
-    Register PtrToUVal =
-        MRI->createGenericVirtualRegister(LLT::scalar(PtrSize));
-    MRI->setRegClass(PtrToUVal, GR.getRegClass(PtrAsIntSpirvType));
-    GR.assignSPIRVTypeToVReg(PtrAsIntSpirvType, PtrToUVal, MIRBuilder.getMF());
-
-    Register PtrCastedToMatchValReg =
-        MRI->createGenericVirtualRegister(LLT::scalar(PtrSize));
-    MRI->setRegClass(PtrCastedToMatchValReg, MRI->getRegClassOrNull(Ptr));
-    SPIRVTypeInst PtrType = GR.getOrCreateSPIRVPointerType(
-        PtrAsIntSpirvType, MIRBuilder,
-        addressSpaceToStorageClass(MemOp.getAddrSpace(), STI));
-    GR.assignSPIRVTypeToVReg(PtrType, PtrCastedToMatchValReg,
-                             MIRBuilder.getMF());
-
-    MIRBuilder.buildInstr(SPIRV::OpBitcast)
-        .addDef(PtrCastedToMatchValReg)
-        .addUse(GR.getSPIRVTypeID(PtrType))
-        .addUse(Ptr)
-        .constrainAllUses(TII, TRI, RBI);
-
-    MIRBuilder.buildInstr(SPIRV::OpAtomicLoad)
-        .addDef(PtrToUVal)
-        .addUse(GR.getSPIRVTypeID(PtrAsIntSpirvType))
-        .addUse(PtrCastedToMatchValReg)
-        .addUse(ScopeReg)
-        .addUse(MemSemReg)
-        .constrainAllUses(TII, TRI, RBI);
-    MIRBuilder.buildInstr(SPIRV::OpConvertUToPtr)
-        .addDef(ResVReg)
-        .addUse(GR.getSPIRVTypeID(ResType))
-        .addUse(PtrToUVal)
-        .constrainAllUses(TII, TRI, RBI);
-    return true;
+    SPIRV::StorageClass::StorageClass SC =
+        addressSpaceToStorageClass(MemOp.getAddrSpace(), STI);
+    return selectAtomicPtrValue(
+        ResVReg, ResType, MIRBuilder, [&](SPIRVTypeInst IntType) {
+          Register CastedPtr = castPtrToPtrToInt(Ptr, SC, MIRBuilder);
+          Register IntResult = createPtrSizedIntReg(MIRBuilder);
+          MIRBuilder.buildInstr(SPIRV::OpAtomicLoad)
+              .addDef(IntResult)
+              .addUse(GR.getSPIRVTypeID(IntType))
+              .addUse(CastedPtr)
+              .addUse(ScopeReg)
+              .addUse(MemSemReg)
+              .constrainAllUses(TII, TRI, RBI);
+          return IntResult;
+        });
   }
   auto AtomicLoad = MIRBuilder.buildInstr(SPIRV::OpAtomicLoad)
                         .addDef(ResVReg)
@@ -2207,38 +2265,21 @@ bool SPIRVInstructionSelector::selectAtomicStore(MachineInstr &I) const {
     // same size as the pointer size using OpConvertPtrToU, bitcast Ptr
     // parameter to pointer to integer type and then generate OpAtomicStore
     // with casted values as required by spec.
-    unsigned PtrSize = GR.getPointerSize();
-    SPIRVTypeInst PtrAsIntSpirvType =
-        GR.getOrCreateSPIRVIntegerType(PtrSize, MIRBuilder);
-
-    Register PtrToUVal =
-        MRI->createGenericVirtualRegister(LLT::scalar(PtrSize));
-    MRI->setRegClass(PtrToUVal, GR.getRegClass(PtrAsIntSpirvType));
-    GR.assignSPIRVTypeToVReg(PtrAsIntSpirvType, PtrToUVal, MIRBuilder.getMF());
-    MIRBuilder.buildInstr(SPIRV::OpConvertPtrToU)
-        .addDef(PtrToUVal)
-        .addUse(GR.getSPIRVTypeID(PtrAsIntSpirvType)) // Result type
-        .addUse(StoreVal)                             // Pointer operand
-        .constrainAllUses(TII, TRI, RBI);
-
-    Register PtrCastedToMatchValReg =
-        MRI->createGenericVirtualRegister(LLT::scalar(PtrSize));
-    MRI->setRegClass(PtrCastedToMatchValReg, MRI->getRegClassOrNull(Ptr));
-    SPIRVTypeInst PtrType = GR.getOrCreateSPIRVPointerType(
-        PtrAsIntSpirvType, MIRBuilder,
-        addressSpaceToStorageClass(MemOp.getAddrSpace(), STI));
-    GR.assignSPIRVTypeToVReg(PtrType, PtrCastedToMatchValReg,
-                             MIRBuilder.getMF());
-
-    MIRBuilder.buildInstr(SPIRV::OpBitcast)
-        .addDef(PtrCastedToMatchValReg)
-        .addUse(GR.getSPIRVTypeID(PtrType))
-        .addUse(Ptr)
-        .constrainAllUses(TII, TRI, RBI);
-
-    StoreVal = PtrToUVal;
-    Ptr = PtrCastedToMatchValReg;
-    PointeeType = PtrAsIntSpirvType;
+    SPIRV::StorageClass::StorageClass SC =
+        addressSpaceToStorageClass(MemOp.getAddrSpace(), STI);
+    return selectAtomicPtrValue(
+        Register(), SPIRVTypeInst(), MIRBuilder, [&](SPIRVTypeInst IntType) {
+          Register ValueAsInt = convertPtrToInt(StoreVal, MIRBuilder);
+          Register CastedPtr = castPtrToPtrToInt(Ptr, SC, MIRBuilder);
+          MIRBuilder.buildInstr(SPIRV::OpAtomicStore)
+              .addUse(CastedPtr)
+              .addUse(ScopeReg)
+              .addUse(MemSemReg)
+              .addUse(ValueAsInt)
+              .constrainAllUses(TII, TRI, RBI);
+          // Stores produce no result, so no OpConvertUToPtr is needed.
+          return Register();
+        });
   }
 
   if (!PointeeType.isTypeIntOrFloat())
@@ -2513,53 +2554,22 @@ bool SPIRVInstructionSelector::selectAtomicRMW(Register ResVReg,
     // converted back to a pointer type using OpConvertUToPtr, similar to atomic
     // load and store.
     MachineIRBuilder MIRBuilder(I);
-    unsigned PtrSize = GR.getPointerSize();
-    SPIRVTypeInst PtrAsIntSpirvType =
-        GR.getOrCreateSPIRVIntegerType(PtrSize, MIRBuilder);
-
-    Register ValueAsIntReg =
-        MRI->createGenericVirtualRegister(LLT::scalar(PtrSize));
-    MRI->setRegClass(ValueAsIntReg, GR.getRegClass(PtrAsIntSpirvType));
-    GR.assignSPIRVTypeToVReg(PtrAsIntSpirvType, ValueAsIntReg,
-                             MIRBuilder.getMF());
-    MIRBuilder.buildInstr(SPIRV::OpConvertPtrToU)
-        .addDef(ValueAsIntReg)
-        .addUse(GR.getSPIRVTypeID(PtrAsIntSpirvType)) // Result type
-        .addUse(ValueReg)                             // Pointer operand
-        .constrainAllUses(TII, TRI, RBI);
-
-    SPIRVTypeInst PtrType = GR.getOrCreateSPIRVPointerType(
-        PtrAsIntSpirvType, MIRBuilder, GR.getPointerStorageClass(Ptr));
-    Register PtrCastedToMatchValReg =
-        MRI->createGenericVirtualRegister(LLT::scalar(PtrSize));
-    MRI->setRegClass(PtrCastedToMatchValReg, GR.getRegClass(PtrType));
-    GR.assignSPIRVTypeToVReg(PtrType, PtrCastedToMatchValReg,
-                             MIRBuilder.getMF());
-    MIRBuilder.buildInstr(SPIRV::OpBitcast)
-        .addDef(PtrCastedToMatchValReg)
-        .addUse(GR.getSPIRVTypeID(PtrType))
-        .addUse(Ptr)
-        .constrainAllUses(TII, TRI, RBI);
-
-    Register ExchangeResReg =
-        MRI->createGenericVirtualRegister(LLT::scalar(PtrSize));
-    MRI->setRegClass(ExchangeResReg, GR.getRegClass(PtrAsIntSpirvType));
-    GR.assignSPIRVTypeToVReg(PtrAsIntSpirvType, ExchangeResReg,
-                             MIRBuilder.getMF());
-    MIRBuilder.buildInstr(SPIRV::OpAtomicExchange)
-        .addDef(ExchangeResReg)
-        .addUse(GR.getSPIRVTypeID(PtrAsIntSpirvType))
-        .addUse(PtrCastedToMatchValReg)
-        .addUse(ScopeReg)
-        .addUse(MemSemReg)
-        .addUse(ValueAsIntReg)
-        .constrainAllUses(TII, TRI, RBI);
-    MIRBuilder.buildInstr(SPIRV::OpConvertUToPtr)
-        .addDef(ResVReg)
-        .addUse(GR.getSPIRVTypeID(ResType))
-        .addUse(ExchangeResReg)
-        .constrainAllUses(TII, TRI, RBI);
-    return true;
+    SPIRV::StorageClass::StorageClass SC = GR.getPointerStorageClass(Ptr);
+    return selectAtomicPtrValue(
+        ResVReg, ResType, MIRBuilder, [&](SPIRVTypeInst IntType) {
+          Register ValueAsInt = convertPtrToInt(ValueReg, MIRBuilder);
+          Register CastedPtr = castPtrToPtrToInt(Ptr, SC, MIRBuilder);
+          Register ExchangeResReg = createPtrSizedIntReg(MIRBuilder);
+          MIRBuilder.buildInstr(SPIRV::OpAtomicExchange)
+              .addDef(ExchangeResReg)
+              .addUse(GR.getSPIRVTypeID(IntType))
+              .addUse(CastedPtr)
+              .addUse(ScopeReg)
+              .addUse(MemSemReg)
+              .addUse(ValueAsInt)
+              .constrainAllUses(TII, TRI, RBI);
+          return ExchangeResReg;
+        });
   }
 
   BuildMI(*I.getParent(), I, I.getDebugLoc(), TII.get(NewOpcode))


        


More information about the llvm-commits mailing list