[Mlir-commits] [mlir] [mlir][affine] Add CollapseNestedIf pattern to collapse nested single-then affine.if ops (PR #213434)

lonely eagle llvmlistbot at llvm.org
Sat Aug 1 04:45:50 PDT 2026


https://github.com/linuxlonelyeagle created https://github.com/llvm/llvm-project/pull/213434

This PR introduces a new canonicalization pattern CollapseNestedIf for AffineIfOp to fold nested single-then affine.if operations into a single affine.if op.

>From 54d6332f9f21cc217b392dd38176dc2839301691 Mon Sep 17 00:00:00 2001
From: linuxlonelyeagle <2020382038 at qq.com>
Date: Sat, 1 Aug 2026 11:43:09 +0000
Subject: [PATCH] update test.

---
 mlir/lib/Dialect/Affine/IR/AffineOps.cpp   | 91 +++++++++++++++++++++-
 mlir/test/Dialect/Affine/canonicalize.mlir | 23 ++++++
 2 files changed, 113 insertions(+), 1 deletion(-)

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)>
 



More information about the Mlir-commits mailing list