diff options
| author | Diego Caballero <diegocaballero@google.com> | 2024-02-28 08:15:47 -0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-02-28 08:15:47 -0800 |
| commit | 4623c114fb51e48b5cbb352adecb60d76077cb0a (patch) | |
| tree | beb1aa91f018a34fb2a0f7d4e7bffc1a67bb7ac3 /mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp | |
| parent | 80cff273906b200eed2f779913f49a64cad8c0c6 (diff) | |
[mlir][Vector] Support vector.insert in bubbling bitcast patterns (#82843)
This PR is adds support for `vector.insert` to the patterns that bubble up and down `vector.bitcat` ops across `vector.extract/extract_slice/insert_slice` ops.
Diffstat (limited to 'mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp')
| -rw-r--r-- | mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp | 74 |
1 files changed, 72 insertions, 2 deletions
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp index 74dd1dfaca0d..a2d4e2166331 100644 --- a/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp +++ b/mlir/lib/Dialect/Vector/Transforms/VectorTransforms.cpp @@ -713,6 +713,76 @@ struct BubbleDownBitCastForStridedSliceExtract // Shuffles vector.bitcast op before vector.insert_strided_slice op. // // This transforms IR like: +// %0 = vector.insert %val, %dst[4] : vector<32xi4> into vector<8x32xi4> +// %1 = vector.bitcast %0 : vector<8x32xi4> to vector<8x16xi8> +// Into: +// %0 = vector.bitcast %val : vector<32xi4> to vector<16xi8> +// %1 = vector.bitcast %dst : vector<8x32xi4> to vector<8x16xi8> +// %2 = vector.insert %0, %1 [4] : vector<16xi8> into vector<8x16xi8> +// +struct BubbleUpBitCastForInsert : public OpRewritePattern<vector::BitCastOp> { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(vector::BitCastOp bitcastOp, + PatternRewriter &rewriter) const override { + VectorType castSrcType = bitcastOp.getSourceVectorType(); + VectorType castDstType = bitcastOp.getResultVectorType(); + + // 0-D and scalable vectors are not supported yet. + if (castSrcType.getRank() == 0 || castSrcType.isScalable() || + castDstType.isScalable()) + return failure(); + + int64_t castSrcLastDim = castSrcType.getShape().back(); + int64_t castDstLastDim = castDstType.getShape().back(); + bool isNumElemsShrink = castSrcLastDim >= castDstLastDim; + int64_t ratio; + if (isNumElemsShrink) { + assert(castSrcLastDim % castDstLastDim == 0); + ratio = castSrcLastDim / castDstLastDim; + } else { + assert(castDstLastDim % castSrcLastDim == 0); + ratio = castDstLastDim / castSrcLastDim; + } + + auto insertOp = bitcastOp.getSource().getDefiningOp<vector::InsertOp>(); + if (!insertOp) + return failure(); + + // Only vector sources are supported for now. + auto insertSrcType = dyn_cast<VectorType>(insertOp.getSourceType()); + if (!insertSrcType) + return failure(); + + // Bitcast the source. + SmallVector<int64_t> srcDims(insertSrcType.getShape()); + srcDims.back() = + isNumElemsShrink ? srcDims.back() / ratio : srcDims.back() * ratio; + VectorType newCastSrcType = + VectorType::get(srcDims, castDstType.getElementType()); + auto newCastSrcOp = rewriter.create<vector::BitCastOp>( + bitcastOp.getLoc(), newCastSrcType, insertOp.getSource()); + + SmallVector<int64_t> dstDims(insertOp.getDestVectorType().getShape()); + dstDims.back() = + isNumElemsShrink ? dstDims.back() / ratio : dstDims.back() * ratio; + VectorType newCastDstType = + VectorType::get(dstDims, castDstType.getElementType()); + + // Bitcast the destination. + auto newCastDstOp = rewriter.create<vector::BitCastOp>( + bitcastOp.getLoc(), newCastDstType, insertOp.getDest()); + + // Generate new insert. + rewriter.replaceOpWithNewOp<vector::InsertOp>( + bitcastOp, newCastSrcOp, newCastDstOp, insertOp.getMixedPosition()); + return success(); + } +}; + +// Shuffles vector.bitcast op before vector.insert_strided_slice op. +// +// This transforms IR like: // %0 = vector.insert_strided_slice %src, %dst { // offsets = [0], strides = [1]} : vector<4xf16> into vector<8xf16> // %1 = vector.bitcast %0: vector<8xf16> to vector<4xf32> @@ -1782,8 +1852,8 @@ void mlir::vector::populateBubbleVectorBitCastOpPatterns( RewritePatternSet &patterns, PatternBenefit benefit) { patterns.add<BubbleDownVectorBitCastForExtract, BubbleDownBitCastForStridedSliceExtract, - BubbleUpBitCastForStridedSliceInsert>(patterns.getContext(), - benefit); + BubbleUpBitCastForInsert, BubbleUpBitCastForStridedSliceInsert>( + patterns.getContext(), benefit); } void mlir::vector::populateBreakDownVectorBitCastOpPatterns( |
