[Mlir-commits] [mlir] [MLIR][NVVM] Add Rubin extensions to tcgen05.commit Op (PR #215125)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Sun Aug 9 11:07:05 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir-llvm
Author: Rajat Bajpai (rajatbajpai)
<details>
<summary>Changes</summary>
This change adds support for 32-bit multicast mask and tracking of only shared-memory reads of A-matrix performed by prior MMA ops.
---
Full diff: https://github.com/llvm/llvm-project/pull/215125.diff
4 Files Affected:
- (modified) mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td (+7-3)
- (modified) mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp (+18-10)
- (added) mlir/test/Target/LLVMIR/nvvm/tcgen05-commit-smem-a-read.mlir (+49)
- (modified) mlir/test/Target/LLVMIR/nvvm/tcgen05-commit.mlir (+27-11)
``````````diff
diff --git a/mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td b/mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td
index ef5d276720be6..0c799b7ac8765 100644
--- a/mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td
+++ b/mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td
@@ -5315,15 +5315,19 @@ def NVVM_Tcgen05CommitOp : NVVM_Op<"tcgen05.commit", [NVVMRequiresSMf<[100, 101,
The multicast variants allow signaling on the *mbarrier objects*
of multiple CTAs within the cluster. Operand `multicastMask`,
when present, specifies the destination CTAs in the cluster such
- that each bit position in the 16-bit `multicastMask` operand
+ that each bit position in the 16-bit or 32-bit `multicastMask` operand
corresponds to the `nvvm.read.ptx.sreg.ctaid` of the destination CTA.
+ When present, the `smem_a_read` attribute restricts tracking to
+ shared-memory reads of matrix A performed by prior `tcgen05.mma`
+ operations.
[For more information, see PTX ISA](https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen-async-sync-operations-commit)
}];
let arguments = (ins
AnyTypeOf<[LLVM_AnyPointer, LLVM_PointerShared]>:$addr,
- Optional<I16>:$multicastMask,
- DefaultValuedAttr<CTAGroupKindAttr, "CTAGroupKind::CTA_1">:$group);
+ Optional<AnyTypeOf<[I16, I32]>>:$multicastMask,
+ DefaultValuedAttr<CTAGroupKindAttr, "CTAGroupKind::CTA_1">:$group,
+ UnitAttr:$smem_a_read);
let assemblyFormat = [{
$addr (`,` `multicast_mask` `=` $multicastMask^)?
diff --git a/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp b/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp
index ab51f1e5fe797..9fe46b8bf6f90 100644
--- a/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp
+++ b/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp
@@ -5002,13 +5002,6 @@ llvm::Intrinsic::ID Tcgen05DeallocOp::getIntrinsicIDAndArgs(
return id;
}
-#define TCGEN05_COMMIT_IMPL(cg, mc) \
- llvm::Intrinsic::nvvm_tcgen05_commit##mc##_##cg
-
-#define GET_TCGEN05_COMMIT_ID(cta_group, has_mc) \
- has_mc ? TCGEN05_COMMIT_IMPL(cta_group, _mc) \
- : TCGEN05_COMMIT_IMPL(cta_group, )
-
llvm::Intrinsic::ID
Tcgen05CommitOp::getIntrinsicIDAndArgs(Operation &op,
LLVM::ModuleTranslation &mt,
@@ -5016,11 +5009,26 @@ Tcgen05CommitOp::getIntrinsicIDAndArgs(Operation &op,
auto curOp = cast<NVVM::Tcgen05CommitOp>(op);
bool hasMulticast = static_cast<bool>(curOp.getMulticastMask());
bool is2CTAMode = curOp.getGroup() == CTAGroupKind::CTA_2;
+ bool hasSmemARead = curOp.getSmemARead();
+ unsigned index = (static_cast<unsigned>(hasSmemARead) << 1) |
+ static_cast<unsigned>(is2CTAMode);
+
+ using namespace llvm::Intrinsic;
+ static constexpr ID IDs[] = {
+ nvvm_tcgen05_commit_cg1,
+ nvvm_tcgen05_commit_cg2,
+ nvvm_tcgen05_commit_smem_a_read_cg1,
+ nvvm_tcgen05_commit_smem_a_read_cg2,
+ };
- llvm::Intrinsic::ID id = is2CTAMode
- ? GET_TCGEN05_COMMIT_ID(cg2, hasMulticast)
- : GET_TCGEN05_COMMIT_ID(cg1, hasMulticast);
+ static constexpr ID multicastIDs[] = {
+ nvvm_tcgen05_commit_mc_cg1,
+ nvvm_tcgen05_commit_mc_cg2,
+ nvvm_tcgen05_commit_smem_a_read_mc_cg1,
+ nvvm_tcgen05_commit_smem_a_read_mc_cg2,
+ };
+ ID id = hasMulticast ? multicastIDs[index] : IDs[index];
// Fill the Intrinsic Args
args.push_back(mt.lookupValue(curOp.getAddr()));
if (hasMulticast)
diff --git a/mlir/test/Target/LLVMIR/nvvm/tcgen05-commit-smem-a-read.mlir b/mlir/test/Target/LLVMIR/nvvm/tcgen05-commit-smem-a-read.mlir
new file mode 100644
index 0000000000000..f483dde87aabe
--- /dev/null
+++ b/mlir/test/Target/LLVMIR/nvvm/tcgen05-commit-smem-a-read.mlir
@@ -0,0 +1,49 @@
+// RUN: mlir-translate -mlir-to-llvmir %s | FileCheck %s
+
+// CHECK-LABEL: @llvm_nvvm_tcgen05_commit_generic_smem_a_read
+llvm.func @llvm_nvvm_tcgen05_commit_generic_smem_a_read(%barrier : !llvm.ptr,
+ %cta_mask : i16,
+ %cta_mask_32 : i32) {
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.smem.a.read.cg1.p0(ptr %{{.*}})
+ nvvm.tcgen05.commit %barrier {smem_a_read} : !llvm.ptr
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.smem.a.read.cg2.p0(ptr %{{.*}})
+ nvvm.tcgen05.commit %barrier {group = #nvvm.cta_group<cta_2>, smem_a_read} : !llvm.ptr
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.smem.a.read.mc.cg1.p0.i16(ptr %{{.*}}, i16 %{{.*}})
+ nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask {smem_a_read} : !llvm.ptr, i16
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.smem.a.read.mc.cg2.p0.i16(ptr %{{.*}}, i16 %{{.*}})
+ nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask {group = #nvvm.cta_group<cta_2>, smem_a_read} : !llvm.ptr, i16
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.smem.a.read.mc.cg1.p0.i32(ptr %{{.*}}, i32 %{{.*}})
+ nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask_32 {smem_a_read} : !llvm.ptr, i32
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.smem.a.read.mc.cg2.p0.i32(ptr %{{.*}}, i32 %{{.*}})
+ nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask_32 {group = #nvvm.cta_group<cta_2>, smem_a_read} : !llvm.ptr, i32
+ llvm.return
+}
+
+// CHECK-LABEL: @llvm_nvvm_tcgen05_commit_shared_smem_a_read
+llvm.func @llvm_nvvm_tcgen05_commit_shared_smem_a_read(%barrier : !llvm.ptr<3>,
+ %cta_mask : i16,
+ %cta_mask_32 : i32) {
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.smem.a.read.cg1.p3(ptr addrspace(3) %{{.*}})
+ nvvm.tcgen05.commit %barrier {smem_a_read} : !llvm.ptr<3>
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.smem.a.read.cg2.p3(ptr addrspace(3) %{{.*}})
+ nvvm.tcgen05.commit %barrier {group = #nvvm.cta_group<cta_2>, smem_a_read} : !llvm.ptr<3>
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.smem.a.read.mc.cg1.p3.i16(ptr addrspace(3) %{{.*}}, i16 %{{.*}})
+ nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask {smem_a_read} : !llvm.ptr<3>, i16
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.smem.a.read.mc.cg2.p3.i16(ptr addrspace(3) %{{.*}}, i16 %{{.*}})
+ nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask {group = #nvvm.cta_group<cta_2>, smem_a_read} : !llvm.ptr<3>, i16
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.smem.a.read.mc.cg1.p3.i32(ptr addrspace(3) %{{.*}}, i32 %{{.*}})
+ nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask_32 {smem_a_read} : !llvm.ptr<3>, i32
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.smem.a.read.mc.cg2.p3.i32(ptr addrspace(3) %{{.*}}, i32 %{{.*}})
+ nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask_32 {group = #nvvm.cta_group<cta_2>, smem_a_read} : !llvm.ptr<3>, i32
+ llvm.return
+}
diff --git a/mlir/test/Target/LLVMIR/nvvm/tcgen05-commit.mlir b/mlir/test/Target/LLVMIR/nvvm/tcgen05-commit.mlir
index 6ef6f9914ffb4..2ab59804ad860 100644
--- a/mlir/test/Target/LLVMIR/nvvm/tcgen05-commit.mlir
+++ b/mlir/test/Target/LLVMIR/nvvm/tcgen05-commit.mlir
@@ -1,33 +1,49 @@
-// RUN: mlir-translate -mlir-to-llvmir %s | FileCheck %s --check-prefix=CHECK-LLVM
+// RUN: mlir-translate -mlir-to-llvmir %s | FileCheck %s
// CHECK-LABEL: @llvm_nvvm_tcgen05_commit_generic
-llvm.func @llvm_nvvm_tcgen05_commit_generic(%barrier : !llvm.ptr, %cta_mask : i16) {
- // CHECK-LLVM: call void @llvm.nvvm.tcgen05.commit.cg1.p0(ptr %{{.*}})
+llvm.func @llvm_nvvm_tcgen05_commit_generic(%barrier : !llvm.ptr,
+ %cta_mask : i16,
+ %cta_mask_32 : i32) {
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.cg1.p0(ptr %{{.*}})
nvvm.tcgen05.commit %barrier : !llvm.ptr
- // CHECK-LLVM: call void @llvm.nvvm.tcgen05.commit.cg2.p0(ptr %{{.*}})
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.cg2.p0(ptr %{{.*}})
nvvm.tcgen05.commit %barrier {group = #nvvm.cta_group<cta_2>} : !llvm.ptr
- // CHECK-LLVM: call void @llvm.nvvm.tcgen05.commit.mc.cg1.p0.i16(ptr %{{.*}}, i16 %{{.*}})
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.mc.cg1.p0.i16(ptr %{{.*}}, i16 %{{.*}})
nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask : !llvm.ptr, i16
- // CHECK-LLVM: call void @llvm.nvvm.tcgen05.commit.mc.cg2.p0.i16(ptr %{{.*}}, i16 %{{.*}})
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.mc.cg2.p0.i16(ptr %{{.*}}, i16 %{{.*}})
nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask {group = #nvvm.cta_group<cta_2>} : !llvm.ptr, i16
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.mc.cg1.p0.i32(ptr %{{.*}}, i32 %{{.*}})
+ nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask_32 : !llvm.ptr, i32
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.mc.cg2.p0.i32(ptr %{{.*}}, i32 %{{.*}})
+ nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask_32 {group = #nvvm.cta_group<cta_2>} : !llvm.ptr, i32
llvm.return
}
// CHECK-LABEL: @llvm_nvvm_tcgen05_commit_shared
-llvm.func @llvm_nvvm_tcgen05_commit_shared(%barrier : !llvm.ptr<3>, %cta_mask : i16) {
- // CHECK-LLVM: call void @llvm.nvvm.tcgen05.commit.cg1.p3(ptr addrspace(3) %{{.*}})
+llvm.func @llvm_nvvm_tcgen05_commit_shared(%barrier : !llvm.ptr<3>,
+ %cta_mask : i16,
+ %cta_mask_32 : i32) {
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.cg1.p3(ptr addrspace(3) %{{.*}})
nvvm.tcgen05.commit %barrier : !llvm.ptr<3>
- // CHECK-LLVM: call void @llvm.nvvm.tcgen05.commit.cg2.p3(ptr addrspace(3) %{{.*}})
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.cg2.p3(ptr addrspace(3) %{{.*}})
nvvm.tcgen05.commit %barrier {group = #nvvm.cta_group<cta_2>} : !llvm.ptr<3>
- // CHECK-LLVM: call void @llvm.nvvm.tcgen05.commit.mc.cg1.p3.i16(ptr addrspace(3) %{{.*}}, i16 %{{.*}})
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.mc.cg1.p3.i16(ptr addrspace(3) %{{.*}}, i16 %{{.*}})
nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask : !llvm.ptr<3>, i16
- // CHECK-LLVM: call void @llvm.nvvm.tcgen05.commit.mc.cg2.p3.i16(ptr addrspace(3) %{{.*}}, i16 %{{.*}})
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.mc.cg2.p3.i16(ptr addrspace(3) %{{.*}}, i16 %{{.*}})
nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask {group = #nvvm.cta_group<cta_2>} : !llvm.ptr<3>, i16
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.mc.cg1.p3.i32(ptr addrspace(3) %{{.*}}, i32 %{{.*}})
+ nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask_32 : !llvm.ptr<3>, i32
+
+ // CHECK: call void @llvm.nvvm.tcgen05.commit.mc.cg2.p3.i32(ptr addrspace(3) %{{.*}}, i32 %{{.*}})
+ nvvm.tcgen05.commit %barrier, multicast_mask = %cta_mask_32 {group = #nvvm.cta_group<cta_2>} : !llvm.ptr<3>, i32
llvm.return
}
``````````
</details>
https://github.com/llvm/llvm-project/pull/215125
More information about the Mlir-commits
mailing list