[llvm] [SPIR-V] Lower OpenCL get_kernel_* queries to OpGetKernel* instructions (PR #227455)
Arseniy Obolenskiy via llvm-commits
llvm-commits at lists.llvm.org
Tue Sep 29 12:49:39 PDT 2026
https://github.com/aobolensk created https://github.com/llvm/llvm-project/pull/227455
Lowers __get_kernel_work_group_size_impl, __get_kernel_preferred_work_group_size_multiple_impl, __get_kernel_sub_group_count_for_ndrange_impl and __get_kernel_max_sub_group_size_for_ndrange_impl.
>From 5cd14db0d9928587e4af24eafe469c92438a999f Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Tue, 29 Sep 2026 21:47:18 +0200
Subject: [PATCH] [SPIR-V] Lower OpenCL get_kernel_* queries to OpGetKernel*
instructions
---
llvm/lib/Target/SPIRV/SPIRVBuiltins.cpp | 94 ++++++++++++-------
llvm/lib/Target/SPIRV/SPIRVBuiltins.td | 4 +
llvm/lib/Target/SPIRV/SPIRVInstrInfo.td | 8 ++
llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.cpp | 36 +++++--
.../SPIRV/transcoding/enqueue_kernel.ll | 40 ++++++++
5 files changed, 140 insertions(+), 42 deletions(-)
diff --git a/llvm/lib/Target/SPIRV/SPIRVBuiltins.cpp b/llvm/lib/Target/SPIRV/SPIRVBuiltins.cpp
index 84196697060e3..8d684993cf0ed 100644
--- a/llvm/lib/Target/SPIRV/SPIRVBuiltins.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVBuiltins.cpp
@@ -2973,6 +2973,58 @@ static bool buildNDRange(const SPIRV::IncomingCall *Call,
.addUse(TmpReg);
}
+static void buildKernelInvokeOperands(
+ const SPIRV::IncomingCall *Call, unsigned InvokeIdx, unsigned ParamIdx,
+ MachineIRBuilder &MIRBuilder, SPIRVGlobalRegistry *GR, Register &InvokeReg,
+ Register &ParamReg, Register &ParamSizeReg, Register &ParamAlignReg) {
+ MachineRegisterInfo *MRI = MIRBuilder.getMRI();
+ const DataLayout &DL = MIRBuilder.getDataLayout();
+
+ // Bypass the addrspacecast so Invoke references the function's <id>.
+ MachineInstr *InvokeGlobalMI =
+ getBlockStructInstr(Call->Arguments[InvokeIdx], MRI);
+ assert(InvokeGlobalMI->getOpcode() == TargetOpcode::G_GLOBAL_VALUE);
+ InvokeReg = InvokeGlobalMI->getOperand(0).getReg();
+ MRI->setRegClass(InvokeReg, &SPIRV::pIDRegClass);
+
+ Register BlockLiteralReg = Call->Arguments[ParamIdx];
+ const SPIRVTypeInst Int8Ty = GR->getOrCreateSPIRVIntegerType(8, MIRBuilder);
+ const SPIRVTypeInst Int8PtrGen = GR->getOrCreateSPIRVPointerType(
+ Int8Ty, MIRBuilder, SPIRV::StorageClass::Generic);
+ Type *PType = const_cast<Type *>(getBlockStructType(BlockLiteralReg, MRI));
+
+ ParamReg = createVirtualRegister(Int8PtrGen, GR, MIRBuilder);
+ MIRBuilder.buildInstr(SPIRV::OpBitcast)
+ .addDef(ParamReg)
+ .addUse(GR->getSPIRVTypeID(Int8PtrGen))
+ .addUse(BlockLiteralReg);
+ // TODO: these numbers should be obtained from block literal structure.
+ ParamSizeReg =
+ buildConstantIntReg32(DL.getTypeStoreSize(PType), MIRBuilder, GR);
+ ParamAlignReg =
+ buildConstantIntReg32(DL.getPrefTypeAlign(PType).value(), MIRBuilder, GR);
+}
+
+static bool buildKernelQuery(const SPIRV::IncomingCall *Call, unsigned Opcode,
+ MachineIRBuilder &MIRBuilder,
+ SPIRVGlobalRegistry *GR) {
+ bool HasNDRange = Call->Arguments.size() == 3;
+ unsigned InvokeIdx = HasNDRange ? 1 : 0;
+ Register InvokeReg, ParamReg, ParamSizeReg, ParamAlignReg;
+ buildKernelInvokeOperands(Call, InvokeIdx, InvokeIdx + 1, MIRBuilder, GR,
+ InvokeReg, ParamReg, ParamSizeReg, ParamAlignReg);
+ auto MIB = MIRBuilder.buildInstr(Opcode)
+ .addDef(Call->ReturnRegister)
+ .addUse(GR->getSPIRVTypeID(Call->ReturnType));
+ if (HasNDRange)
+ MIB.addUse(Call->Arguments[0]);
+ MIB.addUse(InvokeReg)
+ .addUse(ParamReg)
+ .addUse(ParamSizeReg)
+ .addUse(ParamAlignReg);
+ return true;
+}
+
static bool buildEnqueueKernel(const SPIRV::IncomingCall *Call,
MachineIRBuilder &MIRBuilder,
SPIRVGlobalRegistry *GR) {
@@ -2982,7 +3034,6 @@ static bool buildEnqueueKernel(const SPIRV::IncomingCall *Call,
// 3. create a SPIRV operator with arguments.
MachineRegisterInfo *MRI = MIRBuilder.getMRI();
- const DataLayout &DL = MIRBuilder.getDataLayout();
const SPIRVTypeInst Int32Ty = GR->getOrCreateSPIRVIntegerType(32, MIRBuilder);
// 1. prepare call indexes in order we expect them.
@@ -3060,39 +3111,11 @@ static bool buildEnqueueKernel(const SPIRV::IncomingCall *Call,
RetEventReg = NullPtr;
}
- // 2.2 Invoke (Kernel)
- // The Invoke operand of OpEnqueueKernel must be the function's <id>
- // (per SPIR-V spec). The frontend hands us the result of an
- // addrspacecast of @block_invoke_kernel; bypass that cast so the
- // operand references the underlying G_GLOBAL_VALUE register, which
- // selectGlobalValue lowers to a placeholder later rewritten by
- // SPIRVModuleAnalysis to the OpFunction <id>.
- MachineInstr *InvokeGlobalMI =
- getBlockStructInstr(Call->Arguments[InvokeIdx], MRI);
- assert(InvokeGlobalMI->getOpcode() == TargetOpcode::G_GLOBAL_VALUE);
- Register InvokeReg = InvokeGlobalMI->getOperand(0).getReg();
- // OpEnqueueKernel's Invoke operand uses the pID register class.
- MRI->setRegClass(InvokeReg, &SPIRV::pIDRegClass);
-
- // 2.3 Param, Param Size, Param Align
- Register BlockLiteralReg = Call->Arguments[ParamIdx];
- const SPIRVTypeInst Int8Ty = GR->getOrCreateSPIRVIntegerType(8, MIRBuilder);
- const SPIRVTypeInst Int8PtrGen = GR->getOrCreateSPIRVPointerType(
- Int8Ty, MIRBuilder, SPIRV::StorageClass::Generic);
- Type *PType = const_cast<Type *>(getBlockStructType(BlockLiteralReg, MRI));
-
- Register ParamReg = createVirtualRegister(Int8PtrGen, GR, MIRBuilder);
- MIRBuilder.buildInstr(SPIRV::OpBitcast)
- .addDef(ParamReg)
- .addUse(GR->getSPIRVTypeID(Int8PtrGen))
- .addUse(BlockLiteralReg);
- // TODO: these numbers should be obtained from block literal structure.
- Register ParamSizeReg =
- buildConstantIntReg32(DL.getTypeStoreSize(PType), MIRBuilder, GR);
- Register ParamAlignReg =
- buildConstantIntReg32(DL.getPrefTypeAlign(PType).value(), MIRBuilder, GR);
+ Register InvokeReg, ParamReg, ParamSizeReg, ParamAlignReg;
+ buildKernelInvokeOperands(Call, InvokeIdx, ParamIdx, MIRBuilder, GR,
+ InvokeReg, ParamReg, ParamSizeReg, ParamAlignReg);
- // 2.4 Local Size Array
+ // 2.3 Local Size Array
SmallVector<Register, 16> LocalSizes;
if (HasVarArgs) {
Register LocalSizeNumElem = Call->Arguments[LocalSizeNumElemIdx];
@@ -3172,6 +3195,11 @@ static bool generateEnqueueInst(const SPIRV::IncomingCall *Call,
return buildNDRange(Call, MIRBuilder, GR, CB);
case SPIRV::OpEnqueueKernel:
return buildEnqueueKernel(Call, MIRBuilder, GR);
+ case SPIRV::OpGetKernelNDrangeSubGroupCount:
+ case SPIRV::OpGetKernelNDrangeMaxSubGroupSize:
+ case SPIRV::OpGetKernelWorkGroupSize:
+ case SPIRV::OpGetKernelPreferredWorkGroupSizeMultiple:
+ return buildKernelQuery(Call, Opcode, MIRBuilder, GR);
default:
return false;
}
diff --git a/llvm/lib/Target/SPIRV/SPIRVBuiltins.td b/llvm/lib/Target/SPIRV/SPIRVBuiltins.td
index bb8dd1288992f..b401aa564aab7 100644
--- a/llvm/lib/Target/SPIRV/SPIRVBuiltins.td
+++ b/llvm/lib/Target/SPIRV/SPIRVBuiltins.td
@@ -716,6 +716,10 @@ defm : DemangledNativeBuiltin<"__enqueue_kernel_basic_events", OpenCL_std, Enque
defm : DemangledNativeBuiltin<"__enqueue_kernel_varargs", OpenCL_std, Enqueue, 7, 7, OpEnqueueKernel>;
defm : DemangledNativeBuiltin<"__enqueue_kernel_events_varargs", OpenCL_std, Enqueue, 10, 10, OpEnqueueKernel>;
defm : DemangledNativeBuiltin<"__spirv_EnqueueKernel", OpenCL_std, Enqueue, 10, 0, OpEnqueueKernel>;
+defm : DemangledNativeBuiltin<"__get_kernel_work_group_size_impl", OpenCL_std, Enqueue, 2, 2, OpGetKernelWorkGroupSize>;
+defm : DemangledNativeBuiltin<"__get_kernel_preferred_work_group_size_multiple_impl", OpenCL_std, Enqueue, 2, 2, OpGetKernelPreferredWorkGroupSizeMultiple>;
+defm : DemangledNativeBuiltin<"__get_kernel_sub_group_count_for_ndrange_impl", OpenCL_std, Enqueue, 3, 3, OpGetKernelNDrangeSubGroupCount>;
+defm : DemangledNativeBuiltin<"__get_kernel_max_sub_group_size_for_ndrange_impl", OpenCL_std, Enqueue, 3, 3, OpGetKernelNDrangeMaxSubGroupSize>;
defm : DemangledNativeBuiltin<"retain_event", OpenCL_std, Enqueue, 1, 1, OpRetainEvent>;
defm : DemangledNativeBuiltin<"__spirv_RetainEvent", OpenCL_std, Enqueue, 1, 1, OpRetainEvent>;
defm : DemangledNativeBuiltin<"release_event", OpenCL_std, Enqueue, 1, 1, OpReleaseEvent>;
diff --git a/llvm/lib/Target/SPIRV/SPIRVInstrInfo.td b/llvm/lib/Target/SPIRV/SPIRVInstrInfo.td
index cd4faf80c2d97..569da0ebe4b16 100644
--- a/llvm/lib/Target/SPIRV/SPIRVInstrInfo.td
+++ b/llvm/lib/Target/SPIRV/SPIRVInstrInfo.td
@@ -789,6 +789,14 @@ def OpSubgroupMatrixMultiplyAccumulateINTEL: Op<6237, (outs ID:$res),
def OpEnqueueKernel: Op<292, (outs ID:$res), (ins TYPE:$type, ID:$queue, ID:$flags, ID:$NDR, ID:$nevents, ID:$wevents,
ID:$revent, ID:$invoke, ID:$param, ID:$psize, ID:$palign, variable_ops),
"$res = OpEnqueueKernel $type $queue $flags $NDR $nevents $wevents $revent $invoke $param $psize $palign">;
+def OpGetKernelNDrangeSubGroupCount: Op<293, (outs ID:$res), (ins TYPE:$type, ID:$NDR, ID:$invoke, ID:$param, ID:$psize, ID:$palign),
+ "$res = OpGetKernelNDrangeSubGroupCount $type $NDR $invoke $param $psize $palign">;
+def OpGetKernelNDrangeMaxSubGroupSize: Op<294, (outs ID:$res), (ins TYPE:$type, ID:$NDR, ID:$invoke, ID:$param, ID:$psize, ID:$palign),
+ "$res = OpGetKernelNDrangeMaxSubGroupSize $type $NDR $invoke $param $psize $palign">;
+def OpGetKernelWorkGroupSize: Op<295, (outs ID:$res), (ins TYPE:$type, ID:$invoke, ID:$param, ID:$psize, ID:$palign),
+ "$res = OpGetKernelWorkGroupSize $type $invoke $param $psize $palign">;
+def OpGetKernelPreferredWorkGroupSizeMultiple: Op<296, (outs ID:$res), (ins TYPE:$type, ID:$invoke, ID:$param, ID:$psize, ID:$palign),
+ "$res = OpGetKernelPreferredWorkGroupSizeMultiple $type $invoke $param $psize $palign">;
def OpRetainEvent: Op<297, (outs), (ins ID:$event), "OpRetainEvent $event">;
def OpReleaseEvent: Op<298, (outs), (ins ID:$event), "OpReleaseEvent $event">;
def OpCreateUserEvent: Op<299, (outs ID:$res), (ins TYPE:$type),
diff --git a/llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.cpp b/llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.cpp
index 0fc45deb2abbe..75ef439d5f783 100644
--- a/llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.cpp
@@ -332,6 +332,22 @@ static InstrSignature instrToSignature(const MachineInstr &MI,
return Signature;
}
+// Operand index of Invoke in device enqueue instructions, 0 if none.
+static unsigned getInvokeOperandIdx(unsigned Opcode) {
+ switch (Opcode) {
+ case SPIRV::OpEnqueueKernel:
+ return 8;
+ case SPIRV::OpGetKernelNDrangeSubGroupCount:
+ case SPIRV::OpGetKernelNDrangeMaxSubGroupSize:
+ return 3;
+ case SPIRV::OpGetKernelWorkGroupSize:
+ case SPIRV::OpGetKernelPreferredWorkGroupSizeMultiple:
+ return 2;
+ default:
+ return 0;
+ }
+}
+
bool SPIRVModuleAnalysis::isDeclSection(const MachineRegisterInfo &MRI,
const MachineInstr &MI) {
unsigned Opcode = MI.getOpcode();
@@ -351,15 +367,15 @@ bool SPIRVModuleAnalysis::isDeclSection(const MachineRegisterInfo &MRI,
// The OpUndef may be a placeholder for a function reference recorded by
// selectGlobalValue. Skip emitting it if any user consumes it as a
// function-pointer-like operand (OpConstantFunctionPointerINTEL operand 2,
- // or OpEnqueueKernel's Invoke operand at index 8). The rewrite happens
- // in visitFunPtrUse, which aliases the OpUndef's vreg to the function's
- // global <id>.
+ // or the Invoke operand of a device enqueue instruction). The rewrite
+ // happens in visitFunPtrUse, which aliases the OpUndef's vreg to the
+ // function's global <id>.
Register DefReg = MI.getOperand(0).getReg();
if (GR->getFunctionDefinitionByUse(&MI.getOperand(0))) {
for (MachineInstr &UseMI : MRI.use_instructions(DefReg)) {
unsigned UseOp = UseMI.getOpcode();
if (UseOp == SPIRV::OpConstantFunctionPointerINTEL ||
- UseOp == SPIRV::OpEnqueueKernel) {
+ getInvokeOperandIdx(UseOp)) {
MAI.setSkipEmission(&MI);
return false;
}
@@ -587,11 +603,9 @@ void SPIRVModuleAnalysis::collectDeclarations(const Module &M) {
if (DefMO.isReg() && isDeclSection(MRI, MI) &&
!MAI.hasRegisterAlias(MF, DefMO.getReg()))
visitDecl(MRI, SignatureToGReg, GlobalToGReg, MF, MI);
- // OpEnqueueKernel is not a decl, but its Invoke operand may be a
- // function-pointer placeholder OpUndef recorded by selectGlobalValue.
- // Resolve it to the OpFunction's global <id> via visitFunPtrUse.
- if (Opcode == SPIRV::OpEnqueueKernel && MI.getNumOperands() > 8) {
- const MachineOperand &InvokeMO = MI.getOperand(8);
+ // Resolve a function-pointer placeholder Invoke operand.
+ if (unsigned InvokeIdx = getInvokeOperandIdx(Opcode)) {
+ const MachineOperand &InvokeMO = MI.getOperand(InvokeIdx);
if (InvokeMO.isReg()) {
Register InvokeReg = InvokeMO.getReg();
if (!MAI.hasRegisterAlias(MF, InvokeReg)) {
@@ -1755,6 +1769,10 @@ void addInstrRequirements(const MachineInstr &MI,
case SPIRV::OpTypeQueue:
case SPIRV::OpBuildNDRange:
case SPIRV::OpEnqueueKernel:
+ case SPIRV::OpGetKernelNDrangeSubGroupCount:
+ case SPIRV::OpGetKernelNDrangeMaxSubGroupSize:
+ case SPIRV::OpGetKernelWorkGroupSize:
+ case SPIRV::OpGetKernelPreferredWorkGroupSizeMultiple:
Reqs.addCapability(SPIRV::Capability::DeviceEnqueue);
break;
case SPIRV::OpDecorate:
diff --git a/llvm/test/CodeGen/SPIRV/transcoding/enqueue_kernel.ll b/llvm/test/CodeGen/SPIRV/transcoding/enqueue_kernel.ll
index 95a3777f2bb15..592e387db851b 100644
--- a/llvm/test/CodeGen/SPIRV/transcoding/enqueue_kernel.ll
+++ b/llvm/test/CodeGen/SPIRV/transcoding/enqueue_kernel.ll
@@ -124,6 +124,13 @@
; CHECK-DAG: %[[#InvokeKernel4]] = OpFunction %[[#typeVoid]] {{Pure|None}} %[[#typeFnVoidPtrLocal1]]
; CHECK-DAG: %[[#InvokeKernel5]] = OpFunction %[[#typeVoid]] {{Pure|None}} %[[#typeFnVoidPtrLocal3]]
; CHECK-DAG: %[[#InvokeKernel6]] = OpFunction %[[#typeVoid]] {{Pure|None}} %[[#typeFnVoidPtr]]
+
+; CHECK-LABEL: ; -- Begin function kernel_queries
+; CHECK: %[[#]] = OpGetKernelWorkGroupSize %[[#typeInt32]] %[[#QueryKernelPtr:]] %[[#]] %[[#Num16i32]] %[[#Num8i32]]
+; CHECK: %[[#]] = OpGetKernelPreferredWorkGroupSizeMultiple %[[#typeInt32]] %[[#QueryKernelPtr]] %[[#]] %[[#Num16i32]] %[[#Num8i32]]
+; CHECK: %[[#]] = OpGetKernelNDrangeSubGroupCount %[[#typeInt32]] %[[#]] %[[#QueryKernelPtr]] %[[#]] %[[#Num16i32]] %[[#Num8i32]]
+; CHECK: %[[#]] = OpGetKernelNDrangeMaxSubGroupSize %[[#typeInt32]] %[[#]] %[[#QueryKernelPtr]] %[[#]] %[[#Num16i32]] %[[#Num8i32]]
+; CHECK-NOT: _impl
target datalayout = "e-i64:64-v16:16-v24:32-v32:32-v48:64-v96:128-v192:256-v256:256-v512:512-v1024:1024-n8:16:32:64-G1"
target triple = "spirv64-unknown-unknown"
@@ -278,6 +285,39 @@ define internal spir_kernel void @__device_side_enqueue_block_invoke_6_kernel(pt
}
+ at __block_literal_global.3 = internal addrspace(1) constant { i32, i32, ptr addrspace(4) } { i32 16, i32 8, ptr addrspace(4) addrspacecast (ptr @__kernel_queries_block_invoke to ptr addrspace(4)) }, align 8
+
+define spir_kernel void @kernel_queries(ptr addrspace(1) align 4 %out) {
+entry:
+ %nd = alloca %struct.ndrange_t, align 8
+ %wgs = call spir_func i32 @__get_kernel_work_group_size_impl(ptr addrspace(4) addrspacecast (ptr @__kernel_queries_block_invoke_kernel to ptr addrspace(4)), ptr addrspace(4) addrspacecast (ptr addrspace(1) @__block_literal_global.3 to ptr addrspace(4)))
+ store i32 %wgs, ptr addrspace(1) %out, align 4
+ %pwgsm = call spir_func i32 @__get_kernel_preferred_work_group_size_multiple_impl(ptr addrspace(4) addrspacecast (ptr @__kernel_queries_block_invoke_kernel to ptr addrspace(4)), ptr addrspace(4) addrspacecast (ptr addrspace(1) @__block_literal_global.3 to ptr addrspace(4)))
+ %p1 = getelementptr inbounds i32, ptr addrspace(1) %out, i64 1
+ store i32 %pwgsm, ptr addrspace(1) %p1, align 4
+ %sgc = call spir_func i32 @__get_kernel_sub_group_count_for_ndrange_impl(ptr %nd, ptr addrspace(4) addrspacecast (ptr @__kernel_queries_block_invoke_kernel to ptr addrspace(4)), ptr addrspace(4) addrspacecast (ptr addrspace(1) @__block_literal_global.3 to ptr addrspace(4)))
+ %p2 = getelementptr inbounds i32, ptr addrspace(1) %out, i64 2
+ store i32 %sgc, ptr addrspace(1) %p2, align 4
+ %msgs = call spir_func i32 @__get_kernel_max_sub_group_size_for_ndrange_impl(ptr %nd, ptr addrspace(4) addrspacecast (ptr @__kernel_queries_block_invoke_kernel to ptr addrspace(4)), ptr addrspace(4) addrspacecast (ptr addrspace(1) @__block_literal_global.3 to ptr addrspace(4)))
+ %p3 = getelementptr inbounds i32, ptr addrspace(1) %out, i64 3
+ store i32 %msgs, ptr addrspace(1) %p3, align 4
+ ret void
+}
+
+declare spir_func i32 @__get_kernel_work_group_size_impl(ptr addrspace(4), ptr addrspace(4))
+declare spir_func i32 @__get_kernel_preferred_work_group_size_multiple_impl(ptr addrspace(4), ptr addrspace(4))
+declare spir_func i32 @__get_kernel_sub_group_count_for_ndrange_impl(ptr, ptr addrspace(4), ptr addrspace(4))
+declare spir_func i32 @__get_kernel_max_sub_group_size_for_ndrange_impl(ptr, ptr addrspace(4), ptr addrspace(4))
+
+define internal spir_func void @__kernel_queries_block_invoke(ptr addrspace(4) %.block_descriptor) {
+ ret void
+}
+
+define internal spir_kernel void @__kernel_queries_block_invoke_kernel(ptr addrspace(4) %0) {
+ ret void
+}
+
+
!opencl.ocl.version = !{!0}
!0 = !{i32 3, i32 0}
More information about the llvm-commits
mailing list