[Mlir-commits] [mlir] [mlir][SPIRV] Add named-barrier type and OpNamedBarrierInitialize / OpMemoryNamedBarrier (PR #195664)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Mon May 4 08:08:49 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-mlir

Author: Krzysztof Drewniak (krzysz00)

<details>
<summary>Changes</summary>

Adds the SPIR-V named-barrier object (TypeNamedBarrier) along with NamedBarrierInitialize and MemoryNamedBarrier ops, gated on the NamedBarrier capability and SPIR-V 1.1+.

---
Full diff: https://github.com/llvm/llvm-project/pull/195664.diff


11 Files Affected:

- (modified) mlir/include/mlir/Dialect/SPIRV/IR/SPIRVBarrierOps.td (+99) 
- (modified) mlir/include/mlir/Dialect/SPIRV/IR/SPIRVBase.td (+8) 
- (modified) mlir/include/mlir/Dialect/SPIRV/IR/SPIRVTypes.h (+11) 
- (modified) mlir/lib/Dialect/SPIRV/IR/SPIRVDialect.cpp (+9-2) 
- (modified) mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp (+14) 
- (modified) mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp (+17-4) 
- (modified) mlir/lib/Target/SPIRV/Deserialization/DeserializeOps.cpp (+1) 
- (modified) mlir/lib/Target/SPIRV/Deserialization/Deserializer.cpp (+12) 
- (modified) mlir/lib/Target/SPIRV/Deserialization/Deserializer.h (+2) 
- (modified) mlir/lib/Target/SPIRV/Serialization/Serializer.cpp (+5) 
- (modified) mlir/test/Dialect/SPIRV/IR/barrier-ops.mlir (+34) 


``````````diff
diff --git a/mlir/include/mlir/Dialect/SPIRV/IR/SPIRVBarrierOps.td b/mlir/include/mlir/Dialect/SPIRV/IR/SPIRVBarrierOps.td
index 1459330549066..10d8c9993d6e2 100644
--- a/mlir/include/mlir/Dialect/SPIRV/IR/SPIRVBarrierOps.td
+++ b/mlir/include/mlir/Dialect/SPIRV/IR/SPIRVBarrierOps.td
@@ -108,4 +108,103 @@ def SPIRV_MemoryBarrierOp : SPIRV_Op<"MemoryBarrier", []> {
   let assemblyFormat = "$memory_scope `,` $memory_semantics attr-dict";
 }
 
+// -----
+
+def SPIRV_NamedBarrierInitializeOp
+    : SPIRV_Op<"NamedBarrierInitialize", []> {
+  let summary = "Initialize a named barrier with a subgroup count.";
+
+  let description = [{
+    Declare a new named-barrier object.
+
+    Result Type must be the type OpTypeNamedBarrier.
+
+    Subgroup Count must be a 32-bit integer type scalar representing the
+    number of subgroups that must reach the current point of execution.
+
+    Capability: NamedBarrier
+
+    Missing before version 1.1.
+
+    #### Example:
+
+    ```mlir
+    %nb = spirv.NamedBarrierInitialize %subgroup_count : i32 -> !spirv.named_barrier
+    ```
+  }];
+
+  let availability = [
+    MinVersion<SPIRV_V_1_1>,
+    MaxVersion<SPIRV_V_1_6>,
+    Extension<[]>,
+    Capability<[SPIRV_C_NamedBarrier]>
+  ];
+
+  let arguments = (ins
+    SPIRV_Int32:$subgroup_count
+  );
+
+  let results = (outs
+    SPIRV_AnyNamedBarrier:$result
+  );
+
+  let assemblyFormat = [{
+    $subgroup_count attr-dict `:` type($subgroup_count) `->` type($result)
+  }];
+}
+
+// -----
+
+def SPIRV_MemoryNamedBarrierOp
+    : SPIRV_Op<"MemoryNamedBarrier", []> {
+  let summary = [{
+    Wait on a named barrier and apply a memory barrier.
+  }];
+
+  let description = [{
+    Wait for other invocations of this module to reach the current point
+    of execution.
+
+    Named Barrier must be the type OpTypeNamedBarrier.
+
+    If Semantics is not None, this instruction also serves as an
+    OpMemoryBarrier instruction, and also performs and adheres to the
+    description and semantics of an OpMemoryBarrier instruction with the
+    same Memory and Semantics operands. This allows atomically specifying
+    both a control barrier and a memory barrier (that is, without needing
+    two instructions). If Semantics is None, Memory is ignored.
+
+    Capability: NamedBarrier
+
+    Missing before version 1.1.
+
+    #### Example:
+
+    ```mlir
+    spirv.MemoryNamedBarrier %nb, <Workgroup>, <AcquireRelease|WorkgroupMemory>
+        : !spirv.named_barrier
+    ```
+  }];
+
+  let availability = [
+    MinVersion<SPIRV_V_1_1>,
+    MaxVersion<SPIRV_V_1_6>,
+    Extension<[]>,
+    Capability<[SPIRV_C_NamedBarrier]>
+  ];
+
+  let arguments = (ins
+    SPIRV_AnyNamedBarrier:$named_barrier,
+    SPIRV_ScopeAttr:$memory_scope,
+    SPIRV_MemorySemanticsAttr:$memory_semantics
+  );
+
+  let results = (outs);
+
+  let assemblyFormat = [{
+    $named_barrier `,` $memory_scope `,` $memory_semantics attr-dict
+        `:` type($named_barrier)
+  }];
+}
+
 #endif // MLIR_DIALECT_SPIRV_IR_BARRIER_OPS
diff --git a/mlir/include/mlir/Dialect/SPIRV/IR/SPIRVBase.td b/mlir/include/mlir/Dialect/SPIRV/IR/SPIRVBase.td
index b5cc5f6ba8cb5..c13d24942a9e9 100644
--- a/mlir/include/mlir/Dialect/SPIRV/IR/SPIRVBase.td
+++ b/mlir/include/mlir/Dialect/SPIRV/IR/SPIRVBase.td
@@ -4272,6 +4272,7 @@ def SPIRV_IsPtrType : CPred<"::llvm::isa<::mlir::spirv::PointerType>($_self)">;
 def SPIRV_IsRTArrayType : CPred<"::llvm::isa<::mlir::spirv::RuntimeArrayType>($_self)">;
 def SPIRV_IsSampledImageType : CPred<"::llvm::isa<::mlir::spirv::SampledImageType>($_self)">;
 def SPIRV_IsSamplerType : CPred<"::llvm::isa<::mlir::spirv::SamplerType>($_self)">;
+def SPIRV_IsNamedBarrierType : CPred<"::llvm::isa<::mlir::spirv::NamedBarrierType>($_self)">;
 def SPIRV_IsStructType : CPred<"::llvm::isa<::mlir::spirv::StructType>($_self)">;
 def SPIRV_IsTensorArmType : CPred<"::llvm::isa<::mlir::spirv::TensorArmType>($_self)">;
 
@@ -4317,6 +4318,8 @@ def SPIRV_AnySampledImage : DialectType<SPIRV_Dialect, SPIRV_IsSampledImageType,
                                 "any SPIR-V sampled image type">;
 def SPIRV_AnySampler : DialectType<SPIRV_Dialect, SPIRV_IsSamplerType,
                                 "any SPIR-V sampler type">;
+def SPIRV_AnyNamedBarrier : DialectType<SPIRV_Dialect, SPIRV_IsNamedBarrierType,
+                                "any SPIR-V named barrier type">;
 def SPIRV_AnyTensorArm : DialectType<SPIRV_Dialect, SPIRV_IsTensorArmType,
                                  "any SPIR-V tensorArm type">;
 
@@ -4605,6 +4608,9 @@ def SPIRV_OC_OpGroupFMax                      : I32EnumAttrCase<"OpGroupFMax", 2
 def SPIRV_OC_OpGroupUMax                      : I32EnumAttrCase<"OpGroupUMax", 270>;
 def SPIRV_OC_OpGroupSMax                      : I32EnumAttrCase<"OpGroupSMax", 271>;
 def SPIRV_OC_OpNoLine                         : I32EnumAttrCase<"OpNoLine", 317>;
+def SPIRV_OC_OpTypeNamedBarrier               : I32EnumAttrCase<"OpTypeNamedBarrier", 327>;
+def SPIRV_OC_OpNamedBarrierInitialize         : I32EnumAttrCase<"OpNamedBarrierInitialize", 328>;
+def SPIRV_OC_OpMemoryNamedBarrier             : I32EnumAttrCase<"OpMemoryNamedBarrier", 329>;
 def SPIRV_OC_OpModuleProcessed                : I32EnumAttrCase<"OpModuleProcessed", 330>;
 def SPIRV_OC_OpExecutionModeId                : I32EnumAttrCase<"OpExecutionModeId", 331>;
 def SPIRV_OC_OpDecorateId                     : I32EnumAttrCase<"OpDecorateId", 332>;
@@ -4757,6 +4763,8 @@ def SPIRV_OpcodeAttr :
       SPIRV_OC_OpGroupSMin, SPIRV_OC_OpGroupFMax, SPIRV_OC_OpGroupUMax,
       SPIRV_OC_OpGroupSMax, SPIRV_OC_OpNoLine, SPIRV_OC_OpModuleProcessed,
       SPIRV_OC_OpExecutionModeId, SPIRV_OC_OpDecorateId,
+      SPIRV_OC_OpTypeNamedBarrier, SPIRV_OC_OpNamedBarrierInitialize,
+      SPIRV_OC_OpMemoryNamedBarrier,
       SPIRV_OC_OpGroupNonUniformElect, SPIRV_OC_OpGroupNonUniformAll,
       SPIRV_OC_OpGroupNonUniformAny, SPIRV_OC_OpGroupNonUniformAllEqual,
       SPIRV_OC_OpGroupNonUniformBroadcast, SPIRV_OC_OpGroupNonUniformBroadcastFirst,
diff --git a/mlir/include/mlir/Dialect/SPIRV/IR/SPIRVTypes.h b/mlir/include/mlir/Dialect/SPIRV/IR/SPIRVTypes.h
index b7890ff101b2b..70060cfda5fde 100644
--- a/mlir/include/mlir/Dialect/SPIRV/IR/SPIRVTypes.h
+++ b/mlir/include/mlir/Dialect/SPIRV/IR/SPIRVTypes.h
@@ -240,6 +240,17 @@ class SamplerType : public Type::TypeBase<SamplerType, SPIRVType, TypeStorage> {
   static SamplerType get(MLIRContext *context);
 };
 
+// SPIR-V named barrier type (OpTypeNamedBarrier)
+class NamedBarrierType
+    : public Type::TypeBase<NamedBarrierType, SPIRVType, TypeStorage> {
+public:
+  using Base::Base;
+
+  static constexpr StringLiteral name = "spirv.named_barrier";
+
+  static NamedBarrierType get(MLIRContext *context);
+};
+
 /// SPIR-V struct type. Two kinds of struct types are supported:
 /// - Literal: a literal struct type is uniqued by its fields (types + offset
 /// info + decoration info).
diff --git a/mlir/lib/Dialect/SPIRV/IR/SPIRVDialect.cpp b/mlir/lib/Dialect/SPIRV/IR/SPIRVDialect.cpp
index 036d48f0fd637..2b5e7a571f42d 100644
--- a/mlir/lib/Dialect/SPIRV/IR/SPIRVDialect.cpp
+++ b/mlir/lib/Dialect/SPIRV/IR/SPIRVDialect.cpp
@@ -863,6 +863,8 @@ Type SPIRVDialect::parseType(DialectAsmParser &parser) const {
     return parseSampledImageType(*this, parser);
   if (keyword == "sampler")
     return SamplerType::get(getContext());
+  if (keyword == "named_barrier")
+    return NamedBarrierType::get(getContext());
   if (keyword == "struct")
     return parseStructType(*this, parser);
   if (keyword == "matrix")
@@ -911,6 +913,10 @@ static void print(SampledImageType type, DialectAsmPrinter &os) {
 
 static void print(SamplerType type, DialectAsmPrinter &os) { os << "sampler"; }
 
+static void print(NamedBarrierType type, DialectAsmPrinter &os) {
+  os << "named_barrier";
+}
+
 static void print(StructType type, DialectAsmPrinter &os) {
   FailureOr<AsmPrinter::CyclicPrintReset> cyclicPrint;
 
@@ -1005,8 +1011,9 @@ static void print(TensorArmType type, DialectAsmPrinter &os) {
 void SPIRVDialect::printType(Type type, DialectAsmPrinter &os) const {
   TypeSwitch<Type>(type)
       .Case<ArrayType, CooperativeMatrixType, PointerType, RuntimeArrayType,
-            ImageType, SampledImageType, SamplerType, StructType, MatrixType,
-            TensorArmType>([&](auto type) { print(type, os); })
+            ImageType, SampledImageType, SamplerType, NamedBarrierType,
+            StructType, MatrixType, TensorArmType>(
+          [&](auto type) { print(type, os); })
       .DefaultUnreachable("Unhandled SPIR-V type");
 }
 
diff --git a/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp b/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp
index f8aa5bfeba452..84b30208e60dc 100644
--- a/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp
+++ b/mlir/lib/Dialect/SPIRV/IR/SPIRVOps.cpp
@@ -1502,6 +1502,20 @@ LogicalResult spirv::MemoryBarrierOp::verify() {
   return verifyMemorySemantics(getOperation(), getMemorySemantics());
 }
 
+//===----------------------------------------------------------------------===//
+// spirv.MemoryNamedBarrierOp
+//===----------------------------------------------------------------------===//
+
+LogicalResult spirv::MemoryNamedBarrierOp::verify() {
+  return verifyMemorySemantics(getOperation(), getMemorySemantics());
+}
+
+//===----------------------------------------------------------------------===//
+// spirv.NamedBarrierInitializeOp
+//===----------------------------------------------------------------------===//
+
+LogicalResult spirv::NamedBarrierInitializeOp::verify() { return success(); }
+
 //===----------------------------------------------------------------------===//
 // spirv.module
 //===----------------------------------------------------------------------===//
diff --git a/mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp b/mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp
index aafc2180761d0..2ad463674c85e 100644
--- a/mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp
+++ b/mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp
@@ -57,7 +57,7 @@ class TypeExtensionVisitor {
           for (Type elementType : concreteType.getElementTypes())
             add(elementType);
         })
-        .Case<SamplerType>([](auto) { /* no extensions */ })
+        .Case<SamplerType, NamedBarrierType>([](auto) { /* no extensions */ })
         .DefaultUnreachable("Unhandled type");
   }
 
@@ -109,6 +109,11 @@ 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);
+        })
         .DefaultUnreachable("Unhandled type");
   }
 
