[Mlir-commits] [mlir] [MLIR][Vector] Add canonicalization for interleave/deinterleave chain (PR #196979)

Artem Kroviakov llvmlistbot at llvm.org
Tue May 12 01:57:32 PDT 2026


https://github.com/akroviakov updated https://github.com/llvm/llvm-project/pull/196979

>From 4edcfb9665f8a200b661cf7896e20e0a6b524ad8 Mon Sep 17 00:00:00 2001
From: Artem Kroviakov <artem.kroviakov at intel.com>
Date: Mon, 11 May 2026 15:43:02 +0000
Subject: [PATCH 1/2] [MLIR][Vector] Add canonicalization for
 interleave/deinterleave chain

---
 .../mlir/Dialect/Vector/IR/VectorOps.td       |  5 ++-
 mlir/lib/Dialect/Vector/IR/VectorOps.cpp      | 44 +++++++++++++++++++
 mlir/test/Dialect/Vector/canonicalize.mlir    | 10 +++++
 3 files changed, 57 insertions(+), 2 deletions(-)

diff --git a/mlir/include/mlir/Dialect/Vector/IR/VectorOps.td b/mlir/include/mlir/Dialect/Vector/IR/VectorOps.td
index 28a8109cb59c0..5acf2b4ab7649 100644
--- a/mlir/include/mlir/Dialect/Vector/IR/VectorOps.td
+++ b/mlir/include/mlir/Dialect/Vector/IR/VectorOps.td
@@ -584,6 +584,7 @@ def Vector_InterleaveOp :
       return ::llvm::cast<VectorType>(getResult().getType());
     }
   }];
+  let hasCanonicalizer = 1;
 }
 
 class ResultIsHalfSourceVectorType<string result> : TypesMatchWith<
@@ -2560,7 +2561,7 @@ def Vector_TypeCastOp :
 }
 
 def Vector_ConstantMaskOp :
-  Vector_Op<"constant_mask", [Pure, 
+  Vector_Op<"constant_mask", [Pure,
    DeclareOpInterfaceMethods<VectorUnrollOpInterface>
    ]>,
     Arguments<(ins DenseI64ArrayAttr:$mask_dim_sizes)>,
@@ -2620,7 +2621,7 @@ def Vector_ConstantMaskOp :
 }
 
 def Vector_CreateMaskOp :
