summaryrefslogtreecommitdiff
path: root/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp
diff options
context:
space:
mode:
authorBenjamin Maxwell <benjamin.maxwell@arm.com>2024-05-17 17:13:59 +0100
committerGitHub <noreply@github.com>2024-05-17 17:13:59 +0100
commite01ff8238cf62c7149de7b8046bccec9adefbe67 (patch)
treef4e03a13d926b819c0e8fa3fe197ed625c8e25df /mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp
parent8aa6511f4209bba33a74c4ef6e208fda5c0f3d27 (diff)
[mlir][vector] Fix scalability issues in drop innermost unit dims transfer patterns (#92402)
Previously, these rewrites would drop scalable dimensions and treated `[1]` (scalable one dim) as a unit dimension. This patch propagates scalable dimensions and ensures `[1]` is not treated as a unit dimension.
Diffstat (limited to 'mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp')
-rw-r--r--mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp13
1 files changed, 9 insertions, 4 deletions
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp
index 69c497264fd1..f29eba90c3ce 100644
--- a/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp
+++ b/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp
@@ -1237,6 +1237,10 @@ getTransferFoldableInnerUnitDims(MemRefType srcType, VectorType vectorType) {
if (failed(getStridesAndOffset(srcType, srcStrides, srcOffset)))
return failure();
+ auto isUnitDim = [](VectorType type, int dim) {
+ return type.getDimSize(dim) == 1 && !type.getScalableDims()[dim];
+ };
+
// According to vector.transfer_read/write semantics, the vector can be a
// slice. Thus, we have to offset the check index with `rankDiff` in
// `srcStrides` and source dim sizes.
@@ -1247,8 +1251,7 @@ getTransferFoldableInnerUnitDims(MemRefType srcType, VectorType vectorType) {
// It can be folded only if they are 1 and the stride is 1.
int dim = vectorType.getRank() - i - 1;
if (srcStrides[dim + rankDiff] != 1 ||
- srcType.getDimSize(dim + rankDiff) != 1 ||
- vectorType.getDimSize(dim) != 1)
+ srcType.getDimSize(dim + rankDiff) != 1 || !isUnitDim(vectorType, dim))
break;
result++;
}
@@ -1292,7 +1295,8 @@ class DropInnerMostUnitDimsTransferRead
auto resultTargetVecType =
VectorType::get(targetType.getShape().drop_back(dimsToDrop),
- targetType.getElementType());
+ targetType.getElementType(),
+ targetType.getScalableDims().drop_back(dimsToDrop));
auto loc = readOp.getLoc();
SmallVector<OpFoldResult> sizes =
@@ -1378,7 +1382,8 @@ class DropInnerMostUnitDimsTransferWrite
auto resultTargetVecType =
VectorType::get(targetType.getShape().drop_back(dimsToDrop),
- targetType.getElementType());
+ targetType.getElementType(),
+ targetType.getScalableDims().drop_back(dimsToDrop));
Location loc = writeOp.getLoc();
SmallVector<OpFoldResult> sizes =