diff options
| author | Alex Zinenko <zinenko@google.com> | 2022-09-16 16:16:27 +0200 |
|---|---|---|
| committer | Alex Zinenko <zinenko@google.com> | 2022-09-17 08:11:30 +0200 |
| commit | 83df43f3a204c529167caeedc26edae5faebbd31 (patch) | |
| tree | a63e88a15f8a77b0e6bf9e5d4554b37d3c3a8e9a /mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp | |
| parent | f3fae035c7a16e1f4c7d96b115212714375a3d38 (diff) | |
[mlir] use strided layouts in vector transfer on memrefs
One of the vector transformation patterns has been indiscriminately
converting layouts to affine maps. Leverage the strided form when
possible.
Reviewed By: nicolasvasilache, dcaballe
Differential Revision: https://reviews.llvm.org/D134047
Diffstat (limited to 'mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp')
| -rw-r--r-- | mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp | 28 |
1 files changed, 18 insertions, 10 deletions
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp index fded10629b30..18f6c5a154e5 100644 --- a/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp +++ b/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp @@ -24,6 +24,7 @@ #include "mlir/Dialect/Utils/StructuredOpsUtils.h" #include "mlir/Dialect/Vector/IR/VectorOps.h" #include "mlir/Dialect/Vector/Utils/VectorUtils.h" +#include "mlir/IR/BuiltinAttributeInterfaces.h" #include "mlir/IR/BuiltinTypes.h" #include "mlir/IR/ImplicitLocOpBuilder.h" #include "mlir/IR/Matchers.h" @@ -2631,22 +2632,29 @@ class DropInnerMostUnitDims : public OpRewritePattern<vector::TransferReadOp> { targetType.getElementType()); MemRefType resultMemrefType; - if (srcType.getLayout().getAffineMap().isIdentity()) { + MemRefLayoutAttrInterface layout = srcType.getLayout(); + if (layout.isa<AffineMapAttr>() && layout.isIdentity()) { resultMemrefType = MemRefType::get( srcType.getShape().drop_back(dimsToDrop), srcType.getElementType(), - {}, srcType.getMemorySpaceAsInt()); + nullptr, srcType.getMemorySpace()); } else { - AffineMap map = srcType.getLayout().getAffineMap(); - int numSymbols = map.getNumSymbols(); - for (size_t i = 0; i < dimsToDrop; ++i) { - int dim = srcType.getRank() - i - 1; - map = map.replace(rewriter.getAffineDimExpr(dim), - rewriter.getAffineConstantExpr(0), - map.getNumDims() - 1, numSymbols); + MemRefLayoutAttrInterface updatedLayout; + if (auto strided = layout.dyn_cast<StridedLayoutAttr>()) { + auto strides = llvm::to_vector(strided.getStrides().drop_back(dimsToDrop)); + updatedLayout = StridedLayoutAttr::get(strided.getContext(), strided.getOffset(), strides); + } else { + AffineMap map = srcType.getLayout().getAffineMap(); + int numSymbols = map.getNumSymbols(); + for (size_t i = 0; i < dimsToDrop; ++i) { + int dim = srcType.getRank() - i - 1; + map = map.replace(rewriter.getAffineDimExpr(dim), + rewriter.getAffineConstantExpr(0), + map.getNumDims() - 1, numSymbols); + } } resultMemrefType = MemRefType::get( srcType.getShape().drop_back(dimsToDrop), srcType.getElementType(), - map, srcType.getMemorySpaceAsInt()); + updatedLayout, srcType.getMemorySpace()); } auto loc = readOp.getLoc(); |
