[Mlir-commits] [mlir] [mlir][SPIR-V][NFC] Add pushCaps/pushExts helpers to type visitors (PR #195796)

Arseniy Obolenskiy llvmlistbot at llvm.org
Mon May 4 23:13:21 PDT 2026


https://github.com/aobolensk created https://github.com/llvm/llvm-project/pull/195796

Discussed in https://github.com/llvm/llvm-project/pull/195060#discussion_r3183984853 

>From 41698315e550233333c227bb837ce00056da883d Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Tue, 5 May 2026 08:12:28 +0200
Subject: [PATCH] [mlir][SPIR-V][NFC] Add pushCaps/pushExts helpers to type
 visitors

---
 mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp | 114 ++++++++++-------------
 1 file changed, 47 insertions(+), 67 deletions(-)

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



More information about the Mlir-commits mailing list