[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