-  Vector_Op<"create_mask", [Pure, 
+  Vector_Op<"create_mask", [Pure,
    DeclareOpInterfaceMethods<VectorUnrollOpInterface>
    ]>,
     Arguments<(ins Variadic<Index>:$mask_dim_sizes)>,
diff --git a/mlir/lib/Dialect/Vector/IR/VectorOps.cpp b/mlir/lib/Dialect/Vector/IR/VectorOps.cpp
index 51be1e4431e70..83ca718932433 100644
--- a/mlir/lib/Dialect/Vector/IR/VectorOps.cpp
+++ b/mlir/lib/Dialect/Vector/IR/VectorOps.cpp
@@ -8346,6 +8346,50 @@ Value mlir::vector::selectPassthru(OpBuilder &builder, Value mask,
 // InterleaveOp
 //===----------------------------------------------------------------------===//
 
+namespace {
+
+/// This canonicalization folds the following round-trip identity:
+///  interleave(deinterleave(x).even, deinterleave(x).odd) -> x
+struct InterleaveDeinterleaveFolder : public OpRewritePattern<InterleaveOp> {
+  using Base::Base;
+
+  LogicalResult matchAndRewrite(InterleaveOp interleaveOp,
+                                PatternRewriter &rewriter) const override {
+    auto lhsDefOp = interleaveOp.getLhs().getDefiningOp();
+    auto rhsDefOp = interleaveOp.getRhs().getDefiningOp();
+    if (!lhsDefOp || !rhsDefOp)
+      return rewriter.notifyMatchFailure(
+          interleaveOp, "expected both operands to be defined by an op");
+    if (!isa<DeinterleaveOp>(lhsDefOp) || !isa<DeinterleaveOp>(rhsDefOp))
+      return rewriter.notifyMatchFailure(
+          interleaveOp,
+          "expected both operands to be defined by a deinterleave op");
+    if (lhsDefOp != rhsDefOp)
+      return rewriter.notifyMatchFailure(
+          interleaveOp,
+          "expected both operands to be defined by the same deinterleave op");
+    if (auto deinterleaveRes = cast<OpResult>(interleaveOp.getLhs());
+        deinterleaveRes.getResultNumber())
+      return rewriter.notifyMatchFailure(interleaveOp,
+                                         "expected the lhs operand to be the "
+                                         "first result of the deinterleave op");
+    if (auto deinterleaveRes = cast<OpResult>(interleaveOp.getRhs());
+        deinterleaveRes.getResultNumber() != 1)
+      return rewriter.notifyMatchFailure(
+          interleaveOp, "expected the rhs operand to be the second result of "
+                        "the deinterleave op");
+    rewriter.replaceOp(interleaveOp,
+                       cast<DeinterleaveOp>(lhsDefOp).getSource());
+    return success();
+  }
+};
+} // namespace
+
+void InterleaveOp::getCanonicalizationPatterns(RewritePatternSet &results,
+                                               MLIRContext *context) {
+  results.add<InterleaveDeinterleaveFolder>(context);
+}
+
 std::optional<SmallVector<int64_t, 4>> InterleaveOp::getShapeForUnroll() {
   return llvm::to_vector<4>(getResultVectorType().getShape());
 }
diff --git a/mlir/test/Dialect/Vector/canonicalize.mlir b/mlir/test/Dialect/Vector/canonicalize.mlir
index 6aa92ab79a0dd..a43d93dfd1acb 100644
--- a/mlir/test/Dialect/Vector/canonicalize.mlir
+++ b/mlir/test/Dialect/Vector/canonicalize.mlir
@@ -4407,3 +4407,13 @@ func.func @no_fold_alltrue_mask_empty_body_scalar_result(
   %result = vector.mask %all_true, %passthru { vector.yield %val : i32 } : vector<1xi1> -> i32
   return %result : i32
 }
+
+// Fold interleave(deinterleave(x).even, deinterleave(x).odd) -> x
+// CHECK-LABEL: func @interleave_deinterleave_fold
+//  CHECK-SAME: (%[[ARG0:.*]]: vector<4xf32>)
+//       CHECK: return %[[ARG0]]
+func.func @interleave_deinterleave_fold(%arg0: vector<4xf32>) -> vector<4xf32> {
+  %even, %odd = vector.deinterleave %arg0 : vector<4xf32> -> vector<2xf32>
+  %result = vector.interleave %even, %odd : vector<2xf32> -> vector<4xf32>
+  return %result : vector<4xf32>
+}

>From 4d4c161a85ad832defa44a582b114b00c698ed04 Mon Sep 17 00:00:00 2001
From: Artem Kroviakov <artem.kroviakov at intel.com>
Date: Tue, 12 May 2026 08:57:13 +0000
Subject: [PATCH 2/2] Simplify folder

---
 mlir/lib/Dialect/Vector/IR/VectorOps.cpp | 10 +++-------
 1 file changed, 3 insertions(+), 7 deletions(-)

diff --git a/mlir/lib/Dialect/Vector/IR/VectorOps.cpp b/mlir/lib/Dialect/Vector/IR/VectorOps.cpp
index 83ca718932433..716dd8e6bb289 100644
--- a/mlir/lib/Dialect/Vector/IR/VectorOps.cpp
+++ b/mlir/lib/Dialect/Vector/IR/VectorOps.cpp
@@ -8355,12 +8355,9 @@ struct InterleaveDeinterleaveFolder : public OpRewritePattern<InterleaveOp> {
 
   LogicalResult matchAndRewrite(InterleaveOp interleaveOp,
                                 PatternRewriter &rewriter) const override {
-    auto lhsDefOp = interleaveOp.getLhs().getDefiningOp();
-    auto rhsDefOp = interleaveOp.getRhs().getDefiningOp();
+    auto lhsDefOp = interleaveOp.getLhs().getDefiningOp<DeinterleaveOp>();
+    auto rhsDefOp = interleaveOp.getRhs().getDefiningOp<DeinterleaveOp>();
     if (!lhsDefOp || !rhsDefOp)
-      return rewriter.notifyMatchFailure(
-          interleaveOp, "expected both operands to be defined by an op");
-    if (!isa<DeinterleaveOp>(lhsDefOp) || !isa<DeinterleaveOp>(rhsDefOp))
       return rewriter.notifyMatchFailure(
           interleaveOp,
           "expected both operands to be defined by a deinterleave op");
@@ -8378,8 +8375,7 @@ struct InterleaveDeinterleaveFolder : public OpRewritePattern<InterleaveOp> {
       return rewriter.notifyMatchFailure(
           interleaveOp, "expected the rhs operand to be the second result of "
                         "the deinterleave op");
-    rewriter.replaceOp(interleaveOp,
-                       cast<DeinterleaveOp>(lhsDefOp).getSource());
+    rewriter.replaceOp(interleaveOp, lhsDefOp.getSource());
     return success();
   }
 };



More information about the Mlir-commits mailing list