diff options
| author | Benjamin Maxwell <benjamin.maxwell@arm.com> | 2024-05-17 17:13:59 +0100 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-17 17:13:59 +0100 |
| commit | e01ff8238cf62c7149de7b8046bccec9adefbe67 (patch) | |
| tree | f4e03a13d926b819c0e8fa3fe197ed625c8e25df /mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp | |
| parent | 8aa6511f4209bba33a74c4ef6e208fda5c0f3d27 (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.cpp | 13 |
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 = |
