[MLIR][NVVM] Add Rubin extensions to tcgen05.commit Op - #215125
Open
rajatbajpai wants to merge 1 commit into
Open
[MLIR][NVVM] Add Rubin extensions to tcgen05.commit Op#215125rajatbajpai wants to merge 1 commit into
rajatbajpai wants to merge 1 commit into
Conversation
This change adds support for 32-bit multicast mask and tracking of only shared-memory reads of A-matrix performed by prior MMA ops.
|
@llvm/pr-subscribers-mlir @llvm/pr-subscribers-mlir-llvm Author: Rajat Bajpai (rajatbajpai) ChangesThis 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:
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
}
|
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
This change adds support for 32-bit multicast mask and tracking of only shared-memory reads of A-matrix performed by prior MMA ops.