[Mlir-commits] [mlir] 6bb0625 - [MLIR][vector] vector.deinterleave to vector.shuffle decomposition (#177897)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Wed May 6 02:30:29 PDT 2026
Author: Noah Prisament
Date: 2026-05-06T10:30:24+01:00
New Revision: 6bb0625acdd7dac10b06395301407661b5d31478
URL: https://github.com/llvm/llvm-project/commit/6bb0625acdd7dac10b06395301407661b5d31478
DIFF: https://github.com/llvm/llvm-project/commit/6bb0625acdd7dac10b06395301407661b5d31478.diff
LOG: [MLIR][vector] vector.deinterleave to vector.shuffle decomposition (#177897)
This PR adds a rewrite pattern for vector.deinterleave ops that rewrites
them using vector.shuffle ops. This is similar to the existing pattern
for vector.interleave and allows for supporting these ops for lowering
to targets without native deinterleave support. A transform dialect op
is also added to apply this pattern.
---------
Co-authored-by: Andrzej WarzyĆski <andrzej.warzynski at gmail.com>
Added:
mlir/test/Dialect/Vector/vector-interleave-deinterleave-to-shuffle.mlir
Modified:
mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
mlir/lib/Dialect/Vector/Transforms/LowerVectorInterleave.cpp
Removed:
mlir/test/Dialect/Vector/vector-interleave-to-shuffle.mlir
################################################################################
diff --git a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
index dcd5f6ff3ad74..333570d2348ce 100644
--- a/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
+++ b/mlir/include/mlir/Dialect/Vector/TransformOps/VectorTransformOps.td
@@ -418,15 +418,15 @@ def ApplyLowerInterleavePatternsOp : Op<Transform_Dialect,
let assemblyFormat = "attr-dict";
}
-def ApplyInterleaveToShufflePatternsOp : Op<Transform_Dialect,
- "apply_patterns.vector.interleave_to_shuffle",
+def ApplyInterleaveAndDeinterleaveToShufflePatternsOp : Op<Transform_Dialect,
+ "apply_patterns.vector.interleave_and_deinterleave_to_shuffle",
[DeclareOpInterfaceMethods<PatternDescriptorOpInterface>]> {
let description = [{
- Indicates that 1D vector interleave operations should be rewritten as
- vector shuffle operations.
+ Indicates that 1D vector interleave and deinterleave operations should be
+ rewritten as vector shuffle operations.
This is motivated by some current codegen backends not handling vector
- interleave operations.
+ interleave and deinterleave operations.
}];
let assemblyFormat = "attr-dict";
diff --git a/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h b/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
index aa75eff409ef9..d23a4d5c3f5fb 100644
--- a/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
+++ b/mlir/include/mlir/Dialect/Vector/Transforms/LoweringPatterns.h
@@ -290,6 +290,9 @@ void populateVectorInterleaveLoweringPatterns(RewritePatternSet &patterns,
void populateVectorInterleaveToShufflePatterns(RewritePatternSet &patterns,
PatternBenefit benefit = 1);
+void populateVectorDeinterleaveToShufflePatterns(RewritePatternSet &patterns,
+ PatternBenefit benefit = 1);
+
/// Populates the pattern set with the following patterns:
///
/// [UnrollBitCastOp]
diff --git a/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp b/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
index 312bd28ad48cf..8892a8e03ec8e 100644
--- a/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
+++ b/mlir/lib/Dialect/Vector/TransformOps/VectorTransformOps.cpp
@@ -207,9 +207,10 @@ void transform::ApplyLowerInterleavePatternsOp::populatePatterns(
vector::populateVectorInterleaveLoweringPatterns(patterns);
}
-void transform::ApplyInterleaveToShufflePatternsOp::populatePatterns(
- RewritePatternSet &patterns) {
+void transform::ApplyInterleaveAndDeinterleaveToShufflePatternsOp::
+ populatePatterns(RewritePatternSet &patterns) {
vector::populateVectorInterleaveToShufflePatterns(patterns);
+ vector::populateVectorDeinterleaveToShufflePatterns(patterns);
}
void transform::ApplyRewriteNarrowTypePatternsOp::populatePatterns(
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorInterleave.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorInterleave.cpp
index 13ad98de284e2..8f81dbb636325 100644
--- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorInterleave.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorInterleave.cpp
@@ -51,7 +51,7 @@ class UnrollInterleaveOp final : public OpRewritePattern<vector::InterleaveOp> {
public:
UnrollInterleaveOp(int64_t targetRank, MLIRContext *context,
PatternBenefit benefit = 1)
- : OpRewritePattern(context, benefit), targetRank(targetRank){};
+ : OpRewritePattern(context, benefit), targetRank(targetRank) {};
LogicalResult matchAndRewrite(vector::InterleaveOp op,
PatternRewriter &rewriter) const override {
@@ -148,8 +148,9 @@ class UnrollDeinterleaveOp final
private:
int64_t targetRank = 1;
};
+
/// Rewrite vector.interleave op into an equivalent vector.shuffle op, when
-/// applicable: `sourceType` must be 1D and non-scalable.
+/// applicable: `sourceType` must be 0D or 1D, and non-scalable.
///
/// Example:
///
@@ -169,7 +170,7 @@ struct InterleaveToShuffle final : OpRewritePattern<vector::InterleaveOp> {
LogicalResult matchAndRewrite(vector::InterleaveOp op,
PatternRewriter &rewriter) const override {
VectorType sourceType = op.getSourceVectorType();
- if (sourceType.getRank() != 1 || sourceType.isScalable()) {
+ if (sourceType.getRank() > 1 || sourceType.isScalable()) {
return failure();
}
int64_t n = sourceType.getNumElements();
@@ -181,6 +182,45 @@ struct InterleaveToShuffle final : OpRewritePattern<vector::InterleaveOp> {
}
};
+/// Rewrite vector.deinterleave op into two equivalent vector.shuffle ops, when
+/// applicable: `sourceType` must be 1D and non-scalable.
+///
+/// Example:
+///
+/// ```mlir
+/// %evens, %odds = vector.deinterleave %arg0 : vector<4xi32> -> vector<2xi32>
+/// ```
+///
+/// Is rewritten into:
+///
+/// ```mlir
+/// %evens = vector.shuffle %arg0, %arg0 [0, 2] : vector<4xi32>, vector<4xi32>
+/// %odds = vector.shuffle %arg0, %arg0 [1, 3] : vector<4xi32>, vector<4xi32>
+/// ```
+struct DeinterleaveToShuffle final : OpRewritePattern<vector::DeinterleaveOp> {
+ using Base::Base;
+
+ LogicalResult matchAndRewrite(vector::DeinterleaveOp op,
+ PatternRewriter &rewriter) const override {
+ VectorType sourceType = op.getSourceVectorType();
+ if (sourceType.getRank() != 1 || sourceType.isScalable()) {
+ return failure();
+ }
+
+ auto seq = llvm::seq<int64_t>(sourceType.getNumElements() / 2);
+ auto evenZip = llvm::map_to_vector(seq, [](int64_t i) { return i * 2; });
+ auto oddZip = llvm::map_to_vector(evenZip, [](int64_t i) { return i + 1; });
+
+ Value evenResult = vector::ShuffleOp::create(
+ rewriter, op.getLoc(), op.getOperand(), op.getOperand(), evenZip);
+ Value oddResult = vector::ShuffleOp::create(
+ rewriter, op.getLoc(), op.getOperand(), op.getOperand(), oddZip);
+
+ rewriter.replaceOp(op, ValueRange{evenResult, oddResult});
+ return success();
+ }
+};
+
} // namespace
void mlir::vector::populateVectorInterleaveLoweringPatterns(
@@ -193,3 +233,8 @@ void mlir::vector::populateVectorInterleaveToShufflePatterns(
RewritePatternSet &patterns, PatternBenefit benefit) {
patterns.add<InterleaveToShuffle>(patterns.getContext(), benefit);
}
+
+void mlir::vector::populateVectorDeinterleaveToShufflePatterns(
+ RewritePatternSet &patterns, PatternBenefit benefit) {
+ patterns.add<DeinterleaveToShuffle>(patterns.getContext(), benefit);
+}
diff --git a/mlir/test/Dialect/Vector/vector-interleave-deinterleave-to-shuffle.mlir b/mlir/test/Dialect/Vector/vector-interleave-deinterleave-to-shuffle.mlir
new file mode 100644
index 0000000000000..13d1af28f9eb8
--- /dev/null
+++ b/mlir/test/Dialect/Vector/vector-interleave-deinterleave-to-shuffle.mlir
@@ -0,0 +1,60 @@
+// RUN: mlir-opt %s --transform-interpreter | FileCheck %s
+
+// CHECK-LABEL: @vector_interleave_to_shuffle_1d
+func.func @vector_interleave_to_shuffle_1d(%a: vector<7xi16>, %b: vector<7xi16>) -> vector<14xi16> {
+ %0 = vector.interleave %a, %b : vector<7xi16> -> vector<14xi16>
+ return %0 : vector<14xi16>
+}
+// CHECK: vector.shuffle %arg0, %arg1 [0, 7, 1, 8, 2, 9, 3, 10, 4, 11, 5, 12, 6, 13] : vector<7xi16>, vector<7xi16>
+
+// CHECK-LABEL: @vector_interleave_to_shuffle_0d
+func.func @vector_interleave_to_shuffle_0d(%a: vector<f32>, %b: vector<f32>) -> vector<2xf32> {
+ %0 = vector.interleave %a, %b : vector<f32> -> vector<2xf32>
+ return %0 : vector<2xf32>
+}
+// CHECK: vector.shuffle %arg0, %arg1 [0, 1] : vector<f32>, vector<f32>
+
+// CHECK-LABEL: @vector_deinterleave_to_shuffle_1d
+func.func @vector_deinterleave_to_shuffle_1d(%arg0: vector<14xi16>) -> (vector<7xi16>, vector<7xi16>) {
+ %evens, %odds = vector.deinterleave %arg0 : vector<14xi16> -> vector<7xi16>
+ return %evens, %odds : vector<7xi16>, vector<7xi16>
+}
+// CHECK: vector.shuffle %arg0, %arg0 [0, 2, 4, 6, 8, 10, 12] : vector<14xi16>, vector<14xi16>
+// CHECK: vector.shuffle %arg0, %arg0 [1, 3, 5, 7, 9, 11, 13] : vector<14xi16>, vector<14xi16>
+
+// CHECK-LABEL: @vector_deinterleave_size2
+func.func @vector_deinterleave_size2(%arg0: vector<2xi32>) -> (vector<1xi32>, vector<1xi32>) {
+ %evens, %odds = vector.deinterleave %arg0 : vector<2xi32> -> vector<1xi32>
+ return %evens, %odds : vector<1xi32>, vector<1xi32>
+}
+// CHECK: vector.shuffle %arg0, %arg0 [0] : vector<2xi32>, vector<2xi32>
+// CHECK: vector.shuffle %arg0, %arg0 [1] : vector<2xi32>, vector<2xi32>
+
+// CHECK-LABEL: @negative_cases
+// CHECK-NOT: vector.shuffle
+func.func @negative_cases(
+ %a: vector<[4]xi32>, %b: vector<[4]xi32>,
+ %c: vector<2x4xi32>, %d: vector<2x4xi32>,
+ %e: vector<[8]xi32>,
+ %f: vector<2x8xi32>) -> (vector<[8]xi32>, vector<2x8xi32>,
+ vector<[4]xi32>, vector<[4]xi32>,
+ vector<2x4xi32>, vector<2x4xi32>) {
+ %0 = vector.interleave %a, %b : vector<[4]xi32> -> vector<[8]xi32>
+ %1 = vector.interleave %c, %d : vector<2x4xi32> -> vector<2x8xi32>
+ %evens0, %odds0 = vector.deinterleave %e : vector<[8]xi32> -> vector<[4]xi32>
+ %evens1, %odds1 = vector.deinterleave %f : vector<2x8xi32> -> vector<2x4xi32>
+ return %0, %1, %evens0, %odds0, %evens1, %odds1
+ : vector<[8]xi32>, vector<2x8xi32>, vector<[4]xi32>, vector<[4]xi32>, vector<2x4xi32>, vector<2x4xi32>
+}
+
+module attributes {transform.with_named_sequence} {
+ transform.named_sequence @__transform_main(%module_op: !transform.any_op {transform.readonly}) {
+ %f = transform.structured.match ops{["func.func"]} in %module_op
+ : (!transform.any_op) -> !transform.any_op
+
+ transform.apply_patterns to %f {
+ transform.apply_patterns.vector.interleave_and_deinterleave_to_shuffle
+ } : !transform.any_op
+ transform.yield
+ }
+}
diff --git a/mlir/test/Dialect/Vector/vector-interleave-to-shuffle.mlir b/mlir/test/Dialect/Vector/vector-interleave-to-shuffle.mlir
deleted file mode 100644
index d59cd4e6765ba..0000000000000
--- a/mlir/test/Dialect/Vector/vector-interleave-to-shuffle.mlir
+++ /dev/null
@@ -1,20 +0,0 @@
-// RUN: mlir-opt %s --transform-interpreter | FileCheck %s
-
-// CHECK-LABEL: @vector_interleave_to_shuffle
-func.func @vector_interleave_to_shuffle(%a: vector<7xi16>, %b: vector<7xi16>) -> vector<14xi16> {
- %0 = vector.interleave %a, %b : vector<7xi16> -> vector<14xi16>
- return %0 : vector<14xi16>
-}
-// CHECK: vector.shuffle %arg0, %arg1 [0, 7, 1, 8, 2, 9, 3, 10, 4, 11, 5, 12, 6, 13] : vector<7xi16>, vector<7xi16>
-
-module attributes {transform.with_named_sequence} {
- transform.named_sequence @__transform_main(%module_op: !transform.any_op {transform.readonly}) {
- %f = transform.structured.match ops{["func.func"]} in %module_op
- : (!transform.any_op) -> !transform.any_op
-
- transform.apply_patterns to %f {
- transform.apply_patterns.vector.interleave_to_shuffle
- } : !transform.any_op
- transform.yield
- }
-}
More information about the Mlir-commits
mailing list