summaryrefslogtreecommitdiff
path: root/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp
diff options
context:
space:
mode:
authorJacques Pienaar <jpienaar@google.com>2022-03-28 11:24:47 -0700
committerJacques Pienaar <jpienaar@google.com>2022-03-28 11:24:47 -0700
commit7c38fd605ba85657a0ecbea75a8e3a68174d3dff (patch)
tree420972f033748b603360f29cd414847fcaa3bdbd /mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp
parent1066e397fa907629f0da370f9721821c838ed30a (diff)
[mlir] Flip Vector dialect accessors used to prefixed form.
This has been on _Both for a couple of weeks. Flip usages in core with intention to flip flag to _Prefixed in follow up. Needed to add a couple of helper methods in AffineOps and Linalg to facilitate a pure flag flip in follow up as some of these classes are used in templates and so sensitive to Vector dialect changes. Differential Revision: https://reviews.llvm.org/D122151
Diffstat (limited to 'mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp')
-rw-r--r--mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp67
1 files changed, 34 insertions, 33 deletions
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp
index 4308fa6a43be..2a384c3bf785 100644
--- a/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp
@@ -62,7 +62,7 @@ public:
auto srcType = op.getSourceVectorType();
auto dstType = op.getDestVectorType();
- if (op.offsets().getValue().empty())
+ if (op.getOffsets().getValue().empty())
return failure();
auto loc = op.getLoc();
@@ -74,21 +74,21 @@ public:
int64_t rankRest = dstType.getRank() - rankDiff;
// Extract / insert the subvector of matching rank and InsertStridedSlice
// on it.
- Value extracted =
- rewriter.create<ExtractOp>(loc, op.dest(),
- getI64SubArray(op.offsets(), /*dropFront=*/0,
- /*dropBack=*/rankRest));
+ Value extracted = rewriter.create<ExtractOp>(
+ loc, op.getDest(),
+ getI64SubArray(op.getOffsets(), /*dropFront=*/0,
+ /*dropBack=*/rankRest));
// A different pattern will kick in for InsertStridedSlice with matching
// ranks.
auto stridedSliceInnerOp = rewriter.create<InsertStridedSliceOp>(
- loc, op.source(), extracted,
- getI64SubArray(op.offsets(), /*dropFront=*/rankDiff),
- getI64SubArray(op.strides(), /*dropFront=*/0));
+ loc, op.getSource(), extracted,
+ getI64SubArray(op.getOffsets(), /*dropFront=*/rankDiff),
+ getI64SubArray(op.getStrides(), /*dropFront=*/0));
rewriter.replaceOpWithNewOp<InsertOp>(
- op, stridedSliceInnerOp.getResult(), op.dest(),
- getI64SubArray(op.offsets(), /*dropFront=*/0,
+ op, stridedSliceInnerOp.getResult(), op.getDest(),
+ getI64SubArray(op.getOffsets(), /*dropFront=*/0,
/*dropBack=*/rankRest));
return success();
}
@@ -118,7 +118,7 @@ public:
auto srcType = op.getSourceVectorType();
auto dstType = op.getDestVectorType();
- if (op.offsets().getValue().empty())
+ if (op.getOffsets().getValue().empty())
return failure();
int64_t srcRank = srcType.getRank();
@@ -128,18 +128,18 @@ public:
return failure();
if (srcType == dstType) {
- rewriter.replaceOp(op, op.source());
+ rewriter.replaceOp(op, op.getSource());
return success();
}
int64_t offset =
- op.offsets().getValue().front().cast<IntegerAttr>().getInt();
+ op.getOffsets().getValue().front().cast<IntegerAttr>().getInt();
int64_t size = srcType.getShape().front();
int64_t stride =
- op.strides().getValue().front().cast<IntegerAttr>().getInt();
+ op.getStrides().getValue().front().cast<IntegerAttr>().getInt();
auto loc = op.getLoc();
- Value res = op.dest();
+ Value res = op.getDest();
if (srcRank == 1) {
int nSrc = srcType.getShape().front();
@@ -148,8 +148,8 @@ public:
SmallVector<int64_t> offsets(nDest, 0);
for (int64_t i = 0; i < nSrc; ++i)
offsets[i] = i;
- Value scaledSource =
- rewriter.create<ShuffleOp>(loc, op.source(), op.source(), offsets);
+ Value scaledSource = rewriter.create<ShuffleOp>(loc, op.getSource(),
+ op.getSource(), offsets);
// 2. Create a mask where we take the value from scaledSource of dest
// depending on the offset.
@@ -162,7 +162,7 @@ public:
}
// 3. Replace with a ShuffleOp.
- rewriter.replaceOpWithNewOp<ShuffleOp>(op, scaledSource, op.dest(),
+ rewriter.replaceOpWithNewOp<ShuffleOp>(op, scaledSource, op.getDest(),
offsets);
return success();
@@ -172,17 +172,17 @@ public:
for (int64_t off = offset, e = offset + size * stride, idx = 0; off < e;
off += stride, ++idx) {
// 1. extract the proper subvector (or element) from source
- Value extractedSource = extractOne(rewriter, loc, op.source(), idx);
+ Value extractedSource = extractOne(rewriter, loc, op.getSource(), idx);
if (extractedSource.getType().isa<VectorType>()) {
// 2. If we have a vector, extract the proper subvector from destination
// Otherwise we are at the element level and no need to recurse.
- Value extractedDest = extractOne(rewriter, loc, op.dest(), off);
+ Value extractedDest = extractOne(rewriter, loc, op.getDest(), off);
// 3. Reduce the problem to lowering a new InsertStridedSlice op with
// smaller rank.
extractedSource = rewriter.create<InsertStridedSliceOp>(
loc, extractedSource, extractedDest,
- getI64SubArray(op.offsets(), /* dropFront=*/1),
- getI64SubArray(op.strides(), /* dropFront=*/1));
+ getI64SubArray(op.getOffsets(), /* dropFront=*/1),
+ getI64SubArray(op.getStrides(), /* dropFront=*/1));
}
// 4. Insert the extractedSource into the res vector.
res = insertOne(rewriter, loc, extractedSource, res, off);
@@ -212,27 +212,28 @@ public:
PatternRewriter &rewriter) const override {
auto dstType = op.getType();
- assert(!op.offsets().getValue().empty() && "Unexpected empty offsets");
+ assert(!op.getOffsets().getValue().empty() && "Unexpected empty offsets");
int64_t offset =
- op.offsets().getValue().front().cast<IntegerAttr>().getInt();
- int64_t size = op.sizes().getValue().front().cast<IntegerAttr>().getInt();
+ op.getOffsets().getValue().front().cast<IntegerAttr>().getInt();
+ int64_t size =
+ op.getSizes().getValue().front().cast<IntegerAttr>().getInt();
int64_t stride =
- op.strides().getValue().front().cast<IntegerAttr>().getInt();
+ op.getStrides().getValue().front().cast<IntegerAttr>().getInt();
auto loc = op.getLoc();
auto elemType = dstType.getElementType();
assert(elemType.isSignlessIntOrIndexOrFloat());
// Single offset can be more efficiently shuffled.
- if (op.offsets().getValue().size() == 1) {
+ if (op.getOffsets().getValue().size() == 1) {
SmallVector<int64_t, 4> offsets;
offsets.reserve(size);
for (int64_t off = offset, e = offset + size * stride; off < e;
off += stride)
offsets.push_back(off);
- rewriter.replaceOpWithNewOp<ShuffleOp>(op, dstType, op.vector(),
- op.vector(),
+ rewriter.replaceOpWithNewOp<ShuffleOp>(op, dstType, op.getVector(),
+ op.getVector(),
rewriter.getI64ArrayAttr(offsets));
return success();
}
@@ -243,11 +244,11 @@ public:
Value res = rewriter.create<SplatOp>(loc, dstType, zero);
for (int64_t off = offset, e = offset + size * stride, idx = 0; off < e;
off += stride, ++idx) {
- Value one = extractOne(rewriter, loc, op.vector(), off);
+ Value one = extractOne(rewriter, loc, op.getVector(), off);
Value extracted = rewriter.create<ExtractStridedSliceOp>(
- loc, one, getI64SubArray(op.offsets(), /* dropFront=*/1),
- getI64SubArray(op.sizes(), /* dropFront=*/1),
- getI64SubArray(op.strides(), /* dropFront=*/1));
+ loc, one, getI64SubArray(op.getOffsets(), /* dropFront=*/1),
+ getI64SubArray(op.getSizes(), /* dropFront=*/1),
+ getI64SubArray(op.getStrides(), /* dropFront=*/1));
res = insertOne(rewriter, loc, extracted, res, idx);
}
rewriter.replaceOp(op, res);