[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
Wed Aug 12 02:15:17 PDT 2026
================
@@ -858,13 +858,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
+ dag dims_dag = !dag(ins, !listsplat(B32, dim), !foreach(i, !range(dim), "d" # i));
+ string dims_str = !interleave(!foreach(i, !range(dim), "$d" # i), ", ");
+ bit has_multi_dim = !ge(dim, 2);
+
+ // Im2col handling (im2col / im2col_w / im2col_w_128)
+ int num_im2col = !cond(!eq(mode, "im2col"): !add(dim, -2),
+ !eq(mode, "im2col_w"): 2,
+ !eq(mode, "im2col_w_128"): 2,
+ true: 0);
+ dag im2col_dag = !dag(ins, !listsplat(B16, num_im2col),
+ !foreach(i, !range(num_im2col), "im2col" # i));
+ string im2col_str = !interleave(!foreach(i, !range(num_im2col), "$im2col" # i), ", ");
+ 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 = !cond(
+ !eq(mode, "im2col") : "im2col_no_offs",
+ !eq(mode, "im2col_w") : "im2col_no_offs::w",
+ !eq(mode, "tile_scatter4") : "tile::scatter4",
+ true : mode);
+ 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 = !cond(!eq(mode, "im2col") : "im2col_no_offs",
+ !eq(mode, "im2col_w") : "im2col_no_offs::w",
+ true : mode);
----------------
rajatbajpai wrote:
Good point, addressed.
https://github.com/llvm/llvm-project/pull/215503
More information about the llvm-commits
mailing list