[Mlir-commits] [mlir] [mlir] [linalg] Fold reduce(broadcast(x)) max/min (PR #213190)
Chuanqi Xu
llvmlistbot at llvm.org
Thu Jul 30 19:09:46 PDT 2026
https://github.com/ChuanqiXu9 created https://github.com/llvm/llvm-project/pull/213190
This patch tries to implement the optimization for
```
max(broadcast(x)) -> x # (or replace min)
```
We can extend the optimization to other reduce op like add, mul, and, or, xor ... in the future. This patch tries to be small by implementing the simplest case.
AI assisted.
>From b3eeedec377532d4b7210df2af335f9cd3047451 Mon Sep 17 00:00:00 2001
From: "yedeng.yd" <yedeng.yd at alibaba-inc.com>
Date: Fri, 31 Jul 2026 10:05:27 +0800
Subject: [PATCH] [mlir] [linalg] Fold reduce(broadcast(x)) max/min
This patch tries to implement the optimization for
```
max(broadcast(x)) -> x # (or replace min)
```
We can extend the optimization to other reduce op like
add, mul, and, or, xor ... in the future.
AI assisted.
---
.../Dialect/Linalg/IR/LinalgStructuredOps.td | 1 +
mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp | 138 ++++++++++++++++++
mlir/test/Dialect/Linalg/canonicalize.mlir | 83 +++++++++++
3 files changed, 222 insertions(+)
diff --git a/mlir/include/mlir/Dialect/Linalg/IR/LinalgStructuredOps.td b/mlir/include/mlir/Dialect/Linalg/IR/LinalgStructuredOps.td
index fa8125f280db2..fa786f9420247 100644
--- a/mlir/include/mlir/Dialect/Linalg/IR/LinalgStructuredOps.td
+++ b/mlir/include/mlir/Dialect/Linalg/IR/LinalgStructuredOps.td
@@ -388,6 +388,7 @@ def ReduceOp : LinalgStructuredBase_Op<"reduce", [
}];
let hasCustomAssemblyFormat = 1;
+ let hasCanonicalizer = 1;
let hasVerifier = 1;
}
diff --git a/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp b/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp
index 170e1edf8a55d..c5775ca011b75 100644
--- a/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp
+++ b/mlir/lib/Dialect/Linalg/IR/LinalgOps.cpp
@@ -2016,6 +2016,144 @@ LogicalResult ReduceOp::verify() {
return success();
}
+namespace {
+
+enum class BroadcastReduceKind {
+ MaxSI,
+ MaxUI,
+ MinSI,
+ MinUI,
+};
+
+static std::optional<BroadcastReduceKind>
+matchBroadcastReduceBody(ReduceOp reduceOp) {
+ if (reduceOp.getNumDpsInputs() != 1 || reduceOp.getNumDpsInits() != 1 ||
+ !reduceOp.getBody())
+ return std::nullopt;
+
+ // Match the simplest linalg.reduce body. e.g.,
+ //
+ // ^bb0(%in: i32, %acc: i32):
+ // %max = arith.maxsi %in, %acc : i32
+ // linalg.yield %max : i32
+ Block &block = *reduceOp.getBody();
+ if (block.getNumArguments() != 2 ||
+ !llvm::hasSingleElement(block.without_terminator()))
+ return std::nullopt;
+
+ auto yieldOp = dyn_cast<YieldOp>(block.getTerminator());
+ if (!yieldOp || yieldOp.getNumOperands() != 1)
+ return std::nullopt;
+
+ Operation *combineOp = yieldOp.getOperand(0).getDefiningOp();
+ if (!combineOp || combineOp->getNumOperands() != 2 ||
+ combineOp->getOperand(0).getType() != block.getArgument(0).getType() ||
+ combineOp->getOperand(1).getType() != block.getArgument(1).getType())
+ return std::nullopt;
+
+ // Checks that the combine op **only** used the block arguments and
+ // we allow the block arguments to exchange their orders.
+ if (!((combineOp->getOperand(0) == block.getArgument(0) &&
+ combineOp->getOperand(1) == block.getArgument(1)) ||
+ (combineOp->getOperand(0) == block.getArgument(1) &&
+ combineOp->getOperand(1) == block.getArgument(0))))
+ return std::nullopt;
+
+ // TODO: We can extend the list here.
+ return TypeSwitch<Operation *, std::optional<BroadcastReduceKind>>(combineOp)
+ .Case<arith::MaxSIOp>(
+ [](arith::MaxSIOp) { return BroadcastReduceKind::MaxSI; })
+ .Case<arith::MaxUIOp>(
+ [](arith::MaxUIOp) { return BroadcastReduceKind::MaxUI; })
+ .Case<arith::MinSIOp>(
+ [](arith::MinSIOp) { return BroadcastReduceKind::MinSI; })
+ .Case<arith::MinUIOp>(
+ [](arith::MinUIOp) { return BroadcastReduceKind::MinUI; })
+ .Default([](Operation *) -> std::optional<BroadcastReduceKind> {
+ return std::nullopt;
+ });
+}
+
+static bool hasBroadcastReduceIdentity(Value init, BroadcastReduceKind kind) {
+ auto initAttr = getScalarConstantAttrFromDenseSplat(init);
+ if (!initAttr)
+ return false;
+
+ auto integerAttr = dyn_cast<IntegerAttr>(*initAttr);
+ if (!integerAttr)
+ return false;
+
+ const APInt &value = integerAttr.getValue();
+ switch (kind) {
+ case BroadcastReduceKind::MaxSI:
+ return value.isMinSignedValue();
+ case BroadcastReduceKind::MaxUI:
+ return value.isZero();
+ case BroadcastReduceKind::MinSI:
+ return value.isMaxSignedValue();
+ case BroadcastReduceKind::MinUI:
+ return value.isAllOnes();
+ }
+ llvm_unreachable("unknown broadcast reduction kind");
+}
+
+// Fold cases like:
+//
+// maxsi(broadcast(x)) -> x
+// minsi(broadcast(y)) -> y
+//
+// TODO: We can add other op. e.g., add, mul, and, or, xor.
+struct FoldReduceBroadcast : public OpRewritePattern<linalg::ReduceOp> {
+ using OpRewritePattern<linalg::ReduceOp>::OpRewritePattern;
+
+ LogicalResult matchAndRewrite(linalg::ReduceOp reduceOp,
+ PatternRewriter &rewriter) const override {
+ if (reduceOp.getNumResults() != 1)
+ return failure();
+
+ auto broadcastOp =
+ reduceOp.getInputs().front().getDefiningOp<linalg::BroadcastOp>();
+ if (!broadcastOp || broadcastOp.getNumResults() != 1)
+ return failure();
+
+ auto sourceType =
+ dyn_cast<RankedTensorType>(broadcastOp.getInput().getType());
+ auto broadcastType =
+ dyn_cast<RankedTensorType>(broadcastOp.getResult().front().getType());
+ auto resultType =
+ dyn_cast<RankedTensorType>(reduceOp.getResult(0).getType());
+ if (!sourceType || !broadcastType || !resultType ||
+ !sourceType.hasStaticShape() || !broadcastType.hasStaticShape() ||
+ !resultType.hasStaticShape() || sourceType != resultType)
+ return failure();
+
+ ArrayRef<int64_t> broadcastDims = broadcastOp.getDimensions();
+ ArrayRef<int64_t> reduceDims = reduceOp.getDimensions();
+ if (broadcastDims != reduceDims)
+ return failure();
+
+ for (int64_t dimension : broadcastDims)
+ if (broadcastType.getDimSize(dimension) <= 0)
+ return failure();
+
+ std::optional<BroadcastReduceKind> kind =
+ matchBroadcastReduceBody(reduceOp);
+ if (!kind ||
+ !hasBroadcastReduceIdentity(reduceOp.getInits().front(), *kind))
+ return failure();
+
+ rewriter.replaceOp(reduceOp, broadcastOp.getInput());
+ return success();
+ }
+};
+
+} // namespace
+
+void ReduceOp::getCanonicalizationPatterns(RewritePatternSet &results,
+ MLIRContext *context) {
+ results.add<FoldReduceBroadcast>(context);
+}
+
//===----------------------------------------------------------------------===//
// TransposeOp
//===----------------------------------------------------------------------===//
diff --git a/mlir/test/Dialect/Linalg/canonicalize.mlir b/mlir/test/Dialect/Linalg/canonicalize.mlir
index bb11ce0d4dfb8..8670bc7833fb5 100644
--- a/mlir/test/Dialect/Linalg/canonicalize.mlir
+++ b/mlir/test/Dialect/Linalg/canonicalize.mlir
@@ -1,5 +1,88 @@
// RUN: mlir-opt %s -canonicalize="test-convergence" -split-input-file | FileCheck %s
+// CHECK-LABEL: func.func @reduce_broadcast_maxsi
+func.func @reduce_broadcast_maxsi(%input: tensor<3xi32>,
+ %broadcast_init: tensor<3x4xi32>)
+ -> tensor<3xi32> {
+ %reduce_init = arith.constant dense<-2147483648> : tensor<3xi32>
+ %broadcasted = linalg.broadcast
+ ins(%input : tensor<3xi32>)
+ outs(%broadcast_init : tensor<3x4xi32>) dimensions = [1]
+ %result = linalg.reduce
+ ins(%broadcasted : tensor<3x4xi32>)
+ outs(%reduce_init : tensor<3xi32>) dimensions = [1]
+ (%in: i32, %out: i32) {
+ %max = arith.maxsi %in, %out : i32
+ linalg.yield %max : i32
+ }
+ return %result : tensor<3xi32>
+}
+// CHECK-NEXT: return %{{.*}} : tensor<3xi32>
+
+// -----
+
+// CHECK-LABEL: func.func @reduce_broadcast_minui
+func.func @reduce_broadcast_minui(%input: tensor<3xi8>,
+ %broadcast_init: tensor<2x3x4xi8>)
+ -> tensor<3xi8> {
+ %reduce_init = arith.constant dense<255> : tensor<3xi8>
+ %broadcasted = linalg.broadcast
+ ins(%input : tensor<3xi8>)
+ outs(%broadcast_init : tensor<2x3x4xi8>) dimensions = [0, 2]
+ %result = linalg.reduce
+ ins(%broadcasted : tensor<2x3x4xi8>)
+ outs(%reduce_init : tensor<3xi8>) dimensions = [0, 2]
+ (%in: i8, %out: i8) {
+ %min = arith.minui %in, %out : i8
+ linalg.yield %min : i8
+ }
+ return %result : tensor<3xi8>
+}
+// CHECK-NEXT: return %{{.*}} : tensor<3xi8>
+
+// -----
+
+// CHECK-LABEL: func.func @reduce_broadcast_wrong_dimension
+// CHECK: linalg.broadcast
+// CHECK: linalg.reduce
+func.func @reduce_broadcast_wrong_dimension(%input: tensor<3xi32>,
+ %broadcast_init: tensor<3x4xi32>, %reduce_init: tensor<4xi32>)
+ -> tensor<4xi32> {
+ %broadcasted = linalg.broadcast
+ ins(%input : tensor<3xi32>)
+ outs(%broadcast_init : tensor<3x4xi32>) dimensions = [1]
+ %result = linalg.reduce
+ ins(%broadcasted : tensor<3x4xi32>)
+ outs(%reduce_init : tensor<4xi32>) dimensions = [0]
+ (%in: i32, %out: i32) {
+ %max = arith.maxsi %in, %out : i32
+ linalg.yield %max : i32
+ }
+ return %result : tensor<4xi32>
+}
+
+// -----
+
+// CHECK-LABEL: func.func @reduce_broadcast_non_identity_init
+// CHECK: linalg.broadcast
+// CHECK: linalg.reduce
+func.func @reduce_broadcast_non_identity_init(%input: tensor<3xi32>,
+ %broadcast_init: tensor<3x4xi32>)
+ -> tensor<3xi32> {
+ %reduce_init = arith.constant dense<0> : tensor<3xi32>
+ %broadcasted = linalg.broadcast
+ ins(%input : tensor<3xi32>)
+ outs(%broadcast_init : tensor<3x4xi32>) dimensions = [1]
+ %result = linalg.reduce
+ ins(%broadcasted : tensor<3x4xi32>)
+ outs(%reduce_init : tensor<3xi32>) dimensions = [1]
+ (%in: i32, %out: i32) {
+ %max = arith.maxsi %in, %out : i32
+ linalg.yield %max : i32
+ }
+ return %result : tensor<3xi32>
+}
+
// CHECK-LABEL: func @memref_cast(
func.func @memref_cast(%a: index, %b: index) -> memref<?x?xf32> {
%c0 = arith.constant 0 : index
More information about the Mlir-commits
mailing list