Skip to content

Commit d0dd371

Browse files
authored
[MLIR][Canonicalization] Add shape_cast folding patterns (llvm#183061)
### Summary This PR adds two shape_cast-related canonicalization patterns for `vector.to_elements` and `vector.from_elements`. ### Details - Added` ToElements(ShapeCast(X)) -> ToElements(X)` as an in-place fold in `ToElementsOp::fold`. - Added `ShapeCast(FromElements(X)) -> FromElements(X)` as an `OpRewritePattern` — it must be a pattern (not a `fold`) because we have to create new op `FromElementsOp` with updated result type. This cannot be done with a `fold`, because `fold` cannot create new ops and the existing `FromElementsOp` result type differs from the `ShapeCastOp` result type. Mutating the `FromElementsOp` (not root op) would violate the `fold` contract and break other users. - Added lit tests for the both ops (new `vector-to-elements.mlir`, `vector-from-elements.mlir`) --------- Signed-off-by: Alexandra Sidorova <asidorov@amd.com>
1 parent 6b040b0 commit d0dd371

3 files changed

Lines changed: 71 additions & 3 deletions

File tree

‎mlir/lib/Dialect/Vector/IR/VectorOps.cpp‎

Lines changed: 40 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2559,6 +2559,13 @@ LogicalResult ToElementsOp::fold(FoldAdaptor adaptor,
25592559
SmallVectorImpl<OpFoldResult> &results) {
25602560
if (succeeded(foldToElementsFromElements(*this, results)))
25612561
return success();
2562+
2563+
// Y = ToElements(ShapeCast(X)) -> Y = ToElements(X)
2564+
if (auto shapeCast = getSource().getDefiningOp<ShapeCastOp>()) {
2565+
setOperand(shapeCast.getSource());
2566+
return success();
2567+
}
2568+
25622569
return foldToElementsOfBroadcast(*this, results);
25632570
}
25642571

@@ -6723,13 +6730,43 @@ class ShapeCastBroadcastFolder final : public OpRewritePattern<ShapeCastOp> {
67236730
}
67246731
};
67256732

6733+
/// Pattern to rewrite Y = ShapeCast(FromElements(X)) as Y = FromElements(X)
6734+
///
6735+
/// BEFORE:
6736+
/// %1 = vector.from_elements %c1, %c2, %c3 : vector<3xf32>
6737+
/// %2 = vector.shape_cast %1 : vector<3xf32> to vector<1x3xf32>
6738+
/// AFTER:
6739+
/// %2 = vector.from_elements %c1, %c2, %c3 : vector<1x3xf32>
6740+
///
6741+
/// Note: this transformation is implemented as an OpRewritePattern, not as a
6742+
/// fold, because we have to create new op FromElementsOp with updated result
6743+
/// type. This cannot be done with a fold, because fold cannot create new ops
6744+
/// and the existing FromElementsOp result type differs from the ShapeCastOp
6745+
/// result type. Mutating the FromElementsOp (not root op) would violate the
6746+
/// fold contract and break other users.
6747+
class FoldShapeCastOfFromElements final : public OpRewritePattern<ShapeCastOp> {
6748+
public:
6749+
using Base::Base;
6750+
6751+
LogicalResult matchAndRewrite(ShapeCastOp shapeCastOp,
6752+
PatternRewriter &rewriter) const override {
6753+
auto fromElements = shapeCastOp.getSource().getDefiningOp<FromElementsOp>();
6754+
if (!fromElements)
6755+
return failure();
6756+
6757+
rewriter.replaceOpWithNewOp<FromElementsOp>(
6758+
shapeCastOp, shapeCastOp.getResultVectorType(),
6759+
fromElements.getElements());
6760+
return success();
6761+
}
6762+
};
6763+
67266764
} // namespace
67276765

67286766
void ShapeCastOp::getCanonicalizationPatterns(RewritePatternSet &results,
67296767
MLIRContext *context) {
6730-
results
6731-
.add<ShapeCastCreateMaskFolderTrailingOneDim, ShapeCastBroadcastFolder>(
6732-
context);
6768+
results.add<ShapeCastCreateMaskFolderTrailingOneDim, ShapeCastBroadcastFolder,
6769+
FoldShapeCastOfFromElements>(context);
67336770
}
67346771

67356772
//===----------------------------------------------------------------------===//

‎mlir/test/Dialect/Vector/canonicalize.mlir‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1170,6 +1170,19 @@ func.func @dont_fold_expand_collapse(%arg0: vector<1x1x64xf32>) -> vector<8x8xf3
11701170

11711171
// -----
11721172

1173+
// CHECK-LABEL: func.func @fold_shape_cast_from_elements(
1174+
// CHECK-SAME: %[[C1:.*]]: f32, %[[C2:.*]]: f32, %[[C3:.*]]: f32, %[[C4:.*]]: f32
1175+
func.func @fold_shape_cast_from_elements(%c1: f32, %c2: f32, %c3: f32, %c4: f32) -> vector<2x2xf32>{
1176+
// CHECK: %[[VAL_0:.*]] = vector.from_elements %[[C1]], %[[C2]], %[[C3]], %[[C4]] : vector<2x2xf32>
1177+
// CHECK: return %[[VAL_0]] : vector<2x2xf32>
1178+
// CHECK-NOT: vector.shape_cast
1179+
%1 = vector.from_elements %c1, %c2, %c3, %c4 : vector<4xf32>
1180+
%2 = vector.shape_cast %1 : vector<4xf32> to vector<2x2xf32>
1181+
return %2 : vector<2x2xf32>
1182+
}
1183+
1184+
// -----
1185+
11731186
// CHECK-LABEL: func @fold_broadcast_shapecast
11741187
// CHECK-SAME: (%[[V:.+]]: vector<4xf32>)
11751188
// CHECK: return %[[V]]
Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,18 @@
1+
// RUN: mlir-opt %s -canonicalize="test-convergence" -split-input-file -allow-unregistered-dialect | FileCheck %s
2+
3+
// This file contains some tests of folding/canonicalizing vector.to_elements
4+
5+
///===----------------------------------------------===//
6+
/// Tests of `ToElementsOp::fold`
7+
///===----------------------------------------------===//
8+
9+
// CHECK-LABEL: func @to_elements_of_shape_cast_folds
10+
// CHECK-SAME: (%[[VEC:.*]]: vector<4xf32>) -> (f32, f32, f32, f32)
11+
func.func @to_elements_of_shape_cast_folds(%v: vector<4xf32>) -> (f32, f32, f32, f32) {
12+
%sc = vector.shape_cast %v : vector<4xf32> to vector<2x2xf32>
13+
%e:4 = vector.to_elements %sc : vector<2x2xf32>
14+
// CHECK-NOT: vector.shape_cast
15+
// CHECK: %[[E:.*]]:4 = vector.to_elements %[[VEC]] : vector<4xf32>
16+
// CHECK: return %[[E]]#0, %[[E]]#1, %[[E]]#2, %[[E]]#3 : f32, f32, f32, f32
17+
return %e#0, %e#1, %e#2, %e#3 : f32, f32, f32, f32
18+
}

0 commit comments

Comments
 (0)