[Mlir-commits] [mlir] 8ef411a - [MLIR][XeGPU] Avoid chained-reductions in multi_reduction unrolling (#198307)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Wed May 20 07:10:22 PDT 2026
Author: Dmitry Chigarev
Date: 2026-05-20T16:10:16+02:00
New Revision: 8ef411a3bc9d1820144211883f795446fde71386
URL: https://github.com/llvm/llvm-project/commit/8ef411a3bc9d1820144211883f795446fde71386
DIFF: https://github.com/llvm/llvm-project/commit/8ef411a3bc9d1820144211883f795446fde71386.diff
LOG: [MLIR][XeGPU] Avoid chained-reductions in multi_reduction unrolling (#198307)
The PR adds a new unrolling pattern for `vector.multi_reduction` to the
`xegpu-blocking` pass. In comparison with [the upstream reduction
unrolling](https://github.com/llvm/llvm-project/blob/2da84a8307e4ef729458d990b221650a5da22639/mlir/lib/Dialect/Vector/Transforms/VectorUnroll.cpp#L372),
the new pattern performs partial row-wise reductions via elementwise
ops, instead of generating a chain of several multi-reduction ops:
```mlir
// reduction to unroll:
// tile-shape: [8x16]
vector.multi_reduction <add> %vec, %cst : vector<8x48xf32> to vector<8xf32>
// upstream unrolling:
%3 = vector.multi_reduction <add>, %tile_0, %cst [1] : vector<8x16xf32> to vector<8xf32>
%4 = vector.multi_reduction <add>, %tile_1, %3 [1] : vector<8x16xf32> to vector<8xf32>
%5 = vector.multi_reduction <add>, %tile_2, %4 [1] : vector<8x16xf32> to vector<8xf32>
// new xegpu-unrolling
%3 = arith.addf %tile_0, %tile_1 : vector<8x16xf32>
%4 = arith.addf %3, %tile_2 : vector<8x16xf32>
%5 = vector.multi_reduction <add>, %4, %cst [1] : vector<8x16xf32> to vector<8xf32>
```
---------
Signed-off-by: dchigarev <dmitry.chigarev at intel.com>
Added:
Modified:
mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp
mlir/test/Dialect/XeGPU/xegpu-blocking.mlir
mlir/test/Dialect/XeGPU/xegpu-unroll-patterns.mlir
mlir/test/lib/Dialect/XeGPU/TestXeGPUTransforms.cpp
Removed:
################################################################################
diff --git a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp
index a6fd4ce09ebf4..aab36b79845e4 100644
--- a/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp
+++ b/mlir/lib/Dialect/XeGPU/Transforms/XeGPUUnroll.cpp
@@ -12,7 +12,9 @@
//
//===----------------------------------------------------------------------===//
+#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/Utils/IndexingUtils.h"
+#include "mlir/Dialect/Vector/IR/VectorOps.h"
#include "mlir/Dialect/XeGPU/IR/XeGPU.h"
#include "mlir/Dialect/XeGPU/Transforms/Transforms.h"
#include "mlir/Dialect/XeGPU/Transforms/XeGPULayoutImpl.h"
@@ -804,13 +806,152 @@ struct UnrollConvertLayoutOp : public UnrollPattern<xegpu::ConvertLayoutOp> {
}
};
+/// Unrolls vector.multi_reduction by sequentially reducing tiles with
+/// elementwise arith operations first, then a single multi_reduction
+/// per non-reduced tile position. This avoids generating long chains of
+/// multi_reduction ops (as the upstream pattern does) and is more efficient.
+///
+/// Example:
+/// vector.multi_reduction <32x64xf16> to <32xf16> (tile_shape=32, 32)
+/// -- Upstream pattern generates:
+/// %tmp1 = vector.multi_reduction %tile0, %zero_acc <32x32xf16> to <32xf16>
+/// %res = vector.multi_reduction %tmp1, %tile1 <32x32xf16> to <32xf16>
+/// -- This pattern generates:
+/// %tmp1 = arith.reduction %tile0, %tile1 <32x32xf16> -> <32x32xf16> //
+/// elementwise %res = vector.multi_reduction %tmp1, %zero_acc <32x32xf16> to
+/// <32xf16>
+struct UnrollMultiReductionOp
+ : public UnrollPattern<vector::MultiDimReductionOp> {
+ UnrollMultiReductionOp(MLIRContext *context,
+ const xegpu::UnrollOptions &options,
+ PatternBenefit benefit = 2)
+ : UnrollPattern<vector::MultiDimReductionOp>(context, options, benefit) {}
+
+ LogicalResult matchAndRewrite(vector::MultiDimReductionOp reductionOp,
+ PatternRewriter &rewriter) const override {
+ VectorType srcTy = reductionOp.getSourceVectorType();
+ ArrayRef<int64_t> srcShape = srcTy.getShape();
+ int64_t srcRank = srcTy.getRank();
+
+ Location loc = reductionOp.getLoc();
+ Value source = reductionOp.getSource();
+ Value acc = reductionOp.getAcc();
+ vector::CombiningKind kind = reductionOp.getKind();
+
+ // Result must be a vector (not scalar).
+ auto resultType = dyn_cast<VectorType>(reductionOp.getDestType());
+ if (!resultType)
+ return failure();
+
+ std::optional<SmallVector<int64_t>> targetShapeOpt =
+ getTargetShape(reductionOp);
+ if (!targetShapeOpt ||
+ static_cast<int64_t>(targetShapeOpt->size()) != srcRank)
+ return failure();
+
+ SmallVector<int64_t> targetShape = *targetShapeOpt;
+
+ // Check divisibility for all dimensions.
+ for (int64_t i = 0; i < srcRank; ++i) {
+ if (srcShape[i] % targetShape[i] != 0)
+ return failure();
+ }
+
+ SmallVector<bool> reductionMask = reductionOp.getReductionMask();
+ // Identify reduced and kept dimensions from the reduction mask.
+ SmallVector<int64_t> reducedDims, keptDims;
+ for (int64_t i = 0; i < srcRank; ++i) {
+ if (reductionMask[i])
+ reducedDims.push_back(i);
+ else
+ keptDims.push_back(i);
+ }
+
+ // Compute the number of tiles along each reduced dimension and their
+ // product
+ SmallVector<int64_t> numReducedTilesPerDim;
+ for (int64_t d : reducedDims)
+ numReducedTilesPerDim.push_back(srcShape[d] / targetShape[d]);
+
+ // Build kept shapes for iterating over non-reduced dimensions.
+ SmallVector<int64_t> keptShape, keptTileShape;
+ for (int64_t d : keptDims) {
+ keptShape.push_back(srcShape[d]);
+ keptTileShape.push_back(targetShape[d]);
+ }
+
+ // Initialize the result vector for assembly.
+ Value result = arith::ConstantOp::create(rewriter, loc, resultType,
+ rewriter.getZeroAttr(resultType));
+
+ // Iterate over all tile positions in the kept dimensions.
+ // Ex: [off0, off1, _ _ off4]
+ // blanks are offsets for the reduced dims, they will be
+ // generated in the inner loop below
+ for (SmallVector<int64_t> keptOffsets :
+ StaticTileOffsetRange(keptShape, keptTileShape)) {
+
+ // Reconstruct full-rank base offsets with 0 for reduced dims.
+ // Ex: [off0, off1, 0, 0, off4]
+ SmallVector<int64_t> baseOffsets(srcRank, 0);
+ for (auto [idx, dim] : llvm::enumerate(keptDims))
+ baseOffsets[dim] = keptOffsets[idx];
+
+ // Generate the full tile indices for the reduced dimensions.
+ // Ex: if reduceDimShapes = [32, 64] and
+ // reducedDimTargetShapes = [16, 16], then reducedTileCoords:
+ // [(0, 0), (0, 1), (0, 2), (0, 3),
+ // (1, 0), (1, 1), (1, 2), (1, 3)]
+ auto reducedTileCoords = StaticTileOffsetRange(
+ numReducedTilesPerDim, SmallVector<int64_t>(reducedDims.size(), 1));
+
+ // Step 1: Fill "blanks" in the offsets for the reduced dimensions
+ // using 'reducedTileCoords' and extract according tiles.
+ // Ex: tiles = [source[off0, off1, off2_red, off3_red, off4], ...]
+ SmallVector<Value> tiles;
+ for (SmallVector<int64_t> reducedTileIdx : reducedTileCoords) {
+ SmallVector<int64_t> offsets(baseOffsets);
+ for (auto [idx, dim] : llvm::enumerate(reducedDims))
+ offsets[dim] = reducedTileIdx[idx] * targetShape[dim];
+ SmallVector<int64_t> strides(srcRank, 1);
+ Value tile = vector::ExtractStridedSliceOp::create(
+ rewriter, loc, source, offsets, targetShape, strides);
+ tiles.push_back(tile);
+ }
+
+ // Step 2: Sequentially reduce tiles using elementwise arith operations.
+ Value reduced = tiles[0];
+ for (size_t i = 1; i < tiles.size(); ++i)
+ reduced =
+ vector::makeArithReduction(rewriter, loc, kind, reduced, tiles[i]);
+
+ // Step 3: Perform a single multi_reduction with the accumulator slice.
+ SmallVector<int64_t> accStrides(keptTileShape.size(), 1);
+ Value accSlice = vector::ExtractStridedSliceOp::create(
+ rewriter, loc, acc, keptOffsets, keptTileShape, accStrides);
+
+ auto newReduction = vector::MultiDimReductionOp::create(
+ rewriter, loc, reduced, accSlice, reductionMask, kind);
+
+ // Step 4: Insert the reduced result into the output vector.
+ SmallVector<int64_t> dstStrides(keptTileShape.size(), 1);
+ result = vector::InsertStridedSliceOp::create(
+ rewriter, loc, newReduction, result, keptOffsets, dstStrides);
+ }
+
+ rewriter.replaceOp(reductionOp, result);
+ return success();
+ }
+};
+
} // namespace
void mlir::xegpu::populateXeGPUUnrollPatterns(
RewritePatternSet &patterns, const xegpu::UnrollOptions &options) {
- patterns.add<UnrollCreateNdOp, UnrollPrefetchNdOp, UnrollLoadNdOp,
- UnrollStoreNdOp, UnrollDpasOp, UnrollDpasMxOp,
- UnrollLoadMatrixOp, UnrollStoreMatrixOp, UnrollLoadGatherOp,
- UnrollStoreScatterOp, UnrollConvertLayoutOp>(
- patterns.getContext(), options);
+ patterns
+ .add<UnrollCreateNdOp, UnrollPrefetchNdOp, UnrollLoadNdOp,
+ UnrollStoreNdOp, UnrollDpasOp, UnrollDpasMxOp, UnrollLoadMatrixOp,
+ UnrollStoreMatrixOp, UnrollLoadGatherOp, UnrollStoreScatterOp,
+ UnrollConvertLayoutOp, UnrollMultiReductionOp>(patterns.getContext(),
+ options);
}
diff --git a/mlir/test/Dialect/XeGPU/xegpu-blocking.mlir b/mlir/test/Dialect/XeGPU/xegpu-blocking.mlir
index 56a00508eda15..c0ea112edc818 100644
--- a/mlir/test/Dialect/XeGPU/xegpu-blocking.mlir
+++ b/mlir/test/Dialect/XeGPU/xegpu-blocking.mlir
@@ -252,8 +252,10 @@ gpu.module @test_kernel {
%0 = xegpu.create_nd_tdesc %a : memref<512x32xf32> -> !xegpu.tensor_desc<32x128xf32, #l>
%1 = xegpu.load_nd %0[0, 0] {layout = #l}: !xegpu.tensor_desc<32x128xf32, #l> -> vector<32x128xf32>
- // CHECK: vector.multi_reduction <add>, {{.*}}, [[INIT:%[0-9A-Za-z]+]] [1] : vector<16x16xf32> to vector<16xf32>
- // CHECK-COUNT-1: vector.multi_reduction <add>, {{.*}}, [[INIT]] [1] : vector<16x16xf32> to vector<16xf32>
+ // CHECK-COUNT-7: arith.addf {{.*}} : vector<16x16xf32>
+ // CHECK: vector.multi_reduction <add>, {{.*}} [1] : vector<16x16xf32> to vector<16xf32>
+ // CHECK-COUNT-7: arith.addf {{.*}} : vector<16x16xf32>
+ // CHECK: vector.multi_reduction <add>, {{.*}} [1] : vector<16x16xf32> to vector<16xf32>
%2 = vector.multi_reduction <add>, %1, %acc [1]: vector<32x128xf32> to vector<32xf32>
%3 = xegpu.create_nd_tdesc %b : memref<512xf32> -> !xegpu.tensor_desc<32xf32, #r>
diff --git a/mlir/test/Dialect/XeGPU/xegpu-unroll-patterns.mlir b/mlir/test/Dialect/XeGPU/xegpu-unroll-patterns.mlir
index b60700c81f518..49fdf1cbee174 100644
--- a/mlir/test/Dialect/XeGPU/xegpu-unroll-patterns.mlir
+++ b/mlir/test/Dialect/XeGPU/xegpu-unroll-patterns.mlir
@@ -226,5 +226,139 @@ gpu.module @test {
gpu.return
}
+//-----
+ // CHECK-LABEL: multi_reduction_2d_last_dim
+ // CHECK-SAME: [[SRC:%.+]]: vector<32x80xf32>, [[ACC:%.+]]: vector<32xf32>
+ //
+ // Extract column tiles for the first row-tile:
+ // CHECK: [[TILE00:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [0, 0]{{.*}} : vector<32x80xf32> to vector<16x16xf32>
+ // CHECK: [[TILE01:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [0, 16]{{.*}} : vector<32x80xf32> to vector<16x16xf32>
+ // CHECK: [[TILE02:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [0, 32]{{.*}} : vector<32x80xf32> to vector<16x16xf32>
+ // CHECK: [[TILE03:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [0, 48]{{.*}} : vector<32x80xf32> to vector<16x16xf32>
+ // CHECK: [[TILE04:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [0, 64]{{.*}} : vector<32x80xf32> to vector<16x16xf32>
+ //
+ // Perform sequential reduction for first tile (rows 0..15):
+ // CHECK: [[TMP00:%.+]] = arith.addf [[TILE00]], [[TILE01]] : vector<16x16xf32>
+ // CHECK: [[TMP01:%.+]] = arith.addf [[TMP00]], [[TILE02]] : vector<16x16xf32>
+ // CHECK: [[TMP02:%.+]] = arith.addf [[TMP01]], [[TILE03]] : vector<16x16xf32>
+ // CHECK: [[TMP03:%.+]] = arith.addf [[TMP02]], [[TILE04]] : vector<16x16xf32>
+ // CHECK: [[ACC0:%.+]] = vector.extract_strided_slice [[ACC]] {offsets = [0]{{.*}} : vector<32xf32> to vector<16xf32>
+ // CHECK: [[RED0:%.+]] = vector.multi_reduction <add>, [[TMP03]], [[ACC0]] [1] : vector<16x16xf32> to vector<16xf32>
+ // CHECK: [[INS0:%.+]] = vector.insert_strided_slice [[RED0]], {{%.+}} {offsets = [0]{{.*}} : vector<16xf32> into vector<32xf32>
+ //
+ // Extract column tiles for the second row-tile:
+ // CHECK: [[TILE10:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [16, 0]{{.*}} : vector<32x80xf32> to vector<16x16xf32>
+ // CHECK: [[TILE11:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [16, 16]{{.*}} : vector<32x80xf32> to vector<16x16xf32>
+ // CHECK: [[TILE12:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [16, 32]{{.*}} : vector<32x80xf32> to vector<16x16xf32>
+ // CHECK: [[TILE13:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [16, 48]{{.*}} : vector<32x80xf32> to vector<16x16xf32>
+ // CHECK: [[TILE14:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [16, 64]{{.*}} : vector<32x80xf32> to vector<16x16xf32>
+ //
+ // Perform sequential reduction for second tile (rows 16..31):
+ // CHECK: [[TMP10:%.+]] = arith.addf [[TILE10]], [[TILE11]] : vector<16x16xf32>
+ // CHECK: [[TMP11:%.+]] = arith.addf [[TMP10]], [[TILE12]] : vector<16x16xf32>
+ // CHECK: [[TMP12:%.+]] = arith.addf [[TMP11]], [[TILE13]] : vector<16x16xf32>
+ // CHECK: [[TMP13:%.+]] = arith.addf [[TMP12]], [[TILE14]] : vector<16x16xf32>
+ // CHECK: [[ACC1:%.+]] = vector.extract_strided_slice [[ACC]] {offsets = [16]{{.*}} : vector<32xf32> to vector<16xf32>
+ // CHECK: [[RED1:%.+]] = vector.multi_reduction <add>, [[TMP13]], [[ACC1]] [1] : vector<16x16xf32> to vector<16xf32>
+ // CHECK: [[INS1:%.+]] = vector.insert_strided_slice [[RED1]], [[INS0]] {offsets = [16]{{.*}} : vector<16xf32> into vector<32xf32>
+ gpu.func @multi_reduction_2d_last_dim(%src: vector<32x80xf32>, %acc: vector<32xf32>) -> vector<32xf32> {
+ %0 = vector.multi_reduction <add>, %src, %acc {layout_operand_0 = #xegpu.layout<inst_data = [16, 16]>} [1] : vector<32x80xf32> to vector<32xf32>
+ gpu.return %0 : vector<32xf32>
+ }
+
+//-----
+ // Reduction over multiple dimensions [1, 3] in a 4D vector.
+ // source: <4x8x16x32xf32>, target shape: <2x4x16x16xf32>
+ //
+ // CHECK-LABEL: multi_reduction_multi_dim
+ // CHECK-SAME: [[SRC:%.+]]: vector<4x8x16x32xf32>, [[ACC:%.+]]: vector<4x16xf32>
+ //
+ // First kept tile [0, 0]: extract 4 source tiles over reduced dims [1, 3]
+ // CHECK: [[T00:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [0, 0, 0, 0], sizes = [2, 4, 16, 16]{{.*}} : vector<4x8x16x32xf32> to vector<2x4x16x16xf32>
+ // CHECK: [[T01:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [0, 0, 0, 16], sizes = [2, 4, 16, 16]{{.*}} : vector<4x8x16x32xf32> to vector<2x4x16x16xf32>
+ // CHECK: [[T02:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [0, 4, 0, 0], sizes = [2, 4, 16, 16]{{.*}} : vector<4x8x16x32xf32> to vector<2x4x16x16xf32>
+ // CHECK: [[T03:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [0, 4, 0, 16], sizes = [2, 4, 16, 16]{{.*}} : vector<4x8x16x32xf32> to vector<2x4x16x16xf32>
+ // Sequential reduction:
+ // CHECK: [[R00:%.+]] = arith.addf [[T00]], [[T01]] : vector<2x4x16x16xf32>
+ // CHECK: [[R01:%.+]] = arith.addf [[R00]], [[T02]] : vector<2x4x16x16xf32>
+ // CHECK: [[R02:%.+]] = arith.addf [[R01]], [[T03]] : vector<2x4x16x16xf32>
+ // CHECK: [[ACC0:%.+]] = vector.extract_strided_slice [[ACC]] {offsets = [0, 0], sizes = [2, 16]{{.*}} : vector<4x16xf32> to vector<2x16xf32>
+ // CHECK: [[MR0:%.+]] = vector.multi_reduction <add>, [[R02]], [[ACC0]] [1, 3] : vector<2x4x16x16xf32> to vector<2x16xf32>
+ // CHECK: [[INS0:%.+]] = vector.insert_strided_slice [[MR0]], {{%.+}} {offsets = [0, 0]{{.*}} : vector<2x16xf32> into vector<4x16xf32>
+ //
+ // Second kept tile [2, 0]:
+ // CHECK: [[T10:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [2, 0, 0, 0], sizes = [2, 4, 16, 16]{{.*}} : vector<4x8x16x32xf32> to vector<2x4x16x16xf32>
+ // CHECK: [[T11:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [2, 0, 0, 16], sizes = [2, 4, 16, 16]{{.*}} : vector<4x8x16x32xf32> to vector<2x4x16x16xf32>
+ // CHECK: [[T12:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [2, 4, 0, 0], sizes = [2, 4, 16, 16]{{.*}} : vector<4x8x16x32xf32> to vector<2x4x16x16xf32>
+ // CHECK: [[T13:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [2, 4, 0, 16], sizes = [2, 4, 16, 16]{{.*}} : vector<4x8x16x32xf32> to vector<2x4x16x16xf32>
+ // Sequential reduction:
+ // CHECK: [[R10:%.+]] = arith.addf [[T10]], [[T11]] : vector<2x4x16x16xf32>
+ // CHECK: [[R11:%.+]] = arith.addf [[R10]], [[T12]] : vector<2x4x16x16xf32>
+ // CHECK: [[R12:%.+]] = arith.addf [[R11]], [[T13]] : vector<2x4x16x16xf32>
+ // CHECK: [[ACC1:%.+]] = vector.extract_strided_slice [[ACC]] {offsets = [2, 0], sizes = [2, 16]{{.*}} : vector<4x16xf32> to vector<2x16xf32>
+ // CHECK: [[MR1:%.+]] = vector.multi_reduction <add>, [[R12]], [[ACC1]] [1, 3] : vector<2x4x16x16xf32> to vector<2x16xf32>
+ // CHECK: [[INS1:%.+]] = vector.insert_strided_slice [[MR1]], [[INS0]] {offsets = [2, 0]{{.*}} : vector<2x16xf32> into vector<4x16xf32>
+ gpu.func @multi_reduction_multi_dim(%src: vector<4x8x16x32xf32>, %acc: vector<4x16xf32>) -> vector<4x16xf32> {
+ %0 = vector.multi_reduction <add>, %src, %acc {layout_operand_0 = #xegpu.layout<inst_data = [2, 4, 16, 16]>} [1, 3] : vector<4x8x16x32xf32> to vector<4x16xf32>
+ gpu.return %0 : vector<4x16xf32>
+ }
+
+//-----
+ // Reduction over dimension [0] in a 2D vector.
+ // source: <48x32xf32>, target shape: <16x16xf32>
+ //
+ // CHECK-LABEL: multi_reduction_reduce_dim0
+ // CHECK-SAME: [[SRC:%.+]]: vector<48x32xf32>, [[ACC:%.+]]: vector<32xf32>
+ //
+ // First column tile [0]: extract 3 source tiles over reduced dim [0]
+ // CHECK: [[T00:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [0, 0], sizes = [16, 16]{{.*}} : vector<48x32xf32> to vector<16x16xf32>
+ // CHECK: [[T01:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [16, 0], sizes = [16, 16]{{.*}} : vector<48x32xf32> to vector<16x16xf32>
+ // CHECK: [[T02:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [32, 0], sizes = [16, 16]{{.*}} : vector<48x32xf32> to vector<16x16xf32>
+ // Sequential reduction:
+ // CHECK: [[R00:%.+]] = arith.addf [[T00]], [[T01]] : vector<16x16xf32>
+ // CHECK: [[R01:%.+]] = arith.addf [[R00]], [[T02]] : vector<16x16xf32>
+ // CHECK: [[ACC0:%.+]] = vector.extract_strided_slice [[ACC]] {offsets = [0], sizes = [16]{{.*}} : vector<32xf32> to vector<16xf32>
+ // CHECK: [[MR0:%.+]] = vector.multi_reduction <add>, [[R01]], [[ACC0]] [0] : vector<16x16xf32> to vector<16xf32>
+ // CHECK: [[INS0:%.+]] = vector.insert_strided_slice [[MR0]], {{%.+}} {offsets = [0]{{.*}} : vector<16xf32> into vector<32xf32>
+ //
+ // Second column tile [16]:
+ // CHECK: [[T10:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [0, 16], sizes = [16, 16]{{.*}} : vector<48x32xf32> to vector<16x16xf32>
+ // CHECK: [[T11:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [16, 16], sizes = [16, 16]{{.*}} : vector<48x32xf32> to vector<16x16xf32>
+ // CHECK: [[T12:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [32, 16], sizes = [16, 16]{{.*}} : vector<48x32xf32> to vector<16x16xf32>
+ // Sequential reduction:
+ // CHECK: [[R10:%.+]] = arith.addf [[T10]], [[T11]] : vector<16x16xf32>
+ // CHECK: [[R11:%.+]] = arith.addf [[R10]], [[T12]] : vector<16x16xf32>
+ // CHECK: [[ACC1:%.+]] = vector.extract_strided_slice [[ACC]] {offsets = [16], sizes = [16]{{.*}} : vector<32xf32> to vector<16xf32>
+ // CHECK: [[MR1:%.+]] = vector.multi_reduction <add>, [[R11]], [[ACC1]] [0] : vector<16x16xf32> to vector<16xf32>
+ // CHECK: [[INS1:%.+]] = vector.insert_strided_slice [[MR1]], [[INS0]] {offsets = [16]{{.*}} : vector<16xf32> into vector<32xf32>
+ gpu.func @multi_reduction_reduce_dim0(%src: vector<48x32xf32>, %acc: vector<32xf32>) -> vector<32xf32> {
+ %0 = vector.multi_reduction <add>, %src, %acc {layout_operand_0 = #xegpu.layout<inst_data = [16, 16]>} [0] : vector<48x32xf32> to vector<32xf32>
+ gpu.return %0 : vector<32xf32>
+ }
+
+//-----
+ // source: <32x16xf32>, target tile: <16x16xf32>
+ // Verifies that the patterns works correctly when there is
+ // no place for the sequential elementwise 'arith' reduction.
+ //
+ // CHECK-LABEL: multi_reduction_no_elwise
+ // CHECK-SAME: [[SRC:%.+]]: vector<32x16xf32>, [[ACC:%.+]]: vector<32xf32>
+ //
+ // First row tile [0]: single source tile, no arith reduction
+ // CHECK: [[T0:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [0, 0], sizes = [16, 16]{{.*}} : vector<32x16xf32> to vector<16x16xf32>
+ // CHECK: [[ACC0:%.+]] = vector.extract_strided_slice [[ACC]] {offsets = [0], sizes = [16]{{.*}} : vector<32xf32> to vector<16xf32>
+ // CHECK: [[MR0:%.+]] = vector.multi_reduction <add>, [[T0]], [[ACC0]] [1] : vector<16x16xf32> to vector<16xf32>
+ // CHECK: [[INS0:%.+]] = vector.insert_strided_slice [[MR0]], {{%.+}} {offsets = [0]{{.*}} : vector<16xf32> into vector<32xf32>
+ //
+ // Second row tile [16]: single source tile, no arith reduction
+ // CHECK: [[T1:%.+]] = vector.extract_strided_slice [[SRC]] {offsets = [16, 0], sizes = [16, 16]{{.*}} : vector<32x16xf32> to vector<16x16xf32>
+ // CHECK: [[ACC1:%.+]] = vector.extract_strided_slice [[ACC]] {offsets = [16], sizes = [16]{{.*}} : vector<32xf32> to vector<16xf32>
+ // CHECK: [[MR1:%.+]] = vector.multi_reduction <add>, [[T1]], [[ACC1]] [1] : vector<16x16xf32> to vector<16xf32>
+ // CHECK: [[INS1:%.+]] = vector.insert_strided_slice [[MR1]], [[INS0]] {offsets = [16]{{.*}} : vector<16xf32> into vector<32xf32>
+ gpu.func @multi_reduction_no_elwise(%src: vector<32x16xf32>, %acc: vector<32xf32>) -> vector<32xf32> {
+ %0 = vector.multi_reduction <add>, %src, %acc {layout_operand_0 = #xegpu.layout<inst_data = [16, 16]>} [1] : vector<32x16xf32> to vector<32xf32>
+ gpu.return %0 : vector<32xf32>
+ }
+
}
diff --git a/mlir/test/lib/Dialect/XeGPU/TestXeGPUTransforms.cpp b/mlir/test/lib/Dialect/XeGPU/TestXeGPUTransforms.cpp
index 7006051b9c033..581072f5218cd 100644
--- a/mlir/test/lib/Dialect/XeGPU/TestXeGPUTransforms.cpp
+++ b/mlir/test/lib/Dialect/XeGPU/TestXeGPUTransforms.cpp
@@ -94,6 +94,19 @@ struct TestXeGPUUnrollingPatterns
if (isa<xegpu::DpasOp>(op))
return SmallVector<int64_t>{8, 16, 16};
+ // For vector.multi_reduction, read tile shape from the layout attribute
+ // on the source operand (layout_operand_0).
+ if (isa<vector::MultiDimReductionOp>(op)) {
+ xegpu::DistributeLayoutAttr layout =
+ xegpu::getDistributeLayoutAttr(op->getOpOperand(0));
+ if (layout) {
+ auto instData = layout.getEffectiveInstDataAsInt();
+ if (!instData.empty())
+ return instData;
+ }
+ return std::nullopt;
+ }
+
return std::nullopt;
});
More information about the Mlir-commits
mailing list