From 2ec98ffbf12163ee4ff9f4e674eba714bce24ec1 Mon Sep 17 00:00:00 2001 From: Matthias Springer Date: Mon, 19 Dec 2022 10:24:46 +0100 Subject: [mlir][vector] Add scalar vector xfer to memref patterns These patterns devectorize scalar transfers such as vector or vector<1xf32>. Differential Revision: https://reviews.llvm.org/D140215 --- .../Transforms/VectorTransferOpTransforms.cpp | 104 +++++++++++++++++++++ 1 file changed, 104 insertions(+) (limited to 'mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp') diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp index b59b10c43678..727a356210a3 100644 --- a/mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp +++ b/mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp @@ -11,8 +11,10 @@ // //===----------------------------------------------------------------------===// +#include "mlir/Dialect/Affine/IR/AffineOps.h" #include "mlir/Dialect/Arith/IR/Arith.h" #include "mlir/Dialect/MemRef/IR/MemRef.h" +#include "mlir/Dialect/Tensor/IR/Tensor.h" #include "mlir/Dialect/Vector/IR/VectorOps.h" #include "mlir/Dialect/Vector/Transforms/VectorTransforms.h" #include "mlir/Dialect/Vector/Utils/VectorUtils.h" @@ -556,6 +558,101 @@ class FlattenContiguousRowMajorTransferWritePattern } }; +/// Rewrite extractelement(transfer_read) to memref.load. +/// +/// Rewrite only if the extractelement op is the single user of the transfer op. +/// E.g., do not rewrite IR such as: +/// %0 = vector.transfer_read ... : vector<1024xf32> +/// %1 = vector.extractelement %0[%a : index] : vector<1024xf32> +/// %2 = vector.extractelement %0[%b : index] : vector<1024xf32> +/// Rewriting such IR (replacing one vector load with multiple scalar loads) may +/// negatively affect performance. +class FoldScalarExtractOfTransferRead + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(vector::ExtractElementOp extractOp, + PatternRewriter &rewriter) const override { + auto xferOp = extractOp.getVector().getDefiningOp(); + if (!xferOp) + return failure(); + // xfer result must have a single use. Otherwise, it may be better to + // perform a vector load. + if (!extractOp.getVector().hasOneUse()) + return failure(); + // Mask not supported. + if (xferOp.getMask()) + return failure(); + // Map not supported. + if (!xferOp.getPermutationMap().isMinorIdentity()) + return failure(); + // Cannot rewrite if the indices may be out of bounds. The starting point is + // always inbounds, so we don't care in case of 0d transfers. + if (xferOp.hasOutOfBoundsDim() && xferOp.getType().getRank() > 0) + return failure(); + // Construct scalar load. + SmallVector newIndices(xferOp.getIndices().begin(), + xferOp.getIndices().end()); + if (extractOp.getPosition()) { + AffineExpr sym0, sym1; + bindSymbols(extractOp.getContext(), sym0, sym1); + OpFoldResult ofr = makeComposedFoldedAffineApply( + rewriter, extractOp.getLoc(), sym0 + sym1, + {newIndices[newIndices.size() - 1], extractOp.getPosition()}); + if (ofr.is()) { + newIndices[newIndices.size() - 1] = ofr.get(); + } else { + newIndices[newIndices.size() - 1] = + rewriter.create(extractOp.getLoc(), + *getConstantIntValue(ofr)); + } + } + if (xferOp.getSource().getType().isa()) { + rewriter.replaceOpWithNewOp(extractOp, xferOp.getSource(), + newIndices); + } else { + rewriter.replaceOpWithNewOp( + extractOp, xferOp.getSource(), newIndices); + } + return success(); + } +}; + +/// Rewrite scalar transfer_write(broadcast) to memref.store. +class FoldScalarTransferWriteOfBroadcast + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(vector::TransferWriteOp xferOp, + PatternRewriter &rewriter) const override { + // Must be a scalar write. + auto vecType = xferOp.getVectorType(); + if (vecType.getRank() != 0 && + (vecType.getRank() != 1 || vecType.getShape()[0] != 1)) + return failure(); + // Mask not supported. + if (xferOp.getMask()) + return failure(); + // Map not supported. + if (!xferOp.getPermutationMap().isMinorIdentity()) + return failure(); + // Must be a broadcast of a scalar. + auto broadcastOp = xferOp.getVector().getDefiningOp(); + if (!broadcastOp || broadcastOp.getSource().getType().isa()) + return failure(); + // Construct a scalar store. + if (xferOp.getSource().getType().isa()) { + rewriter.replaceOpWithNewOp( + xferOp, broadcastOp.getSource(), xferOp.getSource(), + xferOp.getIndices()); + } else { + rewriter.replaceOpWithNewOp( + xferOp, broadcastOp.getSource(), xferOp.getSource(), + xferOp.getIndices()); + } + return success(); + } +}; } // namespace void mlir::vector::transferOpflowOpt(Operation *rootOp) { @@ -574,6 +671,13 @@ void mlir::vector::transferOpflowOpt(Operation *rootOp) { opt.removeDeadOp(); } +void mlir::vector::populateScalarVectorTransferLoweringPatterns( + RewritePatternSet &patterns, PatternBenefit benefit) { + patterns + .add( + patterns.getContext(), benefit); +} + void mlir::vector::populateVectorTransferDropUnitDimsPatterns( RewritePatternSet &patterns, PatternBenefit benefit) { patterns -- cgit v1.2.3