@@ -816,6 +821,14 @@ SamplerType SamplerType::get(MLIRContext *context) {
   return Base::get(context);
 }
 
+//===----------------------------------------------------------------------===//
+// NamedBarrierType
+//===----------------------------------------------------------------------===//
+
+NamedBarrierType NamedBarrierType::get(MLIRContext *context) {
+  return Base::get(context);
+}
+
 //===----------------------------------------------------------------------===//
 // StructType
 //===----------------------------------------------------------------------===//
@@ -1352,7 +1365,7 @@ TensorArmType::verifyInvariants(function_ref<InFlightDiagnostic()> emitError,
 //===----------------------------------------------------------------------===//
 
 void SPIRVDialect::registerTypes() {
-  addTypes<ArrayType, CooperativeMatrixType, ImageType, MatrixType, PointerType,
-           RuntimeArrayType, SampledImageType, SamplerType, StructType,
-           TensorArmType>();
+  addTypes<ArrayType, CooperativeMatrixType, ImageType, MatrixType,
+           NamedBarrierType, PointerType, RuntimeArrayType, SampledImageType,
+           SamplerType, StructType, TensorArmType>();
 }
diff --git a/mlir/lib/Target/SPIRV/Deserialization/DeserializeOps.cpp b/mlir/lib/Target/SPIRV/Deserialization/DeserializeOps.cpp
index c12647d4255aa..21a8400c8e40f 100644
--- a/mlir/lib/Target/SPIRV/Deserialization/DeserializeOps.cpp
+++ b/mlir/lib/Target/SPIRV/Deserialization/DeserializeOps.cpp
@@ -184,6 +184,7 @@ LogicalResult spirv::Deserializer::processInstruction(
   case spirv::Opcode::OpTypeFunction:
   case spirv::Opcode::OpTypeImage:
   case spirv::Opcode::OpTypeSampler:
+  case spirv::Opcode::OpTypeNamedBarrier:
   case spirv::Opcode::OpTypeSampledImage:
   case spirv::Opcode::OpTypeRuntimeArray:
   case spirv::Opcode::OpTypeStruct:
diff --git a/mlir/lib/Target/SPIRV/Deserialization/Deserializer.cpp b/mlir/lib/Target/SPIRV/Deserialization/Deserializer.cpp
index 440ce262ac594..965af5c1bdc9a 100644
--- a/mlir/lib/Target/SPIRV/Deserialization/Deserializer.cpp
+++ b/mlir/lib/Target/SPIRV/Deserialization/Deserializer.cpp
@@ -1212,6 +1212,8 @@ LogicalResult spirv::Deserializer::processType(spirv::Opcode opcode,
     return processImageType(operands);
   case spirv::Opcode::OpTypeSampler:
     return processSamplerType(operands);
+  case spirv::Opcode::OpTypeNamedBarrier:
+    return processNamedBarrierType(operands);
   case spirv::Opcode::OpTypeSampledImage:
     return processSampledImageType(operands);
   case spirv::Opcode::OpTypeRuntimeArray:
@@ -1705,6 +1707,16 @@ spirv::Deserializer::processSamplerType(ArrayRef<uint32_t> operands) {
   return success();
 }
 
+LogicalResult
+spirv::Deserializer::processNamedBarrierType(ArrayRef<uint32_t> operands) {
+  if (operands.size() != 1)
+    return emitError(unknownLoc,
+                     "OpTypeNamedBarrier must have no parameters");
+
+  typeMap[operands[0]] = spirv::NamedBarrierType::get(context);
+  return success();
+}
+
 //===----------------------------------------------------------------------===//
 // Constant
 //===----------------------------------------------------------------------===//
diff --git a/mlir/lib/Target/SPIRV/Deserialization/Deserializer.h b/mlir/lib/Target/SPIRV/Deserialization/Deserializer.h
index 9725c63deb8c2..f3dab422ef9d9 100644
--- a/mlir/lib/Target/SPIRV/Deserialization/Deserializer.h
+++ b/mlir/lib/Target/SPIRV/Deserialization/Deserializer.h
@@ -324,6 +324,8 @@ class Deserializer {
 
   LogicalResult processSamplerType(ArrayRef<uint32_t> operands);
 
+  LogicalResult processNamedBarrierType(ArrayRef<uint32_t> operands);
+
   LogicalResult processRuntimeArrayType(ArrayRef<uint32_t> operands);
 
   LogicalResult processStructType(ArrayRef<uint32_t> operands);
diff --git a/mlir/lib/Target/SPIRV/Serialization/Serializer.cpp b/mlir/lib/Target/SPIRV/Serialization/Serializer.cpp
index e29a437cca87f..5c9e378b4b7f3 100644
--- a/mlir/lib/Target/SPIRV/Serialization/Serializer.cpp
+++ b/mlir/lib/Target/SPIRV/Serialization/Serializer.cpp
@@ -750,6 +750,11 @@ LogicalResult Serializer::prepareBasicType(
     return success();
   }
 
+  if (isa<spirv::NamedBarrierType>(type)) {
+    typeEnum = spirv::Opcode::OpTypeNamedBarrier;
+    return success();
+  }
+
   if (auto sampledImageType = dyn_cast<spirv::SampledImageType>(type)) {
     typeEnum = spirv::Opcode::OpTypeSampledImage;
     uint32_t imageTypeID = 0;
diff --git a/mlir/test/Dialect/SPIRV/IR/barrier-ops.mlir b/mlir/test/Dialect/SPIRV/IR/barrier-ops.mlir
index e0f935b7ce39d..5e05687887760 100644
--- a/mlir/test/Dialect/SPIRV/IR/barrier-ops.mlir
+++ b/mlir/test/Dialect/SPIRV/IR/barrier-ops.mlir
@@ -48,3 +48,37 @@ func.func @memory_barrier_2() -> () {
   return
 }
 
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.NamedBarrierInitialize
+//===----------------------------------------------------------------------===//
+
+func.func @named_barrier_initialize(%member_count : i32) -> () {
+  // CHECK: %{{.*}} = spirv.NamedBarrierInitialize %[[MEMBER_COUNT:.*]] : i32 -> !spirv.named_barrier
+  %nb = spirv.NamedBarrierInitialize %member_count : i32 -> !spirv.named_barrier
+  return
+}
+
+// -----
+
+//===----------------------------------------------------------------------===//
+// spirv.MemoryNamedBarrier
+//===----------------------------------------------------------------------===//
+
+func.func @memory_named_barrier(%member_count : i32) -> () {
+  %nb = spirv.NamedBarrierInitialize %member_count : i32 -> !spirv.named_barrier
+  // CHECK: spirv.MemoryNamedBarrier %{{.*}}, <Workgroup>, <AcquireRelease|WorkgroupMemory> : !spirv.named_barrier
+  spirv.MemoryNamedBarrier %nb, <Workgroup>, <AcquireRelease|WorkgroupMemory> : !spirv.named_barrier
+  return
+}
+
+// -----
+
+func.func @memory_named_barrier_invalid_semantics(%member_count : i32) -> () {
+  %nb = spirv.NamedBarrierInitialize %member_count : i32 -> !spirv.named_barrier
+  // expected-error @+1 {{expected at most one of these four memory constraints to be set: `Acquire`, `Release`,`AcquireRelease` or `SequentiallyConsistent`}}
+  spirv.MemoryNamedBarrier %nb, <Workgroup>, <Acquire|Release> : !spirv.named_barrier
+  return
+}
+

``````````

</details>


https://github.com/llvm/llvm-project/pull/195664


More information about the Mlir-commits mailing list