diff options
| author | Lei Zhang <antiagainst@google.com> | 2022-11-09 19:37:19 -0500 |
|---|---|---|
| committer | Lei Zhang <antiagainst@google.com> | 2022-11-09 19:42:07 -0500 |
| commit | 39c80656fef68fcaf9707857bf67f643378e6cc8 (patch) | |
| tree | 2c41a883b5454f9ac1e3122c992076c46bd42d83 /mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp | |
| parent | 0e520300580a77f1d7c01ada9a047a7fadb5eb1f (diff) | |
[mlir][vector] Convert extract_strided_slice to extract & insert chain
This is useful for breaking down extract_strided_slice and potentially
cancel with other extract / insert ops before or after.
Reviewed By: ThomasRaoux
Differential Revision: https://reviews.llvm.org/D137471
Diffstat (limited to 'mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp')
| -rw-r--r-- | mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp | 60 |
1 files changed, 58 insertions, 2 deletions
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp index ad6cf85b6253..313a3f9a9c09 100644 --- a/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp +++ b/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp @@ -13,6 +13,7 @@ #include "mlir/Dialect/Vector/Transforms/VectorRewritePatterns.h" #include "mlir/Dialect/Vector/Utils/VectorUtils.h" #include "mlir/IR/BuiltinTypes.h" +#include "mlir/IR/PatternMatch.h" using namespace mlir; using namespace mlir::vector; @@ -231,6 +232,53 @@ public: } }; +/// For a 1-D ExtractStridedSlice, breaks it down into a chain of Extract ops +/// to extract each element from the source, and then a chain of Insert ops +/// to insert to the target vector. +class Convert1DExtractStridedSliceIntoExtractInsertChain final + : public OpRewritePattern<ExtractStridedSliceOp> { +public: + Convert1DExtractStridedSliceIntoExtractInsertChain( + MLIRContext *context, + std::function<bool(ExtractStridedSliceOp)> controlFn, + PatternBenefit benefit) + : OpRewritePattern(context, benefit), controlFn(std::move(controlFn)) {} + + LogicalResult matchAndRewrite(ExtractStridedSliceOp op, + PatternRewriter &rewriter) const override { + if (controlFn && !controlFn(op)) + return failure(); + + // Only handle 1-D cases. + if (op.getOffsets().getValue().size() != 1) + return failure(); + + int64_t offset = + op.getOffsets().getValue().front().cast<IntegerAttr>().getInt(); + int64_t size = + op.getSizes().getValue().front().cast<IntegerAttr>().getInt(); + int64_t stride = + op.getStrides().getValue().front().cast<IntegerAttr>().getInt(); + + Location loc = op.getLoc(); + SmallVector<Value> elements; + elements.reserve(size); + for (int64_t i = offset, e = offset + size * stride; i < e; i += stride) + elements.push_back(rewriter.create<ExtractOp>(loc, op.getVector(), i)); + + Value result = rewriter.create<arith::ConstantOp>( + loc, rewriter.getZeroAttr(op.getType())); + for (int64_t i = 0; i < size; ++i) + result = rewriter.create<InsertOp>(loc, elements[i], result, i); + + rewriter.replaceOp(op, result); + return success(); + } + +private: + std::function<bool(ExtractStridedSliceOp)> controlFn; +}; + /// RewritePattern for ExtractStridedSliceOp where the source vector is n-D. /// For such cases, we can rewrite it to ExtractOp/ExtractElementOp + lower /// rank ExtractStridedSliceOp + InsertOp/InsertElementOp for the n-D case. @@ -285,14 +333,22 @@ public: } }; -void mlir::vector::populateVectorInsertExtractStridedSliceDecompositionPatterns( +void vector::populateVectorInsertExtractStridedSliceDecompositionPatterns( RewritePatternSet &patterns, PatternBenefit benefit) { patterns.add<DecomposeDifferentRankInsertStridedSlice, DecomposeNDExtractStridedSlice>(patterns.getContext(), benefit); } +void vector::populateVectorExtractStridedSliceToExtractInsertChainPatterns( + RewritePatternSet &patterns, + std::function<bool(ExtractStridedSliceOp)> controlFn, + PatternBenefit benefit) { + patterns.add<Convert1DExtractStridedSliceIntoExtractInsertChain>( + patterns.getContext(), std::move(controlFn), benefit); +} + /// Populate the given list with patterns that convert from Vector to LLVM. -void mlir::vector::populateVectorInsertExtractStridedSliceTransforms( +void vector::populateVectorInsertExtractStridedSliceTransforms( RewritePatternSet &patterns, PatternBenefit benefit) { populateVectorInsertExtractStridedSliceDecompositionPatterns(patterns, benefit); |
