[llvm] [SPIR-V] Access-chain aggregate pointers to element 0 for cooperative matrix load/store (PR #202050)

Julian Klappenbach via llvm-commits llvm-commits at lists.llvm.org
Sat Jun 6 08:15:02 PDT 2026


https://github.com/jklappenbach updated https://github.com/llvm/llvm-project/pull/202050

>From 60abd06194db2a9a06f17f38c1ba57a39dc4f2ea Mon Sep 17 00:00:00 2001
From: Julian Klappenbach <julian at twilight.digital>
Date: Sat, 6 Jun 2026 11:14:12 -0400
Subject: [PATCH 1/3] [SPIRV] Cooperative matrix under the Vulkan/Shader flavor
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit

Make SPV_KHR_cooperative_matrix reachable from the Vulkan/Shader flavor (the
OpenCL __spirv_* builtin path is isShader()-gated off), mirroring the texture /
ray-query intrinsic + GlobalISel pattern:

- IntrinsicsSPIRV.td: llvm.spv.cooperative.matrix.{load,store,muladd,splat};
  the opaque matrix is carried as llvm_any_ty (overloaded per concrete shape).
- SPIRVInstructionSelector.cpp: GlobalISel selection for the four ops
  (load/muladd/splat via selectOpWithSrcs, store via selectCoopMatrixStore).
  memory-layout/stride are <id> constant operands; float MulAdd omits the
  integer-only signedness literal; splat = single-scalar OpCompositeConstruct.
- SPIRVModuleAnalysis.{cpp,h}: cooperative matrix under Shader mandates the
  Vulkan memory model (spirv-val rejects Shader + CooperativeMatrixKHR under
  GLSL450). Rather than guess up front, derive the model from the actual
  requirement: after collectReqs, if a Shader module requires CooperativeMatrixKHR
  (RequirementHandler::isCapabilityRequired, added here), upgrade GLSL450 ->
  VulkanKHR (+ VulkanMemoryModelKHR capability + SPV_KHR_vulkan_memory_model
  extension, OpMemoryModel Logical VulkanKHR), unless !spirv.MemoryModel set one
  explicitly. Non-coop shaders keep GLSL450.

The OpTypeCooperativeMatrixKHR opaque type needs no backend change — the
existing BuiltinType machinery lowers target("spirv.CooperativeMatrixKHR", elem,
scope, rows, cols, use) flavor-agnostically.
---
 llvm/include/llvm/IR/IntrinsicsSPIRV.td       | 21 ++++-
 .../Target/SPIRV/SPIRVInstructionSelector.cpp | 33 ++++++++
 llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.cpp | 18 ++++-
 llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.h   |  4 +
 .../cooperative_matrix_kernel.ll              | 79 +++++++++++++++++++
 .../cooperative_matrix_ops_vulkan.ll          | 49 ++++++++++++
 .../cooperative_matrix_type_vulkan.ll         | 33 ++++++++
 7 files changed, 235 insertions(+), 2 deletions(-)
 create mode 100644 llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_kernel.ll
 create mode 100644 llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_ops_vulkan.ll
 create mode 100644 llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_type_vulkan.ll

diff --git a/llvm/include/llvm/IR/IntrinsicsSPIRV.td b/llvm/include/llvm/IR/IntrinsicsSPIRV.td
index 1e758f9f63f49..38868a4449b90 100644
--- a/llvm/include/llvm/IR/IntrinsicsSPIRV.td
+++ b/llvm/include/llvm/IR/IntrinsicsSPIRV.td
@@ -360,5 +360,24 @@ def int_spv_rsqrt : DefaultAttrsIntrinsic<[LLVMMatchType<0>], [llvm_anyfloat_ty]
   def int_spv_unpackhalf2x16 : DefaultAttrsIntrinsic<[llvm_anyfloat_ty], [llvm_i32_ty], [IntrNoMem]>;
   def int_spv_packhalf2x16 : DefaultAttrsIntrinsic<[llvm_anyint_ty], [llvm_anyfloat_ty], [IntrNoMem]>;
 
-
+  // SPV_KHR_cooperative_matrix operations under the Vulkan/Shader flavor. The
+  // KHR builtin (__spirv_*) path is isShader()-gated off (OpenCL-only), so the
+  // Shader flavor reaches these ops through llvm.spv intrinsics + GlobalISel
+  // selection (the texture / ray-query pattern). The matrix value itself is the
+  // opaque target("spirv.CooperativeMatrixKHR", elem, scope, rows, cols, use)
+  // type, carried as llvm_any_ty (overloaded/mangled per concrete shape).
+  //
+  // load:   result = OpCooperativeMatrixLoadKHR  ptr layout stride
+  // store:           OpCooperativeMatrixStoreKHR ptr matrix layout stride
+  // muladd: result = OpCooperativeMatrixMulAddKHR A B C  (float: no operands lit)
+  // splat:  result = OpCompositeConstruct        scalar  (broadcast accumulator)
+  // layout/stride are <id> operands (constants -> OpConstant), not inline imms.
+  def int_spv_cooperative_matrix_load
+    : Intrinsic<[llvm_any_ty], [llvm_anyptr_ty, llvm_i32_ty, llvm_i32_ty]>;
+  def int_spv_cooperative_matrix_store
+    : Intrinsic<[], [llvm_anyptr_ty, llvm_any_ty, llvm_i32_ty, llvm_i32_ty]>;
+  def int_spv_cooperative_matrix_muladd
+    : Intrinsic<[llvm_any_ty], [llvm_any_ty, llvm_any_ty, llvm_any_ty]>;
+  def int_spv_cooperative_matrix_splat
+    : Intrinsic<[llvm_any_ty], [llvm_any_ty]>;
 }
