18#include "llvm/ADT/STLExtras.h"
20#define DEBUG_TYPE "vector-drop-unit-dim"
34 while (!newShape.empty() && newShape.front() == 1 &&
35 !newScalableDims.front()) {
36 newShape = newShape.drop_front(1);
37 newScalableDims = newScalableDims.drop_front(1);
41 if (newShape.empty()) {
42 newShape = oldShape.take_back();
43 newScalableDims = oldType.getScalableDims().take_back();
45 return VectorType::get(newShape, oldType.getElementType(), newScalableDims);
59 return builder.
create(state);
65struct CastAwayExtractStridedSliceLeadingOneDim
69 LogicalResult matchAndRewrite(vector::ExtractStridedSliceOp extractOp,
70 PatternRewriter &rewriter)
const override {
74 VectorType oldSrcType = extractOp.getSourceVectorType();
77 if (newSrcType.getRank() == oldSrcType.getRank())
80 int64_t dropCount = oldSrcType.getRank() - newSrcType.getRank();
82 VectorType oldDstType = extractOp.getType();
83 VectorType newDstType =
84 VectorType::get(oldDstType.getShape().drop_front(dropCount),
85 oldDstType.getElementType(),
86 oldDstType.getScalableDims().drop_front(dropCount));
88 Location loc = extractOp.getLoc();
91 loc, newSrcType, extractOp.getSource());
96 extractOp.getOffsets().getValue().drop_front(dropCount));
98 extractOp.getSizes().getValue().drop_front(dropCount));
100 extractOp.getStrides().getValue().drop_front(dropCount));
102 auto newExtractOp = vector::ExtractStridedSliceOp::create(
103 rewriter, loc, newDstType, newSrcVector, newOffsets, newSizes,
115struct CastAwayInsertStridedSliceLeadingOneDim
119 LogicalResult matchAndRewrite(vector::InsertStridedSliceOp insertOp,
120 PatternRewriter &rewriter)
const override {
121 VectorType oldSrcType = insertOp.getSourceVectorType();
123 VectorType oldDstType = insertOp.getDestVectorType();
126 int64_t srcDropCount = oldSrcType.getRank() - newSrcType.getRank();
127 int64_t dstDropCount = oldDstType.getRank() - newDstType.getRank();
128 if (srcDropCount == 0 && dstDropCount == 0)
132 Location loc = insertOp.getLoc();
134 Value newSrcVector = rewriter.
createOrFold<vector::ShapeCastOp>(
135 loc, newSrcType, insertOp.getValueToStore());
136 Value newDstVector = rewriter.
createOrFold<vector::ShapeCastOp>(
137 loc, newDstType, insertOp.getDest());
140 insertOp.getOffsets().getValue().take_back(newDstType.getRank()));
142 insertOp.getStrides().getValue().take_back(newSrcType.getRank()));
144 auto newInsertOp = vector::InsertStridedSliceOp::create(
145 rewriter, loc, newDstType, newSrcVector, newDstVector, newOffsets,
157struct CastAwayInsertLeadingOneDim :
public OpRewritePattern<vector::InsertOp> {
160 LogicalResult matchAndRewrite(vector::InsertOp insertOp,
161 PatternRewriter &rewriter)
const override {
162 Type oldSrcType = insertOp.getValueToStoreType();
163 Type newSrcType = oldSrcType;
164 int64_t oldSrcRank = 0, newSrcRank = 0;
165 if (
auto type = dyn_cast<VectorType>(oldSrcType)) {
167 oldSrcRank = type.getRank();
168 newSrcRank = cast<VectorType>(newSrcType).getRank();
171 VectorType oldDstType = insertOp.getDestVectorType();
174 int64_t srcDropCount = oldSrcRank - newSrcRank;
175 int64_t dstDropCount = oldDstType.getRank() - newDstType.getRank();
176 if (srcDropCount == 0 && dstDropCount == 0)
180 Location loc = insertOp.getLoc();
182 Value newSrcVector = insertOp.getValueToStore();
183 if (oldSrcRank != 0) {
184 newSrcVector = rewriter.
createOrFold<vector::ShapeCastOp>(
185 loc, cast<VectorType>(newSrcType), insertOp.getValueToStore());
187 Value newDstVector = rewriter.
createOrFold<vector::ShapeCastOp>(
188 loc, newDstType, insertOp.getDest());
194 unsigned oldPosRank = insertOp.getNumIndices();
195 unsigned newPosRank = std::max<int64_t>(0, oldPosRank - dstDropCount);
196 SmallVector<OpFoldResult> oldPosition = insertOp.getMixedPosition();
197 SmallVector<OpFoldResult> newPosition =
198 llvm::to_vector(ArrayRef(oldPosition).take_back(newPosRank));
199 newPosition.resize(newDstType.getRank() - newSrcRank,
202 auto newInsertOp = vector::InsertOp::create(rewriter, loc, newSrcVector,
203 newDstVector, newPosition);
216 return vector::ShapeCastOp::create(
b, loc, newMaskType, mask);
222struct CastAwayTransferReadLeadingOneDim
226 LogicalResult matchAndRewrite(vector::TransferReadOp read,
227 PatternRewriter &rewriter)
const override {
229 if (cast<MaskableOpInterface>(read.getOperation()).isMasked())
232 if (read.getTransferRank() == 0)
234 read,
"Nothing to trim - the transfer itself has rank zero");
236 auto shapedType = cast<ShapedType>(read.getBase().getType());
237 if (shapedType.getElementType() != read.getVectorType().getElementType())
240 VectorType oldType = read.getVectorType();
243 if (newType == oldType)
246 AffineMap oldMap = read.getPermutationMap();
247 ArrayRef<AffineExpr> newResults =
248 oldMap.
getResults().take_back(newType.getRank());
254 if (read.getInBounds())
256 read.getInBoundsAttr().getValue().take_back(newType.getRank()));
258 Value mask = Value();
260 mask = dropUnitDimsFromMask(rewriter, read.getLoc(), read.getMask(),
263 auto newRead = vector::TransferReadOp::create(
264 rewriter, read.getLoc(), newType, read.getBase(), read.getIndices(),
265 AffineMapAttr::get(newMap), read.getPadding(), mask, inBoundsAttr);
275struct CastAwayTransferWriteLeadingOneDim
279 LogicalResult matchAndRewrite(vector::TransferWriteOp write,
280 PatternRewriter &rewriter)
const override {
282 if (cast<MaskableOpInterface>(write.getOperation()).isMasked())
285 if (write.getTransferRank() == 0)
287 write,
"Nothing to trim - the transfer itself has rank zero");
289 auto shapedType = dyn_cast<ShapedType>(write.getBase().getType());
290 if (shapedType.getElementType() != write.getVectorType().getElementType())
293 VectorType oldType = write.getVectorType();
295 if (newType == oldType)
299 ArrayRef<AffineExpr> newResults =
300 oldMap.
getResults().take_back(newType.getRank());
306 if (write.getInBounds())
308 write.getInBoundsAttr().getValue().take_back(newType.getRank()));
310 auto newVector = rewriter.
createOrFold<vector::ShapeCastOp>(
311 write.getLoc(), newType, write.getVector());
313 if (write.getMask()) {
314 Value newMask = dropUnitDimsFromMask(rewriter, write.getLoc(),
315 write.getMask(), newType, newMap);
317 write, newVector, write.getBase(), write.getIndices(),
318 AffineMapAttr::get(newMap), newMask, inBoundsAttr);
323 write, newVector, write.getBase(), write.getIndices(),
324 AffineMapAttr::get(newMap), inBoundsAttr);
332mlir::vector::castAwayContractionLeadingOneDim(vector::ContractionOp contractOp,
333 MaskingOpInterface maskingOp,
335 VectorType oldAccType = dyn_cast<VectorType>(contractOp.getAccType());
336 if (oldAccType ==
nullptr)
338 if (oldAccType.getRank() < 1)
340 if (oldAccType.getShape()[0] != 1)
346 auto oldIndexingMaps = contractOp.getIndexingMapsArray();
349 auto oldIteratorTypes = contractOp.getIteratorTypes();
352 int64_t dimToDrop = oldIndexingMaps[2].getDimPosition(0);
358 for (
const auto &it : llvm::enumerate(oldIteratorTypes)) {
360 if (currDim == dimToDrop)
362 newIteratorTypes.push_back(it.value());
366 contractOp.getAcc()};
368 auto loc = contractOp.getLoc();
370 for (
const auto &it : llvm::enumerate(oldIndexingMaps)) {
373 bool validExtract =
false;
375 auto map = it.value();
376 int64_t orginalZeroDim = it.value().getDimPosition(0);
377 if (orginalZeroDim != dimToDrop) {
383 bool transposeNeeded =
false;
387 for (
int64_t i = 0, e = map.getNumResults(); i < e; ++i) {
388 int64_t currDim = map.getDimPosition(i);
389 if (currDim == dimToDrop) {
390 transposeNeeded =
true;
391 perm.insert(perm.begin(), i);
393 transposeResults.insert(transposeResults.begin(), targetExpr);
397 transposeResults.push_back(targetExpr);
404 bool transposeNonOuterUnitDims =
false;
405 auto operandShape = cast<ShapedType>(operands[it.index()].getType());
406 for (
auto [
index, dim] :
409 operandShape.getDimSize(
index) != 1) {
410 transposeNonOuterUnitDims =
true;
417 if (transposeNeeded) {
419 contractOp.getContext());
420 if (transposeNonOuterUnitDims) {
421 operands[it.index()] = rewriter.
createOrFold<vector::TransposeOp>(
422 loc, operands[it.index()], perm);
430 if (map.getDimPosition(0) == dimToDrop)
433 for (
int64_t i = 0, e = map.getNumResults(); i < e; ++i) {
434 int64_t currDim = map.getDimPosition(i);
435 if (currDim == dimToDrop)
439 currDim < dimToDrop ? currDim : currDim - 1);
440 results.push_back(targetExpr);
442 newIndexingMaps.push_back(
AffineMap::get(map.getNumDims() - 1, 0, results,
443 contractOp.getContext()));
446 newOperands.push_back(validExtract
447 ? vector::ExtractOp::create(rewriter, loc,
448 operands[it.index()],
450 : operands[it.index()]);
455 Operation *newOp = vector::ContractionOp::create(
456 rewriter, loc, newOperands[0], newOperands[1], newOperands[2],
458 rewriter.
getArrayAttr(newIteratorTypes), contractOp.getKind());
461 auto newMask = vector::ExtractOp::create(rewriter, loc, maskingOp.getMask(),
467 return vector::BroadcastOp::create(rewriter, loc,
468 contractOp->getResultTypes()[0],
479struct CastAwayContractionLeadingOneDim
481 using MaskableOpRewritePattern::MaskableOpRewritePattern;
484 matchAndRewriteMaskableOp(vector::ContractionOp contractOp,
485 MaskingOpInterface maskingOp,
486 PatternRewriter &rewriter)
const override {
487 return castAwayContractionLeadingOneDim(contractOp, maskingOp, rewriter);
505 CastAwayElementwiseLeadingOneDim(MLIRContext *context,
506 PatternBenefit benefit = 1)
507 : RewritePattern(MatchAnyOpTypeTag(), benefit, context) {}
509 LogicalResult matchAndRewrite(Operation *op,
510 PatternRewriter &rewriter)
const override {
517 if (newVecType == vecType)
519 int64_t dropDim = vecType.getRank() - newVecType.getRank();
520 SmallVector<Value, 4> newOperands;
522 if (
auto opVecType = dyn_cast<VectorType>(operand.getType())) {
523 newOperands.push_back(vector::ExtractOp::create(
526 newOperands.push_back(operand);
543 auto oldType = cast<VectorType>(operand.
getType());
546 oldType.getScalableDims().take_front(nDropped);
548 llvm::all_of(leadingShape, [](
int64_t d) {
return d == 1; }) &&
549 llvm::none_of(leadingScalable, [](
bool s) {
return s; });
551 return vector::ExtractOp::create(
b, loc, operand,
splatZero(nDropped));
552 VectorType newType = VectorType::get(
553 oldType.getShape().drop_front(nDropped), oldType.getElementType(),
554 oldType.getScalableDims().drop_front(nDropped));
555 return vector::ShapeCastOp::create(
b, loc, newType, operand);
564template <
typename OpTy>
566 using OpRewritePattern<OpTy>::OpRewritePattern;
568 LogicalResult matchAndRewrite(OpTy op,
569 PatternRewriter &rewriter)
const override {
570 VectorType oldResultType = op.getVectorType();
572 if (newResultType == oldResultType)
574 int64_t nDropped = oldResultType.getRank() - newResultType.getRank();
576 Location loc = op.getLoc();
577 SmallVector<Value> newOperands;
578 newOperands.reserve(op->getNumOperands());
579 for (Value operand : op->getOperands()) {
580 if (isa<VectorType>(operand.getType())) {
581 newOperands.push_back(
584 newOperands.push_back(operand);
599template <
typename OpTy>
601 using OpRewritePattern<OpTy>::OpRewritePattern;
603 LogicalResult matchAndRewrite(OpTy op,
604 PatternRewriter &rewriter)
const override {
605 VectorType oldVecType = op.getVectorType();
607 if (newVecType == oldVecType)
609 int64_t nDropped = oldVecType.getRank() - newVecType.getRank();
611 Location loc = op.getLoc();
612 SmallVector<Value> newOperands;
613 newOperands.reserve(op->getNumOperands());
614 for (Value operand : op->getOperands()) {
615 if (isa<VectorType>(operand.getType())) {
616 newOperands.push_back(
619 newOperands.push_back(operand);
632struct CastAwayConstantMaskLeadingOneDim
636 LogicalResult matchAndRewrite(vector::ConstantMaskOp mask,
637 PatternRewriter &rewriter)
const override {
638 VectorType oldType = mask.
getType();
641 if (newType == oldType)
644 int64_t dropDim = oldType.getRank() - newType.getRank();
645 ArrayRef<int64_t> dimSizes = mask.getMaskDimSizes();
649 int64_t flatLeadingSize =
650 llvm::product_of(dimSizes.take_front(dropDim + 1));
651 SmallVector<int64_t> newDimSizes = {flatLeadingSize};
652 newDimSizes.append(dimSizes.begin() + dropDim + 1, dimSizes.end());
654 auto newMask = vector::ConstantMaskOp::create(rewriter, mask.
getLoc(),
655 newType, newDimSizes);
663void mlir::vector::populateCastAwayVectorLeadingOneDimPatterns(
666 .
add<CastAwayExtractStridedSliceLeadingOneDim,
667 CastAwayInsertStridedSliceLeadingOneDim, CastAwayInsertLeadingOneDim,
668 CastAwayConstantMaskLeadingOneDim, CastAwayTransferReadLeadingOneDim,
669 CastAwayTransferWriteLeadingOneDim, CastAwayElementwiseLeadingOneDim,
670 CastAwayContractionLeadingOneDim,
671 CastAwayLoadLikeLeadingOneDim<vector::LoadOp>,
672 CastAwayLoadLikeLeadingOneDim<vector::MaskedLoadOp>,
673 CastAwayLoadLikeLeadingOneDim<vector::ExpandLoadOp>,
674 CastAwayLoadLikeLeadingOneDim<vector::GatherOp>,
675 CastAwayStoreLikeLeadingOneDim<vector::StoreOp>,
676 CastAwayStoreLikeLeadingOneDim<vector::MaskedStoreOp>,
677 CastAwayStoreLikeLeadingOneDim<vector::CompressStoreOp>,
678 CastAwayStoreLikeLeadingOneDim<vector::ScatterOp>>(
static SmallVector< int64_t > splatZero(int64_t rank)
Return a smallVector of size rank containing all zeros.
static Value dropLeadingOneDimsFromOperand(OpBuilder &b, Location loc, Value operand, int64_t nDropped)
static VectorType trimLeadingUnitDims(VectorType oldType)
static Operation * createWithProperties(OpBuilder &builder, Operation *op, ValueRange operands, TypeRange resultTypes)
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
unsigned getNumSymbols() const
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
static AffineMap getPermutationMap(ArrayRef< unsigned > permutation, MLIRContext *context)
Returns an AffineMap representing a permutation.
IntegerAttr getI64IntegerAttr(int64_t value)
AffineExpr getAffineDimExpr(unsigned position)
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
MLIRContext * getContext() const
ArrayAttr getAffineMapArrayAttr(ArrayRef< AffineMap > values)
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
This class helps build Operations.
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
Operation * create(const OperationState &state)
Creates an operation given the fields represented as an OperationState.
Operation is the basic unit of execution within MLIR.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Location getLoc()
The source location the operation was defined or derived from.
Attribute getPropertiesAsAttribute()
Return the properties converted to an attribute.
OperationName getName()
The name of an operation is the key identifier for it.
DictionaryAttr getDiscardableAttrDictionary()
Return all of the discardable attributes on this operation as a DictionaryAttr.
result_type_range getResultTypes()
operand_range getOperands()
Returns an iterator on the underlying Value's.
result_range getResults()
unsigned getNumResults()
Return the number of results held by this operation.
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
RewritePattern is the common base class for all DAG to DAG replacements.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
This class provides an abstraction over the various different ranges of value types.
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
Location getLoc() const
Return the location of this value.
bool hasElementwiseMappableTraits(Operation *op)
Together, Elementwise, Scalarizable, Vectorizable, and Tensorizable provide an easy way for scalar op...
Operation * maskOperation(OpBuilder &builder, Operation *maskableOp, Value mask, Value passthru=Value())
Creates a vector.mask operation around a maskable operation.
VectorType inferTransferOpMaskType(VectorType vecType, AffineMap permMap)
Infers the mask type for a transfer op given its vector type and permutation map.
bool isParallelIterator(Attribute attr)
Returns true if attr has "parallel" iterator type semantics.
Include the generated interface declarations.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
This represents an operation in an abstracted form, suitable for use with the builder APIs.
Attribute propertiesAttr
This Attribute is used to opaquely construct the properties of the operation.
A pattern for ops that implement MaskableOpInterface and that might be masked (i.e.