[Mlir-commits] [mlir] [mlir][SPIR-V] Support literal struct type in spirv.Constant (PR #198414)
Arseniy Obolenskiy
llvmlistbot at llvm.org
Tue May 19 01:56:18 PDT 2026
https://github.com/aobolensk updated https://github.com/llvm/llvm-project/pull/198414
>From e7e398ff1df64e0925cf28a843b363bab644f83e Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Tue, 19 May 2026 00:52:45 +0200
Subject: [PATCH 1/3] [mlir][SPIR-V] Support literal struct type in
spirv.Constant
---
mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp | 24 +++++++++++++--
mlir/test/Dialect/SPIRV/IR/structure-ops.mlir | 30 ++++++++++++++++++-
2 files changed, 50 insertions(+), 4 deletions(-)
diff --git a/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp b/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp
index 90d0895988204..119cdcf0fe4fa 100644
--- a/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp
+++ b/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp
@@ -555,7 +555,7 @@ ParseResult spirv::ConstantOp::parse(OpAsmParser &parser,
void spirv::ConstantOp::print(OpAsmPrinter &printer) {
printer << ' ' << getValue();
- if (isa<spirv::ArrayType>(getType()))
+ if (isa<spirv::ArrayType, spirv::StructType>(getType()))
printer << " : " << getType();
}
@@ -610,10 +610,27 @@ static LogicalResult verifyConstantType(spirv::ConstantOp op, Attribute value,
return success();
}
if (auto arrayAttr = dyn_cast<ArrayAttr>(value)) {
+ if (auto structType = dyn_cast<spirv::StructType>(opType)) {
+ // Identified (possibly recursive) structs are not supported as constants.
+ if (structType.isIdentified())
+ return op.emitOpError(
+ "cannot have an identified struct as a constant type");
+ if (arrayAttr.size() != structType.getNumElements())
+ return op.emitOpError("number of constituents (")
+ << arrayAttr.size()
+ << ") does not match number of struct members ("
+ << structType.getNumElements() << ")";
+ for (auto [idx, element] : llvm::enumerate(arrayAttr.getValue())) {
+ if (failed(verifyConstantType(op, element,
+ structType.getElementType(idx))))
+ return failure();
+ }
+ return success();
+ }
auto arrayType = dyn_cast<spirv::ArrayType>(opType);
if (!arrayType)
return op.emitOpError(
- "must have spirv.array result type for array value");
+ "must have spirv.array or spirv.struct result type for array value");
Type elemType = arrayType.getElementType();
for (Attribute element : arrayAttr.getValue()) {
// Verify array elements recursively.
@@ -638,7 +655,8 @@ bool spirv::ConstantOp::isBuildableWith(Type type) {
return false;
if (isa<SPIRVDialect>(type.getDialect())) {
- // TODO: support constant struct
+ if (auto structType = dyn_cast<spirv::StructType>(type))
+ return !structType.isIdentified();
return isa<spirv::ArrayType>(type);
}
diff --git a/mlir/test/Dialect/SPIRV/IR/structure-ops.mlir b/mlir/test/Dialect/SPIRV/IR/structure-ops.mlir
index 17400a045a7ae..614e4d0a49e53 100644
--- a/mlir/test/Dialect/SPIRV/IR/structure-ops.mlir
+++ b/mlir/test/Dialect/SPIRV/IR/structure-ops.mlir
@@ -66,6 +66,8 @@ func.func @const() -> () {
// CHECK: spirv.Constant dense<4.200000e+00> : !spirv.coopmatrix<16x16xf32, Subgroup, MatrixAcc>
// CHECK: spirv.Constant dense<0> : !spirv.coopmatrix<16x16xi8, Subgroup, MatrixAcc>
// CHECK: spirv.Constant dense<4> : !spirv.coopmatrix<16x16xi8, Subgroup, MatrixAcc>
+ // CHECK: spirv.Constant [1 : i32, 2.000000e+00 : f32] : !spirv.struct<(i32, f32)>
+ // CHECK: spirv.Constant [1 : i32, [dense<2> : vector<2xi32>]] : !spirv.struct<(i32, !spirv.array<1 x vector<2xi32>>)>
%0 = spirv.Constant true
%1 = spirv.Constant 42 : i32
@@ -81,6 +83,8 @@ func.func @const() -> () {
%11 = spirv.Constant dense<4.200000e+00> : !spirv.coopmatrix<16x16xf32, Subgroup, MatrixAcc>
%12 = spirv.Constant dense<0> : !spirv.coopmatrix<16x16xi8, Subgroup, MatrixAcc>
%13 = spirv.Constant dense<4> : !spirv.coopmatrix<16x16xi8, Subgroup, MatrixAcc>
+ %14 = spirv.Constant [1 : i32, 2.0 : f32] : !spirv.struct<(i32, f32)>
+ %15 = spirv.Constant [1 : i32, [dense<2> : vector<2xi32>]] : !spirv.struct<(i32, !spirv.array<1 x vector<2xi32>>)>
return
}
@@ -103,7 +107,7 @@ func.func @array_constant() -> () {
// -----
func.func @array_constant() -> () {
- // expected-error @+1 {{must have spirv.array result type for array value}}
+ // expected-error @+1 {{must have spirv.array or spirv.struct result type for array value}}
%0 = spirv.Constant [dense<3.0> : vector<2xf32>] : !spirv.rtarray<vector<2xf32>>
return
}
@@ -148,6 +152,30 @@ func.func @coop_matrix_const_non_splat() -> () {
// -----
+func.func @struct_constant_wrong_member_count() -> () {
+ // expected-error @+1 {{number of constituents (1) does not match number of struct members (2)}}
+ %0 = spirv.Constant [1 : i32] : !spirv.struct<(i32, f32)>
+ return
+}
+
+// -----
+
+func.func @struct_constant_wrong_member_type() -> () {
+ // expected-error @+1 {{result type ('f32') does not match value type ('i32')}}
+ %0 = spirv.Constant [1 : i32, 2 : i32] : !spirv.struct<(i32, f32)>
+ return
+}
+
+// -----
+
+func.func @struct_constant_identified() -> () {
+ // expected-error @+1 {{cannot have an identified struct as a constant type}}
+ %0 = spirv.Constant [1 : i32] : !spirv.struct<S, (i32)>
+ return
+}
+
+// -----
+
func.func @coop_matrix_const_non_dense() -> () {
// expected-error @+2 {{floating point value not valid for specified type}}
%0 = spirv.Constant 0.000000e+00 : !spirv.coopmatrix<16x16xf32, Subgroup, MatrixAcc>
>From aa8f1604386eac87833f595084a447269aecffb2 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Tue, 19 May 2026 09:15:11 +0200
Subject: [PATCH 2/3] Address comment
---
.../Target/SPIRV/Deserialization/Deserializer.cpp | 2 +-
mlir/lib/Target/SPIRV/Serialization/Serializer.cpp | 7 ++++---
mlir/test/Target/SPIRV/constant.mlir | 14 ++++++++++++++
3 files changed, 19 insertions(+), 4 deletions(-)
diff --git a/mlir/lib/Target/SPIRV/Deserialization/Deserializer.cpp b/mlir/lib/Target/SPIRV/Deserialization/Deserializer.cpp
index cabbecc567241..f27c954a43f3a 100644
--- a/mlir/lib/Target/SPIRV/Deserialization/Deserializer.cpp
+++ b/mlir/lib/Target/SPIRV/Deserialization/Deserializer.cpp
@@ -1912,7 +1912,7 @@ spirv::Deserializer::processConstantComposite(ArrayRef<uint32_t> operands) {
// For normal constants, we just record the attribute (and its type) for
// later materialization at use sites.
constantMap.try_emplace(resultID, attr, shapedType);
- } else if (auto arrayType = dyn_cast<spirv::ArrayType>(resultType)) {
+ } else if (isa<spirv::ArrayType, spirv::StructType>(resultType)) {
auto attr = opBuilder.getArrayAttr(elements);
constantMap.try_emplace(resultID, attr, resultType);
} else {
diff --git a/mlir/lib/Target/SPIRV/Serialization/Serializer.cpp b/mlir/lib/Target/SPIRV/Serialization/Serializer.cpp
index 7a2eaf3df43f1..efd4146e84c78 100644
--- a/mlir/lib/Target/SPIRV/Serialization/Serializer.cpp
+++ b/mlir/lib/Target/SPIRV/Serialization/Serializer.cpp
@@ -1062,9 +1062,10 @@ uint32_t Serializer::prepareArrayConstant(Location loc, Type constType,
uint32_t resultID = getNextID();
SmallVector<uint32_t, 4> operands = {typeID, resultID};
operands.reserve(attr.size() + 2);
- auto elementType = cast<spirv::ArrayType>(constType).getElementType();
- for (Attribute elementAttr : attr) {
- if (auto elementID = prepareConstant(loc, elementType, elementAttr)) {
+ auto compositeType = cast<spirv::CompositeType>(constType);
+ for (auto [idx, elementAttr] : llvm::enumerate(attr)) {
+ if (auto elementID = prepareConstant(loc, compositeType.getElementType(idx),
+ elementAttr)) {
operands.push_back(elementID);
} else {
return 0;
diff --git a/mlir/test/Target/SPIRV/constant.mlir b/mlir/test/Target/SPIRV/constant.mlir
index cc7c93824c6e3..816a6271759a5 100644
--- a/mlir/test/Target/SPIRV/constant.mlir
+++ b/mlir/test/Target/SPIRV/constant.mlir
@@ -348,5 +348,19 @@ spirv.module Logical Vulkan requires #spirv.vce<v1.3,
spirv.ReturnValue %coop : !spirv.coopmatrix<16x16xi8, Subgroup, MatrixAcc>
}
+ // CHECK-LABEL: @struct_const
+ spirv.func @struct_const() -> (!spirv.struct<(i32, f32)>) "None" {
+ // CHECK: spirv.Constant [1 : i32, 2.000000e+00 : f32] : !spirv.struct<(i32, f32)>
+ %0 = spirv.Constant [1 : i32, 2.0 : f32] : !spirv.struct<(i32, f32)>
+ spirv.ReturnValue %0 : !spirv.struct<(i32, f32)>
+ }
+
+ // CHECK-LABEL: @struct_const_nested
+ spirv.func @struct_const_nested() -> (!spirv.struct<(i32, !spirv.array<2 x i32>)>) "None" {
+ // CHECK: spirv.Constant [1 : i32, [2 : i32, 3 : i32]] : !spirv.struct<(i32, !spirv.array<2 x i32>)>
+ %0 = spirv.Constant [1 : i32, [2 : i32, 3 : i32]] : !spirv.struct<(i32, !spirv.array<2 x i32>)>
+ spirv.ReturnValue %0 : !spirv.struct<(i32, !spirv.array<2 x i32>)>
+ }
+
spirv.EntryPoint "GLCompute" @bool_const
}
>From f2d88bf9687be14a02ae935e51fc7299ba847898 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Tue, 19 May 2026 10:56:06 +0200
Subject: [PATCH 3/3] Fix nit
---
mlir/lib/Target/SPIRV/Serialization/Serializer.cpp | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/mlir/lib/Target/SPIRV/Serialization/Serializer.cpp b/mlir/lib/Target/SPIRV/Serialization/Serializer.cpp
index efd4146e84c78..1e3b7aae7fe0b 100644
--- a/mlir/lib/Target/SPIRV/Serialization/Serializer.cpp
+++ b/mlir/lib/Target/SPIRV/Serialization/Serializer.cpp
@@ -1062,7 +1062,7 @@ uint32_t Serializer::prepareArrayConstant(Location loc, Type constType,
uint32_t resultID = getNextID();
SmallVector<uint32_t, 4> operands = {typeID, resultID};
operands.reserve(attr.size() + 2);
- auto compositeType = cast<spirv::CompositeType>(constType);
+ spirv::CompositeType compositeType = cast<spirv::CompositeType>(constType);
for (auto [idx, elementAttr] : llvm::enumerate(attr)) {
if (auto elementID = prepareConstant(loc, compositeType.getElementType(idx),
elementAttr)) {
More information about the Mlir-commits
mailing list