diff options
| author | Benoit Jacob <jacob.benoit.1@gmail.com> | 2024-10-08 11:51:01 -0400 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-10-08 11:51:01 -0400 |
| commit | 10054ba4acbc5378d2e2aa869a5bccd88aa4b59e (patch) | |
| tree | 44ca367f4b03b011142f03889a82e978465e73d8 /mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp | |
| parent | d079743fe67e05697fe55409115a3614e6fe5c45 (diff) | |
[mlir][vector] Add pattern to rewrite contiguous ExtractStridedSlice into Extract (#111541)
Co-authored-by: Jakub Kuderski <kubakuderski@gmail.com>
Diffstat (limited to 'mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp')
| -rw-r--r-- | mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp | 58 |
1 files changed, 58 insertions, 0 deletions
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp index ec2ef3fc7501..c2da9347aadc 100644 --- a/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp +++ b/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp @@ -329,12 +329,70 @@ public: } }; +/// Pattern to rewrite simple cases of N-D extract_strided_slice, where the +/// slice is contiguous, into extract and shape_cast. +class ContiguousExtractStridedSliceToExtract final + : public OpRewritePattern<ExtractStridedSliceOp> { +public: + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(ExtractStridedSliceOp op, + PatternRewriter &rewriter) const override { + if (op.hasNonUnitStrides()) { + return failure(); + } + Value source = op.getOperand(); + auto sourceType = cast<VectorType>(source.getType()); + if (sourceType.isScalable()) { + return failure(); + } + + // Compute the number of offsets to pass to ExtractOp::build. That is the + // difference between the source rank and the desired slice rank. We walk + // the dimensions from innermost out, and stop when the next slice dimension + // is not full-size. + SmallVector<int64_t> sizes = getI64SubArray(op.getSizes()); + int numOffsets; + for (numOffsets = sourceType.getRank(); numOffsets > 0; --numOffsets) { + if (sizes[numOffsets - 1] != sourceType.getDimSize(numOffsets - 1)) { + break; + } + } + + // If not even the inner-most dimension is full-size, this op can't be + // rewritten as an ExtractOp. + if (numOffsets == sourceType.getRank()) { + return failure(); + } + + // Avoid generating slices that have unit outer dimensions. The shape_cast + // op that we create below would take bad generic fallback patterns + // (ShapeCastOpRewritePattern). + while (sizes[numOffsets] == 1 && numOffsets < sourceType.getRank() - 1) { + ++numOffsets; + } + + SmallVector<int64_t> offsets = getI64SubArray(op.getOffsets()); + auto extractOffsets = ArrayRef(offsets).take_front(numOffsets); + Value extract = rewriter.create<vector::ExtractOp>(op->getLoc(), source, + extractOffsets); + rewriter.replaceOpWithNewOp<vector::ShapeCastOp>(op, op.getType(), extract); + return success(); + } +}; + void vector::populateVectorInsertExtractStridedSliceDecompositionPatterns( RewritePatternSet &patterns, PatternBenefit benefit) { patterns.add<DecomposeDifferentRankInsertStridedSlice, DecomposeNDExtractStridedSlice>(patterns.getContext(), benefit); } +void vector::populateVectorContiguousExtractStridedSliceToExtractPatterns( + RewritePatternSet &patterns, PatternBenefit benefit) { + patterns.add<ContiguousExtractStridedSliceToExtract>(patterns.getContext(), + benefit); +} + void vector::populateVectorExtractStridedSliceToExtractInsertChainPatterns( RewritePatternSet &patterns, std::function<bool(ExtractStridedSliceOp)> controlFn, |
