[Mlir-commits] [mlir] [mlir][affine] Add CollapseNestedIf pattern to collapse nested single-then affine.if ops (PR #213434)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Sat Aug 1 04:46:35 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir
@llvm/pr-subscribers-mlir-affine
Author: lonely eagle (linuxlonelyeagle)
<details>
<summary>Changes</summary>
This PR introduces a new canonicalization pattern CollapseNestedIf for AffineIfOp to fold nested single-then affine.if operations into a single affine.if op.
---
Full diff: https://github.com/llvm/llvm-project/pull/213434.diff
2 Files Affected:
- (modified) mlir/lib/Dialect/Affine/IR/AffineOps.cpp (+90-1)
- (modified) mlir/test/Dialect/Affine/canonicalize.mlir (+23)
``````````diff
diff --git a/mlir/lib/Dialect/Affine/IR/AffineOps.cpp b/mlir/lib/Dialect/Affine/IR/AffineOps.cpp
index 98c2be14da5aa..c11d7973552d1 100644
--- a/mlir/lib/Dialect/Affine/IR/AffineOps.cpp
+++ b/mlir/lib/Dialect/Affine/IR/AffineOps.cpp
@@ -3163,6 +3163,95 @@ struct AlwaysTrueOrFalseIf : public OpRewritePattern<AffineIfOp> {
return success();
}
};
+
+// Returns true if both AffineIfOps share the exact same set of Dim operands.
+static bool hasSameDimOperands(AffineIfOp a, AffineIfOp b) {
+ auto dimsA = a.getOperands().take_front(a.getIntegerSet().getNumDims());
+ auto dimsB = b.getOperands().take_front(b.getIntegerSet().getNumDims());
+ return llvm::SmallDenseSet<Value, 4>(dimsA.begin(), dimsA.end()) ==
+ llvm::SmallDenseSet<Value, 4>(dimsB.begin(), dimsB.end());
+}
+
+// Collapses nested single-then AffineIfOps into a single AffineIfOp by merging
+// their constraints when both ops operate on the same set of Dim operands.
+struct CollapseNestedIf : public OpRewritePattern<AffineIfOp> {
+ using OpRewritePattern<AffineIfOp>::OpRewritePattern;
+
+ LogicalResult matchAndRewrite(AffineIfOp innerIf,
+ PatternRewriter &rewriter) const override {
+ // Only combine single-then AffineIfOps where innerIf is the sole statement
+ // in outerIf's then block.
+ if (innerIf.hasElse())
+ return failure();
+
+ auto outerIf = dyn_cast<AffineIfOp>(innerIf->getParentOp());
+ if (!outerIf || outerIf.hasElse())
+ return failure();
+
+ Block *outerThen = outerIf.getThenBlock();
+ if (&outerThen->front() != innerIf.getOperation() ||
+ std::distance(outerThen->begin(), outerThen->end()) != 2)
+ return failure();
+
+ // Only merge when both ops share the same Dim operands (loop IVs). Merging
+ // different dim operands causes the if-statement to lose scheduling
+ // opportunities.
+ if (!hasSameDimOperands(innerIf, outerIf))
+ return failure();
+
+ IntegerSet outerSet = outerIf.getIntegerSet();
+ IntegerSet innerSet = innerIf.getIntegerSet();
+
+ OperandRange outerValues = outerIf.getOperands();
+ OperandRange innerValues = innerIf->getOperands();
+ ValueRange outerDimValues(outerValues.take_front(outerSet.getNumDims()));
+ ValueRange innerDimValues(innerValues.take_front(innerSet.getNumDims()));
+
+ // Map each inner Dim to its position in outer Dims.
+ SmallVector<AffineExpr, 4> dimRepls;
+ for (Value v : innerDimValues) {
+ auto it = llvm::find(outerDimValues, v);
+ dimRepls.push_back(
+ rewriter.getAffineDimExpr(std::distance(outerDimValues.begin(), it)));
+ }
+
+ // Shift inner Symbols after outer Symbols.
+ SmallVector<AffineExpr, 4> symRepls;
+ for (size_t i = 0, e = innerSet.getNumSymbols(),
+ m = outerSet.getNumSymbols();
+ i < e; ++i) {
+ symRepls.push_back(rewriter.getAffineSymbolExpr(m + i));
+ }
+
+ // Remap inner constraints and append to outer constraints.
+ SmallVector<AffineExpr, 8> constraints(outerSet.getConstraints().begin(),
+ outerSet.getConstraints().end());
+ for (AffineExpr e : innerSet.getConstraints())
+ constraints.push_back(e.replaceDimsAndSymbols(dimRepls, symRepls));
+
+ SmallVector<bool, 8> eqFlags(outerSet.getEqFlags().begin(),
+ outerSet.getEqFlags().end());
+ llvm::append_range(eqFlags, innerSet.getEqFlags());
+ auto newSet =
+ IntegerSet::get(outerSet.getNumDims(),
+ outerSet.getNumSymbols() + innerSet.getNumSymbols(),
+ constraints, eqFlags);
+
+ SmallVector<Value, 8> newOperands(outerValues);
+ llvm::append_range(newOperands,
+ innerValues.drop_front(outerSet.getNumDims()));
+
+ rewriter.eraseOp(outerThen->getTerminator());
+ rewriter.mergeBlocks(innerIf.getThenBlock(), outerThen);
+ rewriter.modifyOpInPlace(outerIf, [&]() {
+ outerIf.setIntegerSet(newSet);
+ outerIf->setOperands(newOperands);
+ });
+ rewriter.eraseOp(innerIf);
+ return success();
+ }
+};
+
} // namespace
/// AffineIfOp has two regions -- `then` and `else`. The flow of data should be
@@ -3378,7 +3467,7 @@ LogicalResult AffineIfOp::fold(FoldAdaptor, SmallVectorImpl<OpFoldResult> &) {
void AffineIfOp::getCanonicalizationPatterns(RewritePatternSet &results,
MLIRContext *context) {
- results.add<SimplifyDeadElse, AlwaysTrueOrFalseIf>(context);
+ results.add<SimplifyDeadElse, AlwaysTrueOrFalseIf, CollapseNestedIf>(context);
}
//===----------------------------------------------------------------------===//
diff --git a/mlir/test/Dialect/Affine/canonicalize.mlir b/mlir/test/Dialect/Affine/canonicalize.mlir
index 7d236ef3c2421..cf6c82ffc06bc 100644
--- a/mlir/test/Dialect/Affine/canonicalize.mlir
+++ b/mlir/test/Dialect/Affine/canonicalize.mlir
@@ -654,6 +654,29 @@ func.func @canonicalize_affine_if_compose_apply(%N: index) {
// -----
+// CHECK: #[[$SET:.+]] = affine_set<(d0, d1)[s0, s1] : (d0 - 2 >= 0, -d1 + 99 >= 0, s0 >= 0, d1 - 2 >= 0, -d0 + 99 >= 0, s1 >= 0)>
+
+// CHECK-LABEL: func @collapse_nested_affine_if
+// CHECK-SAME: %[[ARG0:.*]]: index, %[[ARG1:.*]]: index)
+
+func.func @collapse_nested_affine_if(%arg : index, %arg1 : index) {
+ affine.for %i = 0 to 100 {
+ affine.for %j = 0 to 100 {
+ affine.if affine_set<(d0, d1)[s0] : (d0 - 2 >= 0, -d1 + 99 >= 0, s0 >= 0)>(%i, %j)[%arg] {
+ affine.if affine_set<(d0, d1)[s0] : (d0 - 2 >= 0, -d1 + 99 >= 0, s0 >= 0)>(%j, %i)[%arg1] {
+ "test.foo"() : () -> ()
+ }
+ }
+ }
+ }
+ return
+}
+// CHECK: affine.for %[[I:.*]] = 0 to 100
+// CHECK: affine.for %[[J:.*]] = 0 to 100
+// CHECK: affine.if #[[$SET]](%[[I]], %[[J]])[%[[ARG0]], %[[ARG1]]]
+
+// -----
+
// CHECK-DAG: #[[$LBMAP:.*]] = affine_map<()[s0] -> (0, s0)>
// CHECK-DAG: #[[$UBMAP:.*]] = affine_map<()[s0] -> (1024, s0 * 2)>
``````````
</details>
https://github.com/llvm/llvm-project/pull/213434
More information about the Mlir-commits
mailing list