diff options
Diffstat (limited to 'mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp')
| -rw-r--r-- | mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp | 28 |
1 files changed, 23 insertions, 5 deletions
diff --git a/mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp b/mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp index c131fde517f8..4c93d3841bf8 100644 --- a/mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp +++ b/mlir/lib/Dialect/Vector/Transforms/VectorTransferOpTransforms.cpp @@ -568,6 +568,7 @@ namespace { /// memref.collapse_shape on the source so that the resulting /// vector.transfer_read has a 1D source. Requires the source shape to be /// already reduced i.e. without unit dims. +/// /// If `targetVectorBitwidth` is provided, the flattening will only happen if /// the trailing dimension of the vector read is smaller than the provided /// bitwidth. @@ -617,7 +618,7 @@ public: Value collapsedSource = collapseInnerDims(rewriter, loc, source, firstDimToCollapse); MemRefType collapsedSourceType = - dyn_cast<MemRefType>(collapsedSource.getType()); + cast<MemRefType>(collapsedSource.getType()); int64_t collapsedRank = collapsedSourceType.getRank(); assert(collapsedRank == firstDimToCollapse + 1); @@ -658,6 +659,10 @@ private: /// memref.collapse_shape on the source so that the resulting /// vector.transfer_write has a 1D source. Requires the source shape to be /// already reduced i.e. without unit dims. +/// +/// If `targetVectorBitwidth` is provided, the flattening will only happen if +/// the trailing dimension of the vector read is smaller than the provided +/// bitwidth. class FlattenContiguousRowMajorTransferWritePattern : public OpRewritePattern<vector::TransferWriteOp> { public: @@ -674,9 +679,12 @@ public: VectorType vectorType = cast<VectorType>(vector.getType()); Value source = transferWriteOp.getSource(); MemRefType sourceType = dyn_cast<MemRefType>(source.getType()); + + // 0. Check pre-conditions // Contiguity check is valid on tensors only. if (!sourceType) return failure(); + // If this is already 0D/1D, there's nothing to do. if (vectorType.getRank() <= 1) // Already 0D/1D, nothing to do. return failure(); @@ -688,7 +696,6 @@ public: return failure(); if (!vector::isContiguousSlice(sourceType, vectorType)) return failure(); - int64_t firstDimToCollapse = sourceType.getRank() - vectorType.getRank(); // TODO: generalize this pattern, relax the requirements here. if (transferWriteOp.hasOutOfBoundsDim()) return failure(); @@ -697,10 +704,9 @@ public: if (transferWriteOp.getMask()) return failure(); - SmallVector<Value> collapsedIndices = - getCollapsedIndices(rewriter, loc, sourceType.getShape(), - transferWriteOp.getIndices(), firstDimToCollapse); + int64_t firstDimToCollapse = sourceType.getRank() - vectorType.getRank(); + // 1. Collapse the source memref Value collapsedSource = collapseInnerDims(rewriter, loc, source, firstDimToCollapse); MemRefType collapsedSourceType = @@ -708,11 +714,20 @@ public: int64_t collapsedRank = collapsedSourceType.getRank(); assert(collapsedRank == firstDimToCollapse + 1); + // 2. Generate input args for a new vector.transfer_read that will read + // from the collapsed memref. + // 2.1. New dim exprs + affine map SmallVector<AffineExpr, 1> dimExprs{ getAffineDimExpr(firstDimToCollapse, rewriter.getContext())}; auto collapsedMap = AffineMap::get(collapsedRank, 0, dimExprs, rewriter.getContext()); + // 2.2 New indices + SmallVector<Value> collapsedIndices = + getCollapsedIndices(rewriter, loc, sourceType.getShape(), + transferWriteOp.getIndices(), firstDimToCollapse); + + // 3. Create new vector.transfer_write that writes to the collapsed memref VectorType flatVectorType = VectorType::get({vectorType.getNumElements()}, vectorType.getElementType()); Value flatVector = @@ -721,6 +736,9 @@ public: rewriter.create<vector::TransferWriteOp>( loc, flatVector, collapsedSource, collapsedIndices, collapsedMap); flatWrite.setInBoundsAttr(rewriter.getBoolArrayAttr({true})); + + // 4. Replace the old transfer_write with the new one writing the + // collapsed shape rewriter.eraseOp(transferWriteOp); return success(); } |
