[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