diff options
| author | Jacques Pienaar <jpienaar@google.com> | 2022-03-28 11:24:47 -0700 |
|---|---|---|
| committer | Jacques Pienaar <jpienaar@google.com> | 2022-03-28 11:24:47 -0700 |
| commit | 7c38fd605ba85657a0ecbea75a8e3a68174d3dff (patch) | |
| tree | 420972f033748b603360f29cd414847fcaa3bdbd /mlir/lib/Dialect/Vector/Transforms/VectorInsertExtractStridedSliceRewritePatterns.cpp | |
| parent | 1066e397fa907629f0da370f9721821c838ed30a (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.cpp | 67 |
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); |
