[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