summaryrefslogtreecommitdiff
path: root/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp
diff options
context:
space:
mode:
authorAlex Zinenko <zinenko@google.com>2022-09-16 16:16:27 +0200
committerAlex Zinenko <zinenko@google.com>2022-09-17 08:11:30 +0200
commit83df43f3a204c529167caeedc26edae5faebbd31 (patch)
treea63e88a15f8a77b0e6bf9e5d4554b37d3c3a8e9a /mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp
parentf3fae035c7a16e1f4c7d96b115212714375a3d38 (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.cpp28
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();