[llvm] [llvm][SPIRV] Allow casts between AS9 and other AS (PR #194650)

Alex Duran via llvm-commits llvm-commits at lists.llvm.org
Tue Apr 28 08:02:00 PDT 2026


https://github.com/adurang created https://github.com/llvm/llvm-project/pull/194650

None

>From 509443f1add930bad786a1a7f2f186d9dfba6357 Mon Sep 17 00:00:00 2001
From: "Duran, Alex" <alejandro.duran at intel.com>
Date: Tue, 28 Apr 2026 07:16:22 -0700
Subject: [PATCH 1/2] [llvm][spirv] Add support for casts to/from
 CodeSectionINTEL

---
 .../Target/SPIRV/SPIRVInstructionSelector.cpp | 78 +++++++++++++++++++
 llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp  | 29 +++++--
 2 files changed, 102 insertions(+), 5 deletions(-)

diff --git a/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp b/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
index aee3a29c6e42b..0401ca3bcb94b 100644
--- a/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
@@ -2598,6 +2598,36 @@ bool SPIRVInstructionSelector::selectAddrSpaceCast(Register ResVReg,
                        isGenericCastablePtr(DstSC)
                    ? static_cast<uint32_t>(SPIRV::Opcode::GenericCastToPtr)
                    : 0);
+
+    if (STI.canUseExtension(SPIRV::Extension::SPV_INTEL_function_pointers) &&
+        !SpecOpcode) {
+      if (SrcSC == SPIRV::StorageClass::CodeSectionINTEL &&
+          DstSC == SPIRV::StorageClass::Generic) {
+        SpecOpcode = static_cast<uint32_t>(SPIRV::Opcode::PtrCastToGeneric);
+      } else if (SrcSC == SPIRV::StorageClass::Generic &&
+                 DstSC == SPIRV::StorageClass::CodeSectionINTEL) {
+        SpecOpcode = static_cast<uint32_t>(SPIRV::Opcode::GenericCastToPtr);
+      } else if ((SrcSC == SPIRV::StorageClass::CodeSectionINTEL &&
+                  DstSC == SPIRV::StorageClass::CrossWorkgroup) ||
+                 (SrcSC == SPIRV::StorageClass::CrossWorkgroup &&
+                  DstSC == SPIRV::StorageClass::CodeSectionINTEL)) {
+        // For P9 <-> P1, cast through Generic via two OpSpecConstantOp
+        unsigned FirstOpcode = (SrcSC == SPIRV::StorageClass::CodeSectionINTEL)
+            ? static_cast<uint32_t>(SPIRV::Opcode::PtrCastToGeneric)
+            : static_cast<uint32_t>(SPIRV::Opcode::PtrCastToGeneric);
+        unsigned SecondOpcode = static_cast<uint32_t>(SPIRV::Opcode::GenericCastToPtr);
+
+        Register GenericTypeReg = getUcharPtrTypeReg(I, SPIRV::StorageClass::Generic);
+        Register IntermediateReg = MRI->createVirtualRegister(&SPIRV::IDRegClass);
+
+        buildSpecConstantOp(I, IntermediateReg, SrcPtr, GenericTypeReg, FirstOpcode)
+            .constrainAllUses(TII, TRI, RBI);
+        buildSpecConstantOp(I, ResVReg, IntermediateReg, getUcharPtrTypeReg(I, DstSC), SecondOpcode)
+            .constrainAllUses(TII, TRI, RBI);
+        return true;
+      }
+    }
+
     // TODO: OpConstantComposite expects i8*, so we are forced to forget a
     // correct value of ResType and use general i8* instead. Maybe this should
     // be addressed in the emit-intrinsic step to infer a correct
@@ -2664,6 +2694,54 @@ bool SPIRVInstructionSelector::selectAddrSpaceCast(Register ResVReg,
   if (SrcSC == SPIRV::StorageClass::Generic && isUSMStorageClass(DstSC))
     return selectUnOp(ResVReg, ResType, I, SPIRV::OpGenericCastToPtr);
 
+  // Handle function pointer casts with SPV_INTEL_function_pointers extension.
+  if (STI.canUseExtension(SPIRV::Extension::SPV_INTEL_function_pointers)) {
+    // P4 <-> P9 can be handled directly
+    if (SrcSC == SPIRV::StorageClass::CodeSectionINTEL &&
+        DstSC == SPIRV::StorageClass::Generic)
+      return selectUnOp(ResVReg, ResType, I, SPIRV::OpPtrCastToGeneric);
+    if (SrcSC == SPIRV::StorageClass::Generic &&
+        DstSC == SPIRV::StorageClass::CodeSectionINTEL)
+      return selectUnOp(ResVReg, ResType, I, SPIRV::OpGenericCastToPtr);
+
+    // P1 <-> P9, cast through Generic as intermediary
+    if (SrcSC == SPIRV::StorageClass::CodeSectionINTEL &&
+        DstSC == SPIRV::StorageClass::CrossWorkgroup) {
+      SPIRVTypeInst GenericPtrTy =
+          GR.changePointerStorageClass(SrcPtrTy, SPIRV::StorageClass::Generic, I);
+      Register Tmp = createVirtualRegister(GenericPtrTy, &GR, MRI, MRI->getMF());
+      BuildMI(BB, I, DL, TII.get(SPIRV::OpPtrCastToGeneric))
+          .addDef(Tmp)
+          .addUse(GR.getSPIRVTypeID(GenericPtrTy))
+          .addUse(SrcPtr)
+          .constrainAllUses(TII, TRI, RBI);
+      BuildMI(BB, I, DL, TII.get(SPIRV::OpGenericCastToPtr))
+          .addDef(ResVReg)
+          .addUse(GR.getSPIRVTypeID(ResType))
+          .addUse(Tmp)
+          .constrainAllUses(TII, TRI, RBI);
+      return true;
+    }
+
+    if (SrcSC == SPIRV::StorageClass::CrossWorkgroup &&
+        DstSC == SPIRV::StorageClass::CodeSectionINTEL) {
+      SPIRVTypeInst GenericPtrTy =
+          GR.changePointerStorageClass(SrcPtrTy, SPIRV::StorageClass::Generic, I);
+      Register Tmp = createVirtualRegister(GenericPtrTy, &GR, MRI, MRI->getMF());
+      BuildMI(BB, I, DL, TII.get(SPIRV::OpPtrCastToGeneric))
+          .addDef(Tmp)
+          .addUse(GR.getSPIRVTypeID(GenericPtrTy))
+          .addUse(SrcPtr)
+          .constrainAllUses(TII, TRI, RBI);
+      BuildMI(BB, I, DL, TII.get(SPIRV::OpGenericCastToPtr))
+          .addDef(ResVReg)
+          .addUse(GR.getSPIRVTypeID(ResType))
+          .addUse(Tmp)
+          .constrainAllUses(TII, TRI, RBI);
+      return true;
+    }
+  }
+
   // Bitcast for pointers requires that the address spaces must match
   return false;
 }
diff --git a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
index 47ffecc4085ab..685e4c994063f 100644
--- a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
@@ -299,11 +299,30 @@ SPIRVLegalizerInfo::SPIRVLegalizerInfo(const SPIRVSubtarget &ST) {
       .unsupportedIf(typeIs(0, p9))
       .legalIf(all(typeInSet(0, allPtrs), typeInSet(1, allIntScalars)));
 
-  getActionDefinitionsBuilder(G_ADDRSPACE_CAST)
-      .unsupportedIf(
-          LegalityPredicates::any(all(typeIs(0, p9), typeIsNot(1, p9)),
-                                  all(typeIsNot(0, p9), typeIs(1, p9))))
-      .legalForCartesianProduct(allPtrs, allPtrs);
+  if (ST.canUseExtension(SPIRV::Extension::SPV_INTEL_function_pointers)) {
+    // With function pointer extension: allow casts between p1 (CrossWorkgroup),
+    // p4 (Generic), and p9 (CodeSectionINTEL)
+    getActionDefinitionsBuilder(G_ADDRSPACE_CAST)
+        .unsupportedIf(LegalityPredicates::any(
+            // Disallow p9 <-> other pointers except p1, p4, p9
+            all(typeIs(0, p9),
+                LegalityPredicates::any(typeIs(1, p0), typeIs(1, p2),
+                                        typeIs(1, p3), typeIs(1, p5),
+                                        typeIs(1, p6), typeIs(1, p7),
+                                        typeIs(1, p8))),
+            all(typeIs(1, p9),
+                LegalityPredicates::any(typeIs(0, p0), typeIs(0, p2),
+                                        typeIs(0, p3), typeIs(0, p5),
+                                        typeIs(0, p6), typeIs(0, p7),
+                                        typeIs(0, p8)))))
+        .legalForCartesianProduct(allPtrs, allPtrs);
+  } else {
+    getActionDefinitionsBuilder(G_ADDRSPACE_CAST)
+        .unsupportedIf(
+            LegalityPredicates::any(all(typeIs(0, p9), typeIsNot(1, p9)),
+                                    all(typeIsNot(0, p9), typeIs(1, p9))))
+        .legalForCartesianProduct(allPtrs, allPtrs);
+  }
 
   // Should we be legalizing bad scalar sizes like s5 here instead
   // of handling them in the instruction selector?

>From 7cdc623c33c45ca9c4aa3ffcfbd0fd5b9dce3768 Mon Sep 17 00:00:00 2001
From: "Duran, Alex" <alejandro.duran at intel.com>
Date: Tue, 28 Apr 2026 07:59:06 -0700
Subject: [PATCH 2/2] allow other namespaces

---
 .../Target/SPIRV/SPIRVInstructionSelector.cpp | 41 +++++++++++--------
 llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp  | 17 ++------
 2 files changed, 28 insertions(+), 30 deletions(-)

diff --git a/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp b/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
index 0401ca3bcb94b..ed0dc19703add 100644
--- a/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
@@ -2599,6 +2599,7 @@ bool SPIRVInstructionSelector::selectAddrSpaceCast(Register ResVReg,
                    ? static_cast<uint32_t>(SPIRV::Opcode::GenericCastToPtr)
                    : 0);
 
+    // Handle CodeSectionINTEL in constant context with SPV_INTEL_function_pointers.
     if (STI.canUseExtension(SPIRV::Extension::SPV_INTEL_function_pointers) &&
         !SpecOpcode) {
       if (SrcSC == SPIRV::StorageClass::CodeSectionINTEL &&
@@ -2607,22 +2608,30 @@ bool SPIRVInstructionSelector::selectAddrSpaceCast(Register ResVReg,
       } else if (SrcSC == SPIRV::StorageClass::Generic &&
                  DstSC == SPIRV::StorageClass::CodeSectionINTEL) {
         SpecOpcode = static_cast<uint32_t>(SPIRV::Opcode::GenericCastToPtr);
-      } else if ((SrcSC == SPIRV::StorageClass::CodeSectionINTEL &&
-                  DstSC == SPIRV::StorageClass::CrossWorkgroup) ||
-                 (SrcSC == SPIRV::StorageClass::CrossWorkgroup &&
-                  DstSC == SPIRV::StorageClass::CodeSectionINTEL)) {
-        // For P9 <-> P1, cast through Generic via two OpSpecConstantOp
-        unsigned FirstOpcode = (SrcSC == SPIRV::StorageClass::CodeSectionINTEL)
-            ? static_cast<uint32_t>(SPIRV::Opcode::PtrCastToGeneric)
-            : static_cast<uint32_t>(SPIRV::Opcode::PtrCastToGeneric);
-        unsigned SecondOpcode = static_cast<uint32_t>(SPIRV::Opcode::GenericCastToPtr);
+      } else if (SrcSC == SPIRV::StorageClass::CodeSectionINTEL &&
+                 DstSC != SPIRV::StorageClass::CodeSectionINTEL) {
+        // For P9 -> other address spaces, cast through Generic via two OpSpecConstantOp.
+        Register GenericTypeReg = getUcharPtrTypeReg(I, SPIRV::StorageClass::Generic);
+        Register IntermediateReg = MRI->createVirtualRegister(&SPIRV::IDRegClass);
 
+        buildSpecConstantOp(I, IntermediateReg, SrcPtr, GenericTypeReg,
+                            static_cast<uint32_t>(SPIRV::Opcode::PtrCastToGeneric))
+            .constrainAllUses(TII, TRI, RBI);
+        buildSpecConstantOp(I, ResVReg, IntermediateReg, getUcharPtrTypeReg(I, DstSC),
+                            static_cast<uint32_t>(SPIRV::Opcode::GenericCastToPtr))
+            .constrainAllUses(TII, TRI, RBI);
+        return true;
+      } else if (DstSC == SPIRV::StorageClass::CodeSectionINTEL &&
+                 SrcSC != SPIRV::StorageClass::CodeSectionINTEL) {
+        // For other address spaces -> P9, cast through Generic via two OpSpecConstantOp.
         Register GenericTypeReg = getUcharPtrTypeReg(I, SPIRV::StorageClass::Generic);
         Register IntermediateReg = MRI->createVirtualRegister(&SPIRV::IDRegClass);
 
-        buildSpecConstantOp(I, IntermediateReg, SrcPtr, GenericTypeReg, FirstOpcode)
+        buildSpecConstantOp(I, IntermediateReg, SrcPtr, GenericTypeReg,
+                            static_cast<uint32_t>(SPIRV::Opcode::PtrCastToGeneric))
             .constrainAllUses(TII, TRI, RBI);
-        buildSpecConstantOp(I, ResVReg, IntermediateReg, getUcharPtrTypeReg(I, DstSC), SecondOpcode)
+        buildSpecConstantOp(I, ResVReg, IntermediateReg, getUcharPtrTypeReg(I, DstSC),
+                            static_cast<uint32_t>(SPIRV::Opcode::GenericCastToPtr))
             .constrainAllUses(TII, TRI, RBI);
         return true;
       }
@@ -2696,7 +2705,7 @@ bool SPIRVInstructionSelector::selectAddrSpaceCast(Register ResVReg,
 
   // Handle function pointer casts with SPV_INTEL_function_pointers extension.
   if (STI.canUseExtension(SPIRV::Extension::SPV_INTEL_function_pointers)) {
-    // P4 <-> P9 can be handled directly
+    // P4 <-> P9 can be handled directly.
     if (SrcSC == SPIRV::StorageClass::CodeSectionINTEL &&
         DstSC == SPIRV::StorageClass::Generic)
       return selectUnOp(ResVReg, ResType, I, SPIRV::OpPtrCastToGeneric);
@@ -2704,9 +2713,9 @@ bool SPIRVInstructionSelector::selectAddrSpaceCast(Register ResVReg,
         DstSC == SPIRV::StorageClass::CodeSectionINTEL)
       return selectUnOp(ResVReg, ResType, I, SPIRV::OpGenericCastToPtr);
 
-    // P1 <-> P9, cast through Generic as intermediary
+    // Any address space <-> P9, cast through Generic as intermediary.
     if (SrcSC == SPIRV::StorageClass::CodeSectionINTEL &&
-        DstSC == SPIRV::StorageClass::CrossWorkgroup) {
+        DstSC != SPIRV::StorageClass::CodeSectionINTEL) {
       SPIRVTypeInst GenericPtrTy =
           GR.changePointerStorageClass(SrcPtrTy, SPIRV::StorageClass::Generic, I);
       Register Tmp = createVirtualRegister(GenericPtrTy, &GR, MRI, MRI->getMF());
@@ -2723,8 +2732,8 @@ bool SPIRVInstructionSelector::selectAddrSpaceCast(Register ResVReg,
       return true;
     }
 
-    if (SrcSC == SPIRV::StorageClass::CrossWorkgroup &&
-        DstSC == SPIRV::StorageClass::CodeSectionINTEL) {
+    if (DstSC == SPIRV::StorageClass::CodeSectionINTEL &&
+        SrcSC != SPIRV::StorageClass::CodeSectionINTEL) {
       SPIRVTypeInst GenericPtrTy =
           GR.changePointerStorageClass(SrcPtrTy, SPIRV::StorageClass::Generic, I);
       Register Tmp = createVirtualRegister(GenericPtrTy, &GR, MRI, MRI->getMF());
diff --git a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
index 685e4c994063f..0a8047dfbc1ef 100644
--- a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
@@ -300,23 +300,12 @@ SPIRVLegalizerInfo::SPIRVLegalizerInfo(const SPIRVSubtarget &ST) {
       .legalIf(all(typeInSet(0, allPtrs), typeInSet(1, allIntScalars)));
 
   if (ST.canUseExtension(SPIRV::Extension::SPV_INTEL_function_pointers)) {
-    // With function pointer extension: allow casts between p1 (CrossWorkgroup),
-    // p4 (Generic), and p9 (CodeSectionINTEL)
+    // With function pointer extension: allow casts between any address space
+    // and p9 (CodeSectionINTEL), with p4 (Generic) as intermediary.
     getActionDefinitionsBuilder(G_ADDRSPACE_CAST)
-        .unsupportedIf(LegalityPredicates::any(
-            // Disallow p9 <-> other pointers except p1, p4, p9
-            all(typeIs(0, p9),
-                LegalityPredicates::any(typeIs(1, p0), typeIs(1, p2),
-                                        typeIs(1, p3), typeIs(1, p5),
-                                        typeIs(1, p6), typeIs(1, p7),
-                                        typeIs(1, p8))),
-            all(typeIs(1, p9),
-                LegalityPredicates::any(typeIs(0, p0), typeIs(0, p2),
-                                        typeIs(0, p3), typeIs(0, p5),
-                                        typeIs(0, p6), typeIs(0, p7),
-                                        typeIs(0, p8)))))
         .legalForCartesianProduct(allPtrs, allPtrs);
   } else {
+    // Without extension: disallow all casts involving p9.
     getActionDefinitionsBuilder(G_ADDRSPACE_CAST)
         .unsupportedIf(
             LegalityPredicates::any(all(typeIs(0, p9), typeIsNot(1, p9)),



More information about the llvm-commits mailing list