diff --git a/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp b/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
index c3f21fe025bd5..18592de2a4a26 100644
--- a/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
@@ -444,6 +444,7 @@ class SPIRVInstructionSelector : public InstructionSelector {
   bool selectGatherIntrinsic(Register &ResVReg, SPIRVTypeInst ResType,
                              MachineInstr &I) const;
   bool selectImageWriteIntrinsic(MachineInstr &I) const;
+  bool selectCoopMatrixStore(MachineInstr &I) const;
   bool selectResourceGetPointer(Register &ResVReg, SPIRVTypeInst ResType,
                                 MachineInstr &I) const;
   bool selectPushConstantGetPointer(Register &ResVReg, SPIRVTypeInst ResType,
@@ -1592,6 +1593,18 @@ bool SPIRVInstructionSelector::selectSincos(Register ResVReg,
   return false;
 }
 
+bool SPIRVInstructionSelector::selectCoopMatrixStore(MachineInstr &I) const {
+  // Void side-effecting G_INTRINSIC: operand 0 = intrinsic id, operands 1.. =
+  // pointer, matrix, memory_layout (<id> const), stride (<id> const).
+  // OpCooperativeMatrixStoreKHR has no result/result-type.
+  auto MIB = BuildMI(*I.getParent(), I, I.getDebugLoc(),
+                     TII.get(SPIRV::OpCooperativeMatrixStoreKHR));
+  for (unsigned i = 1; i < I.getNumOperands(); ++i)
+    MIB.addUse(I.getOperand(i).getReg());
+  MIB.constrainAllUses(TII, TRI, RBI);
+  return true;
+}
+
 bool SPIRVInstructionSelector::selectOpWithSrcs(Register ResVReg,
                                                 SPIRVTypeInst ResType,
                                                 MachineInstr &I,
@@ -4708,6 +4721,26 @@ bool SPIRVInstructionSelector::selectIntrinsic(Register ResVReg,
       report_fatal_error("incompatible result and operand types in a bitcast");
     return selectOpWithSrcs(ResVReg, ResType, I, {OpReg}, SPIRV::OpBitcast);
   }
+  case Intrinsic::spv_cooperative_matrix_load:
+    // result = OpCooperativeMatrixLoadKHR ptr memory_layout stride
+    return selectOpWithSrcs(ResVReg, ResType, I,
+                            {I.getOperand(2).getReg(), I.getOperand(3).getReg(),
+                             I.getOperand(4).getReg()},
+                            SPIRV::OpCooperativeMatrixLoadKHR);
+  case Intrinsic::spv_cooperative_matrix_store:
+    return selectCoopMatrixStore(I);
+  case Intrinsic::spv_cooperative_matrix_muladd:
+    // result = OpCooperativeMatrixMulAddKHR A B C  (no operands literal: the
+    // signedness mask is for integer matrices; float matmul omits it, valid).
+    return selectOpWithSrcs(ResVReg, ResType, I,
+                            {I.getOperand(2).getReg(), I.getOperand(3).getReg(),
+                             I.getOperand(4).getReg()},
+                            SPIRV::OpCooperativeMatrixMulAddKHR);
+  case Intrinsic::spv_cooperative_matrix_splat:
+    // result = OpCompositeConstruct scalar  (single-scalar construct broadcasts
+    // across the matrix — the zero/identity accumulator).
+    return selectOpWithSrcs(ResVReg, ResType, I, {I.getOperand(2).getReg()},
+                            SPIRV::OpCompositeConstruct);
   case Intrinsic::spv_unref_global:
   case Intrinsic::spv_init_global: {
     MachineInstr *MI = MRI->getVRegDef(I.getOperand(1).getReg());
diff --git a/llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.cpp b/llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.cpp
index 49d9dc95603ca..7e12c21c06540 100644
--- a/llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.cpp
@@ -155,7 +155,9 @@ void SPIRVModuleAnalysis::setBaseInfo(const Module &M) {
     MAI.Mem =
         static_cast<SPIRV::MemoryModel::MemoryModel>(getMetadataUInt(MemMD, 1));
   } else {
-    // TODO: Add support for VulkanMemoryModel.
+    // Default; a Shader module may be upgraded to VulkanKHR after requirements
+    // are collected (see runOnModule), since some capabilities mandate it.
+    // TODO: Add general support for VulkanMemoryModel.
     MAI.Mem = ST->isShader() ? SPIRV::MemoryModel::GLSL450
                              : SPIRV::MemoryModel::OpenCL;
     if (MAI.Mem == SPIRV::MemoryModel::OpenCL) {
@@ -3023,6 +3025,20 @@ bool SPIRVModuleAnalysis::runOnModule(Module &M) {
   collectReqs(M, MAI, MMI, *ST);
   collectDeclarations(M);
 
+  // CooperativeMatrixKHR in a Shader module mandates the Vulkan memory model
+  // (spirv-val rejects Shader + CooperativeMatrixKHR under GLSL450). Now that
+  // requirements are collected, derive the model from the capability rather
+  // than a default, unless one was set explicitly via !spirv.MemoryModel
+  // metadata.
+  if (ST->isShader() && MAI.Mem == SPIRV::MemoryModel::GLSL450 &&
+      !M.getNamedMetadata("spirv.MemoryModel") &&
+      MAI.Reqs.isCapabilityRequired(SPIRV::Capability::CooperativeMatrixKHR)) {
+    MAI.Mem = SPIRV::MemoryModel::VulkanKHR;
+    MAI.Reqs.getAndAddRequirements(SPIRV::OperandCategory::MemoryModelOperand,
+                                   MAI.Mem, *ST);
+    MAI.Reqs.addExtension(SPIRV::Extension::SPV_KHR_vulkan_memory_model);
+  }
+
   // Number rest of registers from N+1 onwards.
   numberRegistersGlobally(M);
 
diff --git a/llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.h b/llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.h
index 6559b5cfc7457..1fb16ad1af03c 100644
--- a/llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.h
+++ b/llvm/lib/Target/SPIRV/SPIRVModuleAnalysis.h
@@ -119,6 +119,10 @@ struct RequirementHandler {
   bool isCapabilityAvailable(Capability::Capability Cap) const {
     return AvailableCaps.contains(Cap);
   }
+  // True if Cap has been added to the module's required capabilities.
+  bool isCapabilityRequired(Capability::Capability Cap) const {
+    return AllCaps.contains(Cap);
+  }
 
   // Remove capability ToRemove, but only if IfPresent is present.
   void removeCapabilityIf(const Capability::Capability ToRemove,
diff --git a/llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_kernel.ll b/llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_kernel.ll
new file mode 100644
index 0000000000000..fadf9288f8e22
--- /dev/null
+++ b/llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_kernel.ll
@@ -0,0 +1,79 @@
+; cajeta-gpu cooperative-matrix increment CM3: a spirv-val-clean GLCompute kernel
+; that runs a cooperative-matrix matmul tile against DESCRIPTOR-BOUND storage
+; buffers under the Vulkan flavor.
+;
+; Unlike the emission tests (CM1/CM2), which pass matrices by value / use bare
+; function pointers, here A and B are read-only StorageBuffers and C a writable
+; StorageBuffer, each a VulkanBuffer handle decorated DescriptorSet/Binding and
+; accessed through the standard llvm.spv.resource.handlefrombinding + getpointer
+; path (exactly how Cajeta binds Buffer<T>). The kernel loads A (use 0) and B
+; (use 1) as cooperative matrices, mul-adds into a zero accumulator C (use 2), and
+; stores C. The whole module must pass `spirv-val --target-env vulkan1.3` — the
+; end-to-end validity proof before the Cajeta surface (CM4).
+
+; RUN: llc -O0 -verify-machineinstrs -mtriple=spirv-unknown-vulkan1.3-compute --spirv-ext=+SPV_KHR_cooperative_matrix,+SPV_KHR_vulkan_memory_model %s -o - | FileCheck %s
+; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv-unknown-vulkan1.3-compute --spirv-ext=+SPV_KHR_cooperative_matrix,+SPV_KHR_vulkan_memory_model %s -o - -filetype=obj | spirv-val --target-env vulkan1.3 %}
+
+; CHECK-DAG: OpCapability CooperativeMatrixKHR
+; CHECK-DAG: OpExtension "SPV_KHR_cooperative_matrix"
+; CHECK-DAG: OpEntryPoint GLCompute %[[#entry:]] "main"
+; CHECK-DAG: OpDecorate %[[#A:]] DescriptorSet 0
+; CHECK-DAG: OpDecorate %[[#A]] Binding 0
+; CHECK-DAG: OpDecorate %[[#B:]] DescriptorSet 0
+; CHECK-DAG: OpDecorate %[[#B]] Binding 1
+; CHECK-DAG: OpDecorate %[[#C:]] DescriptorSet 0
+; CHECK-DAG: OpDecorate %[[#C]] Binding 2
+; CHECK: %[[#MA:]] = OpCooperativeMatrixLoadKHR
+; CHECK: %[[#MB:]] = OpCooperativeMatrixLoadKHR
+; CHECK: %[[#MC0:]] = OpCompositeConstruct
+; CHECK: %[[#MC:]] = OpCooperativeMatrixMulAddKHR %[[#]] %[[#MA]] %[[#MB]] %[[#MC0]]
+; CHECK: OpCooperativeMatrixStoreKHR %[[#]] %[[#MC]]
+
+ at .str.a = private unnamed_addr constant [2 x i8] c"a\00", align 1
+ at .str.b = private unnamed_addr constant [2 x i8] c"b\00", align 1
+ at .str.c = private unnamed_addr constant [2 x i8] c"c\00", align 1
+
+define void @main() local_unnamed_addr #0 {
+entry:
+  ; A: read-only StorageBuffer, binding 0 -> MatrixA (use 0)
+  %ha = tail call target("spirv.VulkanBuffer", [0 x float], 12, 0)
+      @llvm.spv.resource.handlefrombinding.tspirv.VulkanBuffer_a0f32_12_0t(
+          i32 0, i32 0, i32 1, i32 0, ptr nonnull @.str.a)
+  %pa = tail call ptr addrspace(11)
+      @llvm.spv.resource.getpointer.p11.tspirv.VulkanBuffer_a0f32_12_0t(
+          target("spirv.VulkanBuffer", [0 x float], 12, 0) %ha, i32 0)
+  %a = call target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 0)
+      @llvm.spv.cooperative.matrix.load(ptr addrspace(11) %pa, i32 0, i32 16)
+
+  ; B: read-only StorageBuffer, binding 1 -> MatrixB (use 1)
+  %hb = tail call target("spirv.VulkanBuffer", [0 x float], 12, 0)
+      @llvm.spv.resource.handlefrombinding.tspirv.VulkanBuffer_a0f32_12_0t(
+          i32 0, i32 1, i32 1, i32 0, ptr nonnull @.str.b)
+  %pb = tail call ptr addrspace(11)
+      @llvm.spv.resource.getpointer.p11.tspirv.VulkanBuffer_a0f32_12_0t(
+          target("spirv.VulkanBuffer", [0 x float], 12, 0) %hb, i32 0)
+  %b = call target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 1)
+      @llvm.spv.cooperative.matrix.load(ptr addrspace(11) %pb, i32 0, i32 16)
+
+  ; C = A*B + 0  (accumulator, use 2), stored to writable StorageBuffer binding 2
+  %c0 = call target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 2)
+      @llvm.spv.cooperative.matrix.splat(float 0.000000e+00)
+  %c = call target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 2)
+      @llvm.spv.cooperative.matrix.muladd(
+          target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 0) %a,
+          target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 1) %b,
+          target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 2) %c0)
+  %hc = tail call target("spirv.VulkanBuffer", [0 x float], 12, 1)
+      @llvm.spv.resource.handlefrombinding.tspirv.VulkanBuffer_a0f32_12_1t(
+          i32 0, i32 2, i32 1, i32 0, ptr nonnull @.str.c)
+  %pc = tail call ptr addrspace(11)
+      @llvm.spv.resource.getpointer.p11.tspirv.VulkanBuffer_a0f32_12_1t(
+          target("spirv.VulkanBuffer", [0 x float], 12, 1) %hc, i32 0)
+  call void @llvm.spv.cooperative.matrix.store(
+          ptr addrspace(11) %pc,
+          target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 2) %c,
+          i32 0, i32 16)
+  ret void
+}
+
+attributes #0 = { "hlsl.numthreads"="1,1,1" "hlsl.shader"="compute" }
diff --git a/llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_ops_vulkan.ll b/llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_ops_vulkan.ll
new file mode 100644
index 0000000000000..8a47020e2f616
--- /dev/null
+++ b/llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_ops_vulkan.ll
@@ -0,0 +1,49 @@
+; cajeta-gpu cooperative-matrix increment CM2: the cooperative-matrix OPERATIONS
+; (load / store / mul-add / splat-construct) under the Vulkan/Shader flavor, reached
+; via llvm.spv.cooperative.matrix.* intrinsics + GlobalISel selection. The OpenCL
+; __spirv_CooperativeMatrix* builtin path is isShader()-gated off, so the Shader
+; flavor Cajeta emits cannot use it; these intrinsics are the texture / ray-query
+; pattern that runs for every flavor.
+;
+; A single matmul tile fragment: C = A * B + 0. A is MatrixA (use 0), B is MatrixB
+; (use 1), C the accumulator (use 2); all 16x16 f32, Subgroup scope (3). memory
+; layout 0 = RowMajorKHR, stride 16.
+;
+; Text-emission check (FileCheck only); the spirv-val-clean, descriptor-bound
+; compute kernel is increment CM3.
+
+; RUN: llc -verify-machineinstrs -O0 -mtriple=spirv-unknown-vulkan1.3-compute --spirv-ext=+SPV_KHR_cooperative_matrix,+SPV_KHR_vulkan_memory_model %s -o - | FileCheck %s
+
+; CHECK-DAG: OpCapability CooperativeMatrixKHR
+; CHECK-DAG: OpExtension "SPV_KHR_cooperative_matrix"
+; CHECK-DAG: %[[#F32:]] = OpTypeFloat 32
+; The three matrix types (A use 0, B use 1, C accumulator use 2) all lower to
+; OpTypeCooperativeMatrixKHR over the f32 component type.
+; CHECK-DAG: OpTypeCooperativeMatrixKHR %[[#F32]]
+; Data flow: mul-add consumes both loaded operands plus the splat-constructed
+; accumulator; the store consumes the mul-add result.
+; CHECK: %[[#A:]] = OpCooperativeMatrixLoadKHR
+; CHECK: %[[#B:]] = OpCooperativeMatrixLoadKHR
+; CHECK: %[[#C0:]] = OpCompositeConstruct
+; CHECK: %[[#C:]] = OpCooperativeMatrixMulAddKHR %[[#]] %[[#A]] %[[#B]] %[[#C0]]
+; CHECK: OpCooperativeMatrixStoreKHR %[[#]] %[[#C]]
+
+define spir_func void @matmul_tile(ptr %pa, ptr %pb, ptr %pc) {
+entry:
+  %a = call target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 0)
+       @llvm.spv.cooperative.matrix.load(ptr %pa, i32 0, i32 16)
+  %b = call target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 1)
+       @llvm.spv.cooperative.matrix.load(ptr %pb, i32 0, i32 16)
+  %c0 = call target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 2)
+        @llvm.spv.cooperative.matrix.splat(float 0.000000e+00)
+  %c = call target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 2)
+       @llvm.spv.cooperative.matrix.muladd(
+         target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 0) %a,
+         target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 1) %b,
+         target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 2) %c0)
+  call void @llvm.spv.cooperative.matrix.store(
+         ptr %pc,
+         target("spirv.CooperativeMatrixKHR", float, 3, 16, 16, 2) %c,
+         i32 0, i32 16)
+  ret void
+}
diff --git a/llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_type_vulkan.ll b/llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_type_vulkan.ll
new file mode 100644
index 0000000000000..442001c55f237
--- /dev/null
+++ b/llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_type_vulkan.ll
@@ -0,0 +1,33 @@
+; cajeta-gpu cooperative-matrix increment CM1: SPV_KHR_cooperative_matrix opaque
+; CooperativeMatrix TYPE under the Vulkan/Shader flavor. Proves the backend lowers
+; the parameterized target("spirv.CooperativeMatrixKHR", elem, scope, rows, cols, use)
+; to OpTypeCooperativeMatrixKHR under spirv-unknown-vulkan1.3-compute (the flavor
+; Cajeta emits), gated behind the CooperativeMatrixKHR capability + the
+; SPV_KHR_cooperative_matrix extension — and errors cleanly without it.
+;
+; Like the ray-query opaque types, the existing BuiltinType machinery lowers the
+; type flavor-agnostically, so no backend change is needed for the TYPE; this test
+; locks that in under Vulkan. The cooperative-matrix OPERATIONS (load/store/muladd/
+; length) reach the Shader flavor via llvm.spv.cooperative.matrix.* intrinsics in
+; increment CM2 (the OpenCL __spirv_* builtin path is isShader()-gated off).
+;
+; Text-emission check only: a type-only module forces the type via a by-value
+; parameter (which pulls in the Linkage capability), so it is intentionally not run
+; through spirv-val here. A spirv-val-clean, matrix-USING compute kernel arrives in
+; increment CM3.
+
+; RUN: not llc -O0 -mtriple=spirv-unknown-vulkan1.3-compute %s -o /dev/null 2>&1 | FileCheck %s --check-prefix=CHECK-ERROR
+; RUN: llc -verify-machineinstrs -O0 -mtriple=spirv-unknown-vulkan1.3-compute --spirv-ext=+SPV_KHR_cooperative_matrix,+SPV_KHR_vulkan_memory_model %s -o - | FileCheck %s
+
+; CHECK-ERROR: LLVM ERROR: OpTypeCooperativeMatrixKHR type requires the following SPIR-V extension: SPV_KHR_cooperative_matrix
+
+; CHECK-DAG: OpCapability CooperativeMatrixKHR
+; CHECK-DAG: OpExtension "SPV_KHR_cooperative_matrix"
+; CHECK-DAG: {{%[0-9]+}} = OpTypeCooperativeMatrixKHR
+
+; A by-value parameter of the opaque type forces OpTypeCooperativeMatrixKHR into the
+; module (referenced by OpTypeFunction — cannot be eliminated like an unused local).
+define spir_func void @use_coopmat(target("spirv.CooperativeMatrixKHR", i32, 3, 12, 12, 2) %m) {
+entry:
+  ret void
+}

>From 5907efed4fb86d9a71e7f62d69828be09ffae522 Mon Sep 17 00:00:00 2001
From: Julian Klappenbach <julian at twilight.digital>
Date: Sat, 6 Jun 2026 11:14:10 -0400
Subject: [PATCH 2/3] [SPIR-V] Deduce pointee type for all global variables,
 not only initialized ones
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit

SPIRVEmitIntrinsics::processGlobalValue only recorded a global variable's
element type (deduceElementTypeHelper) when hasInitializer() was true. For an
undef-initialized, non-constant aggregate global — e.g. a Workgroup `[N x T]`
shared tile (`@g = internal addrspace(3) global [N x T] undef`) — hasInitializer
returns false, so the variable's type was never recorded from its declared value
type. Its pointee type was then inferred lazily from a flat element-typed GEP use
(`getelementptr T, ptr @g, %i`) as the scalar `T`.

With the variable typed as a scalar pointer, the array-to-pointer-decay rewrite
in visitGetElementPtrInst (which prepends a 0 index so a Logical-SPIR-V access
chain can index the array) is skipped, because its DeducedPointeeTy is not an
ArrayType. The GEP stays a byte-offset pointer-arithmetic form Logical SPIR-V
cannot express, and the dynamic index is dropped: every invocation accesses
element 0. A Workgroup array indexed by a loop variable thus collapsed to a
scalar variable, silently corrupting workgroup-shared reductions/staging on
Vulkan (the emitted module passes spirv-val but computes wrong results).

A global variable's pointee type is concrete and authoritative, so record it for
every global in processGlobalValue (which runs before the use-based forward
pass). The deduction result is still ignored at the call site — TypedPointerType
isn't expressible in general LLVM IR; it is stored in the Global Registry.

(cherry picked from commit 2849c532820328544bee3ea3d8acd25289d3f457)
---
 llvm/lib/Target/SPIRV/SPIRVEmitIntrinsics.cpp | 19 ++++++++--
 .../type-deduce-global-array-undef.ll         | 37 +++++++++++++++++++
 2 files changed, 52 insertions(+), 4 deletions(-)
 create mode 100644 llvm/test/CodeGen/SPIRV/pointers/type-deduce-global-array-undef.ll

diff --git a/llvm/lib/Target/SPIRV/SPIRVEmitIntrinsics.cpp b/llvm/lib/Target/SPIRV/SPIRVEmitIntrinsics.cpp
index 97fa49d8836fb..89a25456830e4 100644
--- a/llvm/lib/Target/SPIRV/SPIRVEmitIntrinsics.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVEmitIntrinsics.cpp
@@ -2553,12 +2553,23 @@ void SPIRVEmitIntrinsics::processGlobalValue(GlobalVariable &GV,
   if (!shouldEmitIntrinsicsForGlobalValue(GVUsers, GV, CurrF))
     return;
 
+  // Deduce and record the variable's pointee type from its declared value type,
+  // for EVERY global — not only hasInitializer() ones. A global variable's type
+  // is concrete and authoritative; recording it here (this runs before the
+  // use-based forward pass) keeps an undef-initialized, non-constant aggregate
+  // — e.g. a Workgroup `[N x T]` shared tile — from being collapsed to its
+  // element type by flat element-typed GEP accesses. Without this,
+  // hasInitializer() excludes such globals (undef + non-constant), the
+  // variable's type is inferred from a `getelementptr T, ...` use as scalar
+  // `T`, the array-to-pointer-decay GEP rewrite in visitGetElementPtrInst is
+  // skipped (its ArrayType check fails), and under Logical SPIR-V the dynamic
+  // index is dropped — every invocation then accesses element 0. (Result
+  // ignored: TypedPointerType isn't expressible in general LLVM IR; it is
+  // stored in the Global Registry.)
+  deduceElementTypeHelper(&GV, false);
+
   Constant *Init = nullptr;
   if (hasInitializer(&GV)) {
-    // Deduce element type and store results in Global Registry.
-    // Result is ignored, because TypedPointerType is not supported
-    // by llvm IR general logic.
-    deduceElementTypeHelper(&GV, false);
     Init = GV.getInitializer();
     Value *InitOp = Init;
     if (isa<UndefValue>(Init) && Init->getType()->isAggregateType()) {
diff --git a/llvm/test/CodeGen/SPIRV/pointers/type-deduce-global-array-undef.ll b/llvm/test/CodeGen/SPIRV/pointers/type-deduce-global-array-undef.ll
new file mode 100644
index 0000000000000..8aa546a4e2a51
--- /dev/null
+++ b/llvm/test/CodeGen/SPIRV/pointers/type-deduce-global-array-undef.ll
@@ -0,0 +1,37 @@
+; A global variable's pointee type is concrete (its value type) and must be used
+; even for an undef-initialized, NON-constant aggregate — e.g. a Workgroup
+; `[N x T]` shared tile. Otherwise the variable's type is inferred from a flat
+; element-typed GEP use as the scalar element, the array-to-pointer-decay rewrite
+; is skipped, and under Logical SPIR-V the dynamic index is dropped (every
+; invocation accesses element 0). The variable must keep its array type and be
+; indexed with an OpAccessChain carrying the dynamic index.
+
+; RUN: llc -verify-machineinstrs -O0 -mtriple=spirv-unknown-vulkan1.3-compute %s -o - | FileCheck %s
+; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv-unknown-vulkan1.3-compute %s -o - -filetype=obj | spirv-val %}
+
+ at tile = internal addrspace(3) global [64 x i32] undef, align 16
+
+; CHECK-DAG: %[[#U32:]] = OpTypeInt 32 0
+; CHECK-DAG: %[[#Arr:]] = OpTypeArray %[[#U32]] %[[#]]
+; CHECK-DAG: %[[#ArrPtr:]] = OpTypePointer Workgroup %[[#Arr]]
+; CHECK-DAG: %[[#EltPtr:]] = OpTypePointer Workgroup %[[#U32]]
+; CHECK-DAG: %[[#Tile:]] = OpVariable %[[#ArrPtr]] Workgroup
+
+; The (dynamic) index survives as an OpAccessChain into the array tile; it is not
+; dropped, and the variable is not collapsed to a scalar `OpTypePointer Workgroup
+; %u32` with no index.
+; CHECK: %[[#Id:]] = OpCompositeExtract %[[#U32]] %[[#]] 0
+; CHECK: %[[#Ptr:]] = OpAccessChain %[[#EltPtr]] %[[#Tile]] %[[#Id]]
+; CHECK: OpStore %[[#Ptr]] %[[#]]
+
+define void @store_dynamic_index() #0 {
+entry:
+  %id = call i32 @llvm.spv.thread.id.in.group.i32(i32 0)
+  %p = getelementptr i32, ptr addrspace(3) @tile, i32 %id
+  store i32 42, ptr addrspace(3) %p, align 4
+  ret void
+}
+
+declare i32 @llvm.spv.thread.id.in.group.i32(i32)
+
+attributes #0 = { "hlsl.shader"="compute" "hlsl.numthreads"="64,1,1" }

>From fddd8513f3fb36c98246d0c00f1edd8cce4ace6a Mon Sep 17 00:00:00 2001
From: Julian Klappenbach <julian at twilight.digital>
Date: Sat, 6 Jun 2026 11:14:52 -0400
Subject: [PATCH 3/3] [SPIR-V] Access-chain aggregate pointers to element 0 for
 cooperative matrix load/store
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit

OpCooperativeMatrixLoadKHR / OpCooperativeMatrixStoreKHR require the Pointer to
point to a scalar or vector — the tile's element type. When the source is a
workgroup-shared array tile, however, the pointer reaching the selector is a
pointer to the whole `[N x T]` array: in opaque-pointer IR `&arr[0]` is the same
SSA value as `&arr`, and a zero-index element GEP is simplified back to the array
base during SPIRVEmitIntrinsics. The cooperative-matrix op was then emitted with
an array pointer, which spirv-val rejects ("Pointer's Type must be a scalar or
vector type").

Index such a pointer to its first element with an OpAccessChain before the
cooperative-matrix op, so it receives the required scalar pointer. Index 0 selects
the same address the MemoryLayout/Stride operands then walk from, so semantics are
unchanged. Pointers that are already element-typed (e.g. a StorageBuffer access
chain, or a dynamic-offset Workgroup access chain) are not aggregates and pass
through untouched.

Verified with llc + spirv-val on a Workgroup-tile cooperative-matrix load, and a
full LDS-staged GEMM (CoopStage copy -> Barrier -> load(Shared) -> mma) computing
bit-exact results on RADV / gfx1151.
---
 .../Target/SPIRV/SPIRVInstructionSelector.cpp | 56 ++++++++++++++++++-
 .../cooperative_matrix_workgroup_source.ll    | 42 ++++++++++++++
 2 files changed, 97 insertions(+), 1 deletion(-)
 create mode 100644 llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_workgroup_source.ll

diff --git a/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp b/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
index 18592de2a4a26..c797395b095dc 100644
--- a/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVInstructionSelector.cpp
@@ -445,6 +445,7 @@ class SPIRVInstructionSelector : public InstructionSelector {
                              MachineInstr &I) const;
   bool selectImageWriteIntrinsic(MachineInstr &I) const;
   bool selectCoopMatrixStore(MachineInstr &I) const;
+  Register coopMatrixElementPtr(Register PtrReg, MachineInstr &I) const;
   bool selectResourceGetPointer(Register &ResVReg, SPIRVTypeInst ResType,
                                 MachineInstr &I) const;
   bool selectPushConstantGetPointer(Register &ResVReg, SPIRVTypeInst ResType,
@@ -1593,13 +1594,58 @@ bool SPIRVInstructionSelector::selectSincos(Register ResVReg,
   return false;
 }
 
+// OpCooperativeMatrixLoad/StoreKHR require the Pointer to point to a scalar or
+// vector (the tile's element type). A workgroup/shared array tile, however,
+// reaches the selector as a pointer to the whole `[N x T]` array: in opaque-
+// pointer IR `&arr[0]` is the same SSA value as `&arr`, and a zero-index element
+// GEP is simplified back to the array base during SPIRVEmitIntrinsics — so by
+// selection the pointer's pointee type is the array. (A dynamic-offset access is
+// already an OpAccessChain to an element and is unaffected.) Index the array to
+// its first element so the cooperative-matrix op gets the scalar pointer it
+// requires; the access chain's index 0 selects the same address the layout/stride
+// operands then walk from. Returns PtrReg unchanged when it is not an aggregate.
+Register SPIRVInstructionSelector::coopMatrixElementPtr(Register PtrReg,
+                                                        MachineInstr &I) const {
+  SPIRVTypeInst PtrType = GR.getSPIRVTypeForVReg(PtrReg);
+  if (!PtrType)
+    return PtrReg;
+  SPIRVTypeInst PointeeType = GR.getPointeeType(PtrType);
+  if (!PointeeType || PointeeType->getOpcode() != SPIRV::OpTypeArray)
+    return PtrReg;
+  SPIRVTypeInst ElemType =
+      GR.getSPIRVTypeForVReg(PointeeType->getOperand(1).getReg());
+  if (!ElemType)
+    return PtrReg;
+  SPIRV::StorageClass::StorageClass SC = GR.getPointerStorageClass(PtrReg);
+  MachineIRBuilder MIRBuilder(I);
+  SPIRVTypeInst ElemPtrType =
+      GR.getOrCreateSPIRVPointerType(ElemType, MIRBuilder, SC);
+  SPIRVTypeInst I32Type = GR.getOrCreateSPIRVIntegerType(32, I, TII);
+  Register Zero = buildZerosVal(I32Type, I);
+  Register NewPtr = MRI->createVirtualRegister(GR.getRegClass(ElemPtrType));
+  GR.assignSPIRVTypeToVReg(ElemPtrType, NewPtr, *I.getParent()->getParent());
+  unsigned Opcode =
+      STI.isLogicalSPIRV() ? SPIRV::OpAccessChain : SPIRV::OpPtrAccessChain;
+  BuildMI(*I.getParent(), I, I.getDebugLoc(), TII.get(Opcode))
+      .addDef(NewPtr)
+      .addUse(GR.getSPIRVTypeID(ElemPtrType))
+      .addUse(PtrReg)
+      .addUse(Zero)
+      .constrainAllUses(TII, TRI, RBI);
+  return NewPtr;
+}
+
 bool SPIRVInstructionSelector::selectCoopMatrixStore(MachineInstr &I) const {
   // Void side-effecting G_INTRINSIC: operand 0 = intrinsic id, operands 1.. =
   // pointer, matrix, memory_layout (<id> const), stride (<id> const).
   // OpCooperativeMatrixStoreKHR has no result/result-type.
+  // Build the (possible) element access chain BEFORE the store, so its result is
+  // defined before the store that consumes it.
+  Register Ptr = coopMatrixElementPtr(I.getOperand(1).getReg(), I);
   auto MIB = BuildMI(*I.getParent(), I, I.getDebugLoc(),
                      TII.get(SPIRV::OpCooperativeMatrixStoreKHR));
-  for (unsigned i = 1; i < I.getNumOperands(); ++i)
+  MIB.addUse(Ptr);
+  for (unsigned i = 2; i < I.getNumOperands(); ++i)
     MIB.addUse(I.getOperand(i).getReg());
   MIB.constrainAllUses(TII, TRI, RBI);
   return true;
@@ -4723,10 +4769,18 @@ bool SPIRVInstructionSelector::selectIntrinsic(Register ResVReg,
   }
   case Intrinsic::spv_cooperative_matrix_load:
     // result = OpCooperativeMatrixLoadKHR ptr memory_layout stride
+<<<<<<< HEAD
     return selectOpWithSrcs(ResVReg, ResType, I,
                             {I.getOperand(2).getReg(), I.getOperand(3).getReg(),
                              I.getOperand(4).getReg()},
                             SPIRV::OpCooperativeMatrixLoadKHR);
+=======
+    return selectOpWithSrcs(
+        ResVReg, ResType, I,
+        {coopMatrixElementPtr(I.getOperand(2).getReg(), I),
+         I.getOperand(3).getReg(), I.getOperand(4).getReg()},
+        SPIRV::OpCooperativeMatrixLoadKHR);
+>>>>>>> 6114125dc940 ([SPIR-V] Access-chain aggregate pointers to element 0 for cooperative matrix load/store)
   case Intrinsic::spv_cooperative_matrix_store:
     return selectCoopMatrixStore(I);
   case Intrinsic::spv_cooperative_matrix_muladd:
diff --git a/llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_workgroup_source.ll b/llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_workgroup_source.ll
new file mode 100644
index 0000000000000..8ba3257d28bfa
--- /dev/null
+++ b/llvm/test/CodeGen/SPIRV/extensions/SPV_KHR_cooperative_matrix/cooperative_matrix_workgroup_source.ll
@@ -0,0 +1,42 @@
+; A cooperative-matrix load/store whose source is a workgroup-shared array tile.
+;
+; In opaque-pointer IR the source pointer is the whole `[N x T]` Workgroup array
+; (`&arr[0]` is the same SSA value as `&arr`, and a zero-index element GEP folds to
+; the base in SPIRVEmitIntrinsics). OpCooperativeMatrixLoad/StoreKHR require the
+; Pointer to point to a scalar/vector, so the selector must access-chain the array
+; to element 0 before the cooperative-matrix op. (Regression test for the
+; cooperative-matrix aggregate-pointer access-chain fix.)
+
+; RUN: llc -verify-machineinstrs -O0 -mtriple=spirv-unknown-vulkan1.3-compute --spirv-ext=+SPV_KHR_cooperative_matrix,+SPV_KHR_vulkan_memory_model %s -o - | FileCheck %s
+; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv-unknown-vulkan1.3-compute --spirv-ext=+SPV_KHR_cooperative_matrix,+SPV_KHR_vulkan_memory_model %s -o - -filetype=obj | spirv-val %}
+
+ at tile = internal addrspace(3) global [256 x half] undef, align 16
+
+; CHECK-DAG: %[[#Half:]] = OpTypeFloat 16
+; CHECK-DAG: %[[#U32:]] = OpTypeInt 32 0
+; The undef, non-constant Workgroup global keeps its array type (it is not
+; collapsed to a scalar pointer).
+; CHECK-DAG: %[[#Arr:]] = OpTypeArray %[[#Half]] %[[#]]
+; CHECK-DAG: %[[#ArrPtr:]] = OpTypePointer Workgroup %[[#Arr]]
+; CHECK-DAG: %[[#EltPtr:]] = OpTypePointer Workgroup %[[#Half]]
+; CHECK-DAG: %[[#Tile:]] = OpVariable %[[#ArrPtr]] Workgroup
+; CHECK-DAG: %[[#Zero:]] = OpConstant %[[#U32]] 0
+
+; The array tile is access-chained to element 0 before the load, and that scalar
+; pointer is what the cooperative-matrix load consumes.
+; CHECK: %[[#LdPtr:]] = OpAccessChain %[[#EltPtr]] %[[#Tile]] %[[#Zero]]
+; CHECK: %[[#Mat:]] = OpCooperativeMatrixLoadKHR %[[#]] %[[#LdPtr]]
+; CHECK: %[[#StPtr:]] = OpAccessChain %[[#EltPtr]] %[[#Tile]] %[[#Zero]]
+; CHECK: OpCooperativeMatrixStoreKHR %[[#StPtr]] %[[#Mat]]
+
+define void @coop_matrix_workgroup_source() #0 {
+entry:
+  %m = call target("spirv.CooperativeMatrixKHR", half, 3, 16, 16, 0)
+       @llvm.spv.cooperative.matrix.load(ptr addrspace(3) @tile, i32 0, i32 16)
+  call void
+       @llvm.spv.cooperative.matrix.store(ptr addrspace(3) @tile,
+       target("spirv.CooperativeMatrixKHR", half, 3, 16, 16, 0) %m, i32 0, i32 16)
+  ret void
+}
+
+attributes #0 = { "hlsl.shader"="compute" "hlsl.numthreads"="1,1,1" }



More information about the llvm-commits mailing list