[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