From 10054ba4acbc5378d2e2aa869a5bccd88aa4b59e Mon Sep 17 00:00:00 2001 From: Benoit Jacob Date: Tue, 8 Oct 2024 11:51:01 -0400 Subject: [mlir][vector] Add pattern to rewrite contiguous ExtractStridedSlice into Extract (#111541) Co-authored-by: Jakub Kuderski --- ...torInsertExtractStridedSliceRewritePatterns.cpp | 58 ++++++++++++++++++++++ 1 file changed, 58 insertions(+) (limited to 'mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp') 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 { +public: + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(ExtractStridedSliceOp op, + PatternRewriter &rewriter) const override { + if (op.hasNonUnitStrides()) { + return failure(); + } + Value source = op.getOperand(); + auto sourceType = cast(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 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 offsets = getI64SubArray(op.getOffsets()); + auto extractOffsets = ArrayRef(offsets).take_front(numOffsets); + Value extract = rewriter.create(op->getLoc(), source, + extractOffsets); + rewriter.replaceOpWithNewOp(op, op.getType(), extract); + return success(); + } +}; + void vector::populateVectorInsertExtractStridedSliceDecompositionPatterns( RewritePatternSet &patterns, PatternBenefit benefit) { patterns.add(patterns.getContext(), benefit); } +void vector::populateVectorContiguousExtractStridedSliceToExtractPatterns( + RewritePatternSet &patterns, PatternBenefit benefit) { + patterns.add(patterns.getContext(), + benefit); +} + void vector::populateVectorExtractStridedSliceToExtractInsertChainPatterns( RewritePatternSet &patterns, std::function controlFn, -- cgit v1.2.3