18#include "llvm/ADT/STLExtras.h"
20#define DEBUG_TYPE "vector-drop-unit-dim"
28 bool trimOnlyOneDim =
false,
29 bool allowRank0 =
false) {
36 while (!newShape.empty() && newShape.front() == 1 &&
37 !newScalableDims.front()) {
38 newShape = newShape.drop_front(1);
39 newScalableDims = newScalableDims.drop_front(1);
46 if (newShape.empty() && !allowRank0) {
47 newShape = oldShape.take_back();
48 newScalableDims = oldType.getScalableDims().take_back();
50 return VectorType::get(newShape, oldType.getElementType(), newScalableDims);
64 return builder.
create(state);
70struct CastAwayExtractStridedSliceLeadingOneDim
74 LogicalResult matchAndRewrite(vector::ExtractStridedSliceOp extractOp,
75 PatternRewriter &rewriter)
const override {
79 VectorType oldSrcType = extractOp.getSourceVectorType();
82 if (newSrcType.getRank() == oldSrcType.getRank())
85 int64_t dropCount = oldSrcType.getRank() - newSrcType.getRank();
87 VectorType oldDstType = extractOp.getType();
88 VectorType newDstType =
89 VectorType::get(oldDstType.getShape().drop_front(dropCount),
90 oldDstType.getElementType(),
91 oldDstType.getScalableDims().drop_front(dropCount));
93 Location loc = extractOp.getLoc();
96 loc, newSrcType, extractOp.getSource());
101 extractOp.getOffsets().getValue().drop_front(dropCount));
103 extractOp.getSizes().getValue().drop_front(dropCount));
105 extractOp.getStrides().getValue().drop_front(dropCount));
107 auto newExtractOp = vector::ExtractStridedSliceOp::create(
108 rewriter, loc, newDstType, newSrcVector, newOffsets, newSizes,
120struct CastAwayInsertStridedSliceLeadingOneDim
124 LogicalResult matchAndRewrite(vector::InsertStridedSliceOp insertOp,
125 PatternRewriter &rewriter)
const override {
126 VectorType oldSrcType = insertOp.getSourceVectorType();
128 VectorType oldDstType = insertOp.getDestVectorType();
131 int64_t srcDropCount = oldSrcType.getRank() - newSrcType.getRank();
132 int64_t dstDropCount = oldDstType.getRank() - newDstType.getRank();
133 if (srcDropCount == 0 && dstDropCount == 0)
137 Location loc = insertOp.getLoc();
139 Value newSrcVector = rewriter.
createOrFold<vector::ShapeCastOp>(
140 loc, newSrcType, insertOp.getValueToStore());
141 Value newDstVector = rewriter.
createOrFold<vector::ShapeCastOp>(
142 loc, newDstType, insertOp.getDest());
145 insertOp.getOffsets().getValue().take_back(newDstType.getRank()));
147 insertOp.getStrides().getValue().take_back(newSrcType.getRank()));
149 auto newInsertOp = vector::InsertStridedSliceOp::create(
150 rewriter, loc, newDstType, newSrcVector, newDstVector, newOffsets,
162struct CastAwayInsertLeadingOneDim :
public OpRewritePattern<vector::InsertOp> {
165 LogicalResult matchAndRewrite(vector::InsertOp insertOp,
166 PatternRewriter &rewriter)
const override {
167 Type oldSrcType = insertOp.getValueToStoreType();
168 Type newSrcType = oldSrcType;
169 int64_t oldSrcRank = 0, newSrcRank = 0;
170 if (
auto type = dyn_cast<VectorType>(oldSrcType)) {
172 oldSrcRank = type.getRank();
173 newSrcRank = cast<VectorType>(newSrcType).getRank();
176 VectorType oldDstType = insertOp.getDestVectorType();
179 int64_t srcDropCount = oldSrcRank - newSrcRank;
180 int64_t dstDropCount = oldDstType.getRank() - newDstType.getRank();
181 if (srcDropCount == 0 && dstDropCount == 0)
185 Location loc = insertOp.getLoc();
187 Value newSrcVector = insertOp.getValueToStore();
188 if (oldSrcRank != 0) {
189 newSrcVector = rewriter.
createOrFold<vector::ShapeCastOp>(
190 loc, cast<VectorType>(newSrcType), insertOp.getValueToStore());
192 Value newDstVector = rewriter.
createOrFold<vector::ShapeCastOp>(
193 loc, newDstType, insertOp.getDest());
199 unsigned oldPosRank = insertOp.getNumIndices();
200 unsigned newPosRank = std::max<int64_t>(0, oldPosRank - dstDropCount);
201 SmallVector<OpFoldResult> oldPosition = insertOp.getMixedPosition();
202 SmallVector<OpFoldResult> newPosition =
203 llvm::to_vector(ArrayRef(oldPosition).take_back(newPosRank));
204 newPosition.resize(newDstType.getRank() - newSrcRank,
207 auto newInsertOp = vector::InsertOp::create(rewriter, loc, newSrcVector,
208 newDstVector, newPosition);
221 return vector::ShapeCastOp::create(
b, loc, newMaskType, mask);
227struct CastAwayTransferReadLeadingOneDim
231 LogicalResult matchAndRewrite(vector::TransferReadOp read,
232 PatternRewriter &rewriter)
const override {
234 if (cast<MaskableOpInterface>(read.getOperation()).isMasked())
237 if (read.getTransferRank() == 0)
239 read,
"Nothing to trim - the transfer itself has rank zero");
241 auto shapedType = cast<ShapedType>(read.getBase().getType());
242 if (shapedType.getElementType() != read.getVectorType().getElementType())
245 VectorType oldType = read.getVectorType();
248 if (newType == oldType)
251 AffineMap oldMap = read.getPermutationMap();
252 ArrayRef<AffineExpr> newResults =
253 oldMap.
getResults().take_back(newType.getRank());
259 if (read.getInBounds())
261 read.getInBoundsAttr().getValue().take_back(newType.getRank()));
263 Value mask = Value();
265 mask = dropUnitDimsFromMask(rewriter, read.getLoc(), read.getMask(),
268 auto newRead = vector::TransferReadOp::create(
269 rewriter, read.getLoc(), newType, read.getBase(), read.getIndices(),
270 AffineMapAttr::get(newMap), read.getPadding(), mask, inBoundsAttr);
280struct CastAwayTransferWriteLeadingOneDim
284 LogicalResult matchAndRewrite(vector::TransferWriteOp write,
285 PatternRewriter &rewriter)
const override {
287 if (cast<MaskableOpInterface>(write.getOperation()).isMasked())
290 if (write.getTransferRank() == 0)
292 write,
"Nothing to trim - the transfer itself has rank zero");
294 auto shapedType = dyn_cast<ShapedType>(write.getBase().getType());
295 if (shapedType.getElementType() != write.getVectorType().getElementType())
298 VectorType oldType = write.getVectorType();
300 if (newType == oldType)
304 ArrayRef<AffineExpr> newResults =
305 oldMap.
getResults().take_back(newType.getRank());
311 if (write.getInBounds())
313 write.getInBoundsAttr().getValue().take_back(newType.getRank()));
315 auto newVector = rewriter.
createOrFold<vector::ShapeCastOp>(
316 write.getLoc(), newType, write.getVector());
318 if (write.getMask()) {
319 Value newMask = dropUnitDimsFromMask(rewriter, write.getLoc(),
320 write.getMask(), newType, newMap);
322 write, newVector, write.getBase(), write.getIndices(),
323 AffineMapAttr::get(newMap), newMask, inBoundsAttr);
328 write, newVector, write.getBase(), write.getIndices(),
329 AffineMapAttr::get(newMap), inBoundsAttr);
343 auto oldValTy = cast<VectorType>(oldVal.
getType());
344 if (oldValTy.getRank() == 1) {
345 return rewriter.
createOrFold<ExtractOp>(loc, oldVal, 0);
364 if (!isa<VectorType>(oldVal.
getType())) {
365 return rewriter.
createOrFold<BroadcastOp>(loc, newTy, oldVal);
368 return rewriter.
createOrFold<ShapeCastOp>(loc, newTy, oldVal);
372mlir::vector::castAwayContractionLeadingOneDim(vector::ContractionOp contractOp,
373 MaskingOpInterface maskingOp,
375 VectorType oldAccType = dyn_cast<VectorType>(contractOp.getAccType());
376 if (oldAccType ==
nullptr)
378 if (oldAccType.getRank() < 1)
380 if (oldAccType.getShape()[0] != 1 || oldAccType.getScalableDims()[0])
383 auto oldIndexingMaps = contractOp.getIndexingMapsArray();
386 auto oldIteratorTypes = contractOp.getIteratorTypes();
390 int64_t dimToDrop = oldIndexingMaps[2].getDimPosition(0);
396 for (
const auto &it : llvm::enumerate(oldIteratorTypes)) {
398 if (currDim == dimToDrop)
400 newIteratorTypes.push_back(it.value());
404 contractOp.getAcc()};
406 auto loc = contractOp.getLoc();
408 for (
const auto &it : llvm::enumerate(oldIndexingMaps)) {
411 bool validExtract =
false;
413 auto map = it.value();
414 int64_t orginalZeroDim = it.value().getDimPosition(0);
415 if (orginalZeroDim != dimToDrop) {
421 bool transposeNeeded =
false;
425 for (
int64_t i = 0, e = map.getNumResults(); i < e; ++i) {
426 int64_t currDim = map.getDimPosition(i);
427 if (currDim == dimToDrop) {
428 transposeNeeded =
true;
429 perm.insert(perm.begin(), i);
431 transposeResults.insert(transposeResults.begin(), targetExpr);
435 transposeResults.push_back(targetExpr);
442 bool transposeNonOuterUnitDims =
false;
443 auto operandShape = cast<ShapedType>(operands[it.index()].getType());
444 for (
auto [
index, dim] :
447 operandShape.getDimSize(
index) != 1) {
448 transposeNonOuterUnitDims =
true;
455 if (transposeNeeded) {
457 contractOp.getContext());
458 if (transposeNonOuterUnitDims) {
466 operands[it.index()] = rewriter.
createOrFold<vector::TransposeOp>(
467 loc, operands[it.index()], perm);
475 if (map.getDimPosition(0) == dimToDrop)
478 for (
int64_t i = 0, e = map.getNumResults(); i < e; ++i) {
479 int64_t currDim = map.getDimPosition(i);
480 if (currDim == dimToDrop)
484 currDim < dimToDrop ? currDim : currDim - 1);
485 results.push_back(targetExpr);
487 newIndexingMaps.push_back(
AffineMap::get(map.getNumDims() - 1, 0, results,
488 contractOp.getContext()));
491 auto oldVal = operands[it.index()];
492 newOperands.push_back(
500 Operation *newOp = vector::ContractionOp::create(
501 rewriter, loc, newOperands[0], newOperands[1], newOperands[2],
503 rewriter.
getArrayAttr(newIteratorTypes), contractOp.getKind());
510 maskingOp.getMask());
516 rewriter, loc, newOp->
getResult(0), contractOp->getResultTypes()[0]);
530struct CastAwayContractionLeadingOneDim
532 using MaskableOpRewritePattern::MaskableOpRewritePattern;
535 matchAndRewriteMaskableOp(vector::ContractionOp contractOp,
536 MaskingOpInterface maskingOp,
537 PatternRewriter &rewriter)
const override {
538 return castAwayContractionLeadingOneDim(contractOp, maskingOp, rewriter);
556 CastAwayElementwiseLeadingOneDim(MLIRContext *context,
557 PatternBenefit benefit = 1)
558 : RewritePattern(MatchAnyOpTypeTag(), benefit, context) {}
560 LogicalResult matchAndRewrite(Operation *op,
561 PatternRewriter &rewriter)
const override {
568 if (newVecType == vecType)
570 int64_t dropDim = vecType.getRank() - newVecType.getRank();
571 SmallVector<Value, 4> newOperands;
573 if (
auto opVecType = dyn_cast<VectorType>(operand.getType())) {
574 newOperands.push_back(vector::ExtractOp::create(
577 newOperands.push_back(operand);
594 auto oldType = cast<VectorType>(operand.
getType());
597 oldType.getScalableDims().take_front(nDropped);
599 llvm::all_of(leadingShape, [](
int64_t d) {
return d == 1; }) &&
600 llvm::none_of(leadingScalable, [](
bool s) {
return s; });
602 return vector::ExtractOp::create(
b, loc, operand,
splatZero(nDropped));
603 VectorType newType = VectorType::get(
604 oldType.getShape().drop_front(nDropped), oldType.getElementType(),
605 oldType.getScalableDims().drop_front(nDropped));
606 return vector::ShapeCastOp::create(
b, loc, newType, operand);
615template <
typename OpTy>
617 using OpRewritePattern<OpTy>::OpRewritePattern;
619 LogicalResult matchAndRewrite(OpTy op,
620 PatternRewriter &rewriter)
const override {
621 VectorType oldResultType = op.getVectorType();
623 if (newResultType == oldResultType)
625 int64_t nDropped = oldResultType.getRank() - newResultType.getRank();
627 Location loc = op.getLoc();
628 SmallVector<Value> newOperands;
629 newOperands.reserve(op->getNumOperands());
630 for (Value operand : op->getOperands()) {
631 if (isa<VectorType>(operand.getType())) {
632 newOperands.push_back(
635 newOperands.push_back(operand);
650template <
typename OpTy>
652 using OpRewritePattern<OpTy>::OpRewritePattern;
654 LogicalResult matchAndRewrite(OpTy op,
655 PatternRewriter &rewriter)
const override {
656 VectorType oldVecType = op.getVectorType();
658 if (newVecType == oldVecType)
660 int64_t nDropped = oldVecType.getRank() - newVecType.getRank();
662 Location loc = op.getLoc();
663 SmallVector<Value> newOperands;
664 newOperands.reserve(op->getNumOperands());
665 for (Value operand : op->getOperands()) {
666 if (isa<VectorType>(operand.getType())) {
667 newOperands.push_back(
670 newOperands.push_back(operand);
683struct CastAwayConstantMaskLeadingOneDim
687 LogicalResult matchAndRewrite(vector::ConstantMaskOp mask,
688 PatternRewriter &rewriter)
const override {
689 VectorType oldType = mask.
getType();
692 if (newType == oldType)
695 int64_t dropDim = oldType.getRank() - newType.getRank();
696 ArrayRef<int64_t> dimSizes = mask.getMaskDimSizes();
700 int64_t flatLeadingSize =
701 llvm::product_of(dimSizes.take_front(dropDim + 1));
702 SmallVector<int64_t> newDimSizes = {flatLeadingSize};
703 newDimSizes.append(dimSizes.begin() + dropDim + 1, dimSizes.end());
705 auto newMask = vector::ConstantMaskOp::create(rewriter, mask.
getLoc(),
706 newType, newDimSizes);
714void mlir::vector::populateCastAwayVectorLeadingOneDimPatterns(
717 .
add<CastAwayExtractStridedSliceLeadingOneDim,
718 CastAwayInsertStridedSliceLeadingOneDim, CastAwayInsertLeadingOneDim,
719 CastAwayConstantMaskLeadingOneDim, CastAwayTransferReadLeadingOneDim,
720 CastAwayTransferWriteLeadingOneDim, CastAwayElementwiseLeadingOneDim,
721 CastAwayContractionLeadingOneDim,
722 CastAwayLoadLikeLeadingOneDim<vector::LoadOp>,
723 CastAwayLoadLikeLeadingOneDim<vector::MaskedLoadOp>,
724 CastAwayLoadLikeLeadingOneDim<vector::ExpandLoadOp>,
725 CastAwayLoadLikeLeadingOneDim<vector::GatherOp>,
726 CastAwayStoreLikeLeadingOneDim<vector::StoreOp>,
727 CastAwayStoreLikeLeadingOneDim<vector::MaskedStoreOp>,
728 CastAwayStoreLikeLeadingOneDim<vector::CompressStoreOp>,
729 CastAwayStoreLikeLeadingOneDim<vector::ScatterOp>>(
static VectorType trimLeadingUnitDims(VectorType oldType, bool trimOnlyOneDim=false, bool allowRank0=false)
static SmallVector< int64_t > splatZero(int64_t rank)
Return a smallVector of size rank containing all zeros.
static Value restoreLeadingUnitDimViaShapeCastOrBcast(RewriterBase &rewriter, Location loc, mlir::Value oldVal, mlir::Type newTy)
static Value dropLeadingUnitDimViaShapeCastOrExtract(RewriterBase &rewriter, Location loc, mlir::Value oldVal)
static Value dropLeadingOneDimsFromOperand(OpBuilder &b, Location loc, Value operand, int64_t nDropped)
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.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
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.