diff options
| author | Andrzej WarzyĆski <andrzej.warzynski@arm.com> | 2023-12-04 16:56:43 +0000 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-12-04 16:56:43 +0000 |
| commit | bbd2b08b95fe76bea138c1b03c1cd42ed3ee04df (patch) | |
| tree | 8eddda7b522f76c5b344d1f0138c6cdcd09da31a /mlir/lib/Dialect/Vector/Transforms/LowerVectorTranspose.cpp | |
| parent | 74e59e7752ac80607d817308a63d0574a6e8cbe1 (diff) | |
[mlir][vector] Make `TransposeOpLowering` configurable (#73915)
Following the discussion here:
* https://github.com/llvm/llvm-project/pull/72105
this patch makes the `TransposeOpLowering` configurable so that one can select
whether to favour `vector.shape_cast` over `vector.transpose`.
As per the discussion in #72105, using `vector.shape_cast` is very beneficial
and desirable when targeting `LLVM IR` (CPU lowering), but won't work when
targeting `SPIR-V` today (GPU lowering). Hence the need for a mechanism to be
able to disable/enable the pattern introduced in #72105. This patch proposes one
such mechanism.
While this should solve the problem that we are facing today, it's understood to
be a temporary workaround. It should be removed once support for lowering
`vector.shape_cast` to SPIR-V is added. Also, (once implemented) the following
proposal might make this workaround redundant:
* https://discourse.llvm.org/t/improving-handling-of-unit-dimensions-in-the-vector-dialect/
Diffstat (limited to 'mlir/lib/Dialect/Vector/Transforms/LowerVectorTranspose.cpp')
| -rw-r--r-- | mlir/lib/Dialect/Vector/Transforms/LowerVectorTranspose.cpp | 34 |
1 files changed, 18 insertions, 16 deletions
diff --git a/mlir/lib/Dialect/Vector/Transforms/LowerVectorTranspose.cpp b/mlir/lib/Dialect/Vector/Transforms/LowerVectorTranspose.cpp index 97f6caca1b25..4d43a76c4a4e 100644 --- a/mlir/lib/Dialect/Vector/Transforms/LowerVectorTranspose.cpp +++ b/mlir/lib/Dialect/Vector/Transforms/LowerVectorTranspose.cpp @@ -334,22 +334,24 @@ public: return rewriter.notifyMatchFailure( op, "Options specifies lowering to shuffle"); - // Replace: - // vector.transpose %0, [1, 0] : vector<nx1x<eltty>> to - // vector<1xnxelty> - // with: - // vector.shape_cast %0 : vector<nx1x<eltty>> to vector<1xnxelty> - // - // Source with leading unit dim (inverse) is also replaced. Unit dim must - // be fixed. Non-unit can be scalable. - if (resType.getRank() == 2 && - ((resType.getShape().front() == 1 && - !resType.getScalableDims().front()) || - (resType.getShape().back() == 1 && - !resType.getScalableDims().back())) && - transp == ArrayRef<int64_t>({1, 0})) { - rewriter.replaceOpWithNewOp<vector::ShapeCastOp>(op, resType, input); - return success(); + if (vectorTransformOptions.useShapeCast) { + // Replace: + // vector.transpose %0, [1, 0] : vector<nx1x<eltty>> to + // vector<1xnxelty> + // with: + // vector.shape_cast %0 : vector<nx1x<eltty>> to vector<1xnxelty> + // + // Source with leading unit dim (inverse) is also replaced. Unit dim must + // be fixed. Non-unit can be scalable. + if (resType.getRank() == 2 && + ((resType.getShape().front() == 1 && + !resType.getScalableDims().front()) || + (resType.getShape().back() == 1 && + !resType.getScalableDims().back())) && + transp == ArrayRef<int64_t>({1, 0})) { + rewriter.replaceOpWithNewOp<vector::ShapeCastOp>(op, resType, input); + return success(); + } } if (inputType.isScalable()) |
