[Mlir-commits] [mlir] [mlir][vector] Generalize multi_reduction innerparallel unrolling to N dimensions (PR #182301)
Erick Ochoa Lopez
llvmlistbot at llvm.org
Thu Feb 19 07:50:38 PST 2026
https://github.com/amd-eochoalo updated https://github.com/llvm/llvm-project/pull/182301
>From 03213d82c26512cf3268a99c7daef7fd9b8b4c89 Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Wed, 18 Feb 2026 11:55:25 -0500
Subject: [PATCH 01/19] Add tests
---
.../Vector/TransformOps/VectorTransformOps.td | 23 ++
.../TransformOps/VectorTransformOps.cpp | 8 +
.../vector-multi-reduction-lowering.mlir | 255 ------------------
...vector-multi-reduction-outer-lowering.mlir | 192 -------------
.../vector-multi-reduction-unrolling.mlir | 156 +++++++++++
.../python/dialects/transform_vector_ext.py | 8 +
6 files changed, 195 insertions(+), 447 deletions(-)
delete mode 100644 mlir/test/Dialect/Vector/vector-multi-reduction-lowering.mlir
delete mode 100644 mlir/test/Dialect/Vector/vector-multi-reduction-outer-lowering.mlir
create mode 100644 mlir/test/Dialect/Vector/vector-multi-reduction-unrolling.mlir
diff --git a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
index 6eb96e2a8fdab..685c88c17e556 100644
--- a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
+++ b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
@@ -283,6 +283,29 @@ def ApplyMultiReductionFlatteningPatternsOp: Op<Transform_Dialect,
}];
}
+def ApplyMultiReductionUnrollingPatternsOp: Op<Transform_Dialect,
+ "apply_patterns.vector.multi_reduction_unrolling",
+ [DeclareOpInterfaceMethods<PatternDescriptorOpInterface>]> {
+ let description = [{
+ Indicates that 2-D vector multi_reduction operations should be unrolled
+ into either a sequence of vector.reduction ops (innerreduction) or
+ element-wise arith ops (innerparallel).
+
+ This populates the patterns from
+ `populateVectorMultiReductionUnrollingPatterns`, i.e.:
+ * `TwoDimMultiReductionToReduction` (innerreduction)
+ * `TwoDimMultiReductionToElementWise` (innerparallel)
+ }];
+
+ let arguments = (ins DefaultValuedAttr<VectorMultiReductionLoweringAttr,
+ "vector::VectorMultiReductionLowering::InnerParallel">:$lowering_strategy
+ );
+
+ let assemblyFormat = [{
+ (`lowering_strategy` `=` $lowering_strategy^)? attr-dict
+ }];
+}
+
def ApplyLowerOuterProductPatternsOp : Op<Transform_Dialect,
"apply_patterns.vector.lower_outerproduct",
[DeclareOpInterfaceMethods<PatternDescriptorOpInterface>]> {
diff --git a/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp b/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
index f3529ac26523f..6c97c6501a23e 100644
--- a/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
+++ b/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
@@ -154,6 +154,14 @@ void transform::ApplyMultiReductionFlatteningPatternsOp::populatePatterns(
patterns, vectorTransformOptions.vectorMultiReductionLowering);
}
+void transform::ApplyMultiReductionUnrollingPatternsOp::populatePatterns(
+ RewritePatternSet &patterns) {
+ vector::VectorTransformsOptions vectorTransformOptions;
+ vectorTransformOptions.setVectorMultiReductionLowering(getLoweringStrategy());
+ vector::populateVectorMultiReductionUnrollingPatterns(
+ patterns, vectorTransformOptions.vectorMultiReductionLowering);
+}
+
void transform::ApplyLowerOuterProductPatternsOp::populatePatterns(
RewritePatternSet &patterns) {
populateVectorOuterProductLoweringPatterns(patterns);
diff --git a/mlir/test/Dialect/Vector/vector-multi-reduction-lowering.mlir b/mlir/test/Dialect/Vector/vector-multi-reduction-lowering.mlir
deleted file mode 100644
index 6b79a78e6a42a..0000000000000
--- a/mlir/test/Dialect/Vector/vector-multi-reduction-lowering.mlir
+++ /dev/null
@@ -1,255 +0,0 @@
-// RUN: mlir-opt %s --transform-interpreter | FileCheck %s
-
-// Patterns applied:
-// * TwoDimMultiReductionToReduction from populateVectorMultiReductionUnrollingPatterns
-func.func @vector_multi_reduction(%arg0: vector<2x4xf32>, %acc: vector<2xf32>) -> vector<2xf32> {
- %0 = vector.multi_reduction <mul>, %arg0, %acc [1] : vector<2x4xf32> to vector<2xf32>
- return %0 : vector<2xf32>
-}
-// CHECK-LABEL: func @vector_multi_reduction
-// CHECK-SAME: %[[INPUT:.+]]: vector<2x4xf32>, %[[ACC:.*]]: vector<2xf32>)
-// CHECK-DAG: %[[RESULT_VEC_0:.+]] = arith.constant dense<{{.*}}> : vector<2xf32>
-// CHECK: %[[V0:.+]] = vector.extract %[[INPUT]][0]
-// CHECK: %[[ACC0:.+]] = vector.extract %[[ACC]][0]
-// CHECK: %[[RV0:.+]] = vector.reduction <mul>, %[[V0]], %[[ACC0]] : vector<4xf32> into f32
-// CHECK: %[[RESULT_VEC_1:.+]] = vector.insert %[[RV0:.+]], %[[RESULT_VEC_0]] [0] : f32 into vector<2xf32>
-// CHECK: %[[V1:.+]] = vector.extract %[[INPUT]][1]
-// CHECK: %[[ACC1:.+]] = vector.extract %[[ACC]][1]
-// CHECK: %[[RV1:.+]] = vector.reduction <mul>, %[[V1]], %[[ACC1]] : vector<4xf32> into f32
-// CHECK: %[[RESULT_VEC:.+]] = vector.insert %[[RV1:.+]], %[[RESULT_VEC_1]] [1] : f32 into vector<2xf32>
-// CHECK: return %[[RESULT_VEC]]
-
-// Patterns applied:
-// * ReduceMultiDimReductionRank from populateVectorMultiReductionFlatteningPatterns
-// * TwoDimMultiReductionToReduction from populateVectorMultiReductionUnrollingPatterns
-func.func @vector_reduction_inner(%arg0: vector<2x3x4x5xi32>, %acc: vector<2x3xi32>) -> vector<2x3xi32> {
- %0 = vector.multi_reduction <add>, %arg0, %acc [2, 3] : vector<2x3x4x5xi32> to vector<2x3xi32>
- return %0 : vector<2x3xi32>
-}
-// CHECK-LABEL: func @vector_reduction_inner
-// CHECK-SAME: %[[INPUT:.+]]: vector<2x3x4x5xi32>, %[[ACC:.*]]: vector<2x3xi32>
-// CHECK-DAG: %[[FLAT_RESULT_VEC_0:.+]] = arith.constant dense<0> : vector<6xi32>
-// CHECK: %[[RESHAPED_INPUT:.+]] = vector.shape_cast %[[INPUT]] : vector<2x3x4x5xi32> to vector<6x20xi32>
-// CHECK: %[[V0:.+]] = vector.extract %[[RESHAPED_INPUT]][0] : vector<20xi32> from vector<6x20xi32>
-// CHECK: %[[ACC0:.+]] = vector.extract %[[ACC]][0, 0] : i32 from vector<2x3xi32>
-// CHECK: %[[V0R:.+]] = vector.reduction <add>, %[[V0]], %[[ACC0]] : vector<20xi32> into i32
-// CHECK: %[[FLAT_RESULT_VEC_1:.+]] = vector.insert %[[V0R]], %[[FLAT_RESULT_VEC_0]] [0] : i32 into vector<6xi32>
-// CHECK: %[[V1:.+]] = vector.extract %[[RESHAPED_INPUT]][1] : vector<20xi32> from vector<6x20xi32>
-// CHECK: %[[ACC1:.+]] = vector.extract %[[ACC]][0, 1] : i32 from vector<2x3xi32>
-// CHECK: %[[V1R:.+]] = vector.reduction <add>, %[[V1]], %[[ACC1]] : vector<20xi32> into i32
-// CHECK: %[[FLAT_RESULT_VEC_2:.+]] = vector.insert %[[V1R]], %[[FLAT_RESULT_VEC_1]] [1] : i32 into vector<6xi32>
-// CHECK: %[[V2:.+]] = vector.extract %[[RESHAPED_INPUT]][2] : vector<20xi32> from vector<6x20xi32>
-// CHECK: %[[ACC2:.+]] = vector.extract %[[ACC]][0, 2] : i32 from vector<2x3xi32>
-// CHECK: %[[V2R:.+]] = vector.reduction <add>, %[[V2]], %[[ACC2]] : vector<20xi32> into i32
-// CHECK: %[[FLAT_RESULT_VEC_3:.+]] = vector.insert %[[V2R]], %[[FLAT_RESULT_VEC_2]] [2] : i32 into vector<6xi32>
-// CHECK: %[[V3:.+]] = vector.extract %[[RESHAPED_INPUT]][3] : vector<20xi32> from vector<6x20xi32>
-// CHECK: %[[ACC3:.+]] = vector.extract %[[ACC]][1, 0] : i32 from vector<2x3xi32>
-// CHECK: %[[V3R:.+]] = vector.reduction <add>, %[[V3]], %[[ACC3]] : vector<20xi32> into i32
-// CHECK: %[[FLAT_RESULT_VEC_4:.+]] = vector.insert %[[V3R]], %[[FLAT_RESULT_VEC_3]] [3] : i32 into vector<6xi32>
-// CHECK: %[[V4:.+]] = vector.extract %[[RESHAPED_INPUT]][4] : vector<20xi32> from vector<6x20xi32>
-// CHECK: %[[ACC4:.+]] = vector.extract %[[ACC]][1, 1] : i32 from vector<2x3xi32>
-// CHECK: %[[V4R:.+]] = vector.reduction <add>, %[[V4]], %[[ACC4]] : vector<20xi32> into i32
-// CHECK: %[[FLAT_RESULT_VEC_5:.+]] = vector.insert %[[V4R]], %[[FLAT_RESULT_VEC_4]] [4] : i32 into vector<6xi32>
-// CHECK: %[[V5:.+]] = vector.extract %[[RESHAPED_INPUT]][5] : vector<20xi32> from vector<6x20xi32>
-// CHECK: %[[ACC5:.+]] = vector.extract %[[ACC]][1, 2] : i32 from vector<2x3xi32>
-// CHECK: %[[V5R:.+]] = vector.reduction <add>, %[[V5]], %[[ACC5]] : vector<20xi32> into i32
-// CHECK: %[[FLAT_RESULT_VEC:.+]] = vector.insert %[[V5R]], %[[FLAT_RESULT_VEC_5]] [5] : i32 into vector<6xi32>
-// CHECK: %[[RESULT:.+]] = vector.shape_cast %[[FLAT_RESULT_VEC]] : vector<6xi32> to vector<2x3xi32>
-// CHECK: return %[[RESULT]]
-
-// Patterns applied:
-// * TwoDimMultiReductionToReduction from populateVectorMultiReductionUnrollingPatterns
-func.func @vectorize_dynamic_reduction(%arg0: tensor<?x?xf32>, %arg1: tensor<?xf32>) -> tensor<?xf32> {
- %c0 = arith.constant 0 : index
- %dim = tensor.dim %arg0, %c0 : tensor<?x?xf32>
- %c1 = arith.constant 1 : index
- %dim_0 = tensor.dim %arg0, %c1 : tensor<?x?xf32>
- %c0_1 = arith.constant 0 : index
- %cst = arith.constant 0.000000e+00 : f32
- %0 = vector.create_mask %dim, %dim_0 : vector<4x8xi1>
- %1 = vector.mask %0 { vector.transfer_read %arg0[%c0_1, %c0_1], %cst {in_bounds = [true, true]} : tensor<?x?xf32>, vector<4x8xf32> } : vector<4x8xi1> -> vector<4x8xf32>
- %cst_2 = arith.constant 0.000000e+00 : f32
- %2 = vector.create_mask %dim : vector<4xi1>
- %3 = vector.mask %2 { vector.transfer_read %arg1[%c0_1], %cst_2 {in_bounds = [true]} : tensor<?xf32>, vector<4xf32> } : vector<4xi1> -> vector<4xf32>
- %4 = vector.mask %0 { vector.multi_reduction <add>, %1, %3 [1] : vector<4x8xf32> to vector<4xf32> } : vector<4x8xi1> -> vector<4xf32>
- %c0_3 = arith.constant 0 : index
- %5 = vector.mask %2 { vector.transfer_write %4, %arg1[%c0_3] {in_bounds = [true]} : vector<4xf32>, tensor<?xf32> } : vector<4xi1> -> tensor<?xf32>
- return %5 : tensor<?xf32>
-}
-
-// Verify that the original 2-D mask is sliced and propagated properly to the
-// vector.reduction instances.
-
-// CHECK-LABEL: func.func @vectorize_dynamic_reduction
-// CHECK: %[[VAL_8:.*]] = tensor.dim
-// CHECK: %[[VAL_9:.*]] = tensor.dim
-// CHECK: %[[VAL_10:.*]] = vector.create_mask %[[VAL_8]], %[[VAL_9]] : vector<4x8xi1>
-
-// CHECK: %[[VAL_16:.*]] = vector.extract %[[VAL_10]][0] : vector<8xi1> from vector<4x8xi1>
-// CHECK: %[[VAL_17:.*]] = vector.mask %[[VAL_16]] { vector.reduction <add>, %{{.*}} : vector<8xf32> into f32 } : vector<8xi1> -> f32
-// CHECK: %[[VAL_18:.*]] = vector.insert
-
-// CHECK: %[[VAL_21:.*]] = vector.extract %[[VAL_10]][1] : vector<8xi1> from vector<4x8xi1>
-// CHECK: %[[VAL_22:.*]] = vector.mask %[[VAL_21]] { vector.reduction <add>, %{{.*}} : vector<8xf32> into f32 } : vector<8xi1> -> f32
-// CHECK: %[[VAL_23:.*]] = vector.insert
-
-// CHECK: %[[VAL_26:.*]] = vector.extract %[[VAL_10]][2] : vector<8xi1> from vector<4x8xi1>
-// CHECK: %[[VAL_27:.*]] = vector.mask %[[VAL_26]] { vector.reduction <add>, %{{.*}} : vector<8xf32> into f32 } : vector<8xi1> -> f32
-// CHECK: %[[VAL_28:.*]] = vector.insert
-
-// CHECK: %[[VAL_31:.*]] = vector.extract %[[VAL_10]][3] : vector<8xi1> from vector<4x8xi1>
-// CHECK: %[[VAL_32:.*]] = vector.mask %[[VAL_31]] { vector.reduction <add>, %{{.*}} : vector<8xf32> into f32 } : vector<8xi1> -> f32
-// CHECK: %[[VAL_33:.*]] = vector.insert
-
-// Patterns applied:
-// * OneDimMultiReductionToTwoDim from populateVectorMultiReductionTransformationPatterns
-// * TwoDimMultiReductionToReduction from populateVectorMultiReductionUnrollingPatterns
-func.func @vectorize_1d_dynamic_reduction(%arg0: tensor<?xf32>) -> f32 {
- %c0 = arith.constant 0 : index
- %dim = tensor.dim %arg0, %c0 : tensor<?xf32>
- %c0_1 = arith.constant 0 : index
- %cst = arith.constant 0.000000e+00 : f32
- %0 = vector.create_mask %dim : vector<8xi1>
- %1 = vector.mask %0 { vector.transfer_read %arg0[%c0_1], %cst {in_bounds = [true]} : tensor<?xf32>, vector<8xf32> } : vector<8xi1> -> vector<8xf32>
- %4 = vector.mask %0 { vector.multi_reduction <add>, %1, %cst [0] : vector<8xf32> to f32 } : vector<8xi1> -> f32
- return %4 : f32
-}
-
-// Verify that a 1-D vector.multi_reduction is transformed into a vector.reduction.
-// This transform expands 1-D vectors into 2-D.
-
-// CHECK-LABEL: func.func @vectorize_1d_dynamic_reduction(
-// CHECK: %[[VAL_5:.*]] = vector.create_mask {{.*}} : vector<8xi1>
-// CHECK: %[[VAL_7:.*]] = vector.mask %[[VAL_5]] { vector.reduction <add>, %{{.*}} : vector<8xf32> into f32 } : vector<8xi1> -> f32
-
-
-// Patterns applied:
-// * InnerOuterDimReductionConversion from populateVectorMultiReductionTransformationPatterns
-// * ReduceMultiDimReductionRank from populateVectorMultiReductionFlatteningPatterns
-// * TwoDimMultiReductionToReduction from populateVectorMultiReductionUnrollingPatterns
-func.func @vectorize_dynamic_transpose_reduction(%arg0: tensor<?x?x?xf32>, %arg1: tensor<?x?xf32>) -> tensor<?x?xf32> {
- %c0 = arith.constant 0 : index
- %dim = tensor.dim %arg0, %c0 : tensor<?x?x?xf32>
- %c1 = arith.constant 1 : index
- %dim_0 = tensor.dim %arg0, %c1 : tensor<?x?x?xf32>
- %c2 = arith.constant 2 : index
- %dim_1 = tensor.dim %arg0, %c2 : tensor<?x?x?xf32>
- %c0_2 = arith.constant 0 : index
- %cst = arith.constant 0.000000e+00 : f32
- %0 = vector.create_mask %dim, %dim_0, %dim_1 : vector<4x8x16xi1>
- %1 = vector.mask %0 { vector.transfer_read %arg0[%c0_2, %c0_2, %c0_2], %cst {in_bounds = [true, true, true]} : tensor<?x?x?xf32>, vector<4x8x16xf32> } : vector<4x8x16xi1> -> vector<4x8x16xf32>
- %cst_3 = arith.constant 0.000000e+00 : f32
- %2 = vector.create_mask %dim_1, %dim_0 : vector<16x8xi1>
- %3 = vector.mask %2 { vector.transfer_read %arg1[%c0_2, %c0_2], %cst_3 {in_bounds = [true, true], permutation_map = affine_map<(d0, d1) -> (d1, d0)>} : tensor<?x?xf32>, vector<8x16xf32> } : vector<16x8xi1> -> vector<8x16xf32>
- %4 = vector.mask %0 { vector.multi_reduction <add>, %1, %3 [0] : vector<4x8x16xf32> to vector<8x16xf32> } : vector<4x8x16xi1> -> vector<8x16xf32>
- %c0_4 = arith.constant 0 : index
- %5 = vector.mask %2 { vector.transfer_write %4, %arg1[%c0_4, %c0_4] {in_bounds = [true, true], permutation_map = affine_map<(d0, d1) -> (d1, d0)>} : vector<8x16xf32>, tensor<?x?xf32> } : vector<16x8xi1> -> tensor<?x?xf32>
- return %5 : tensor<?x?xf32>
-}
-
-// CHECK-LABEL: func.func @vectorize_dynamic_transpose_reduction
-// CHECK: %[[VAL_6:.*]] = tensor.dim
-// CHECK: %[[VAL_7:.*]] = tensor.dim
-// CHECK: %[[VAL_8:.*]] = tensor.dim
-// CHECK: %[[VAL_135:.*]] = vector.create_mask %{{.*}}, %{{.*}}, %{{.*}} : vector<4x8x16xi1>
-// CHECK: %[[VAL_139:.*]] = vector.transpose %[[VAL_135]], [1, 2, 0] : vector<4x8x16xi1> to vector<8x16x4xi1>
-
-// Just checking a few instances to make sure the vector mask is properly propagated:
-
-// CHECK: %[[VAL_143:.*]] = vector.extract %[[VAL_139]][0, 0] : vector<4xi1> from vector<8x16x4xi1>
-// CHECK: %[[VAL_144:.*]] = vector.mask %[[VAL_143]] { vector.reduction <add>
-// CHECK: %[[VAL_145:.*]] = vector.insert %[[VAL_144]]
-
-// CHECK: %[[VAL_148:.*]] = vector.extract %[[VAL_139]][0, 1] : vector<4xi1> from vector<8x16x4xi1>
-// CHECK: %[[VAL_149:.*]] = vector.mask %[[VAL_148]] { vector.reduction <add>
-// CHECK: %[[VAL_150:.*]] = vector.insert %[[VAL_149]]
-
-// CHECK: %[[VAL_153:.*]] = vector.extract %[[VAL_139]][0, 2] : vector<4xi1> from vector<8x16x4xi1>
-// CHECK: %[[VAL_154:.*]] = vector.mask %[[VAL_153]] { vector.reduction <add>
-// CHECK: %[[VAL_155:.*]] = vector.insert %[[VAL_154]]
-
-// CHECK: %[[VAL_158:.*]] = vector.extract %[[VAL_139]][0, 3] : vector<4xi1> from vector<8x16x4xi1>
-// CHECK: %[[VAL_159:.*]] = vector.mask %[[VAL_158]] { vector.reduction <add>
-// CHECK: %[[VAL_160:.*]] = vector.insert %[[VAL_159]]
-
-// Patterns applied:
-// * ReduceMultiDimReductionRank from populateVectorMultiReductionFlatteningPatterns
-// * TwoDimMultiReductionToReduction from populateVectorMultiReductionUnrollingPatterns
-func.func private @vector_multi_reduction_non_scalable_dim(%A : vector<8x[4]x2xf32>, %B: vector<8x[4]xf32>) -> vector<8x[4]xf32> {
- %0 = vector.multi_reduction <add>, %A, %B [2] : vector<8x[4]x2xf32> to vector<8x[4]xf32>
- return %0 : vector<8x[4]xf32>
-}
-// CHECK-LABEL: func.func private @vector_multi_reduction_non_scalable_dim(
-// CHECK-SAME: %[[VAL_0:.*]]: vector<8x[4]x2xf32>,
-// CHECK-SAME: %[[VAL_1:.*]]: vector<8x[4]xf32>) -> vector<8x[4]xf32> {
-// CHECK-DAG: %[[VAL_2:.*]] = arith.constant dense<0.000000e+00> : vector<[32]xf32>
-
-// CHECK: %[[VAL_35:.*]] = vector.extract %[[VAL_0]][0, 0] : vector<2xf32> from vector<8x[4]x2xf32>
-// CHECK: %[[VAL_36:.*]] = vector.extract %[[VAL_1]][0, 0] : f32 from vector<8x[4]xf32>
-// CHECK: %[[VAL_37:.*]] = vector.reduction <add>, %[[VAL_35]], %[[VAL_36]] : vector<2xf32> into f32
-// CHECK: %[[VAL_38:.*]] = vector.insert %[[VAL_37]], %[[VAL_2]] [0] : f32 into vector<[32]xf32>
-
-// CHECK: %[[VAL_39:.*]] = vector.extract %[[VAL_0]][0, 1] : vector<2xf32> from vector<8x[4]x2xf32>
-// CHECK: %[[VAL_40:.*]] = vector.extract %[[VAL_1]][0, 1] : f32 from vector<8x[4]xf32>
-// CHECK: %[[VAL_41:.*]] = vector.reduction <add>, %[[VAL_39]], %[[VAL_40]] : vector<2xf32> into f32
-// CHECK: %[[VAL_42:.*]] = vector.insert %[[VAL_41]], %[[VAL_38]] [1] : f32 into vector<[32]xf32>
-
-// (...)
-
-// CHECK: %[[VAL_159:.*]] = vector.extract %[[VAL_0]][7, 3] : vector<2xf32> from vector<8x[4]x2xf32>
-// CHECK: %[[VAL_160:.*]] = vector.extract %[[VAL_1]][7, 3] : f32 from vector<8x[4]xf32>
-// CHECK: %[[VAL_161:.*]] = vector.reduction <add>, %[[VAL_159]], %[[VAL_160]] : vector<2xf32> into f32
-// CHECK: %[[VAL_162:.*]] = vector.insert %[[VAL_161]], %{{.*}} [31] : f32 into vector<[32]xf32>
-
-// CHECK: %[[VAL_163:.*]] = vector.shape_cast %[[VAL_162]] : vector<[32]xf32> to vector<8x[4]xf32>
-// CHECK: return %[[VAL_163]] : vector<8x[4]xf32>
-
-// Check that OneDimMultiReductionToTwoDim handles scalable dim
-// Patterns applied:
-// * OneDimMultiReductionToTwoDim from populateVectorMultiReductionTransformationPatterns
-// * TwoDimMultiReductionToReduction from populateVectorMultiReductionUnrollingPatterns
-func.func @vector_multi_reduction_scalable_dim_1d(%A: vector<[4]xf32>, %B: f32, %C: vector<[4]xi1>) -> f32 {
- %0 = vector.mask %C { vector.multi_reduction <add>, %A, %B [0] : vector<[4]xf32> to f32 } : vector<[4]xi1> -> f32
- return %0 : f32
-}
-
-// CHECK-LABEL: func.func @vector_multi_reduction_scalable_dim_1d(
-// CHECK-SAME: %[[ARG_0:.*]]: vector<[4]xf32>,
-// CHECK-SAME: %[[ARG_1:.*]]: f32,
-// CHECK-SAME: %[[ARG_2:.*]]: vector<[4]xi1>) -> f32 {
-// CHECK: %[[VAL_2:.*]] = vector.mask %[[ARG_2]] { vector.reduction <add>, %[[ARG_0]], %[[ARG_1]] : vector<[4]xf32> into f32 } : vector<[4]xi1> -> f32
-// CHECK: return %[[VAL_2]] : f32
-
-// Patterns applied:
-// * TwoDimMultiReductionToReduction from populateVectorMultiReductionUnrollingPatterns
-func.func @vector_multi_reduction_scalable_dim_2d(%A: vector<2x[4]xf32>, %B: vector<2xf32>, %C: vector<2x[4]xi1>) -> vector<2xf32> {
- %0 = vector.mask %C { vector.multi_reduction <add>, %A, %B [1] : vector<2x[4]xf32> to vector<2xf32> } : vector<2x[4]xi1> -> vector<2xf32>
- return %0 : vector<2xf32>
-}
-
-// CHECK-LABEL: func.func @vector_multi_reduction_scalable_dim_2d(
-// CHECK-SAME: %[[ARG_0:.*]]: vector<2x[4]xf32>,
-// CHECK-SAME: %[[ARG_1:.*]]: vector<2xf32>,
-// CHECK-SAME: %[[ARG_2:.*]]: vector<2x[4]xi1>) -> vector<2xf32> {
-// CHECK-DAG: %[[C0_2xf32:.*]] = arith.constant dense<0.000000e+00> : vector<2xf32>
-// CHECK: %[[ARG0_0:.*]] = vector.extract %[[ARG_0]][0] : vector<[4]xf32> from vector<2x[4]xf32>
-// CHECK: %[[ARG1_0:.*]] = vector.extract %[[ARG_1]][0] : f32 from vector<2xf32>
-// CHECK: %[[ARG2_0:.*]] = vector.extract %[[ARG_2]][0] : vector<[4]xi1> from vector<2x[4]xi1>
-// CHECK: %[[REDUCE_0:.*]] = vector.mask %[[ARG2_0]] { vector.reduction <add>, %[[ARG0_0]], %[[ARG1_0]] : vector<[4]xf32> into f32 } : vector<[4]xi1> -> f32
-// CHECK: %[[INSERT_0:.*]] = vector.insert %[[REDUCE_0]], %[[C0_2xf32]] [0] : f32 into vector<2xf32>
-// CHECK: %[[ARG0_1:.*]] = vector.extract %[[ARG_0]][1] : vector<[4]xf32> from vector<2x[4]xf32>
-// CHECK: %[[ARG1_1:.*]] = vector.extract %[[ARG_1]][1] : f32 from vector<2xf32>
-// CHECK: %[[ARG2_1:.*]] = vector.extract %[[ARG_2]][1] : vector<[4]xi1> from vector<2x[4]xi1>
-// CHECK: %[[REDUCE_1:.*]] = vector.mask %[[ARG2_1]] { vector.reduction <add>, %[[ARG0_1]], %[[ARG1_1]] : vector<[4]xf32> into f32 } : vector<[4]xi1> -> f32
-// CHECK: %[[INSERT_1:.*]] = vector.insert %[[REDUCE_1]], %[[INSERT_0]] [1] : f32 into vector<2xf32>
-// CHECK: return %[[INSERT_1]] : vector<2xf32>
-
-module attributes {transform.with_named_sequence} {
- transform.named_sequence @__transform_main(%root : !transform.any_op {transform.readonly}) {
- %func_op = transform.structured.match ops{["func.func"]} in %root : (!transform.any_op) -> !transform.op<"func.func">
- transform.apply_patterns to %func_op {
- transform.apply_patterns.vector.lower_multi_reduction lowering_strategy = "innerreduction"
- } : !transform.op<"func.func">
- transform.yield
- }
-}
diff --git a/mlir/test/Dialect/Vector/vector-multi-reduction-outer-lowering.mlir b/mlir/test/Dialect/Vector/vector-multi-reduction-outer-lowering.mlir
deleted file mode 100644
index d0ab71e3f400f..0000000000000
--- a/mlir/test/Dialect/Vector/vector-multi-reduction-outer-lowering.mlir
+++ /dev/null
@@ -1,192 +0,0 @@
-// RUN: mlir-opt %s --transform-interpreter | FileCheck %s
-
-func.func @vector_multi_reduction(%arg0: vector<2x4xf32>, %acc: vector<2xf32>) -> vector<2xf32> {
- %0 = vector.multi_reduction <mul>, %arg0, %acc [1] : vector<2x4xf32> to vector<2xf32>
- return %0 : vector<2xf32>
-}
-
-// CHECK-LABEL: func @vector_multi_reduction
-// CHECK-SAME: %[[INPUT:.+]]: vector<2x4xf32>, %[[ACC:.*]]: vector<2xf32>
-// CHECK: %[[TRANSPOSED:.+]] = vector.transpose %[[INPUT]], [1, 0] : vector<2x4xf32> to vector<4x2xf32>
-// CHECK: %[[V0:.+]] = vector.extract %[[TRANSPOSED]][0] : vector<2xf32> from vector<4x2xf32>
-// CHECK: %[[RV0:.+]] = arith.mulf %[[V0]], %[[ACC]] : vector<2xf32>
-// CHECK: %[[V1:.+]] = vector.extract %[[TRANSPOSED]][1] : vector<2xf32> from vector<4x2xf32>
-// CHECK: %[[RV01:.+]] = arith.mulf %[[V1]], %[[RV0]] : vector<2xf32>
-// CHECK: %[[V2:.+]] = vector.extract %[[TRANSPOSED]][2] : vector<2xf32> from vector<4x2xf32>
-// CHECK: %[[RV012:.+]] = arith.mulf %[[V2]], %[[RV01]] : vector<2xf32>
-// CHECK: %[[V3:.+]] = vector.extract %[[TRANSPOSED]][3] : vector<2xf32> from vector<4x2xf32>
-// CHECK: %[[RESULT_VEC:.+]] = arith.mulf %[[V3]], %[[RV012]] : vector<2xf32>
-// CHECK: return %[[RESULT_VEC]] : vector<2xf32>
-
-func.func @vector_multi_reduction_min(%arg0: vector<2x4xf32>, %acc: vector<2xf32>) -> vector<2xf32> {
- %0 = vector.multi_reduction <minnumf>, %arg0, %acc [1] : vector<2x4xf32> to vector<2xf32>
- return %0 : vector<2xf32>
-}
-
-// CHECK-LABEL: func @vector_multi_reduction_min
-// CHECK-SAME: %[[INPUT:.+]]: vector<2x4xf32>, %[[ACC:.*]]: vector<2xf32>
-// CHECK: %[[TRANSPOSED:.+]] = vector.transpose %[[INPUT]], [1, 0] : vector<2x4xf32> to vector<4x2xf32>
-// CHECK: %[[V0:.+]] = vector.extract %[[TRANSPOSED]][0] : vector<2xf32> from vector<4x2xf32>
-// CHECK: %[[RV0:.+]] = arith.minnumf %[[V0]], %[[ACC]] : vector<2xf32>
-// CHECK: %[[V1:.+]] = vector.extract %[[TRANSPOSED]][1] : vector<2xf32> from vector<4x2xf32>
-// CHECK: %[[RV01:.+]] = arith.minnumf %[[V1]], %[[RV0]] : vector<2xf32>
-// CHECK: %[[V2:.+]] = vector.extract %[[TRANSPOSED]][2] : vector<2xf32> from vector<4x2xf32>
-// CHECK: %[[RV012:.+]] = arith.minnumf %[[V2]], %[[RV01]] : vector<2xf32>
-// CHECK: %[[V3:.+]] = vector.extract %[[TRANSPOSED]][3] : vector<2xf32> from vector<4x2xf32>
-// CHECK: %[[RESULT_VEC:.+]] = arith.minnumf %[[V3]], %[[RV012]] : vector<2xf32>
-// CHECK: return %[[RESULT_VEC]] : vector<2xf32>
-
-func.func @vector_multi_reduction_max(%arg0: vector<2x4xf32>, %acc: vector<2xf32>) -> vector<2xf32> {
- %0 = vector.multi_reduction <maxnumf>, %arg0, %acc [1] : vector<2x4xf32> to vector<2xf32>
- return %0 : vector<2xf32>
-}
-
-// CHECK-LABEL: func @vector_multi_reduction_max
-// CHECK-SAME: %[[INPUT:.+]]: vector<2x4xf32>, %[[ACC:.*]]: vector<2xf32>
-// CHECK: %[[TRANSPOSED:.+]] = vector.transpose %[[INPUT]], [1, 0] : vector<2x4xf32> to vector<4x2xf32>
-// CHECK: %[[V0:.+]] = vector.extract %[[TRANSPOSED]][0] : vector<2xf32> from vector<4x2xf32>
-// CHECK: %[[RV0:.+]] = arith.maxnumf %[[V0]], %[[ACC]] : vector<2xf32>
-// CHECK: %[[V1:.+]] = vector.extract %[[TRANSPOSED]][1] : vector<2xf32> from vector<4x2xf32>
-// CHECK: %[[RV01:.+]] = arith.maxnumf %[[V1]], %[[RV0]] : vector<2xf32>
-// CHECK: %[[V2:.+]] = vector.extract %[[TRANSPOSED]][2] : vector<2xf32> from vector<4x2xf32>
-// CHECK: %[[RV012:.+]] = arith.maxnumf %[[V2]], %[[RV01]] : vector<2xf32>
-// CHECK: %[[V3:.+]] = vector.extract %[[TRANSPOSED]][3] : vector<2xf32> from vector<4x2xf32>
-// CHECK: %[[RESULT_VEC:.+]] = arith.maxnumf %[[V3]], %[[RV012]] : vector<2xf32>
-// CHECK: return %[[RESULT_VEC]] : vector<2xf32>
-
-func.func @vector_multi_reduction_and(%arg0: vector<2x4xi32>, %acc: vector<2xi32>) -> vector<2xi32> {
- %0 = vector.multi_reduction <and>, %arg0, %acc [1] : vector<2x4xi32> to vector<2xi32>
- return %0 : vector<2xi32>
-}
-
-// CHECK-LABEL: func @vector_multi_reduction_and
-// CHECK-SAME: %[[INPUT:.+]]: vector<2x4xi32>, %[[ACC:.*]]: vector<2xi32>
-// CHECK: %[[TRANSPOSED:.+]] = vector.transpose %[[INPUT]], [1, 0] : vector<2x4xi32> to vector<4x2xi32>
-// CHECK: %[[V0:.+]] = vector.extract %[[TRANSPOSED]][0] : vector<2xi32> from vector<4x2xi32>
-// CHECK: %[[RV0:.+]] = arith.andi %[[V0]], %[[ACC]] : vector<2xi32>
-// CHECK: %[[V1:.+]] = vector.extract %[[TRANSPOSED]][1] : vector<2xi32> from vector<4x2xi32>
-// CHECK: %[[RV01:.+]] = arith.andi %[[V1]], %[[RV0]] : vector<2xi32>
-// CHECK: %[[V2:.+]] = vector.extract %[[TRANSPOSED]][2] : vector<2xi32> from vector<4x2xi32>
-// CHECK: %[[RV012:.+]] = arith.andi %[[V2]], %[[RV01]] : vector<2xi32>
-// CHECK: %[[V3:.+]] = vector.extract %[[TRANSPOSED]][3] : vector<2xi32> from vector<4x2xi32>
-// CHECK: %[[RESULT_VEC:.+]] = arith.andi %[[V3]], %[[RV012]] : vector<2xi32>
-// CHECK: return %[[RESULT_VEC]] : vector<2xi32>
-
-func.func @vector_multi_reduction_or(%arg0: vector<2x4xi32>, %acc: vector<2xi32>) -> vector<2xi32> {
- %0 = vector.multi_reduction <or>, %arg0, %acc [1] : vector<2x4xi32> to vector<2xi32>
- return %0 : vector<2xi32>
-}
-
-// CHECK-LABEL: func @vector_multi_reduction_or
-// CHECK-SAME: %[[INPUT:.+]]: vector<2x4xi32>, %[[ACC:.*]]: vector<2xi32>
-// CHECK: %[[TRANSPOSED:.+]] = vector.transpose %[[INPUT]], [1, 0] : vector<2x4xi32> to vector<4x2xi32>
-// CHECK: %[[V0:.+]] = vector.extract %[[TRANSPOSED]][0] : vector<2xi32> from vector<4x2xi32>
-// CHECK: %[[RV0:.+]] = arith.ori %[[V0]], %[[ACC]] : vector<2xi32>
-// CHECK: %[[V1:.+]] = vector.extract %[[TRANSPOSED]][1] : vector<2xi32> from vector<4x2xi32>
-// CHECK: %[[RV01:.+]] = arith.ori %[[V1]], %[[RV0]] : vector<2xi32>
-// CHECK: %[[V2:.+]] = vector.extract %[[TRANSPOSED]][2] : vector<2xi32> from vector<4x2xi32>
-// CHECK: %[[RV012:.+]] = arith.ori %[[V2]], %[[RV01]] : vector<2xi32>
-// CHECK: %[[V3:.+]] = vector.extract %[[TRANSPOSED]][3] : vector<2xi32> from vector<4x2xi32>
-// CHECK: %[[RESULT_VEC:.+]] = arith.ori %[[V3]], %[[RV012]] : vector<2xi32>
-// CHECK: return %[[RESULT_VEC]] : vector<2xi32>
-
-func.func @vector_multi_reduction_xor(%arg0: vector<2x4xi32>, %acc: vector<2xi32>) -> vector<2xi32> {
- %0 = vector.multi_reduction <xor>, %arg0, %acc [1] : vector<2x4xi32> to vector<2xi32>
- return %0 : vector<2xi32>
-}
-
-// CHECK-LABEL: func @vector_multi_reduction_xor
-// CHECK-SAME: %[[INPUT:.+]]: vector<2x4xi32>, %[[ACC:.*]]: vector<2xi32>
-// CHECK: %[[TRANSPOSED:.+]] = vector.transpose %[[INPUT]], [1, 0] : vector<2x4xi32> to vector<4x2xi32>
-// CHECK: %[[V0:.+]] = vector.extract %[[TRANSPOSED]][0] : vector<2xi32> from vector<4x2xi32>
-// CHECK: %[[RV0:.+]] = arith.xori %[[V0]], %[[ACC]] : vector<2xi32>
-// CHECK: %[[V1:.+]] = vector.extract %[[TRANSPOSED]][1] : vector<2xi32> from vector<4x2xi32>
-// CHECK: %[[RV01:.+]] = arith.xori %[[V1]], %[[RV0]] : vector<2xi32>
-// CHECK: %[[V2:.+]] = vector.extract %[[TRANSPOSED]][2] : vector<2xi32> from vector<4x2xi32>
-// CHECK: %[[RV012:.+]] = arith.xori %[[V2]], %[[RV01]] : vector<2xi32>
-// CHECK: %[[V3:.+]] = vector.extract %[[TRANSPOSED]][3] : vector<2xi32> from vector<4x2xi32>
-// CHECK: %[[RESULT_VEC:.+]] = arith.xori %[[V3]], %[[RV012]] : vector<2xi32>
-// CHECK: return %[[RESULT_VEC]] : vector<2xi32>
-
-
-func.func @vector_reduction_outer(%arg0: vector<2x3x4x5xi32>, %acc: vector<2x3xi32>) -> vector<2x3xi32> {
- %0 = vector.multi_reduction <add>, %arg0, %acc [2, 3] : vector<2x3x4x5xi32> to vector<2x3xi32>
- return %0 : vector<2x3xi32>
-}
-
-// CHECK-LABEL: func @vector_reduction_outer
-// CHECK-SAME: %[[INPUT:.+]]: vector<2x3x4x5xi32>, %[[ACC:.*]]: vector<2x3xi32>
-// CHECK: %[[TRANSPOSED:.+]] = vector.transpose %[[INPUT]], [2, 3, 0, 1] : vector<2x3x4x5xi32> to vector<4x5x2x3xi32>
-// CHECK: %[[RESHAPED:.+]] = vector.shape_cast %[[TRANSPOSED]] : vector<4x5x2x3xi32> to vector<20x6xi32>
-// CHECK: %[[FACC:.+]] = vector.shape_cast %[[ACC]] : vector<2x3xi32> to vector<6xi32>
-// CHECK: %[[V0:.+]] = vector.extract %[[RESHAPED]][0] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R:.+]] = arith.addi %[[V0]], %[[FACC]] : vector<6xi32>
-// CHECK: %[[V1:.+]] = vector.extract %[[RESHAPED]][1] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R0:.+]] = arith.addi %[[V1]], %[[R]] : vector<6xi32>
-// CHECK: %[[V2:.+]] = vector.extract %[[RESHAPED]][2] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R1:.+]] = arith.addi %[[V2]], %[[R0]] : vector<6xi32>
-// CHECK: %[[V3:.+]] = vector.extract %[[RESHAPED]][3] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R2:.+]] = arith.addi %[[V3]], %[[R1]] : vector<6xi32>
-// CHECK: %[[V4:.+]] = vector.extract %[[RESHAPED]][4] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R3:.+]] = arith.addi %[[V4]], %[[R2]] : vector<6xi32>
-// CHECK: %[[V5:.+]] = vector.extract %[[RESHAPED]][5] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R4:.+]] = arith.addi %[[V5]], %[[R3]] : vector<6xi32>
-// CHECK: %[[V6:.+]] = vector.extract %[[RESHAPED]][6] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R5:.+]] = arith.addi %[[V6]], %[[R4]] : vector<6xi32>
-// CHECK: %[[V7:.+]] = vector.extract %[[RESHAPED]][7] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R6:.+]] = arith.addi %[[V7]], %[[R5]] : vector<6xi32>
-// CHECK: %[[V8:.+]] = vector.extract %[[RESHAPED]][8] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R7:.+]] = arith.addi %[[V8]], %[[R6]] : vector<6xi32>
-// CHECK: %[[V9:.+]] = vector.extract %[[RESHAPED]][9] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R8:.+]] = arith.addi %[[V9]], %[[R7]] : vector<6xi32>
-// CHECK: %[[V10:.+]] = vector.extract %[[RESHAPED]][10] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R9:.+]] = arith.addi %[[V10]], %[[R8]] : vector<6xi32>
-// CHECK: %[[V11:.+]] = vector.extract %[[RESHAPED]][11] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R10:.+]] = arith.addi %[[V11]], %[[R9]] : vector<6xi32>
-// CHECK: %[[V12:.+]] = vector.extract %[[RESHAPED]][12] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R11:.+]] = arith.addi %[[V12]], %[[R10]] : vector<6xi32>
-// CHECK: %[[V13:.+]] = vector.extract %[[RESHAPED]][13] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R12:.+]] = arith.addi %[[V13]], %[[R11]] : vector<6xi32>
-// CHECK: %[[V14:.+]] = vector.extract %[[RESHAPED]][14] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R13:.+]] = arith.addi %[[V14]], %[[R12]] : vector<6xi32>
-// CHECK: %[[V15:.+]] = vector.extract %[[RESHAPED]][15] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R14:.+]] = arith.addi %[[V15]], %[[R13]] : vector<6xi32>
-// CHECK: %[[V16:.+]] = vector.extract %[[RESHAPED]][16] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R15:.+]] = arith.addi %[[V16]], %[[R14]] : vector<6xi32>
-// CHECK: %[[V17:.+]] = vector.extract %[[RESHAPED]][17] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R16:.+]] = arith.addi %[[V17]], %[[R15]] : vector<6xi32>
-// CHECK: %[[V18:.+]] = vector.extract %[[RESHAPED]][18] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R17:.+]] = arith.addi %[[V18]], %[[R16]] : vector<6xi32>
-// CHECK: %[[V19:.+]] = vector.extract %[[RESHAPED]][19] : vector<6xi32> from vector<20x6xi32>
-// CHECK: %[[R18:.+]] = arith.addi %[[V19]], %[[R17]] : vector<6xi32>
-// CHECK: %[[RESULT_VEC:.+]] = vector.shape_cast %[[R18]] : vector<6xi32> to vector<2x3xi32>
-// CHECK: return %[[RESULT_VEC]] : vector<2x3xi32>
-
-func.func @vector_multi_reduction_parallel_middle(%arg0: vector<3x4x5xf32>, %acc: vector<4xf32>) -> vector<4xf32> {
- %0 = vector.multi_reduction <add>, %arg0, %acc [0, 2] : vector<3x4x5xf32> to vector<4xf32>
- return %0 : vector<4xf32>
-}
-
-// CHECK-LABEL: func @vector_multi_reduction_parallel_middle
-// CHECK-SAME: %[[INPUT:.+]]: vector<3x4x5xf32>, %[[ACC:.+]]: vector<4xf32>
-// CHECK: vector.transpose %[[INPUT]], [0, 2, 1] : vector<3x4x5xf32> to vector<3x5x4xf32>
-
-// This test is mainly to catch a bug that running
-// `InnerOuterDimReductionConversion` on this function results in an
-// infinite loop. So just check that some value is returned.
-func.func @vector_reduction_1D(%arg0 : vector<2xf32>, %acc: f32) -> f32 {
- %0 = vector.multi_reduction #vector.kind<maxnumf>, %arg0, %acc [0] : vector<2xf32> to f32
- return %0 : f32
-}
-// CHECK-LABEL: func @vector_reduction_1D
-// CHECK: return %{{.+}}
-
-module attributes {transform.with_named_sequence} {
- transform.named_sequence @__transform_main(%root : !transform.any_op {transform.readonly}) {
- %func_op = transform.structured.match ops{["func.func"]} in %root : (!transform.any_op) -> !transform.op<"func.func">
- transform.apply_patterns to %func_op {
- transform.apply_patterns.vector.lower_multi_reduction lowering_strategy = "innerparallel"
- } : !transform.op<"func.func">
- transform.yield
- }
-}
diff --git a/mlir/test/Dialect/Vector/vector-multi-reduction-unrolling.mlir b/mlir/test/Dialect/Vector/vector-multi-reduction-unrolling.mlir
new file mode 100644
index 0000000000000..d4fb79a1d4668
--- /dev/null
+++ b/mlir/test/Dialect/Vector/vector-multi-reduction-unrolling.mlir
@@ -0,0 +1,156 @@
+// RUN: mlir-opt %s --transform-interpreter='entry-point=innerreduction' | FileCheck %s --check-prefixes=INNER_REDUCTION,ALL
+// RUN: mlir-opt %s --transform-interpreter='entry-point=innerparallel' | FileCheck %s --check-prefixes=INNER_PARALLEL,ALL
+
+// ALL-LABEL: func @negative_rank1_and_rank3
+func.func @negative_rank1_and_rank3(
+ %rank1: vector<8xf32>, %rank1_acc: f32,
+ %rank3: vector<2x3x4xf32>, %rank3_acc: vector<2x3xf32>) -> (f32, vector<2x3xf32>) {
+ // ALL: vector.multi_reduction <add>, {{.+}} [0] : vector<8xf32> to f32
+ %0 = vector.multi_reduction <add>, %rank1, %rank1_acc [0] : vector<8xf32> to f32
+ // ALL: vector.multi_reduction <add>, {{.+}} [2] : vector<2x3x4xf32> to vector<2x3xf32>
+ %1 = vector.multi_reduction <add>, %rank3, %rank3_acc [2] : vector<2x3x4xf32> to vector<2x3xf32>
+ return %0, %1 : f32, vector<2x3xf32>
+}
+
+// ALL-LABEL: func @inner_reduction_2d
+// ALL-SAME: %[[INPUT:.+]]: vector<2x4xf32>, %[[ACC:.+]]: vector<2xf32>
+func.func @inner_reduction_2d(%arg0: vector<2x4xf32>, %acc: vector<2xf32>) -> vector<2xf32> {
+ // INNER_REDUCTION: %[[RESULT_VEC_0:.+]] = arith.constant dense<{{.+}}> : vector<2xf32>
+ // INNER_REDUCTION: %[[V0:.+]] = vector.extract %[[INPUT]][0]
+ // INNER_REDUCTION: %[[ACC0:.+]] = vector.extract %[[ACC]][0]
+ // INNER_REDUCTION: %[[RV0:.+]] = vector.reduction <mul>, %[[V0]], %[[ACC0]] : vector<4xf32> into f32
+ // INNER_REDUCTION: %[[RESULT_VEC_1:.+]] = vector.insert %[[RV0]], %[[RESULT_VEC_0]] [0] : f32 into vector<2xf32>
+ // INNER_REDUCTION: %[[V1:.+]] = vector.extract %[[INPUT]][1]
+ // INNER_REDUCTION: %[[ACC1:.+]] = vector.extract %[[ACC]][1]
+ // INNER_REDUCTION: %[[RV1:.+]] = vector.reduction <mul>, %[[V1]], %[[ACC1]] : vector<4xf32> into f32
+ // INNER_REDUCTION: %[[RESULT:.+]] = vector.insert %[[RV1]], %[[RESULT_VEC_1]] [1] : f32 into vector<2xf32>
+
+ // INNER_PARALLEL: %[[RESULT:.+]] = vector.multi_reduction <mul>, %[[INPUT]], %[[ACC]] [1]
+ %0 = vector.multi_reduction <mul>, %arg0, %acc [1] : vector<2x4xf32> to vector<2xf32>
+ // ALL: return %[[RESULT]]
+ return %0 : vector<2xf32>
+}
+
+func.func @inner_reduction_2d_masked_dynamic(%arg0: tensor<?x?xf32>, %arg1: tensor<?xf32>) -> tensor<?xf32> {
+ %c0 = arith.constant 0 : index
+ %dim = tensor.dim %arg0, %c0 : tensor<?x?xf32>
+ %c1 = arith.constant 1 : index
+ %dim_0 = tensor.dim %arg0, %c1 : tensor<?x?xf32>
+ %c0_1 = arith.constant 0 : index
+ %cst = arith.constant 0.000000e+00 : f32
+ %0 = vector.create_mask %dim, %dim_0 : vector<4x8xi1>
+ %1 = vector.mask %0 { vector.transfer_read %arg0[%c0_1, %c0_1], %cst {in_bounds = [true, true]} : tensor<?x?xf32>, vector<4x8xf32> } : vector<4x8xi1> -> vector<4x8xf32>
+ %cst_2 = arith.constant 0.000000e+00 : f32
+ %2 = vector.create_mask %dim : vector<4xi1>
+ %3 = vector.mask %2 { vector.transfer_read %arg1[%c0_1], %cst_2 {in_bounds = [true]} : tensor<?xf32>, vector<4xf32> } : vector<4xi1> -> vector<4xf32>
+ %4 = vector.mask %0 { vector.multi_reduction <add>, %1, %3 [1] : vector<4x8xf32> to vector<4xf32> } : vector<4x8xi1> -> vector<4xf32>
+ %c0_3 = arith.constant 0 : index
+ %5 = vector.mask %2 { vector.transfer_write %4, %arg1[%c0_3] {in_bounds = [true]} : vector<4xf32>, tensor<?xf32> } : vector<4xi1> -> tensor<?xf32>
+ return %5 : tensor<?xf32>
+}
+
+// ALL-LABEL: func @inner_reduction_2d_masked_dynamic
+// INNER_REDUCTION: %[[DIM_0:.+]] = tensor.dim
+// INNER_REDUCTION: %[[DIM_1:.+]] = tensor.dim
+// INNER_REDUCTION: %[[MASK_2D:.+]] = vector.create_mask %[[DIM_0]], %[[DIM_1]] : vector<4x8xi1>
+//
+// INNER_REDUCTION: %[[MASK_SLICE_0:.+]] = vector.extract %[[MASK_2D]][0] : vector<8xi1> from vector<4x8xi1>
+// INNER_REDUCTION: %[[REDUCE_0:.+]] = vector.mask %[[MASK_SLICE_0]] { vector.reduction <add>, %{{.+}} : vector<8xf32> into f32 } : vector<8xi1> -> f32
+// INNER_REDUCTION: %[[INSERT_0:.+]] = vector.insert
+//
+// INNER_REDUCTION: %[[MASK_SLICE_1:.+]] = vector.extract %[[MASK_2D]][1] : vector<8xi1> from vector<4x8xi1>
+// INNER_REDUCTION: %[[REDUCE_1:.+]] = vector.mask %[[MASK_SLICE_1]] { vector.reduction <add>, %{{.+}} : vector<8xf32> into f32 } : vector<8xi1> -> f32
+// INNER_REDUCTION: %[[INSERT_1:.+]] = vector.insert
+//
+// INNER_REDUCTION: %[[MASK_SLICE_2:.+]] = vector.extract %[[MASK_2D]][2] : vector<8xi1> from vector<4x8xi1>
+// INNER_REDUCTION: %[[REDUCE_2:.+]] = vector.mask %[[MASK_SLICE_2]] { vector.reduction <add>, %{{.+}} : vector<8xf32> into f32 } : vector<8xi1> -> f32
+// INNER_REDUCTION: %[[INSERT_2:.+]] = vector.insert
+//
+// INNER_REDUCTION: %[[MASK_SLICE_3:.+]] = vector.extract %[[MASK_2D]][3] : vector<8xi1> from vector<4x8xi1>
+// INNER_REDUCTION: %[[REDUCE_3:.+]] = vector.mask %[[MASK_SLICE_3]] { vector.reduction <add>, %{{.+}} : vector<8xf32> into f32 } : vector<8xi1> -> f32
+// INNER_REDUCTION: %[[INSERT_3:.+]] = vector.insert
+//
+// INNER_PARALLEL: vector.multi_reduction <add>
+
+// ALL-LABEL: func @inner_reduction_2d_scalable
+// ALL-SAME: %[[INPUT:.+]]: vector<2x[4]xf32>
+// ALL-SAME: %[[ACC:.+]]: vector<2xf32>
+// ALL-SAME: %[[MASK:.+]]: vector<2x[4]xi1>
+func.func @inner_reduction_2d_scalable(%input: vector<2x[4]xf32>, %acc: vector<2xf32>, %mask: vector<2x[4]xi1>) -> vector<2xf32> {
+ // INNER_REDUCTION: %[[INIT:.+]] = arith.constant dense<0.000000e+00> : vector<2xf32>
+ // INNER_REDUCTION: %[[INPUT_0:.+]] = vector.extract %[[INPUT]][0] : vector<[4]xf32> from vector<2x[4]xf32>
+ // INNER_REDUCTION: %[[ACC_0:.+]] = vector.extract %[[ACC]][0] : f32 from vector<2xf32>
+ // INNER_REDUCTION: %[[MASK_0:.+]] = vector.extract %[[MASK]][0] : vector<[4]xi1> from vector<2x[4]xi1>
+ // INNER_REDUCTION: %[[REDUCE_0:.+]] = vector.mask %[[MASK_0]] { vector.reduction <add>, %[[INPUT_0]], %[[ACC_0]] : vector<[4]xf32> into f32 } : vector<[4]xi1> -> f32
+ // INNER_REDUCTION: %[[INSERT_0:.+]] = vector.insert %[[REDUCE_0]], %[[INIT]] [0] : f32 into vector<2xf32>
+ // INNER_REDUCTION: %[[INPUT_1:.+]] = vector.extract %[[INPUT]][1] : vector<[4]xf32> from vector<2x[4]xf32>
+ // INNER_REDUCTION: %[[ACC_1:.+]] = vector.extract %[[ACC]][1] : f32 from vector<2xf32>
+ // INNER_REDUCTION: %[[MASK_1:.+]] = vector.extract %[[MASK]][1] : vector<[4]xi1> from vector<2x[4]xi1>
+ // INNER_REDUCTION: %[[REDUCE_1:.+]] = vector.mask %[[MASK_1]] { vector.reduction <add>, %[[INPUT_1]], %[[ACC_1]] : vector<[4]xf32> into f32 } : vector<[4]xi1> -> f32
+ // INNER_REDUCTION: %[[RESULT:.+]] = vector.insert %[[REDUCE_1]], %[[INSERT_0]] [1] : f32 into vector<2xf32>
+
+ // INNER_PARALLEL: %[[RESULT:.+]] = vector.mask %[[MASK]] { vector.multi_reduction <add>, %[[INPUT]], %[[ACC]] [1] {{.+}} } : vector<2x[4]xi1> -> vector<2xf32>
+ // ALL: return %[[RESULT]] : vector<2xf32>
+ %0 = vector.mask %mask { vector.multi_reduction <add>, %input, %acc [1] : vector<2x[4]xf32> to vector<2xf32> } : vector<2x[4]xi1> -> vector<2xf32>
+ return %0 : vector<2xf32>
+}
+
+// ALL-LABEL: func @inner_parallel_2d
+// ALL-SAME: %[[INPUT:.+]]: vector<4x2xf32>, %[[ACC:.+]]: vector<2xf32>
+func.func @inner_parallel_2d(%arg0: vector<4x2xf32>, %acc: vector<2xf32>) -> vector<2xf32> {
+ // INNER_PARALLEL: %[[V0:.+]] = vector.extract %[[INPUT]][0] : vector<2xf32> from vector<4x2xf32>
+ // INNER_PARALLEL: %[[RV0:.+]] = arith.mulf %[[V0]], %[[ACC]] : vector<2xf32>
+ // INNER_PARALLEL: %[[V1:.+]] = vector.extract %[[INPUT]][1] : vector<2xf32> from vector<4x2xf32>
+ // INNER_PARALLEL: %[[RV1:.+]] = arith.mulf %[[V1]], %[[RV0]] : vector<2xf32>
+ // INNER_PARALLEL: %[[V2:.+]] = vector.extract %[[INPUT]][2] : vector<2xf32> from vector<4x2xf32>
+ // INNER_PARALLEL: %[[RV2:.+]] = arith.mulf %[[V2]], %[[RV1]] : vector<2xf32>
+ // INNER_PARALLEL: %[[V3:.+]] = vector.extract %[[INPUT]][3] : vector<2xf32> from vector<4x2xf32>
+ // INNER_PARALLEL: %[[RESULT:.+]] = arith.mulf %[[V3]], %[[RV2]] : vector<2xf32>
+ // INNER_REDUCTION: %[[RESULT:.+]] = vector.multi_reduction <mul>, %[[INPUT]], %[[ACC]] [0]
+ // ALL: return %[[RESULT]] : vector<2xf32>
+ %0 = vector.multi_reduction <mul>, %arg0, %acc [0] : vector<4x2xf32> to vector<2xf32>
+ return %0 : vector<2xf32>
+}
+
+// ALL-LABEL: func @inner_parallel_2d_masked
+// ALL-SAME: %[[INPUT:.+]]: vector<4x2xf32>, %[[ACC:.+]]: vector<2xf32>, %[[MASK:.+]]: vector<4x2xi1>
+func.func @inner_parallel_2d_masked(%arg0: vector<4x2xf32>, %acc: vector<2xf32>, %mask: vector<4x2xi1>) -> vector<2xf32> {
+ // INNER_PARALLEL: %[[V0:.+]] = vector.extract %[[INPUT]][0] : vector<2xf32> from vector<4x2xf32>
+ // INNER_PARALLEL: %[[M0:.+]] = vector.extract %[[MASK]][0] : vector<2xi1> from vector<4x2xi1>
+ // INNER_PARALLEL: %[[RED0:.+]] = arith.mulf %[[V0]], %[[ACC]] : vector<2xf32>
+ // INNER_PARALLEL: %[[RV0:.+]] = arith.select %[[M0]], %[[RED0]], %[[ACC]] : vector<2xi1>, vector<2xf32>
+ // INNER_PARALLEL: %[[V1:.+]] = vector.extract %[[INPUT]][1] : vector<2xf32> from vector<4x2xf32>
+ // INNER_PARALLEL: %[[M1:.+]] = vector.extract %[[MASK]][1] : vector<2xi1> from vector<4x2xi1>
+ // INNER_PARALLEL: %[[RED1:.+]] = arith.mulf %[[V1]], %[[RV0]] : vector<2xf32>
+ // INNER_PARALLEL: %[[RV1:.+]] = arith.select %[[M1]], %[[RED1]], %[[RV0]] : vector<2xi1>, vector<2xf32>
+ // INNER_PARALLEL: %[[V2:.+]] = vector.extract %[[INPUT]][2] : vector<2xf32> from vector<4x2xf32>
+ // INNER_PARALLEL: %[[M2:.+]] = vector.extract %[[MASK]][2] : vector<2xi1> from vector<4x2xi1>
+ // INNER_PARALLEL: %[[RED2:.+]] = arith.mulf %[[V2]], %[[RV1]] : vector<2xf32>
+ // INNER_PARALLEL: %[[RV2:.+]] = arith.select %[[M2]], %[[RED2]], %[[RV1]] : vector<2xi1>, vector<2xf32>
+ // INNER_PARALLEL: %[[V3:.+]] = vector.extract %[[INPUT]][3] : vector<2xf32> from vector<4x2xf32>
+ // INNER_PARALLEL: %[[M3:.+]] = vector.extract %[[MASK]][3] : vector<2xi1> from vector<4x2xi1>
+ // INNER_PARALLEL: %[[RED3:.+]] = arith.mulf %[[V3]], %[[RV2]] : vector<2xf32>
+ // INNER_PARALLEL: %[[RESULT:.+]] = arith.select %[[M3]], %[[RED3]], %[[RV2]] : vector<2xi1>, vector<2xf32>
+ // INNER_REDUCTION: %[[RESULT:.+]] = vector.mask %[[MASK]] { vector.multi_reduction <mul>, %[[INPUT]], %[[ACC]] [0] {{.+}} } : vector<4x2xi1> -> vector<2xf32>
+ // ALL: return %[[RESULT]] : vector<2xf32>
+ %0 = vector.mask %mask { vector.multi_reduction <mul>, %arg0, %acc [0] : vector<4x2xf32> to vector<2xf32> } : vector<4x2xi1> -> vector<2xf32>
+ return %0 : vector<2xf32>
+}
+
+module attributes {transform.with_named_sequence} {
+ transform.named_sequence @innerreduction(%root : !transform.any_op {transform.readonly}) {
+ %func_op = transform.structured.match ops{["func.func"]} in %root : (!transform.any_op) -> !transform.op<"func.func">
+ transform.apply_patterns to %func_op {
+ transform.apply_patterns.vector.multi_reduction_unrolling lowering_strategy = "innerreduction"
+ } : !transform.op<"func.func">
+ transform.yield
+ }
+
+ transform.named_sequence @innerparallel(%root : !transform.any_op {transform.readonly}) {
+ %func_op = transform.structured.match ops{["func.func"]} in %root : (!transform.any_op) -> !transform.op<"func.func">
+ transform.apply_patterns to %func_op {
+ transform.apply_patterns.vector.multi_reduction_unrolling lowering_strategy = "innerparallel"
+ } : !transform.op<"func.func">
+ transform.yield
+ }
+}
diff --git a/mlir/test/python/dialects/transform_vector_ext.py b/mlir/test/python/dialects/transform_vector_ext.py
index 29ce8ba63cd53..76e39864fea9f 100644
--- a/mlir/test/python/dialects/transform_vector_ext.py
+++ b/mlir/test/python/dialects/transform_vector_ext.py
@@ -116,6 +116,14 @@ def enum_configurable_patterns():
lowering_strategy=vector.VectorMultiReductionLowering.InnerReduction
)
+ # CHECK: transform.apply_patterns.vector.multi_reduction_unrolling
+ vector.ApplyMultiReductionUnrollingPatternsOp()
+ # CHECK: transform.apply_patterns.vector.multi_reduction_unrolling
+ # CHECK-SAME: lowering_strategy = innerreduction
+ vector.ApplyMultiReductionUnrollingPatternsOp(
+ lowering_strategy=vector.VectorMultiReductionLowering.InnerReduction
+ )
+
# CHECK: transform.apply_patterns.vector.lower_transpose
vector.ApplyLowerTransposePatternsOp()
# CHECK: transform.apply_patterns.vector.lower_transpose
>From b67d56e11d3384b26f2ebe237c9f429969991002 Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Fri, 30 Jan 2026 11:13:52 -0500
Subject: [PATCH 02/19] [mlir][vector] rank reduce unrolling for
vector.multi_reduction.
---
.../Transforms/LowerVectorMultiReduction.cpp | 199 +++++++++++++++++-
.../vector-multi-reduction-pass-lowering.mlir | 6 +-
2 files changed, 198 insertions(+), 7 deletions(-)
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index fec04c967c9e1..90fed006327e2 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -486,6 +486,195 @@ struct OneDimMultiReductionToTwoDim
}
};
+/// Unrolls outermost dimension for vector.multi_reduction.
+/// This patterns matches operations which reduce the outermost dimension,
+/// it does not transform operations for which the outermost dimension is not
+/// a reduction dimension.
+///
+/// There are two cases to consider:
+/// 1. The base case is when the outermost dimension is the only reduction
+/// dimension.
+/// 2. The general case is when the outermost dimension is not the only
+/// reduction dimension.
+///
+/// The base case transformation:
+///
+/// ```mlir
+/// %res = vector.multi_reduction <add> %src, %acc [0] : vector<NxMx...xf32> to
+/// vector<Mx...xf32>
+/// ```
+///
+/// will extract N vectors from %src and then perform elementwise operations.
+///
+/// ```mlir
+/// %0 = vector.extract %src[0] : vector<Mx...xf32> from vector<NxMx...xf32>
+/// ...
+/// %Nminus1 = vector.extract %src[ [[N-1]] ] : vector<Mx...x.f32> from
+/// vector<NxMx...xf32>
+///
+/// %res0 = arith.addf %0, %acc : vector<Mx...xf32>
+/// ...
+/// %res = arith.addf %Nminus1, %resNminus2 : vector<Mx...xf32>
+/// ```
+///
+/// For the general case, we still extract N vectors, but produce N
+/// vector.multi_reduction instead of elementwise operations.
+///
+/// ```mlir
+/// %res = vector.multi_reduction <add> %src, %acc [0, [[REDUCTION_DIMS]] ] :
+/// vector<NxMx...xf32> to vector<Ix...xf32>
+///
+/// ```mlir
+/// %0 = vector.extract %src[0] : vector<Mx...xf32> from vector<NxMx...xf32>
+/// ...
+/// %Nminus1 = vector.extract %src[ [[N-1]] ] : vector<Mx...x.f32> from
+/// vector<NxMx...xf32>
+///
+/// %red0 = vector.multi_reduction %0, %acc [ [[REDUCTION_DIMS]] ] :
+/// vector<Mx...xf32> to vector<Ix...xf32>
+/// ...
+/// %res = vector.multi_reduction %Nminus1, %redNminus2 [ [[REDUCTION_DIMS]] ] :
+/// vector<Mx...xf32> to vector<Ix...xf32>
+/// ```
+struct UnrollMultiReductionOuterBaseCase
+ : public OpRewritePattern<vector::MultiDimReductionOp> {
+ using Base::Base;
+
+ LogicalResult matchAndRewrite(vector::MultiDimReductionOp multiReductionOp,
+ PatternRewriter &rewriter) const override {
+ if (!multiReductionOp.isReducedDim(0))
+ return rewriter.notifyMatchFailure(
+ multiReductionOp,
+ "expected outermost dimension to be reduced dimension.");
+
+ Type elementType = getElementTypeOrSelf(multiReductionOp.getDestType());
+ if (!elementType.isIntOrIndexOrFloat())
+ return rewriter.notifyMatchFailure(
+ multiReductionOp, "expected integer or float element type.");
+
+ ArrayRef<int64_t> reductionDims = multiReductionOp.getReductionDims();
+ if (reductionDims.size() > 1)
+ return rewriter.notifyMatchFailure(
+ multiReductionOp, "expected only one reduction dimension.");
+
+ Location loc = multiReductionOp.getLoc();
+ Value source = multiReductionOp.getSource();
+
+ ArrayRef<int64_t> srcShape =
+ multiReductionOp.getSourceVectorType().getShape();
+ int64_t numElementwiseOps = srcShape.front();
+
+ OpBuilder::InsertionGuard guard(rewriter);
+ auto maskableOp =
+ cast<vector::MaskableOpInterface>(multiReductionOp.getOperation());
+ bool isMasked = maskableOp.isMasked();
+ Operation *rootOp;
+ Value mask = nullptr;
+ if (isMasked) {
+ rewriter.setInsertionPoint(maskableOp.getMaskingOp());
+ rootOp = maskableOp.getMaskingOp();
+ mask = maskableOp.getMaskingOp().getMask();
+ } else {
+ rootOp = multiReductionOp;
+ }
+
+ SmallVector<Value> vectors;
+ for (int64_t i = 0; i < numElementwiseOps; ++i)
+ vectors.push_back(vector::ExtractOp::create(rewriter, loc, source, i));
+
+ SmallVector<Value> masks;
+ for (int64_t i = 0; i < numElementwiseOps; ++i)
+ if (isMasked)
+ masks.push_back(vector::ExtractOp::create(rewriter, loc, mask, i));
+ else
+ masks.push_back(nullptr);
+
+ Value result = multiReductionOp.getAcc();
+ for (auto [innerVector, innerMask] : llvm::zip(vectors, masks))
+ result = makeArithReduction(rewriter, loc, multiReductionOp.getKind(),
+ innerVector, result, /*fastmath=*/nullptr,
+ innerMask);
+
+ rewriter.replaceOp(rootOp, result);
+ return success();
+ }
+};
+
+struct UnrollMultiReductionOuterGeneralCase
+ : public OpRewritePattern<vector::MultiDimReductionOp> {
+ using Base::Base;
+
+ LogicalResult matchAndRewrite(vector::MultiDimReductionOp multiReductionOp,
+ PatternRewriter &rewriter) const override {
+ if (!multiReductionOp.isReducedDim(0))
+ return rewriter.notifyMatchFailure(
+ multiReductionOp,
+ "expected outermost dimension to be reduced dimension.");
+
+ Type elementType = getElementTypeOrSelf(multiReductionOp.getDestType());
+ if (!elementType.isIntOrIndexOrFloat())
+ return rewriter.notifyMatchFailure(
+ multiReductionOp, "expected integer or float element type.");
+
+ ArrayRef<int64_t> reductionDims = multiReductionOp.getReductionDims();
+ if (reductionDims.size() <= 1)
+ return rewriter.notifyMatchFailure(
+ multiReductionOp, "expected more than one reduction dimension.");
+
+ Location loc = multiReductionOp.getLoc();
+ Value source = multiReductionOp.getSource();
+
+ ArrayRef<int64_t> srcShape =
+ multiReductionOp.getSourceVectorType().getShape();
+ int64_t numElementwiseOps = srcShape.front();
+
+ OpBuilder::InsertionGuard guard(rewriter);
+ auto maskableOp =
+ cast<vector::MaskableOpInterface>(multiReductionOp.getOperation());
+ bool isMasked = maskableOp.isMasked();
+ Operation *rootOp;
+ Value mask = nullptr;
+ if (isMasked) {
+ rewriter.setInsertionPoint(maskableOp.getMaskingOp());
+ rootOp = maskableOp.getMaskingOp();
+ mask = maskableOp.getMaskingOp().getMask();
+ } else {
+ rootOp = multiReductionOp;
+ }
+
+ SmallVector<Value> vectors;
+ for (int64_t i = 0; i < numElementwiseOps; ++i)
+ vectors.push_back(vector::ExtractOp::create(rewriter, loc, source, i));
+
+ SmallVector<Value> masks;
+ for (int64_t i = 0; i < numElementwiseOps; ++i)
+ if (isMasked)
+ masks.push_back(vector::ExtractOp::create(rewriter, loc, mask, i));
+ else
+ masks.push_back(nullptr);
+
+ ArrayRef<bool> reductionMask =
+ ArrayRef<bool>(multiReductionOp.getReductionMask()).drop_front();
+ Value result = multiReductionOp.getAcc();
+ for (auto [innerVector, innerMask] : llvm::zip(vectors, masks)) {
+
+ auto reductionOp = vector::MultiDimReductionOp::create(
+ rewriter, loc, innerVector, result, reductionMask,
+ multiReductionOp.getKind());
+
+ if (isMasked) {
+ auto maskOp = vector::maskOperation(rewriter, reductionOp, innerMask);
+ result = maskOp->getResult(0);
+ } else {
+ result = reductionOp.getResult();
+ }
+ }
+
+ rewriter.replaceOp(rootOp, result);
+ return success();
+ }
+};
+
struct LowerVectorMultiReductionPass
: public vector::impl::LowerVectorMultiReductionBase<
LowerVectorMultiReductionPass> {
@@ -541,12 +730,14 @@ void mlir::vector::populateVectorMultiReductionFlatteningPatterns(
void mlir::vector::populateVectorMultiReductionUnrollingPatterns(
RewritePatternSet &patterns, VectorMultiReductionLowering options,
PatternBenefit benefit) {
- if (options == VectorMultiReductionLowering ::InnerReduction)
+ if (options == VectorMultiReductionLowering ::InnerReduction) {
patterns.add<TwoDimMultiReductionToReduction>(patterns.getContext(),
benefit);
- else
- patterns.add<TwoDimMultiReductionToElementWise>(patterns.getContext(),
- benefit);
+ } else {
+ patterns.add<UnrollMultiReductionOuterBaseCase,
+ UnrollMultiReductionOuterGeneralCase>(patterns.getContext(),
+ benefit);
+ }
}
void mlir::vector::populateVectorMultiReductionLoweringPatterns(
diff --git a/mlir/test/Dialect/Vector/vector-multi-reduction-pass-lowering.mlir b/mlir/test/Dialect/Vector/vector-multi-reduction-pass-lowering.mlir
index ddbc5c7bdb2c0..e01bf446eb83c 100644
--- a/mlir/test/Dialect/Vector/vector-multi-reduction-pass-lowering.mlir
+++ b/mlir/test/Dialect/Vector/vector-multi-reduction-pass-lowering.mlir
@@ -21,12 +21,12 @@ func.func @vector_multi_reduction(%arg0: vector<2x4xf32>, %acc: vector<2xf32>) -
// INNER-PARALLEL: %[[TRANSPOSED:.+]] = vector.transpose %[[INPUT]], [1, 0] : vector<2x4xf32> to vector<4x2xf32>
// INNER-PARALLEL: %[[V0:.+]] = vector.extract %[[TRANSPOSED]][0] : vector<2xf32> from vector<4x2xf32>
-// INNER-PARALLEL: %[[RV0:.+]] = arith.mulf %[[V0]], %[[ACC]] : vector<2xf32>
// INNER-PARALLEL: %[[V1:.+]] = vector.extract %[[TRANSPOSED]][1] : vector<2xf32> from vector<4x2xf32>
-// INNER-PARALLEL: %[[RV01:.+]] = arith.mulf %[[V1]], %[[RV0]] : vector<2xf32>
// INNER-PARALLEL: %[[V2:.+]] = vector.extract %[[TRANSPOSED]][2] : vector<2xf32> from vector<4x2xf32>
-// INNER-PARALLEL: %[[RV012:.+]] = arith.mulf %[[V2]], %[[RV01]] : vector<2xf32>
// INNER-PARALLEL: %[[V3:.+]] = vector.extract %[[TRANSPOSED]][3] : vector<2xf32> from vector<4x2xf32>
+// INNER-PARALLEL: %[[RV0:.+]] = arith.mulf %[[V0]], %[[ACC]] : vector<2xf32>
+// INNER-PARALLEL: %[[RV01:.+]] = arith.mulf %[[V1]], %[[RV0]] : vector<2xf32>
+// INNER-PARALLEL: %[[RV012:.+]] = arith.mulf %[[V2]], %[[RV01]] : vector<2xf32>
// INNER-PARALLEL: %[[RESULT_VEC:.+]] = arith.mulf %[[V3]], %[[RV012]] : vector<2xf32>
// INNER-PARALLEL: return %[[RESULT_VEC]] : vector<2xf32>
>From d9326dc9b5e43e8a7bd8c649a6d7aa97f93fa581 Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Fri, 30 Jan 2026 12:07:30 -0500
Subject: [PATCH 03/19] [mlir][vector] Add tests for multi_reduction unrolling.
---
.../Vector/TransformOps/VectorTransformOps.td | 20 ++++
.../Vector/Transforms/LoweringPatterns.h | 24 +++++
.../TransformOps/VectorTransformOps.cpp | 8 ++
.../Transforms/LowerVectorMultiReduction.cpp | 16 ++++
.../Vector/td/unroll-multi-reduction.mlir | 24 +++++
.../Vector/unroll-vector-multi-reduction.mlir | 92 +++++++++++++++++++
6 files changed, 184 insertions(+)
create mode 100644 mlir/test/Dialect/Vector/td/unroll-multi-reduction.mlir
create mode 100644 mlir/test/Dialect/Vector/unroll-vector-multi-reduction.mlir
diff --git a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
index 685c88c17e556..01fb33274828a 100644
--- a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
+++ b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
@@ -306,6 +306,26 @@ def ApplyMultiReductionUnrollingPatternsOp: Op<Transform_Dialect,
}];
}
+def ApplyUnrollMultiReductionPatternsOp : Op<Transform_Dialect,
+ "apply_patterns.vector.unroll_multi_reduction",
+ [DeclareOpInterfaceMethods<PatternDescriptorOpInterface>]> {
+ let description = [{
+ Unrolls vector.multi_reduction operations by progressively reducing rank
+ along the outermost dimension.
+
+ This is an alternative to the flattening-based lowering that preserves
+ the n-D structure during progressive lowering.
+ }];
+
+ let arguments = (ins DefaultValuedAttr<VectorMultiReductionLoweringAttr,
+ "vector::VectorMultiReductionLowering::InnerParallel">:$lowering_strategy
+ );
+
+ let assemblyFormat = [{
+ (`lowering_strategy` `=` $lowering_strategy^)? attr-dict
+ }];
+}
+
def ApplyLowerOuterProductPatternsOp : Op<Transform_Dialect,
"apply_patterns.vector.lower_outerproduct",
[DeclareOpInterfaceMethods<PatternDescriptorOpInterface>]> {
diff --git a/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h b/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
index 33487a9d8d6e0..a0334afa05dad 100644
--- a/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
+++ b/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
@@ -118,6 +118,30 @@ void populateVectorMultiReductionLoweringPatterns(
RewritePatternSet &patterns, VectorMultiReductionLowering options,
PatternBenefit benefit = 1);
+/// Collect a set of patterns to unroll vector.multi_reduction ops by
+/// progressively reducing rank along the outermost dimension.
+///
+/// For OuterReduction (outermost dim is reduction):
+/// [UnrollMultiReductionOuterBaseCase]
+/// When the outermost dimension is the only reduction dimension, unroll to
+/// produce elementwise arithmetic operations.
+///
+/// [UnrollMultiReductionOuterGeneralCase]
+/// When the outermost dimension is one of multiple reduction dimensions,
+/// unroll to produce smaller multi_reduction operations.
+///
+/// For InnerReduction (innermost dim is reduction):
+/// [UnrollMultiReductionInnerBaseCase]
+/// When the innermost dimension is the only reduction dimension, unroll along
+/// the outermost parallel dimension.
+///
+/// [UnrollMultiReductionInnerGeneralCase]
+/// When the innermost dimension is one of multiple reduction dimensions,
+/// unroll along the outermost parallel dimension.
+void populateVectorUnrollMultiReduction(RewritePatternSet &patterns,
+ VectorMultiReductionLowering options,
+ PatternBenefit benefit = 1);
+
/// Populate the pattern set with the following patterns:
///
/// [TransferReadToVectorLoadLowering]
diff --git a/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp b/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
index 6c97c6501a23e..fca1ca6bade92 100644
--- a/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
+++ b/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
@@ -162,6 +162,14 @@ void transform::ApplyMultiReductionUnrollingPatternsOp::populatePatterns(
patterns, vectorTransformOptions.vectorMultiReductionLowering);
}
+void transform::ApplyUnrollMultiReductionPatternsOp::populatePatterns(
+ RewritePatternSet &patterns) {
+ vector::VectorTransformsOptions vectorTransformOptions;
+ vectorTransformOptions.setVectorMultiReductionLowering(getLoweringStrategy());
+ vector::populateVectorUnrollMultiReduction(
+ patterns, vectorTransformOptions.vectorMultiReductionLowering);
+}
+
void transform::ApplyLowerOuterProductPatternsOp::populatePatterns(
RewritePatternSet &patterns) {
populateVectorOuterProductLoweringPatterns(patterns);
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index 90fed006327e2..5ac13dd9fe4d5 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -740,6 +740,22 @@ void mlir::vector::populateVectorMultiReductionUnrollingPatterns(
}
}
+void mlir::vector::populateVectorUnrollMultiReduction(
+ RewritePatternSet &patterns, VectorMultiReductionLowering options,
+ PatternBenefit benefit) {
+ if (options == VectorMultiReductionLowering::InnerReduction) {
+ // TODO: Add UnrollMultiReductionInnerBaseCase and
+ // UnrollMultiReductionInnerGeneralCase patterns here once implemented.
+ // For now, fall back to the existing 2-D based lowering.
+ patterns.add<TwoDimMultiReductionToReduction>(patterns.getContext(),
+ benefit);
+ } else {
+ patterns.add<UnrollMultiReductionOuterBaseCase,
+ UnrollMultiReductionOuterGeneralCase>(patterns.getContext(),
+ benefit);
+ }
+}
+
void mlir::vector::populateVectorMultiReductionLoweringPatterns(
RewritePatternSet &patterns, VectorMultiReductionLowering options,
PatternBenefit benefit) {
diff --git a/mlir/test/Dialect/Vector/td/unroll-multi-reduction.mlir b/mlir/test/Dialect/Vector/td/unroll-multi-reduction.mlir
new file mode 100644
index 0000000000000..96a68723266d3
--- /dev/null
+++ b/mlir/test/Dialect/Vector/td/unroll-multi-reduction.mlir
@@ -0,0 +1,24 @@
+module attributes {transform.with_named_sequence} {
+ transform.named_sequence @unroll_multi_reduction(%module_op: !transform.any_op {transform.readonly}) {
+
+ %func_op = transform.structured.match ops{["func.func"]} in %module_op
+ : (!transform.any_op) -> !transform.any_op
+ transform.apply_patterns to %func_op {
+ // Test patterns
+ transform.apply_patterns.vector.unroll_multi_reduction
+ } : !transform.any_op
+
+ transform.yield
+ }
+ transform.named_sequence @unroll_multi_reduction_inner(%module_op: !transform.any_op {transform.readonly}) {
+
+ %func_op = transform.structured.match ops{["func.func"]} in %module_op
+ : (!transform.any_op) -> !transform.any_op
+ transform.apply_patterns to %func_op {
+ // Test patterns
+ transform.apply_patterns.vector.unroll_multi_reduction lowering_strategy = "innerreduction"
+ } : !transform.any_op
+
+ transform.yield
+ }
+}
diff --git a/mlir/test/Dialect/Vector/unroll-vector-multi-reduction.mlir b/mlir/test/Dialect/Vector/unroll-vector-multi-reduction.mlir
new file mode 100644
index 0000000000000..79086e2b0b9ad
--- /dev/null
+++ b/mlir/test/Dialect/Vector/unroll-vector-multi-reduction.mlir
@@ -0,0 +1,92 @@
+// RUN: mlir-opt --split-input-file %s -transform-preload-library='transform-library-paths=%p/td/unroll-multi-reduction.mlir' \
+// RUN: -transform-interpreter=entry-point=unroll_multi_reduction | FileCheck %s
+
+//===----------------------------------------------------------------------===//
+// Test UnrollVectorMultiReduction
+//===----------------------------------------------------------------------===//
+
+// CHECK-LABEL: func @unroll_vector_multi_reduction(
+// CHECK-SAME: %[[SOURCE:.+]]: vector<2x3x5xf32>,
+// CHECK-SAME: %[[ACC:.+]]: vector<3x5xf32>
+func.func @unroll_vector_multi_reduction(%source: vector<2x3x5xf32>, %acc: vector<3x5xf32>) -> (vector<3x5xf32>) {
+ // CHECK-DAG: %[[VEC_0:.+]] = vector.extract %[[SOURCE]][0] : vector<3x5xf32> from vector<2x3x5xf32>
+ // CHECK-DAG: %[[VEC_1:.+]] = vector.extract %[[SOURCE]][1] : vector<3x5xf32> from vector<2x3x5xf32>
+
+ // CHECK: %[[RES_0:.+]] = arith.addf %[[VEC_0]], %[[ACC]] : vector<3x5xf32>
+ // CHECK: %[[RES_1:.+]] = arith.addf %[[VEC_1]], %[[RES_0]] : vector<3x5xf32>
+ %1 = vector.multi_reduction <add>, %source, %acc [0] : vector<2x3x5xf32> to vector<3x5xf32>
+
+ // CHECK: return %[[RES_1]]
+ return %1 : vector<3x5xf32>
+}
+
+// -----
+
+// CHECK-LABEL: func @unroll_vector_multi_reduction_masked(
+// CHECK-SAME: %[[SOURCE:.+]]: vector<2x3x5xf32>,
+// CHECK-SAME: %[[MASK:.+]]: vector<2x3x5xi1>,
+// CHECK-SAME: %[[ACC:.+]]: vector<3x5xf32>
+func.func @unroll_vector_multi_reduction_masked(%source: vector<2x3x5xf32>, %mask: vector<2x3x5xi1>, %acc: vector<3x5xf32>) -> (vector<3x5xf32>) {
+ // CHECK-DAG: %[[VEC_0:.+]] = vector.extract %[[SOURCE]][0] : vector<3x5xf32> from vector<2x3x5xf32>
+ // CHECK-DAG: %[[VEC_1:.+]] = vector.extract %[[SOURCE]][1] : vector<3x5xf32> from vector<2x3x5xf32>
+
+ // CHECK-DAG: %[[MASK_0:.+]] = vector.extract %[[MASK]][0] : vector<3x5xi1> from vector<2x3x5xi1>
+ // CHECK-DAG: %[[MASK_1:.+]] = vector.extract %[[MASK]][1] : vector<3x5xi1> from vector<2x3x5xi1>
+
+ // CHECK: %[[RES_0:.+]] = arith.addf %[[VEC_0]], %[[ACC]] : vector<3x5xf32>
+ // CHECK: %[[RES_MASKED_0:.+]] = arith.select %[[MASK_0]], %[[RES_0]], %[[ACC]] : vector<3x5xi1>, vector<3x5xf32>
+
+ // CHECK: %[[RES_1:.+]] = arith.addf %[[VEC_1]], %[[RES_MASKED_0]] : vector<3x5xf32>
+ // CHECK: %[[RES_MASKED_1:.+]] = arith.select %[[MASK_1]], %[[RES_1]], %[[RES_MASKED_0]] : vector<3x5xi1>, vector<3x5xf32>
+
+ %0 = vector.mask %mask {
+ %1 = vector.multi_reduction <add>, %source, %acc [0] : vector<2x3x5xf32> to vector<3x5xf32>
+ } : vector<2x3x5xi1> -> vector<3x5xf32>
+
+ // CHECK: return %[[RES_MASKED_1]]
+ return %0 : vector<3x5xf32>
+}
+
+// -----
+
+// CHECK-LABEL: func @unroll_vector_multi_reduction_general(
+// CHECK-SAME: %[[SOURCE:.+]]: vector<2x3x5xf32>,
+// CHECK-SAME: %[[ACC:.+]]: vector<3xf32>
+func.func @unroll_vector_multi_reduction_general(%source: vector<2x3x5xf32>, %acc: vector<3xf32>) -> (vector<3xf32>) {
+
+ // CHECK-DAG: %[[VEC_0:.+]] = vector.extract %[[SOURCE]][0] : vector<3x5xf32> from vector<2x3x5xf32>
+ // CHECK-DAG: %[[VEC_1:.+]] = vector.extract %[[SOURCE]][1] : vector<3x5xf32> from vector<2x3x5xf32>
+
+ // CHECK: %[[RES_0:.+]] = vector.multi_reduction <add>, %[[VEC_0]], %[[ACC]] [1] : vector<3x5xf32> to vector<3xf32>
+ // CHECK: %[[RES_1:.+]] = vector.multi_reduction <add>, %[[VEC_1]], %[[RES_0]] [1] : vector<3x5xf32> to vector<3xf32>
+
+ %1 = vector.multi_reduction <add>, %source, %acc [0, 2] : vector<2x3x5xf32> to vector<3xf32>
+
+ // CHECK: return %[[RES_1]]
+ return %1 : vector<3xf32>
+}
+
+// -----
+
+// CHECK-LABEL: func @unroll_vector_multi_reduction_general_masked(
+// CHECK-SAME: %[[SOURCE:.+]]: vector<2x3x5xf32>,
+// CHECK-SAME: %[[MASK:.+]]: vector<2x3x5xi1>,
+// CHECK-SAME: %[[ACC:.+]]: vector<3xf32>
+func.func @unroll_vector_multi_reduction_general_masked(%source: vector<2x3x5xf32>, %mask: vector<2x3x5xi1>, %acc: vector<3xf32>) -> (vector<3xf32>) {
+
+ // CHECK-DAG: %[[VEC_0:.+]] = vector.extract %[[SOURCE]][0] : vector<3x5xf32> from vector<2x3x5xf32>
+ // CHECK-DAG: %[[VEC_1:.+]] = vector.extract %[[SOURCE]][1] : vector<3x5xf32> from vector<2x3x5xf32>
+
+ // CHECK-DAG: %[[MASK_0:.+]] = vector.extract %[[MASK]][0] : vector<3x5xi1> from vector<2x3x5xi1>
+ // CHECK-DAG: %[[MASK_1:.+]] = vector.extract %[[MASK]][1] : vector<3x5xi1> from vector<2x3x5xi1>
+
+ // CHECK: %[[RES_0:.+]] = vector.mask %[[MASK_0]] { vector.multi_reduction <add>, %[[VEC_0]], %[[ACC]] [1] : vector<3x5xf32> to vector<3xf32> } : vector<3x5xi1> -> vector<3xf32>
+ // CHECK: %[[RES_1:.+]] = vector.mask %[[MASK_1]] { vector.multi_reduction <add>, %[[VEC_1]], %[[RES_0]] [1] : vector<3x5xf32> to vector<3xf32> } : vector<3x5xi1> -> vector<3xf32>
+
+ %0 = vector.mask %mask {
+ %1 = vector.multi_reduction <add>, %source, %acc [0, 2] : vector<2x3x5xf32> to vector<3xf32>
+ } : vector<2x3x5xi1> -> vector<3xf32>
+
+ // CHECK: return %[[RES_1]]
+ return %0 : vector<3xf32>
+}
>From ad7f27abd5f45fefc3e4319a3a0f173daf78a778 Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Wed, 18 Feb 2026 16:20:55 -0500
Subject: [PATCH 04/19] Use already existing infrastructure
---
.../Vector/TransformOps/VectorTransformOps.td | 20 ----
.../Vector/Transforms/LoweringPatterns.h | 24 -----
.../TransformOps/VectorTransformOps.cpp | 8 --
.../Transforms/LowerVectorMultiReduction.cpp | 16 ----
.../Vector/unroll-vector-multi-reduction.mlir | 92 -------------------
5 files changed, 160 deletions(-)
delete mode 100644 mlir/test/Dialect/Vector/unroll-vector-multi-reduction.mlir
diff --git a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
index 01fb33274828a..685c88c17e556 100644
--- a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
+++ b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
@@ -306,26 +306,6 @@ def ApplyMultiReductionUnrollingPatternsOp: Op<Transform_Dialect,
}];
}
-def ApplyUnrollMultiReductionPatternsOp : Op<Transform_Dialect,
- "apply_patterns.vector.unroll_multi_reduction",
- [DeclareOpInterfaceMethods<PatternDescriptorOpInterface>]> {
- let description = [{
- Unrolls vector.multi_reduction operations by progressively reducing rank
- along the outermost dimension.
-
- This is an alternative to the flattening-based lowering that preserves
- the n-D structure during progressive lowering.
- }];
-
- let arguments = (ins DefaultValuedAttr<VectorMultiReductionLoweringAttr,
- "vector::VectorMultiReductionLowering::InnerParallel">:$lowering_strategy
- );
-
- let assemblyFormat = [{
- (`lowering_strategy` `=` $lowering_strategy^)? attr-dict
- }];
-}
-
def ApplyLowerOuterProductPatternsOp : Op<Transform_Dialect,
"apply_patterns.vector.lower_outerproduct",
[DeclareOpInterfaceMethods<PatternDescriptorOpInterface>]> {
diff --git a/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h b/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
index a0334afa05dad..33487a9d8d6e0 100644
--- a/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
+++ b/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
@@ -118,30 +118,6 @@ void populateVectorMultiReductionLoweringPatterns(
RewritePatternSet &patterns, VectorMultiReductionLowering options,
PatternBenefit benefit = 1);
-/// Collect a set of patterns to unroll vector.multi_reduction ops by
-/// progressively reducing rank along the outermost dimension.
-///
-/// For OuterReduction (outermost dim is reduction):
-/// [UnrollMultiReductionOuterBaseCase]
-/// When the outermost dimension is the only reduction dimension, unroll to
-/// produce elementwise arithmetic operations.
-///
-/// [UnrollMultiReductionOuterGeneralCase]
-/// When the outermost dimension is one of multiple reduction dimensions,
-/// unroll to produce smaller multi_reduction operations.
-///
-/// For InnerReduction (innermost dim is reduction):
-/// [UnrollMultiReductionInnerBaseCase]
-/// When the innermost dimension is the only reduction dimension, unroll along
-/// the outermost parallel dimension.
-///
-/// [UnrollMultiReductionInnerGeneralCase]
-/// When the innermost dimension is one of multiple reduction dimensions,
-/// unroll along the outermost parallel dimension.
-void populateVectorUnrollMultiReduction(RewritePatternSet &patterns,
- VectorMultiReductionLowering options,
- PatternBenefit benefit = 1);
-
/// Populate the pattern set with the following patterns:
///
/// [TransferReadToVectorLoadLowering]
diff --git a/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp b/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
index fca1ca6bade92..6c97c6501a23e 100644
--- a/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
+++ b/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
@@ -162,14 +162,6 @@ void transform::ApplyMultiReductionUnrollingPatternsOp::populatePatterns(
patterns, vectorTransformOptions.vectorMultiReductionLowering);
}
-void transform::ApplyUnrollMultiReductionPatternsOp::populatePatterns(
- RewritePatternSet &patterns) {
- vector::VectorTransformsOptions vectorTransformOptions;
- vectorTransformOptions.setVectorMultiReductionLowering(getLoweringStrategy());
- vector::populateVectorUnrollMultiReduction(
- patterns, vectorTransformOptions.vectorMultiReductionLowering);
-}
-
void transform::ApplyLowerOuterProductPatternsOp::populatePatterns(
RewritePatternSet &patterns) {
populateVectorOuterProductLoweringPatterns(patterns);
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index 5ac13dd9fe4d5..90fed006327e2 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -740,22 +740,6 @@ void mlir::vector::populateVectorMultiReductionUnrollingPatterns(
}
}
-void mlir::vector::populateVectorUnrollMultiReduction(
- RewritePatternSet &patterns, VectorMultiReductionLowering options,
- PatternBenefit benefit) {
- if (options == VectorMultiReductionLowering::InnerReduction) {
- // TODO: Add UnrollMultiReductionInnerBaseCase and
- // UnrollMultiReductionInnerGeneralCase patterns here once implemented.
- // For now, fall back to the existing 2-D based lowering.
- patterns.add<TwoDimMultiReductionToReduction>(patterns.getContext(),
- benefit);
- } else {
- patterns.add<UnrollMultiReductionOuterBaseCase,
- UnrollMultiReductionOuterGeneralCase>(patterns.getContext(),
- benefit);
- }
-}
-
void mlir::vector::populateVectorMultiReductionLoweringPatterns(
RewritePatternSet &patterns, VectorMultiReductionLowering options,
PatternBenefit benefit) {
diff --git a/mlir/test/Dialect/Vector/unroll-vector-multi-reduction.mlir b/mlir/test/Dialect/Vector/unroll-vector-multi-reduction.mlir
deleted file mode 100644
index 79086e2b0b9ad..0000000000000
--- a/mlir/test/Dialect/Vector/unroll-vector-multi-reduction.mlir
+++ /dev/null
@@ -1,92 +0,0 @@
-// RUN: mlir-opt --split-input-file %s -transform-preload-library='transform-library-paths=%p/td/unroll-multi-reduction.mlir' \
-// RUN: -transform-interpreter=entry-point=unroll_multi_reduction | FileCheck %s
-
-//===----------------------------------------------------------------------===//
-// Test UnrollVectorMultiReduction
-//===----------------------------------------------------------------------===//
-
-// CHECK-LABEL: func @unroll_vector_multi_reduction(
-// CHECK-SAME: %[[SOURCE:.+]]: vector<2x3x5xf32>,
-// CHECK-SAME: %[[ACC:.+]]: vector<3x5xf32>
-func.func @unroll_vector_multi_reduction(%source: vector<2x3x5xf32>, %acc: vector<3x5xf32>) -> (vector<3x5xf32>) {
- // CHECK-DAG: %[[VEC_0:.+]] = vector.extract %[[SOURCE]][0] : vector<3x5xf32> from vector<2x3x5xf32>
- // CHECK-DAG: %[[VEC_1:.+]] = vector.extract %[[SOURCE]][1] : vector<3x5xf32> from vector<2x3x5xf32>
-
- // CHECK: %[[RES_0:.+]] = arith.addf %[[VEC_0]], %[[ACC]] : vector<3x5xf32>
- // CHECK: %[[RES_1:.+]] = arith.addf %[[VEC_1]], %[[RES_0]] : vector<3x5xf32>
- %1 = vector.multi_reduction <add>, %source, %acc [0] : vector<2x3x5xf32> to vector<3x5xf32>
-
- // CHECK: return %[[RES_1]]
- return %1 : vector<3x5xf32>
-}
-
-// -----
-
-// CHECK-LABEL: func @unroll_vector_multi_reduction_masked(
-// CHECK-SAME: %[[SOURCE:.+]]: vector<2x3x5xf32>,
-// CHECK-SAME: %[[MASK:.+]]: vector<2x3x5xi1>,
-// CHECK-SAME: %[[ACC:.+]]: vector<3x5xf32>
-func.func @unroll_vector_multi_reduction_masked(%source: vector<2x3x5xf32>, %mask: vector<2x3x5xi1>, %acc: vector<3x5xf32>) -> (vector<3x5xf32>) {
- // CHECK-DAG: %[[VEC_0:.+]] = vector.extract %[[SOURCE]][0] : vector<3x5xf32> from vector<2x3x5xf32>
- // CHECK-DAG: %[[VEC_1:.+]] = vector.extract %[[SOURCE]][1] : vector<3x5xf32> from vector<2x3x5xf32>
-
- // CHECK-DAG: %[[MASK_0:.+]] = vector.extract %[[MASK]][0] : vector<3x5xi1> from vector<2x3x5xi1>
- // CHECK-DAG: %[[MASK_1:.+]] = vector.extract %[[MASK]][1] : vector<3x5xi1> from vector<2x3x5xi1>
-
- // CHECK: %[[RES_0:.+]] = arith.addf %[[VEC_0]], %[[ACC]] : vector<3x5xf32>
- // CHECK: %[[RES_MASKED_0:.+]] = arith.select %[[MASK_0]], %[[RES_0]], %[[ACC]] : vector<3x5xi1>, vector<3x5xf32>
-
- // CHECK: %[[RES_1:.+]] = arith.addf %[[VEC_1]], %[[RES_MASKED_0]] : vector<3x5xf32>
- // CHECK: %[[RES_MASKED_1:.+]] = arith.select %[[MASK_1]], %[[RES_1]], %[[RES_MASKED_0]] : vector<3x5xi1>, vector<3x5xf32>
-
- %0 = vector.mask %mask {
- %1 = vector.multi_reduction <add>, %source, %acc [0] : vector<2x3x5xf32> to vector<3x5xf32>
- } : vector<2x3x5xi1> -> vector<3x5xf32>
-
- // CHECK: return %[[RES_MASKED_1]]
- return %0 : vector<3x5xf32>
-}
-
-// -----
-
-// CHECK-LABEL: func @unroll_vector_multi_reduction_general(
-// CHECK-SAME: %[[SOURCE:.+]]: vector<2x3x5xf32>,
-// CHECK-SAME: %[[ACC:.+]]: vector<3xf32>
-func.func @unroll_vector_multi_reduction_general(%source: vector<2x3x5xf32>, %acc: vector<3xf32>) -> (vector<3xf32>) {
-
- // CHECK-DAG: %[[VEC_0:.+]] = vector.extract %[[SOURCE]][0] : vector<3x5xf32> from vector<2x3x5xf32>
- // CHECK-DAG: %[[VEC_1:.+]] = vector.extract %[[SOURCE]][1] : vector<3x5xf32> from vector<2x3x5xf32>
-
- // CHECK: %[[RES_0:.+]] = vector.multi_reduction <add>, %[[VEC_0]], %[[ACC]] [1] : vector<3x5xf32> to vector<3xf32>
- // CHECK: %[[RES_1:.+]] = vector.multi_reduction <add>, %[[VEC_1]], %[[RES_0]] [1] : vector<3x5xf32> to vector<3xf32>
-
- %1 = vector.multi_reduction <add>, %source, %acc [0, 2] : vector<2x3x5xf32> to vector<3xf32>
-
- // CHECK: return %[[RES_1]]
- return %1 : vector<3xf32>
-}
-
-// -----
-
-// CHECK-LABEL: func @unroll_vector_multi_reduction_general_masked(
-// CHECK-SAME: %[[SOURCE:.+]]: vector<2x3x5xf32>,
-// CHECK-SAME: %[[MASK:.+]]: vector<2x3x5xi1>,
-// CHECK-SAME: %[[ACC:.+]]: vector<3xf32>
-func.func @unroll_vector_multi_reduction_general_masked(%source: vector<2x3x5xf32>, %mask: vector<2x3x5xi1>, %acc: vector<3xf32>) -> (vector<3xf32>) {
-
- // CHECK-DAG: %[[VEC_0:.+]] = vector.extract %[[SOURCE]][0] : vector<3x5xf32> from vector<2x3x5xf32>
- // CHECK-DAG: %[[VEC_1:.+]] = vector.extract %[[SOURCE]][1] : vector<3x5xf32> from vector<2x3x5xf32>
-
- // CHECK-DAG: %[[MASK_0:.+]] = vector.extract %[[MASK]][0] : vector<3x5xi1> from vector<2x3x5xi1>
- // CHECK-DAG: %[[MASK_1:.+]] = vector.extract %[[MASK]][1] : vector<3x5xi1> from vector<2x3x5xi1>
-
- // CHECK: %[[RES_0:.+]] = vector.mask %[[MASK_0]] { vector.multi_reduction <add>, %[[VEC_0]], %[[ACC]] [1] : vector<3x5xf32> to vector<3xf32> } : vector<3x5xi1> -> vector<3xf32>
- // CHECK: %[[RES_1:.+]] = vector.mask %[[MASK_1]] { vector.multi_reduction <add>, %[[VEC_1]], %[[RES_0]] [1] : vector<3x5xf32> to vector<3xf32> } : vector<3x5xi1> -> vector<3xf32>
-
- %0 = vector.mask %mask {
- %1 = vector.multi_reduction <add>, %source, %acc [0, 2] : vector<2x3x5xf32> to vector<3xf32>
- } : vector<2x3x5xi1> -> vector<3xf32>
-
- // CHECK: return %[[RES_1]]
- return %0 : vector<3xf32>
-}
>From 209e1d513ce7ced8702e63207ea02582d36f4953 Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Wed, 18 Feb 2026 16:34:54 -0500
Subject: [PATCH 05/19] fix tests
---
.../Transforms/LowerVectorMultiReduction.cpp | 5 +++++
.../vector-multi-reduction-unrolling.mlir | 20 ++++++++++---------
2 files changed, 16 insertions(+), 9 deletions(-)
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index 90fed006327e2..6df55232c605d 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -542,6 +542,11 @@ struct UnrollMultiReductionOuterBaseCase
LogicalResult matchAndRewrite(vector::MultiDimReductionOp multiReductionOp,
PatternRewriter &rewriter) const override {
+ auto srcRank = multiReductionOp.getSourceVectorType().getRank();
+ if (srcRank < 2)
+ return rewriter.notifyMatchFailure(multiReductionOp,
+ "expected source rank >= 2.");
+
if (!multiReductionOp.isReducedDim(0))
return rewriter.notifyMatchFailure(
multiReductionOp,
diff --git a/mlir/test/Dialect/Vector/vector-multi-reduction-unrolling.mlir b/mlir/test/Dialect/Vector/vector-multi-reduction-unrolling.mlir
index d4fb79a1d4668..3eefb8c53f92c 100644
--- a/mlir/test/Dialect/Vector/vector-multi-reduction-unrolling.mlir
+++ b/mlir/test/Dialect/Vector/vector-multi-reduction-unrolling.mlir
@@ -99,12 +99,12 @@ func.func @inner_reduction_2d_scalable(%input: vector<2x[4]xf32>, %acc: vector<2
// ALL-SAME: %[[INPUT:.+]]: vector<4x2xf32>, %[[ACC:.+]]: vector<2xf32>
func.func @inner_parallel_2d(%arg0: vector<4x2xf32>, %acc: vector<2xf32>) -> vector<2xf32> {
// INNER_PARALLEL: %[[V0:.+]] = vector.extract %[[INPUT]][0] : vector<2xf32> from vector<4x2xf32>
- // INNER_PARALLEL: %[[RV0:.+]] = arith.mulf %[[V0]], %[[ACC]] : vector<2xf32>
// INNER_PARALLEL: %[[V1:.+]] = vector.extract %[[INPUT]][1] : vector<2xf32> from vector<4x2xf32>
- // INNER_PARALLEL: %[[RV1:.+]] = arith.mulf %[[V1]], %[[RV0]] : vector<2xf32>
// INNER_PARALLEL: %[[V2:.+]] = vector.extract %[[INPUT]][2] : vector<2xf32> from vector<4x2xf32>
- // INNER_PARALLEL: %[[RV2:.+]] = arith.mulf %[[V2]], %[[RV1]] : vector<2xf32>
// INNER_PARALLEL: %[[V3:.+]] = vector.extract %[[INPUT]][3] : vector<2xf32> from vector<4x2xf32>
+ // INNER_PARALLEL: %[[RV0:.+]] = arith.mulf %[[V0]], %[[ACC]] : vector<2xf32>
+ // INNER_PARALLEL: %[[RV1:.+]] = arith.mulf %[[V1]], %[[RV0]] : vector<2xf32>
+ // INNER_PARALLEL: %[[RV2:.+]] = arith.mulf %[[V2]], %[[RV1]] : vector<2xf32>
// INNER_PARALLEL: %[[RESULT:.+]] = arith.mulf %[[V3]], %[[RV2]] : vector<2xf32>
// INNER_REDUCTION: %[[RESULT:.+]] = vector.multi_reduction <mul>, %[[INPUT]], %[[ACC]] [0]
// ALL: return %[[RESULT]] : vector<2xf32>
@@ -116,19 +116,21 @@ func.func @inner_parallel_2d(%arg0: vector<4x2xf32>, %acc: vector<2xf32>) -> vec
// ALL-SAME: %[[INPUT:.+]]: vector<4x2xf32>, %[[ACC:.+]]: vector<2xf32>, %[[MASK:.+]]: vector<4x2xi1>
func.func @inner_parallel_2d_masked(%arg0: vector<4x2xf32>, %acc: vector<2xf32>, %mask: vector<4x2xi1>) -> vector<2xf32> {
// INNER_PARALLEL: %[[V0:.+]] = vector.extract %[[INPUT]][0] : vector<2xf32> from vector<4x2xf32>
+ // INNER_PARALLEL: %[[V1:.+]] = vector.extract %[[INPUT]][1] : vector<2xf32> from vector<4x2xf32>
+ // INNER_PARALLEL: %[[V2:.+]] = vector.extract %[[INPUT]][2] : vector<2xf32> from vector<4x2xf32>
+ // INNER_PARALLEL: %[[V3:.+]] = vector.extract %[[INPUT]][3] : vector<2xf32> from vector<4x2xf32>
+
// INNER_PARALLEL: %[[M0:.+]] = vector.extract %[[MASK]][0] : vector<2xi1> from vector<4x2xi1>
+ // INNER_PARALLEL: %[[M1:.+]] = vector.extract %[[MASK]][1] : vector<2xi1> from vector<4x2xi1>
+ // INNER_PARALLEL: %[[M2:.+]] = vector.extract %[[MASK]][2] : vector<2xi1> from vector<4x2xi1>
+ // INNER_PARALLEL: %[[M3:.+]] = vector.extract %[[MASK]][3] : vector<2xi1> from vector<4x2xi1>
+
// INNER_PARALLEL: %[[RED0:.+]] = arith.mulf %[[V0]], %[[ACC]] : vector<2xf32>
// INNER_PARALLEL: %[[RV0:.+]] = arith.select %[[M0]], %[[RED0]], %[[ACC]] : vector<2xi1>, vector<2xf32>
- // INNER_PARALLEL: %[[V1:.+]] = vector.extract %[[INPUT]][1] : vector<2xf32> from vector<4x2xf32>
- // INNER_PARALLEL: %[[M1:.+]] = vector.extract %[[MASK]][1] : vector<2xi1> from vector<4x2xi1>
// INNER_PARALLEL: %[[RED1:.+]] = arith.mulf %[[V1]], %[[RV0]] : vector<2xf32>
// INNER_PARALLEL: %[[RV1:.+]] = arith.select %[[M1]], %[[RED1]], %[[RV0]] : vector<2xi1>, vector<2xf32>
- // INNER_PARALLEL: %[[V2:.+]] = vector.extract %[[INPUT]][2] : vector<2xf32> from vector<4x2xf32>
- // INNER_PARALLEL: %[[M2:.+]] = vector.extract %[[MASK]][2] : vector<2xi1> from vector<4x2xi1>
// INNER_PARALLEL: %[[RED2:.+]] = arith.mulf %[[V2]], %[[RV1]] : vector<2xf32>
// INNER_PARALLEL: %[[RV2:.+]] = arith.select %[[M2]], %[[RED2]], %[[RV1]] : vector<2xi1>, vector<2xf32>
- // INNER_PARALLEL: %[[V3:.+]] = vector.extract %[[INPUT]][3] : vector<2xf32> from vector<4x2xf32>
- // INNER_PARALLEL: %[[M3:.+]] = vector.extract %[[MASK]][3] : vector<2xi1> from vector<4x2xi1>
// INNER_PARALLEL: %[[RED3:.+]] = arith.mulf %[[V3]], %[[RV2]] : vector<2xf32>
// INNER_PARALLEL: %[[RESULT:.+]] = arith.select %[[M3]], %[[RED3]], %[[RV2]] : vector<2xi1>, vector<2xf32>
// INNER_REDUCTION: %[[RESULT:.+]] = vector.mask %[[MASK]] { vector.multi_reduction <mul>, %[[INPUT]], %[[ACC]] [0] {{.+}} } : vector<4x2xi1> -> vector<2xf32>
>From 8e55a24b14eecdf46e80444423126db6345b78d2 Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 08:45:44 -0500
Subject: [PATCH 06/19] Delete's old implementation and renames new
implementation
---
.../Vector/TransformOps/VectorTransformOps.td | 3 +-
.../Vector/Transforms/LoweringPatterns.h | 11 ++--
.../Transforms/LowerVectorMultiReduction.cpp | 65 ++-----------------
3 files changed, 14 insertions(+), 65 deletions(-)
diff --git a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
index 685c88c17e556..a7de823de7705 100644
--- a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
+++ b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
@@ -294,7 +294,8 @@ def ApplyMultiReductionUnrollingPatternsOp: Op<Transform_Dialect,
This populates the patterns from
`populateVectorMultiReductionUnrollingPatterns`, i.e.:
* `TwoDimMultiReductionToReduction` (innerreduction)
- * `TwoDimMultiReductionToElementWise` (innerparallel)
+ * `UnrollMultiReductionInnerParallelBaseCase`
+ * `UnrollMultiReductionInnerParallelGeneralCase`
}];
let arguments = (ins DefaultValuedAttr<VectorMultiReductionLoweringAttr,
diff --git a/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h b/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
index 33487a9d8d6e0..ecc6c420e3a82 100644
--- a/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
+++ b/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
@@ -89,10 +89,13 @@ void populateVectorMultiReductionFlatteningPatterns(
/// Populate the pattern set with the following patterns:
///
-/// [TwoDimMultiReductionToElementWise]
-/// Once in 2-D vector.multi_reduction form, with an **outermost** reduction
-/// dimension, unroll the outer dimension to obtain a sequence of 1-D vector
-/// ops. This also has an opportunity for tree-reduction (in the future).
+/// [UnrollMultiReductionInnerParallelBaseCase]
+/// Rank reducing unrolling for inner-parallel case, when there is only one
+/// reduction dimension and it is the outermost one.
+///
+/// [UnrollMultiReductionInnerParallelGeneralCase]
+/// Rank reducing unrolling for inner-parallel general case, when there is
+/// more than one reduction and it is the outermost one.
///
/// [TwoDimMultiReductionToReduction]
/// Once in 2-D vector.multi_reduction form, with an **innermost** reduction
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index 6df55232c605d..a68300c8b9088 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -301,61 +301,6 @@ class ReduceMultiDimReductionRank
const bool useInnerDimsForReduction;
};
-/// Unrolls vector.multi_reduction with outermost reductions
-/// and combines results
-struct TwoDimMultiReductionToElementWise
- : public OpRewritePattern<vector::MultiDimReductionOp> {
- using Base::Base;
-
- LogicalResult matchAndRewrite(vector::MultiDimReductionOp multiReductionOp,
- PatternRewriter &rewriter) const override {
- auto srcRank = multiReductionOp.getSourceVectorType().getRank();
- // Rank-2 ["parallel", "reduce"] or bail.
- if (srcRank != 2)
- return failure();
-
- if (multiReductionOp.isReducedDim(1) || !multiReductionOp.isReducedDim(0))
- return failure();
-
- auto loc = multiReductionOp.getLoc();
- ArrayRef<int64_t> srcShape =
- multiReductionOp.getSourceVectorType().getShape();
-
- Type elementType = getElementTypeOrSelf(multiReductionOp.getDestType());
- if (!elementType.isIntOrIndexOrFloat())
- return failure();
-
- OpBuilder::InsertionGuard guard(rewriter);
- auto maskableOp =
- cast<vector::MaskableOpInterface>(multiReductionOp.getOperation());
- Operation *rootOp;
- Value mask = nullptr;
- if (maskableOp.isMasked()) {
- rewriter.setInsertionPoint(maskableOp.getMaskingOp());
- rootOp = maskableOp.getMaskingOp();
- mask = maskableOp.getMaskingOp().getMask();
- } else {
- rootOp = multiReductionOp;
- }
-
- Value result = multiReductionOp.getAcc();
- for (int64_t i = 0; i < srcShape[0]; i++) {
- auto operand = vector::ExtractOp::create(rewriter, loc,
- multiReductionOp.getSource(), i);
- Value extractMask = nullptr;
- if (mask) {
- extractMask = vector::ExtractOp::create(rewriter, loc, mask, i);
- }
- result =
- makeArithReduction(rewriter, loc, multiReductionOp.getKind(), operand,
- result, /*fastmath=*/nullptr, extractMask);
- }
-
- rewriter.replaceOp(rootOp, result);
- return success();
- }
-};
-
/// Converts 2d vector.multi_reduction with inner most reduction dimension into
/// a sequence of vector.reduction ops.
struct TwoDimMultiReductionToReduction
@@ -536,7 +481,7 @@ struct OneDimMultiReductionToTwoDim
/// %res = vector.multi_reduction %Nminus1, %redNminus2 [ [[REDUCTION_DIMS]] ] :
/// vector<Mx...xf32> to vector<Ix...xf32>
/// ```
-struct UnrollMultiReductionOuterBaseCase
+struct UnrollMultiReductionInnerParallelBaseCase
: public OpRewritePattern<vector::MultiDimReductionOp> {
using Base::Base;
@@ -605,7 +550,7 @@ struct UnrollMultiReductionOuterBaseCase
}
};
-struct UnrollMultiReductionOuterGeneralCase
+struct UnrollMultiReductionInnerParallelGeneralCase
: public OpRewritePattern<vector::MultiDimReductionOp> {
using Base::Base;
@@ -739,9 +684,9 @@ void mlir::vector::populateVectorMultiReductionUnrollingPatterns(
patterns.add<TwoDimMultiReductionToReduction>(patterns.getContext(),
benefit);
} else {
- patterns.add<UnrollMultiReductionOuterBaseCase,
- UnrollMultiReductionOuterGeneralCase>(patterns.getContext(),
- benefit);
+ patterns.add<UnrollMultiReductionInnerParallelBaseCase,
+ UnrollMultiReductionInnerParallelGeneralCase>(
+ patterns.getContext(), benefit);
}
}
>From 59fc900a9cc371b26f04c46d50e63a321460f55f Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 08:47:55 -0500
Subject: [PATCH 07/19] Remove unnecessary file after re-structuring
---
.../Vector/td/unroll-multi-reduction.mlir | 24 -------------------
1 file changed, 24 deletions(-)
delete mode 100644 mlir/test/Dialect/Vector/td/unroll-multi-reduction.mlir
diff --git a/mlir/test/Dialect/Vector/td/unroll-multi-reduction.mlir b/mlir/test/Dialect/Vector/td/unroll-multi-reduction.mlir
deleted file mode 100644
index 96a68723266d3..0000000000000
--- a/mlir/test/Dialect/Vector/td/unroll-multi-reduction.mlir
+++ /dev/null
@@ -1,24 +0,0 @@
-module attributes {transform.with_named_sequence} {
- transform.named_sequence @unroll_multi_reduction(%module_op: !transform.any_op {transform.readonly}) {
-
- %func_op = transform.structured.match ops{["func.func"]} in %module_op
- : (!transform.any_op) -> !transform.any_op
- transform.apply_patterns to %func_op {
- // Test patterns
- transform.apply_patterns.vector.unroll_multi_reduction
- } : !transform.any_op
-
- transform.yield
- }
- transform.named_sequence @unroll_multi_reduction_inner(%module_op: !transform.any_op {transform.readonly}) {
-
- %func_op = transform.structured.match ops{["func.func"]} in %module_op
- : (!transform.any_op) -> !transform.any_op
- transform.apply_patterns to %func_op {
- // Test patterns
- transform.apply_patterns.vector.unroll_multi_reduction lowering_strategy = "innerreduction"
- } : !transform.any_op
-
- transform.yield
- }
-}
>From 567d6d6a9131dd1d47f23168db3bafa4aa214ef2 Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 09:05:08 -0500
Subject: [PATCH 08/19] Add new case for general case
---
.../vector-multi-reduction-unrolling.mlir | 25 ++++++++++++++++---
1 file changed, 21 insertions(+), 4 deletions(-)
diff --git a/mlir/test/Dialect/Vector/vector-multi-reduction-unrolling.mlir b/mlir/test/Dialect/Vector/vector-multi-reduction-unrolling.mlir
index 3eefb8c53f92c..2c7fada2bd57d 100644
--- a/mlir/test/Dialect/Vector/vector-multi-reduction-unrolling.mlir
+++ b/mlir/test/Dialect/Vector/vector-multi-reduction-unrolling.mlir
@@ -95,9 +95,9 @@ func.func @inner_reduction_2d_scalable(%input: vector<2x[4]xf32>, %acc: vector<2
return %0 : vector<2xf32>
}
-// ALL-LABEL: func @inner_parallel_2d
+// ALL-LABEL: func @inner_parallel_base
// ALL-SAME: %[[INPUT:.+]]: vector<4x2xf32>, %[[ACC:.+]]: vector<2xf32>
-func.func @inner_parallel_2d(%arg0: vector<4x2xf32>, %acc: vector<2xf32>) -> vector<2xf32> {
+func.func @inner_parallel_base(%arg0: vector<4x2xf32>, %acc: vector<2xf32>) -> vector<2xf32> {
// INNER_PARALLEL: %[[V0:.+]] = vector.extract %[[INPUT]][0] : vector<2xf32> from vector<4x2xf32>
// INNER_PARALLEL: %[[V1:.+]] = vector.extract %[[INPUT]][1] : vector<2xf32> from vector<4x2xf32>
// INNER_PARALLEL: %[[V2:.+]] = vector.extract %[[INPUT]][2] : vector<2xf32> from vector<4x2xf32>
@@ -112,9 +112,26 @@ func.func @inner_parallel_2d(%arg0: vector<4x2xf32>, %acc: vector<2xf32>) -> vec
return %0 : vector<2xf32>
}
-// ALL-LABEL: func @inner_parallel_2d_masked
+// ALL-LABEL: func @inner_parallel_general
+// ALL-SAME: %[[INPUT:.+]]: vector<4x2x3xf32>, %[[ACC:.+]]: vector<2xf32>
+func.func @inner_parallel_general(%arg0: vector<4x2x3xf32>, %acc: vector<2xf32>) -> vector<2xf32> {
+ // INNER_PARALLEL: %[[V0:.+]] = vector.extract %[[INPUT]][0] : vector<2x3xf32> from vector<4x2x3xf32>
+ // INNER_PARALLEL: %[[V1:.+]] = vector.extract %[[INPUT]][1] : vector<2x3xf32> from vector<4x2x3xf32>
+ // INNER_PARALLEL: %[[V2:.+]] = vector.extract %[[INPUT]][2] : vector<2x3xf32> from vector<4x2x3xf32>
+ // INNER_PARALLEL: %[[V3:.+]] = vector.extract %[[INPUT]][3] : vector<2x3xf32> from vector<4x2x3xf32>
+ // INNER_PARALLEL: %[[RV0:.+]] = vector.multi_reduction <mul>, %[[V0]], %[[ACC]] [1] : vector<2x3xf32>
+ // INNER_PARALLEL: %[[RV1:.+]] = vector.multi_reduction <mul>, %[[V1]], %[[RV0]] [1] : vector<2x3xf32>
+ // INNER_PARALLEL: %[[RV2:.+]] = vector.multi_reduction <mul>, %[[V2]], %[[RV1]] [1] : vector<2x3xf32>
+ // INNER_PARALLEL: %[[RESULT:.+]] = vector.multi_reduction <mul>, %[[V3]], %[[RV2]] [1] : vector<2x3xf32>
+ // INNER_REDUCTION: %[[RESULT:.+]] = vector.multi_reduction <mul>, %[[INPUT]], %[[ACC]] [0, 2]
+ %0 = vector.multi_reduction <mul>, %arg0, %acc [0, 2] : vector<4x2x3xf32> to vector<2xf32>
+ // ALL: return %[[RESULT]]
+ return %0 : vector<2xf32>
+}
+
+// ALL-LABEL: func @inner_parallel_base_masked
// ALL-SAME: %[[INPUT:.+]]: vector<4x2xf32>, %[[ACC:.+]]: vector<2xf32>, %[[MASK:.+]]: vector<4x2xi1>
-func.func @inner_parallel_2d_masked(%arg0: vector<4x2xf32>, %acc: vector<2xf32>, %mask: vector<4x2xi1>) -> vector<2xf32> {
+func.func @inner_parallel_base_masked(%arg0: vector<4x2xf32>, %acc: vector<2xf32>, %mask: vector<4x2xi1>) -> vector<2xf32> {
// INNER_PARALLEL: %[[V0:.+]] = vector.extract %[[INPUT]][0] : vector<2xf32> from vector<4x2xf32>
// INNER_PARALLEL: %[[V1:.+]] = vector.extract %[[INPUT]][1] : vector<2xf32> from vector<4x2xf32>
// INNER_PARALLEL: %[[V2:.+]] = vector.extract %[[INPUT]][2] : vector<2xf32> from vector<4x2xf32>
>From d2bd3e69b82571a16052aadee5f1a185e9f51583 Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 09:19:11 -0500
Subject: [PATCH 09/19] Use MaskableOpRewritePattern
---
.../Transforms/LowerVectorMultiReduction.cpp | 61 ++++++-------------
1 file changed, 19 insertions(+), 42 deletions(-)
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index a68300c8b9088..fcbbd06a4716a 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -482,11 +482,13 @@ struct OneDimMultiReductionToTwoDim
/// vector<Mx...xf32> to vector<Ix...xf32>
/// ```
struct UnrollMultiReductionInnerParallelBaseCase
- : public OpRewritePattern<vector::MultiDimReductionOp> {
- using Base::Base;
+ : public vector::MaskableOpRewritePattern<vector::MultiDimReductionOp> {
+ using MaskableOpRewritePattern::MaskableOpRewritePattern;
- LogicalResult matchAndRewrite(vector::MultiDimReductionOp multiReductionOp,
- PatternRewriter &rewriter) const override {
+ FailureOr<Value>
+ matchAndRewriteMaskableOp(vector::MultiDimReductionOp multiReductionOp,
+ vector::MaskingOpInterface maskingOp,
+ PatternRewriter &rewriter) const override {
auto srcRank = multiReductionOp.getSourceVectorType().getRank();
if (srcRank < 2)
return rewriter.notifyMatchFailure(multiReductionOp,
@@ -514,19 +516,7 @@ struct UnrollMultiReductionInnerParallelBaseCase
multiReductionOp.getSourceVectorType().getShape();
int64_t numElementwiseOps = srcShape.front();
- OpBuilder::InsertionGuard guard(rewriter);
- auto maskableOp =
- cast<vector::MaskableOpInterface>(multiReductionOp.getOperation());
- bool isMasked = maskableOp.isMasked();
- Operation *rootOp;
- Value mask = nullptr;
- if (isMasked) {
- rewriter.setInsertionPoint(maskableOp.getMaskingOp());
- rootOp = maskableOp.getMaskingOp();
- mask = maskableOp.getMaskingOp().getMask();
- } else {
- rootOp = multiReductionOp;
- }
+ Value mask = maskingOp ? maskingOp.getMask() : nullptr;
SmallVector<Value> vectors;
for (int64_t i = 0; i < numElementwiseOps; ++i)
@@ -534,7 +524,7 @@ struct UnrollMultiReductionInnerParallelBaseCase
SmallVector<Value> masks;
for (int64_t i = 0; i < numElementwiseOps; ++i)
- if (isMasked)
+ if (mask)
masks.push_back(vector::ExtractOp::create(rewriter, loc, mask, i));
else
masks.push_back(nullptr);
@@ -545,17 +535,18 @@ struct UnrollMultiReductionInnerParallelBaseCase
innerVector, result, /*fastmath=*/nullptr,
innerMask);
- rewriter.replaceOp(rootOp, result);
- return success();
+ return result;
}
};
struct UnrollMultiReductionInnerParallelGeneralCase
- : public OpRewritePattern<vector::MultiDimReductionOp> {
- using Base::Base;
+ : public vector::MaskableOpRewritePattern<vector::MultiDimReductionOp> {
+ using MaskableOpRewritePattern::MaskableOpRewritePattern;
- LogicalResult matchAndRewrite(vector::MultiDimReductionOp multiReductionOp,
- PatternRewriter &rewriter) const override {
+ FailureOr<Value>
+ matchAndRewriteMaskableOp(vector::MultiDimReductionOp multiReductionOp,
+ vector::MaskingOpInterface maskingOp,
+ PatternRewriter &rewriter) const override {
if (!multiReductionOp.isReducedDim(0))
return rewriter.notifyMatchFailure(
multiReductionOp,
@@ -578,19 +569,7 @@ struct UnrollMultiReductionInnerParallelGeneralCase
multiReductionOp.getSourceVectorType().getShape();
int64_t numElementwiseOps = srcShape.front();
- OpBuilder::InsertionGuard guard(rewriter);
- auto maskableOp =
- cast<vector::MaskableOpInterface>(multiReductionOp.getOperation());
- bool isMasked = maskableOp.isMasked();
- Operation *rootOp;
- Value mask = nullptr;
- if (isMasked) {
- rewriter.setInsertionPoint(maskableOp.getMaskingOp());
- rootOp = maskableOp.getMaskingOp();
- mask = maskableOp.getMaskingOp().getMask();
- } else {
- rootOp = multiReductionOp;
- }
+ Value mask = maskingOp ? maskingOp.getMask() : nullptr;
SmallVector<Value> vectors;
for (int64_t i = 0; i < numElementwiseOps; ++i)
@@ -598,7 +577,7 @@ struct UnrollMultiReductionInnerParallelGeneralCase
SmallVector<Value> masks;
for (int64_t i = 0; i < numElementwiseOps; ++i)
- if (isMasked)
+ if (mask)
masks.push_back(vector::ExtractOp::create(rewriter, loc, mask, i));
else
masks.push_back(nullptr);
@@ -607,12 +586,11 @@ struct UnrollMultiReductionInnerParallelGeneralCase
ArrayRef<bool>(multiReductionOp.getReductionMask()).drop_front();
Value result = multiReductionOp.getAcc();
for (auto [innerVector, innerMask] : llvm::zip(vectors, masks)) {
-
auto reductionOp = vector::MultiDimReductionOp::create(
rewriter, loc, innerVector, result, reductionMask,
multiReductionOp.getKind());
- if (isMasked) {
+ if (innerMask) {
auto maskOp = vector::maskOperation(rewriter, reductionOp, innerMask);
result = maskOp->getResult(0);
} else {
@@ -620,8 +598,7 @@ struct UnrollMultiReductionInnerParallelGeneralCase
}
}
- rewriter.replaceOp(rootOp, result);
- return success();
+ return result;
}
};
>From 42b5dec83812e89476979eedc8434c4b63fe9acf Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 09:27:21 -0500
Subject: [PATCH 10/19] Remove notifyMatchFailure in complementary patterns
---
.../Vector/Transforms/LowerVectorMultiReduction.cpp | 10 +++-------
1 file changed, 3 insertions(+), 7 deletions(-)
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index fcbbd06a4716a..d33370fe2d2d0 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -548,19 +548,15 @@ struct UnrollMultiReductionInnerParallelGeneralCase
vector::MaskingOpInterface maskingOp,
PatternRewriter &rewriter) const override {
if (!multiReductionOp.isReducedDim(0))
- return rewriter.notifyMatchFailure(
- multiReductionOp,
- "expected outermost dimension to be reduced dimension.");
+ return failure();
Type elementType = getElementTypeOrSelf(multiReductionOp.getDestType());
if (!elementType.isIntOrIndexOrFloat())
- return rewriter.notifyMatchFailure(
- multiReductionOp, "expected integer or float element type.");
+ return failure();
ArrayRef<int64_t> reductionDims = multiReductionOp.getReductionDims();
if (reductionDims.size() <= 1)
- return rewriter.notifyMatchFailure(
- multiReductionOp, "expected more than one reduction dimension.");
+ return failure();
Location loc = multiReductionOp.getLoc();
Value source = multiReductionOp.getSource();
>From deb59951d06eb48822079a41b8682f3e9e1d88ee Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 09:32:04 -0500
Subject: [PATCH 11/19] Remove unnecessary check
---
.../Vector/Transforms/LowerVectorMultiReduction.cpp | 9 ---------
1 file changed, 9 deletions(-)
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index d33370fe2d2d0..4821d40172c70 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -499,11 +499,6 @@ struct UnrollMultiReductionInnerParallelBaseCase
multiReductionOp,
"expected outermost dimension to be reduced dimension.");
- Type elementType = getElementTypeOrSelf(multiReductionOp.getDestType());
- if (!elementType.isIntOrIndexOrFloat())
- return rewriter.notifyMatchFailure(
- multiReductionOp, "expected integer or float element type.");
-
ArrayRef<int64_t> reductionDims = multiReductionOp.getReductionDims();
if (reductionDims.size() > 1)
return rewriter.notifyMatchFailure(
@@ -550,10 +545,6 @@ struct UnrollMultiReductionInnerParallelGeneralCase
if (!multiReductionOp.isReducedDim(0))
return failure();
- Type elementType = getElementTypeOrSelf(multiReductionOp.getDestType());
- if (!elementType.isIntOrIndexOrFloat())
- return failure();
-
ArrayRef<int64_t> reductionDims = multiReductionOp.getReductionDims();
if (reductionDims.size() <= 1)
return failure();
>From 665efcb14bdc4b2116c336662db5b115d0c30521 Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 09:49:44 -0500
Subject: [PATCH 12/19] Improve loops
---
.../Transforms/LowerVectorMultiReduction.cpp | 20 ++++++++-----------
1 file changed, 8 insertions(+), 12 deletions(-)
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index 4821d40172c70..7e6f90a128114 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -513,16 +513,14 @@ struct UnrollMultiReductionInnerParallelBaseCase
Value mask = maskingOp ? maskingOp.getMask() : nullptr;
- SmallVector<Value> vectors;
+ SmallVector<Value> vectors(numElementwiseOps);
for (int64_t i = 0; i < numElementwiseOps; ++i)
- vectors.push_back(vector::ExtractOp::create(rewriter, loc, source, i));
+ vectors[i] = vector::ExtractOp::create(rewriter, loc, source, i);
- SmallVector<Value> masks;
+ SmallVector<Value> masks(numElementwiseOps);
for (int64_t i = 0; i < numElementwiseOps; ++i)
if (mask)
- masks.push_back(vector::ExtractOp::create(rewriter, loc, mask, i));
- else
- masks.push_back(nullptr);
+ masks[i] = vector::ExtractOp::create(rewriter, loc, mask, i);
Value result = multiReductionOp.getAcc();
for (auto [innerVector, innerMask] : llvm::zip(vectors, masks))
@@ -558,16 +556,14 @@ struct UnrollMultiReductionInnerParallelGeneralCase
Value mask = maskingOp ? maskingOp.getMask() : nullptr;
- SmallVector<Value> vectors;
+ SmallVector<Value> vectors(numElementwiseOps);
for (int64_t i = 0; i < numElementwiseOps; ++i)
- vectors.push_back(vector::ExtractOp::create(rewriter, loc, source, i));
+ vectors[i] = vector::ExtractOp::create(rewriter, loc, source, i);
- SmallVector<Value> masks;
+ SmallVector<Value> masks(numElementwiseOps);
for (int64_t i = 0; i < numElementwiseOps; ++i)
if (mask)
- masks.push_back(vector::ExtractOp::create(rewriter, loc, mask, i));
- else
- masks.push_back(nullptr);
+ masks[i] = vector::ExtractOp::create(rewriter, loc, mask, i);
ArrayRef<bool> reductionMask =
ArrayRef<bool>(multiReductionOp.getReductionMask()).drop_front();
>From 12e15a12d36750d228931c835999d6d9a8a4ee49 Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 10:11:34 -0500
Subject: [PATCH 13/19] sink loop and fix dangling array ref
---
.../Vector/Transforms/LowerVectorMultiReduction.cpp | 11 ++++++-----
1 file changed, 6 insertions(+), 5 deletions(-)
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index 7e6f90a128114..6c9b89bb62ae7 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -518,8 +518,8 @@ struct UnrollMultiReductionInnerParallelBaseCase
vectors[i] = vector::ExtractOp::create(rewriter, loc, source, i);
SmallVector<Value> masks(numElementwiseOps);
- for (int64_t i = 0; i < numElementwiseOps; ++i)
- if (mask)
+ if (mask)
+ for (int64_t i = 0; i < numElementwiseOps; ++i)
masks[i] = vector::ExtractOp::create(rewriter, loc, mask, i);
Value result = multiReductionOp.getAcc();
@@ -561,12 +561,13 @@ struct UnrollMultiReductionInnerParallelGeneralCase
vectors[i] = vector::ExtractOp::create(rewriter, loc, source, i);
SmallVector<Value> masks(numElementwiseOps);
- for (int64_t i = 0; i < numElementwiseOps; ++i)
- if (mask)
+ if (mask)
+ for (int64_t i = 0; i < numElementwiseOps; ++i)
masks[i] = vector::ExtractOp::create(rewriter, loc, mask, i);
+ SmallVector<bool> fullReductionMask = multiReductionOp.getReductionMask();
ArrayRef<bool> reductionMask =
- ArrayRef<bool>(multiReductionOp.getReductionMask()).drop_front();
+ ArrayRef<bool>(fullReductionMask).drop_front();
Value result = multiReductionOp.getAcc();
for (auto [innerVector, innerMask] : llvm::zip(vectors, masks)) {
auto reductionOp = vector::MultiDimReductionOp::create(
>From c4bfc072d8ff02c175b72d9b809c46fe32d67c4e Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 10:14:23 -0500
Subject: [PATCH 14/19] Fix documentation
---
.../mlir/Dialect/Vector/TransformOps/VectorTransformOps.td | 2 +-
mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h | 2 +-
.../lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp | 1 +
3 files changed, 3 insertions(+), 2 deletions(-)
diff --git a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
index a7de823de7705..10950e701faa5 100644
--- a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
+++ b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
@@ -287,7 +287,7 @@ def ApplyMultiReductionUnrollingPatternsOp: Op<Transform_Dialect,
"apply_patterns.vector.multi_reduction_unrolling",
[DeclareOpInterfaceMethods<PatternDescriptorOpInterface>]> {
let description = [{
- Indicates that 2-D vector multi_reduction operations should be unrolled
+ Indicates that vector multi_reduction operations should be unrolled
into either a sequence of vector.reduction ops (innerreduction) or
element-wise arith ops (innerparallel).
diff --git a/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h b/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
index ecc6c420e3a82..7a1823c26a46e 100644
--- a/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
+++ b/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
@@ -95,7 +95,7 @@ void populateVectorMultiReductionFlatteningPatterns(
///
/// [UnrollMultiReductionInnerParallelGeneralCase]
/// Rank reducing unrolling for inner-parallel general case, when there is
-/// more than one reduction and it is the outermost one.
+/// more than one reduction dimension and it is the outermost one.
///
/// [TwoDimMultiReductionToReduction]
/// Once in 2-D vector.multi_reduction form, with an **innermost** reduction
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index 6c9b89bb62ae7..b7a2c21f356e6 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -468,6 +468,7 @@ struct OneDimMultiReductionToTwoDim
/// ```mlir
/// %res = vector.multi_reduction <add> %src, %acc [0, [[REDUCTION_DIMS]] ] :
/// vector<NxMx...xf32> to vector<Ix...xf32>
+/// ```
///
/// ```mlir
/// %0 = vector.extract %src[0] : vector<Mx...xf32> from vector<NxMx...xf32>
>From 6030d6235d44515a14c231f23c87bd5fd105afcd Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 10:19:37 -0500
Subject: [PATCH 15/19] Rename variable
---
.../Vector/Transforms/LowerVectorMultiReduction.cpp | 10 +++++-----
1 file changed, 5 insertions(+), 5 deletions(-)
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index b7a2c21f356e6..a6b3e8bba06d7 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -553,17 +553,17 @@ struct UnrollMultiReductionInnerParallelGeneralCase
ArrayRef<int64_t> srcShape =
multiReductionOp.getSourceVectorType().getShape();
- int64_t numElementwiseOps = srcShape.front();
+ int64_t outerDimSize = srcShape.front();
Value mask = maskingOp ? maskingOp.getMask() : nullptr;
- SmallVector<Value> vectors(numElementwiseOps);
- for (int64_t i = 0; i < numElementwiseOps; ++i)
+ SmallVector<Value> vectors(outerDimSize);
+ for (int64_t i = 0; i < outerDimSize; ++i)
vectors[i] = vector::ExtractOp::create(rewriter, loc, source, i);
- SmallVector<Value> masks(numElementwiseOps);
+ SmallVector<Value> masks(outerDimSize);
if (mask)
- for (int64_t i = 0; i < numElementwiseOps; ++i)
+ for (int64_t i = 0; i < outerDimSize; ++i)
masks[i] = vector::ExtractOp::create(rewriter, loc, mask, i);
SmallVector<bool> fullReductionMask = multiReductionOp.getReductionMask();
>From 939a6865fd24f43a48d325c9a2ff45f9cda9654d Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 10:22:05 -0500
Subject: [PATCH 16/19] Splits documentation
---
.../Transforms/LowerVectorMultiReduction.cpp | 53 ++++++++-----------
1 file changed, 22 insertions(+), 31 deletions(-)
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index a6b3e8bba06d7..38ae3caa89d07 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -432,17 +432,8 @@ struct OneDimMultiReductionToTwoDim
};
/// Unrolls outermost dimension for vector.multi_reduction.
-/// This patterns matches operations which reduce the outermost dimension,
-/// it does not transform operations for which the outermost dimension is not
-/// a reduction dimension.
-///
-/// There are two cases to consider:
-/// 1. The base case is when the outermost dimension is the only reduction
+/// Matches when the outermost dimension is the only reduction
/// dimension.
-/// 2. The general case is when the outermost dimension is not the only
-/// reduction dimension.
-///
-/// The base case transformation:
///
/// ```mlir
/// %res = vector.multi_reduction <add> %src, %acc [0] : vector<NxMx...xf32> to
@@ -461,27 +452,6 @@ struct OneDimMultiReductionToTwoDim
/// ...
/// %res = arith.addf %Nminus1, %resNminus2 : vector<Mx...xf32>
/// ```
-///
-/// For the general case, we still extract N vectors, but produce N
-/// vector.multi_reduction instead of elementwise operations.
-///
-/// ```mlir
-/// %res = vector.multi_reduction <add> %src, %acc [0, [[REDUCTION_DIMS]] ] :
-/// vector<NxMx...xf32> to vector<Ix...xf32>
-/// ```
-///
-/// ```mlir
-/// %0 = vector.extract %src[0] : vector<Mx...xf32> from vector<NxMx...xf32>
-/// ...
-/// %Nminus1 = vector.extract %src[ [[N-1]] ] : vector<Mx...x.f32> from
-/// vector<NxMx...xf32>
-///
-/// %red0 = vector.multi_reduction %0, %acc [ [[REDUCTION_DIMS]] ] :
-/// vector<Mx...xf32> to vector<Ix...xf32>
-/// ...
-/// %res = vector.multi_reduction %Nminus1, %redNminus2 [ [[REDUCTION_DIMS]] ] :
-/// vector<Mx...xf32> to vector<Ix...xf32>
-/// ```
struct UnrollMultiReductionInnerParallelBaseCase
: public vector::MaskableOpRewritePattern<vector::MultiDimReductionOp> {
using MaskableOpRewritePattern::MaskableOpRewritePattern;
@@ -533,6 +503,27 @@ struct UnrollMultiReductionInnerParallelBaseCase
}
};
+/// Unrolls outermost dimension for vector.multi_reduction.
+/// Matches when the outermost dimension is not the only
+/// reduction dimension.
+///
+/// ```mlir
+/// %res = vector.multi_reduction <add> %src, %acc [0, [[REDUCTION_DIMS]] ] :
+/// vector<NxMx...xf32> to vector<Ix...xf32>
+/// ```
+///
+/// ```mlir
+/// %0 = vector.extract %src[0] : vector<Mx...xf32> from vector<NxMx...xf32>
+/// ...
+/// %Nminus1 = vector.extract %src[ [[N-1]] ] : vector<Mx...x.f32> from
+/// vector<NxMx...xf32>
+///
+/// %red0 = vector.multi_reduction %0, %acc [ [[REDUCTION_DIMS]] ] :
+/// vector<Mx...xf32> to vector<Ix...xf32>
+/// ...
+/// %res = vector.multi_reduction %Nminus1, %redNminus2 [ [[REDUCTION_DIMS]] ] :
+/// vector<Mx...xf32> to vector<Ix...xf32>
+/// ```
struct UnrollMultiReductionInnerParallelGeneralCase
: public vector::MaskableOpRewritePattern<vector::MultiDimReductionOp> {
using MaskableOpRewritePattern::MaskableOpRewritePattern;
>From a87b3ef0645330460f6b6fa5a62cacb75d8bcda0 Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 10:29:14 -0500
Subject: [PATCH 17/19] Make type explicit
---
.../Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp | 3 ++-
1 file changed, 2 insertions(+), 1 deletion(-)
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index 38ae3caa89d07..3756d6d471ed7 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -567,7 +567,8 @@ struct UnrollMultiReductionInnerParallelGeneralCase
multiReductionOp.getKind());
if (innerMask) {
- auto maskOp = vector::maskOperation(rewriter, reductionOp, innerMask);
+ Operation *maskOp =
+ vector::maskOperation(rewriter, reductionOp, innerMask);
result = maskOp->getResult(0);
} else {
result = reductionOp.getResult();
>From e7381a4b8a9ce9a6a91133bc94d01b935e5a2b77 Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 10:46:00 -0500
Subject: [PATCH 18/19] Use zip_equal
---
.../Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp | 4 ++--
1 file changed, 2 insertions(+), 2 deletions(-)
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index 3756d6d471ed7..8b817c74f72a5 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -494,7 +494,7 @@ struct UnrollMultiReductionInnerParallelBaseCase
masks[i] = vector::ExtractOp::create(rewriter, loc, mask, i);
Value result = multiReductionOp.getAcc();
- for (auto [innerVector, innerMask] : llvm::zip(vectors, masks))
+ for (auto [innerVector, innerMask] : llvm::zip_equal(vectors, masks))
result = makeArithReduction(rewriter, loc, multiReductionOp.getKind(),
innerVector, result, /*fastmath=*/nullptr,
innerMask);
@@ -561,7 +561,7 @@ struct UnrollMultiReductionInnerParallelGeneralCase
ArrayRef<bool> reductionMask =
ArrayRef<bool>(fullReductionMask).drop_front();
Value result = multiReductionOp.getAcc();
- for (auto [innerVector, innerMask] : llvm::zip(vectors, masks)) {
+ for (auto [innerVector, innerMask] : llvm::zip_equal(vectors, masks)) {
auto reductionOp = vector::MultiDimReductionOp::create(
rewriter, loc, innerVector, result, reductionMask,
multiReductionOp.getKind());
>From 0b7d094c51805a3a9ab3b6e9403b03036d0ec50a Mon Sep 17 00:00:00 2001
From: Erick Ochoa <erick.ochoalopez at amd.com>
Date: Thu, 19 Feb 2026 10:49:21 -0500
Subject: [PATCH 19/19] clarify comment
---
.../Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp | 3 +++
1 file changed, 3 insertions(+)
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
index 8b817c74f72a5..28257c2ec39a5 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorMultiReduction.cpp
@@ -435,6 +435,9 @@ struct OneDimMultiReductionToTwoDim
/// Matches when the outermost dimension is the only reduction
/// dimension.
///
+/// In this case [0] refers to rank at position N, so it is the outermost
+/// dimension.
+///
/// ```mlir
/// %res = vector.multi_reduction <add> %src, %acc [0] : vector<NxMx...xf32> to
/// vector<Mx...xf32>
More information about the Mlir-commits
mailing list