summaryrefslogtreecommitdiff
path: root/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp
diff options
context:
space:
mode:
authorLei Zhang <antiagainst@google.com>2022-11-09 19:37:19 -0500
committerLei Zhang <antiagainst@google.com>2022-11-09 19:42:07 -0500
commit39c80656fef68fcaf9707857bf67f643378e6cc8 (patch)
tree2c41a883b5454f9ac1e3122c992076c46bd42d83 /mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp
parent0e520300580a77f1d7c01ada9a047a7fadb5eb1f (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.cpp60
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);