[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