[llvm] [mlir] [NVVM][NVPTX] Change TMA Tensor reduction ops to use flag for reduction ops (PR #213638)
Durgadoss R via llvm-commits
llvm-commits at lists.llvm.org
Mon Aug 3 05:15:05 PDT 2026
================
@@ -856,38 +856,59 @@ def TMAReductionFlags : Operand<i32> {
let PrintMethod = "printTmaReductionMode";
}
+def tma_reduction_imm :
+ TImmLeaf<i32, [{ return Imm >= 0 && Imm < 8; }]>;
+
// TMA Copy from Shared to Global memory with Reduction
-multiclass CP_ASYNC_BULK_TENSOR_REDUCE_INTR<int dim, bit shared32, string mode> {
+multiclass CP_ASYNC_BULK_TENSOR_REDUCE_INTR<int dim, string mode> {
defvar dims_dag = TMA_DIMS_UTIL<dim>.ins_dag;
defvar dims_str = TMA_DIMS_UTIL<dim>.base_str;
defvar asm_str = " [$tmap, {{" # dims_str # "}}], [$src]";
- defvar rc = !if(shared32, B32, B64);
// For im2col mode, the actual asm_str is "im2col_no_offs"
defvar mode_asm_str = !if(!eq(mode, "im2col"),
"im2col_no_offs", mode);
- defvar prefix = "cp.reduce.async.bulk.tensor" # "." # dim # "d" # ".global.shared::cta";
+ defvar prefix = "cp.reduce.async.bulk.tensor"
+ # "." # dim # "d"
+ # ".global.shared::cta";
defvar suffix = "." # mode_asm_str # ".bulk_group";
+ defvar intr = !cast<Intrinsic>(
+ "int_nvvm_cp_async_bulk_tensor_reduce_" # mode # "_" # dim # "d"
+ );
+
+ defvar intr_dag = !con(
+ (intr addr:$src, B64:$tmap),
+ !setdagop(dims_dag, intr),
+ (intr (i64 srcvalue), tma_reduction_imm:$red_op, 0)
----------------
durga4github wrote:
I was expecting this to the same as line 889. could you please help me understand what `srcvalue` is here?
https://github.com/llvm/llvm-project/pull/213638
More information about the llvm-commits
mailing list