[llvm] [NVVM][NVPTX] Add tensor map override support in S2G and Reduce intrinsics (PR #215503)

Rajat Bajpai via llvm-commits llvm-commits at lists.llvm.org
Thu Aug 13 02:58:01 PDT 2026


================
@@ -858,23 +860,187 @@ foreach dim = 3...5 in {
 defm TMA_S2G_TILE_SCATTER4_2D : TMA_TENSOR_S2G_INTR<5, "tile_scatter4",
                                 [callSubtarget<"hasTMABlackwellSupport">]>;
 
+// Base class for TMA tensor override intrinsics
+class TMA_TENSOR_OVERRIDE_BASE<int dim, string mode, string override> {
+  // Dimension handling.
+  defvar dims_util = TMA_DIMS_UTIL<dim>;
+  dag dims_dag = dims_util.ins_dag;
+  string dims_str = dims_util.base_str;
+  bit has_multi_dim = !ge(dim, 2);
+
+  // Im2col handling (im2col / im2col_w / im2col_w_128).
+  defvar im2col_util = TMA_IM2COL_UTIL<dim, mode>;
+  int num_im2col = im2col_util.offsets;
+  dag im2col_dag = im2col_util.ins_dag;
+  string im2col_str = im2col_util.base_str;
+  string im2col_component = !if(num_im2col, ", {{" # im2col_str # "}}", "");
+
+  // Tensor size handling
+  bit is_addr_dim_stride = !eq(override, "override_addr_dim_stride");
+  dag tensor_size_dag = !if(is_addr_dim_stride,
+                          !dag(ins, !listsplat(B16, dim), !foreach(i, !range(dim), "ts" # i)),
+                          (ins));
+  string tensor_size_str = !interleave(!foreach(i, !range(dim), "$ts" # i), ", ");
+
+  // Stride handling for multi-dimensional tensors
+  bit need_strides = !and(is_addr_dim_stride, has_multi_dim);
+  dag lower_stride_dag = !if(need_strides,
+                           !dag(ins, !listsplat(B32, !sub(dim, 1)),
+                                !foreach(i, !range(!sub(dim, 1)), "stride" # i)),
+                           (ins));
+  string lower_stride_str = !interleave(!foreach(i, !range(!sub(dim, 1)), "$stride" # i), ", ");
+  dag upper_stride_dag = !if(need_strides, (ins B16:$upper_stride), (ins));
+
+  // Compose assembly string components
+  string stride_component = !if(has_multi_dim,
+                              "{{" # lower_stride_str # "}}, $upper_stride, ",
+                              "");
+  string tensor_size_component = !if(is_addr_dim_stride,
+                                  "{{" # tensor_size_str # "}}, " # stride_component,
+                                  "");
+
+  // Common suffix components
+  string dim_suffix = !if(has_multi_dim,
+                         ".override::global_dim_stride",
+                         ".override::global_dim");
+  string override_suffix = ".override::global_address" # !if(is_addr_dim_stride, dim_suffix, "");
+}
+
+// TMA Copy from Shared to Global memory (with Override Global Address)
+multiclass NVPTX_CP_ASYNC_BULK_TENSOR_S2G_OVERRIDE_INTR<int dim, string mode,
+                                                    string override> {
+  defvar base = TMA_TENSOR_OVERRIDE_BASE<dim, mode, override>;
+  defvar is_scatter4 = !eq(mode, "tile_scatter4");
+  defvar mode_asm_str = TMA_MODE_ASM_UTIL<mode>.ret;
+  defvar inst_name = "cp.async.bulk.tensor." # !if(is_scatter4, 2, dim) #
+                     "d.global.shared::cta." # mode_asm_str # ".bulk_group";
+  defvar suffix = base.override_suffix;
+  defvar asm_operands = " [$tmap, $override_addr, " # base.tensor_size_component #
+                        "{{" # base.dims_str # "}}], [$src]";
+
+  defvar is_override_dim_stride_1d =
+      !and(base.is_addr_dim_stride, !eq(dim, 1));
+  defvar intr = !cast<Intrinsic>(
+      "int_nvvm_cp_async_bulk_tensor_s2g_" # mode # "_" #
+      !if(is_override_dim_stride_1d, "override_addr_dim", override) # "_" #
+      !if(is_scatter4, 2, dim) # "d");
+
+  defvar ins_dag_base =
+      !con((ins ADDR:$src, B64:$tmap, B64:$override_addr),
+           base.tensor_size_dag, base.lower_stride_dag, base.upper_stride_dag,
+           base.dims_dag);
+  defvar ins_dag = ins_dag_base;
+  defvar ins_dag_ch = !con(ins_dag_base, (ins B64:$ch));
+
+  defvar intr_dag_base =
+      !con((intr addr:$src, B64:$tmap, B64:$override_addr),
+           !setdagop(base.tensor_size_dag, intr),
+           !setdagop(base.lower_stride_dag, intr),
+           !setdagop(base.upper_stride_dag, intr),
+           !setdagop(base.dims_dag, intr));
+  defvar intr_dag =
+      !con(intr_dag_base, (intr (i64 srcvalue), 0));
+  defvar intr_dag_ch =
+      !con(intr_dag_base, (intr i64:$ch, -1));
+
+  let Predicates = [callSubtarget<"hasRubinFamilySupport">] in {
+    def "" : NVPTXInst<(outs), ins_dag,
+             inst_name # suffix # asm_operands # ";",
+             [intr_dag]>;
+    def _CH : NVPTXInst<(outs), ins_dag_ch,
+              inst_name # suffix # ".L2::cache_hint" # asm_operands # ", $ch;",
+              [intr_dag_ch]>;
+  }
+}
+
 def TMAReductionFlags : Operand<i32> {
   let PrintMethod = "printTmaReductionMode";
 }
 
 def tma_tensor_reduction_imm :
     TImmLeaf<i32, [{ return Imm >= 0 && Imm < 8; }]>;
 
+// TMA Copy from Shared to Global memory with Reduction (with Override Global Address)
+multiclass NVPTX_CP_ASYNC_BULK_TENSOR_REDUCE_OVERRIDE_INTR<int dim, string mode,
+                                                    string override> {
+  defvar base = TMA_TENSOR_OVERRIDE_BASE<dim, mode, override>;
+  defvar mode_asm_str = TMA_MODE_ASM_UTIL<mode>.ret;
+  defvar inst_name = "cp.reduce.async.bulk.tensor." # dim #
+                     "d.global.shared::cta";
+  defvar suffix = "." # mode_asm_str # ".bulk_group" # base.override_suffix;
+  defvar asm_operands = " [$tmap, $override_addr, " # base.tensor_size_component #
+                        "{{" # base.dims_str # "}}], [$src]";
+
+  defvar is_override_dim_stride_1d =
+      !and(base.is_addr_dim_stride, !eq(dim, 1));
+
+  defvar intr = !cast<Intrinsic>(
+      "int_nvvm_cp_async_bulk_tensor_reduce_" # mode # "_" #
+      !if(is_override_dim_stride_1d, "override_addr_dim", override) # "_" #
+      dim # "d");
+
+  defvar ins_dag_base =
+      !con((ins ADDR:$src, B64:$tmap, B64:$override_addr),
+           base.tensor_size_dag, base.lower_stride_dag, base.upper_stride_dag,
+           base.dims_dag);
+  defvar ins_dag = !con(ins_dag_base, (ins TMAReductionFlags:$red_op));
+  defvar ins_dag_ch =
+      !con(ins_dag_base, (ins B64:$ch, TMAReductionFlags:$red_op));
+
+  defvar intr_dag_base =
+      !con((intr addr:$src, B64:$tmap, B64:$override_addr),
+           !setdagop(base.tensor_size_dag, intr),
+           !setdagop(base.lower_stride_dag, intr),
+           !setdagop(base.upper_stride_dag, intr),
+           !setdagop(base.dims_dag, intr));
+  defvar intr_dag =
+      !con(intr_dag_base,
+           (intr (i64 srcvalue), tma_tensor_reduction_imm:$red_op, 0));
+  defvar intr_dag_ch =
+      !con(intr_dag_base,
+           (intr B64:$ch, tma_tensor_reduction_imm:$red_op, -1));
+
+  let Predicates = [callSubtarget<"hasRubinFamilySupport">] in {
+    def "" : NVPTXInst<(outs), ins_dag,
+             inst_name # "${red_op}" # suffix # asm_operands # ";",
+             [intr_dag]>;
+    def _CH : NVPTXInst<(outs), ins_dag_ch,
+              inst_name # "${red_op}" # suffix # ".L2::cache_hint" # asm_operands #
+              ", $ch;",
+              [intr_dag_ch]>;
+  }
+}
+
+foreach dim = 1...5 in {
+  foreach mode = !if(
+      !ge(dim, 3), ["tile", "im2col", "im2col_w"], ["tile"]) in {
+    foreach override = !if(!eq(mode, "tile"),
+                            ["override_addr", "override_addr_dim_stride"],
+                            ["override_addr"]) in {
+      defm S2G_ # dim # "D_" #
+               !toupper(!subst("::", "_", mode) # "_" # override)
+          : NVPTX_CP_ASYNC_BULK_TENSOR_S2G_OVERRIDE_INTR<
+                dim, mode, override>;
+      defm RED_ # dim # "D_" #
+               !toupper(!subst("::", "_", mode) # "_" # override)
+          : NVPTX_CP_ASYNC_BULK_TENSOR_REDUCE_OVERRIDE_INTR<
+                dim, mode, override>;
+    }
+  }
+}
+
+defm S2G_2D_TILE_SCATTER4_OVERRIDE_ADDR
+    : NVPTX_CP_ASYNC_BULK_TENSOR_S2G_OVERRIDE_INTR<
+          5, "tile_scatter4", "override_addr">;
----------------
rajatbajpai wrote:

done, but it is going over 80ish line.

https://github.com/llvm/llvm-project/pull/215503


More information about the llvm-commits mailing list