[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