[llvm] [mlir] [LLVM][NVPTX] Add async bulk copy global to shared extensions (PR #222323)
Rajat Bajpai via llvm-commits
llvm-commits at lists.llvm.org
Wed Sep 9 11:46:36 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,
+ (Intr async_bulk_scope_imm:$scope, tma_validate_data_imm:$validate),
+ (Intr tma_validate_data_imm:$validate));
+ defvar p_no_mc16 = !con((Intr addr:$dst, addr:$mbar, addr:$src, i32:$size,
+ i16:$mask, i64:$ch, 0, !if(has_ch, -1, 0)), p_tail);
+ defvar p_mc16 = !con((Intr addr:$dst, addr:$mbar, addr:$src, i32:$size,
+ i16:$mask, i64:$ch, -1, !if(has_ch, -1, 0)), p_tail);
+ defvar p_no_mc32 = !con((Intr addr:$dst, addr:$mbar, addr:$src, i32:$size,
+ i32:$mask, i64:$ch, 0, !if(has_ch, -1, 0)), p_tail);
+ defvar p_mc32 = !con((Intr addr:$dst, addr:$mbar, addr:$src, i32:$size,
+ i32:$mask, i64:$ch, -1, !if(has_ch, -1, 0)), p_tail);
+
+ defvar args_ch = !if(has_ch, ", $ch", "");
+
+ // 16-bit multicast mask variants
+ def "" : NVPTXInst<(outs), ins_i16,
+ asm_nomc # " [$dst], [$src], $size, [$mbar]" # args_ch # ";",
+ [p_no_mc16]>, Requires<preds>;
+ def _MC : NVPTXInst<(outs), ins_i16,
+ asm_mc16 # " [$dst], [$src], $size, [$mbar], $mask" # args_ch # ";",
+ [p_mc16]>, Requires<preds>;
+
+ // 32-bit multicast mask variants
+ let Predicates = [hasRubinFamilySupport] in {
+ def _NO_MC32 : NVPTXInst<(outs), ins_i32,
+ asm_nomc # " [$dst], [$src], $size, [$mbar]" # args_ch # ";",
+ [p_no_mc32]>;
+ def _MC32 : NVPTXInst<(outs), ins_i32,
+ asm_mc32 # " [$dst], [$src], $size, [$mbar], $mask" # args_ch # ";",
+ [p_mc32]>;
----------------
rajatbajpai wrote:
The sm_107f has an implicit PTX version requirement of 9.4. If we use sm_107f with any PTX version < 9.4, it will emit an error. So, in short, we don’t need a separate PTX version predicate here.
https://github.com/llvm/llvm-project/pull/222323
More information about the llvm-commits
mailing list