[Mlir-commits] [mlir] [mlir][SPIR-V][NFC] Add pushCaps/pushExts helpers to type visitors (PR #195796)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Mon May 4 23:14:04 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir-spirv
Author: Arseniy Obolenskiy (aobolensk)
<details>
<summary>Changes</summary>
Discussed in https://github.com/llvm/llvm-project/pull/195060#discussion_r3183984853
---
Full diff: https://github.com/llvm/llvm-project/pull/195796.diff
1 Files Affected:
- (modified) mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp (+47-67)
``````````diff
diff --git a/mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp b/mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp
index 2ad463674c85e..734f312476871 100644
--- a/mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp
+++ b/mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp
@@ -70,6 +70,12 @@ class TypeExtensionVisitor {
void addConcrete(ScalarType type);
void addConcrete(TensorArmType type);
+ template <Extension... Es>
+ void pushExts() {
+ static constexpr Extension exts[] = {Es...};
+ extensions.push_back(exts);
+ }
+
SPIRVType::ExtensionArrayRefVector &extensions;
std::optional<StorageClass> storage;
llvm::SmallDenseSet<std::pair<Type, std::optional<StorageClass>>> seen;
@@ -109,11 +115,8 @@ class TypeCapabilityVisitor {
add(elementType);
})
.Case<SamplerType>([](auto) { /* no capabilities */ })
- .Case<NamedBarrierType>([this](auto) {
- static const Capability caps[] = {Capability::NamedBarrier};
- ArrayRef<Capability> ref(caps, std::size(caps));
- capabilities.push_back(ref);
- })
+ .Case<NamedBarrierType>(
+ [this](auto) { pushCaps<Capability::NamedBarrier>(); })
.DefaultUnreachable("Unhandled type");
}
@@ -130,6 +133,12 @@ class TypeCapabilityVisitor {
void addConcrete(TensorArmType type);
void addConcrete(VectorType type);
+ template <Capability... Cs>
+ void pushCaps() {
+ static constexpr Capability caps[] = {Cs...};
+ capabilities.push_back(caps);
+ }
+
SPIRVType::CapabilityArrayRefVector &capabilities;
std::optional<StorageClass> storage;
llvm::SmallDenseSet<std::pair<Type, std::optional<StorageClass>>> seen;
@@ -224,10 +233,8 @@ void TypeCapabilityVisitor::addConcrete(VectorType type) {
add(type.getElementType());
int64_t vecSize = type.getNumElements();
- if (vecSize == 8 || vecSize == 16) {
- static constexpr auto cap = Capability::Vector16;
- capabilities.push_back(cap);
- }
+ if (vecSize == 8 || vecSize == 16)
+ pushCaps<Capability::Vector16>();
}
//===----------------------------------------------------------------------===//
@@ -307,23 +314,17 @@ CooperativeMatrixUseKHR CooperativeMatrixType::getUse() const {
void TypeExtensionVisitor::addConcrete(CooperativeMatrixType type) {
add(type.getElementType());
- static constexpr auto ext = Extension::SPV_KHR_cooperative_matrix;
- extensions.push_back(ext);
+ pushExts<Extension::SPV_KHR_cooperative_matrix>();
}
void TypeCapabilityVisitor::addConcrete(CooperativeMatrixType type) {
Type elementType = type.getElementType();
add(elementType);
- static constexpr auto caps = Capability::CooperativeMatrixKHR;
- capabilities.push_back(caps);
- if (elementType.isBF16()) {
- static constexpr auto caps = Capability::BFloat16CooperativeMatrixKHR;
- capabilities.push_back(caps);
- }
- if (elementType.isF8E4M3FN() || elementType.isF8E5M2()) {
- static constexpr auto caps = Capability::Float8CooperativeMatrixEXT;
- capabilities.push_back(caps);
- }
+ pushCaps<Capability::CooperativeMatrixKHR>();
+ if (elementType.isBF16())
+ pushCaps<Capability::BFloat16CooperativeMatrixKHR>();
+ if (elementType.isF8E4M3FN() || elementType.isF8E5M2())
+ pushCaps<Capability::Float8CooperativeMatrixEXT>();
}
//===----------------------------------------------------------------------===//
@@ -536,8 +537,7 @@ unsigned RuntimeArrayType::getArrayStride() const { return getImpl()->stride; }
void TypeCapabilityVisitor::addConcrete(RuntimeArrayType type) {
add(type.getElementType());
- static constexpr auto cap = Capability::Shader;
- capabilities.push_back(cap);
+ pushCaps<Capability::Shader>();
}
//===----------------------------------------------------------------------===//
@@ -565,15 +565,11 @@ bool ScalarType::isValid(IntegerType type) {
}
void TypeExtensionVisitor::addConcrete(ScalarType type) {
- if (type.isBF16()) {
- static constexpr auto ext = Extension::SPV_KHR_bfloat16;
- extensions.push_back(ext);
- }
+ if (type.isBF16())
+ pushExts<Extension::SPV_KHR_bfloat16>();
- if (type.isF8E4M3FN() || type.isF8E5M2()) {
- static constexpr auto ext = Extension::SPV_EXT_float8;
- extensions.push_back(ext);
- }
+ if (type.isF8E4M3FN() || type.isF8E5M2())
+ pushExts<Extension::SPV_EXT_float8>();
// 8- or 16-bit integer/floating-point numbers will require extra extensions
// to appear in interface storage classes. See SPV_KHR_16bit_storage and
@@ -585,17 +581,13 @@ void TypeExtensionVisitor::addConcrete(ScalarType type) {
case StorageClass::PushConstant:
case StorageClass::StorageBuffer:
case StorageClass::Uniform:
- if (type.getIntOrFloatBitWidth() == 8) {
- static constexpr auto ext = Extension::SPV_KHR_8bit_storage;
- extensions.push_back(ext);
- }
+ if (type.getIntOrFloatBitWidth() == 8)
+ pushExts<Extension::SPV_KHR_8bit_storage>();
[[fallthrough]];
case StorageClass::Input:
case StorageClass::Output:
- if (type.getIntOrFloatBitWidth() == 16) {
- static constexpr auto ext = Extension::SPV_KHR_16bit_storage;
- extensions.push_back(ext);
- }
+ if (type.getIntOrFloatBitWidth() == 16)
+ pushExts<Extension::SPV_KHR_16bit_storage>();
break;
default:
break;
@@ -612,13 +604,11 @@ void TypeCapabilityVisitor::addConcrete(ScalarType type) {
#define STORAGE_CASE(storage, cap8, cap16) \
case StorageClass::storage: { \
if (bitwidth == 8) { \
- static constexpr auto cap = Capability::cap8; \
- capabilities.push_back(cap); \
+ pushCaps<Capability::cap8>(); \
return; \
} \
if (bitwidth == 16) { \
- static constexpr auto cap = Capability::cap16; \
- capabilities.push_back(cap); \
+ pushCaps<Capability::cap16>(); \
return; \
} \
/* For 64-bit integers/floats, Int64/Float64 enables support for all */ \
@@ -637,8 +627,7 @@ void TypeCapabilityVisitor::addConcrete(ScalarType type) {
case StorageClass::Input:
case StorageClass::Output: {
if (bitwidth == 16) {
- static constexpr auto cap = Capability::StorageInputOutput16;
- capabilities.push_back(cap);
+ pushCaps<Capability::StorageInputOutput16>();
return;
}
break;
@@ -653,10 +642,9 @@ void TypeCapabilityVisitor::addConcrete(ScalarType type) {
// capabilities for special bitwidths.
#define WIDTH_CASE(type, width) \
- case width: { \
- static constexpr auto cap = Capability::type##width; \
- capabilities.push_back(cap); \
- } break
+ case width: \
+ pushCaps<Capability::type##width>(); \
+ break
if (auto intType = dyn_cast<IntegerType>(type)) {
switch (bitwidth) {
@@ -673,22 +661,17 @@ void TypeCapabilityVisitor::addConcrete(ScalarType type) {
assert(isa<FloatType>(type));
switch (bitwidth) {
case 8: {
- if (type.isF8E4M3FN() || type.isF8E5M2()) {
- static constexpr auto cap = Capability::Float8EXT;
- capabilities.push_back(cap);
- } else {
+ if (type.isF8E4M3FN() || type.isF8E5M2())
+ pushCaps<Capability::Float8EXT>();
+ else
llvm_unreachable("invalid 8-bit float type to getCapabilities");
- }
break;
}
case 16: {
- if (type.isBF16()) {
- static constexpr auto cap = Capability::BFloat16TypeKHR;
- capabilities.push_back(cap);
- } else {
- static constexpr auto cap = Capability::Float16;
- capabilities.push_back(cap);
- }
+ if (type.isBF16())
+ pushCaps<Capability::BFloat16TypeKHR>();
+ else
+ pushCaps<Capability::Float16>();
break;
}
WIDTH_CASE(Float, 64);
@@ -1288,8 +1271,7 @@ ArrayRef<int64_t> MatrixType::getShape() const { return getImpl()->shape; }
void TypeCapabilityVisitor::addConcrete(MatrixType type) {
add(type.getColumnType());
- static constexpr auto cap = Capability::Matrix;
- capabilities.push_back(cap);
+ pushCaps<Capability::Matrix>();
}
//===----------------------------------------------------------------------===//
@@ -1337,14 +1319,12 @@ ArrayRef<int64_t> TensorArmType::getShape() const { return getImpl()->shape; }
void TypeExtensionVisitor::addConcrete(TensorArmType type) {
add(type.getElementType());
- static constexpr auto ext = Extension::SPV_ARM_tensors;
- extensions.push_back(ext);
+ pushExts<Extension::SPV_ARM_tensors>();
}
void TypeCapabilityVisitor::addConcrete(TensorArmType type) {
add(type.getElementType());
- static constexpr auto cap = Capability::TensorsARM;
- capabilities.push_back(cap);
+ pushCaps<Capability::TensorsARM>();
}
LogicalResult
``````````
</details>
https://github.com/llvm/llvm-project/pull/195796
More information about the Mlir-commits
mailing list