[llvm] [mlir] [LLVM][NVPTX] Add async bulk copy global to shared extensions (PR #222323)
Durgadoss R via llvm-commits
llvm-commits at lists.llvm.org
Fri Sep 11 02:00:39 PDT 2026
================
@@ -561,44 +576,127 @@ multiclass CP_ASYNC_BULK_S2G_INTR<bit has_ch> {
defm CP_ASYNC_BULK_S2G : CP_ASYNC_BULK_S2G_INTR<has_ch = 0>;
defm CP_ASYNC_BULK_S2G_CH : CP_ASYNC_BULK_S2G_INTR<has_ch = 1>;
-multiclass CP_ASYNC_BULK_G2S_INTR<bit has_ch> {
- defvar Intr = int_nvvm_cp_async_bulk_global_to_shared_cluster;
+def TMAValidateDataFlags : Operand<i32> {
+ let PrintMethod = "printTMAValidateDataFlags";
+}
- def "" : NVPTXInst<(outs),
- (ins ADDR:$dst, ADDR:$mbar, ADDR:$src,
- B32:$size, B16:$mask, B64:$ch),
- !if(has_ch,
- CpAsyncBulkStr<0, 1>.G2S # " [$dst], [$src], $size, [$mbar], $ch;",
- CpAsyncBulkStr<0, 0>.G2S # " [$dst], [$src], $size, [$mbar];"),
- [(Intr addr:$dst, addr:$mbar, addr:$src, i32:$size, i16:$mask, i64:$ch, 0, !if(has_ch, -1, 0))]>,
- Requires<[PTX80, SM90]>;
+// Matches validate pattern values 0-5. The 1-5 values are Rubin-only
+// validate patterns. 0 denotes disabled validate pattern (default).
+def tma_validate_data_imm : TImmLeaf<i32, [{
+ return Imm == 0 || (Imm >= 1 && Imm <= 5 &&
+ Subtarget->hasFeature(NVPTX::SM107f));
+}]>;
- def _MC : NVPTXInst<(outs),
- (ins ADDR:$dst, ADDR:$mbar, ADDR:$src,
- B32:$size, B16:$mask, B64:$ch),
- !if(has_ch,
- CpAsyncBulkStr<1, 1>.G2S # " [$dst], [$src], $size, [$mbar], $mask, $ch;",
- CpAsyncBulkStr<1, 0>.G2S # " [$dst], [$src], $size, [$mbar], $mask;"),
- [(Intr addr:$dst, addr:$mbar, addr:$src, i32:$size, i16:$mask, i64:$ch, -1, !if(has_ch, -1, 0))]>,
- Requires<[PTX80, SM90]>;
+// Prints the scope for relaxed memory ordering semantics.
+def MemScopeFlags : Operand<i32> {
+ let PrintMethod = "printMemScope";
+}
+
+// Matches scope values 0-3 (cta/cluster/gpu/sys) for relaxed memory ordering semantics.
+def async_bulk_scope_imm : TImmLeaf<i32, [{ return Imm >= 0 && Imm <= 3; }]>;
+
+multiclass CP_ASYNC_BULK_G2S_INTR<bit has_ch, bit is_relaxed,
+ list<Predicate> preds = [hasCpAsyncBulkClusterSupport]> {
+ defvar Intr =
+ !cast<Intrinsic>("int_nvvm_cp_async_bulk_global_to_shared_cluster" #
+ !if(is_relaxed, "_relaxed", ""));
+
+ defvar asm_nomc = CpAsyncBulkStr<0, has_ch, 0, is_relaxed>.G2S;
+ defvar asm_mc16 = CpAsyncBulkStr<1, has_ch, 0, is_relaxed>.G2S;
+ defvar asm_mc32 = CpAsyncBulkStr<2, has_ch, 0, is_relaxed>.G2S;
+
+ // Relaxed form carries the extra scope operand.
+ defvar ins_tail = !if(is_relaxed,
+ (ins MemScopeFlags:$scope, TMAValidateDataFlags:$validate),
+ (ins TMAValidateDataFlags:$validate));
+ defvar ins_i16 = !con((ins ADDR:$dst, ADDR:$mbar, ADDR:$src, B32:$size,
+ B16:$mask, B64:$ch), ins_tail);
+ defvar ins_i32 = !con((ins ADDR:$dst, ADDR:$mbar, ADDR:$src, B32:$size,
+ B32:$mask, B64:$ch), ins_tail);
+
+ // Relaxed form carries the extra scope operand.
+ defvar p_tail = !if(is_relaxed,
----------------
durga4github wrote:
How do I read "p"? I am sure you had some logic behind the naming and seems I can't get it ;-)
Is it for pattern/pat?
https://github.com/llvm/llvm-project/pull/222323
More information about the llvm-commits
mailing list