[Mlir-commits] [mlir] [mlir][SPIR-V] Allow SpecConstantComposite constituents to reference other SpecConstantComposites (PR #193416)
Arseniy Obolenskiy
llvmlistbot at llvm.org
Mon Apr 27 23:09:32 PDT 2026
https://github.com/aobolensk updated https://github.com/llvm/llvm-project/pull/193416
>From 7684bd1c64a024ec68565e12cbf97da2fa488080 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Wed, 22 Apr 2026 08:05:45 +0200
Subject: [PATCH 1/2] [mlir][SPIR-V] Allow SpecConstantComposite constituents
to reference other SpecConstantComposites
The verifier for spirv.SpecConstantComposite previously assumed all constituents were spirv.SpecConstant ops, which caused a crash when referencing nested spirv.SpecConstantComposite ops
Per the SPIR-V spec (s3.3.7, OpSpecConstantComposite), constituents "must be the <id>s of other specialization constants, constant declarations, or an OpUndef", which includes OpSpecConstantComposite
---
mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp | 26 +++++++---
mlir/test/Dialect/SPIRV/IR/structure-ops.mlir | 51 +++++++++++++++++++
2 files changed, 71 insertions(+), 6 deletions(-)
diff --git a/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp b/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp
index 9300483a0f92f..86c06d301b47c 100644
--- a/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp
+++ b/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp
@@ -1851,15 +1851,29 @@ LogicalResult spirv::SpecConstantCompositeOp::verify() {
for (auto index : llvm::seq<uint32_t>(0, constituents.size())) {
auto constituent = cast<FlatSymbolRefAttr>(constituents[index]);
- auto constituentSpecConstOp =
- dyn_cast<spirv::SpecConstantOp>(SymbolTable::lookupNearestSymbolFrom(
- (*this)->getParentOp(), constituent.getAttr()));
+ auto *constituentOp = SymbolTable::lookupNearestSymbolFrom(
+ (*this)->getParentOp(), constituent.getAttr());
+
+ if (!constituentOp)
+ return emitError("unknown constituent symbol ") << constituent.getAttr();
+
+ Type constituentType;
+ if (auto specConstOp = dyn_cast<spirv::SpecConstantOp>(constituentOp)) {
+ constituentType = specConstOp.getDefaultValue().getType();
+ } else if (auto specConstCompositeOp =
+ dyn_cast<spirv::SpecConstantCompositeOp>(constituentOp)) {
+ constituentType = specConstCompositeOp.getType();
+ } else {
+ return emitError("unsupported constituent ")
+ << constituent.getAttr()
+ << ": must reference a spirv.SpecConstant or "
+ "spirv.SpecConstantComposite";
+ }
- if (constituentSpecConstOp.getDefaultValue().getType() !=
- cType.getElementType(index))
+ if (constituentType != cType.getElementType(index))
return emitError("has incorrect types of operands: expected ")
<< cType.getElementType(index) << ", but provided "
- << constituentSpecConstOp.getDefaultValue().getType();
+ << constituentType;
}
return success();
diff --git a/mlir/test/Dialect/SPIRV/IR/structure-ops.mlir b/mlir/test/Dialect/SPIRV/IR/structure-ops.mlir
index 991530458bd26..17400a045a7ae 100644
--- a/mlir/test/Dialect/SPIRV/IR/structure-ops.mlir
+++ b/mlir/test/Dialect/SPIRV/IR/structure-ops.mlir
@@ -968,6 +968,57 @@ spirv.module Logical GLSL450 {
spirv.SpecConstantComposite @scc (@sc1, @sc2, @sc3) : vector<3xf32>
}
+// -----
+
+// Nested composite: array of arrays
+spirv.module Logical GLSL450 {
+ spirv.SpecConstant @sc1 = 1.5 : f32
+ spirv.SpecConstant @sc2 = 2.5 : f32
+ spirv.SpecConstantComposite @scc_inner (@sc1, @sc2) : !spirv.array<2 x f32>
+ // CHECK: spirv.SpecConstantComposite @scc_nested (@scc_inner, @scc_inner) : !spirv.array<2 x !spirv.array<2 x f32>>
+ spirv.SpecConstantComposite @scc_nested (@scc_inner, @scc_inner) : !spirv.array<2 x !spirv.array<2 x f32>>
+}
+
+// -----
+
+// Struct with composite and scalar constituents
+spirv.module Logical GLSL450 {
+ spirv.SpecConstant @sc1 = 1 : i32
+ spirv.SpecConstant @sc2 = 2.5 : f32
+ spirv.SpecConstant @sc3 = 3.5 : f32
+ spirv.SpecConstantComposite @scc_vec (@sc2, @sc3) : vector<2xf32>
+ // CHECK: spirv.SpecConstantComposite @scc_struct (@sc1, @scc_vec) : !spirv.struct<(i32, vector<2xf32>)>
+ spirv.SpecConstantComposite @scc_struct (@sc1, @scc_vec) : !spirv.struct<(i32, vector<2xf32>)>
+}
+
+// -----
+
+// Type mismatch with composite constituent
+spirv.module Logical GLSL450 {
+ spirv.SpecConstant @sc1 = 1.5 : f32
+ spirv.SpecConstant @sc2 = 2.5 : f32
+ spirv.SpecConstantComposite @scc_inner (@sc1, @sc2) : !spirv.array<2 x f32>
+ // expected-error @+1 {{has incorrect types of operands: expected '!spirv.array<3 x f32>', but provided '!spirv.array<2 x f32>'}}
+ spirv.SpecConstantComposite @scc_bad (@scc_inner) : !spirv.array<1 x !spirv.array<3 x f32>>
+}
+
+// -----
+
+// Unsupported constituent (not a SpecConstant or SpecConstantComposite)
+spirv.module Logical GLSL450 {
+ spirv.GlobalVariable @gv : !spirv.ptr<f32, Private>
+ // expected-error @+1 {{unsupported constituent "gv": must reference a spirv.SpecConstant or spirv.SpecConstantComposite}}
+ spirv.SpecConstantComposite @scc (@gv) : !spirv.array<1 x f32>
+}
+
+// -----
+
+// Unknown constituent symbol
+spirv.module Logical GLSL450 {
+ // expected-error @+1 {{unknown constituent symbol "does_not_exist"}}
+ spirv.SpecConstantComposite @scc (@does_not_exist) : !spirv.array<1 x f32>
+}
+
//===----------------------------------------------------------------------===//
// spirv.SpecConstantComposite (spirv.KHR.coopmatrix)
//===----------------------------------------------------------------------===//
>From aeee9dbbfdb0cd20c21059be9667ca7e211ffef0 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Tue, 28 Apr 2026 08:09:17 +0200
Subject: [PATCH 2/2] Address review comment
---
mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp b/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp
index 86c06d301b47c..49a14b5b30f0f 100644
--- a/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp
+++ b/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp
@@ -1851,7 +1851,7 @@ LogicalResult spirv::SpecConstantCompositeOp::verify() {
for (auto index : llvm::seq<uint32_t>(0, constituents.size())) {
auto constituent = cast<FlatSymbolRefAttr>(constituents[index]);
- auto *constituentOp = SymbolTable::lookupNearestSymbolFrom(
+ Operation *constituentOp = SymbolTable::lookupNearestSymbolFrom(
(*this)->getParentOp(), constituent.getAttr());
if (!constituentOp)
More information about the Mlir-commits
mailing list