27 for (
unsigned pos : permutation)
28 newInBoundsValues[pos] =
29 cast<BoolAttr>(attr.getValue()[
index++]).getValue();
37 auto originalVecType = cast<VectorType>(vec.
getType());
39 newShape.append(originalVecType.getShape().begin(),
40 originalVecType.getShape().end());
43 newScalableDims.append(originalVecType.getScalableDims().begin(),
44 originalVecType.getScalableDims().end());
45 VectorType newVecType = VectorType::get(
46 newShape, originalVecType.getElementType(), newScalableDims);
47 return vector::BroadcastOp::create(builder, loc, newVecType, vec);
57 e = cast<VectorType>(broadcasted.
getType()).getRank();
59 permutation.push_back(i);
60 for (
int64_t i = 0; i < addedRank; ++i)
61 permutation.push_back(i);
62 return vector::TransposeOp::create(builder, loc, broadcasted, permutation);
89struct TransferReadPermutationLowering
91 using MaskableOpRewritePattern::MaskableOpRewritePattern;
93 FailureOr<mlir::Value>
94 matchAndRewriteMaskableOp(vector::TransferReadOp op,
95 MaskingOpInterface maskOp,
96 PatternRewriter &rewriter)
const override {
98 if (op.getTransferRank() == 0)
101 SmallVector<unsigned> permutation;
102 AffineMap map = op.getPermutationMap();
107 op,
"map is not permutable to minor identity, apply another pattern");
109 AffineMap permutationMap =
117 AffineMap newMap = permutationMap.
compose(map);
119 ArrayRef<int64_t> originalShape = op.getVectorType().getShape();
120 SmallVector<int64_t> newVectorShape(originalShape.size());
121 ArrayRef<bool> originalScalableDims = op.getVectorType().getScalableDims();
122 SmallVector<bool> newScalableDims(originalShape.size());
123 for (
const auto &pos : llvm::enumerate(permutation)) {
124 newVectorShape[pos.value()] = originalShape[pos.index()];
125 newScalableDims[pos.value()] = originalScalableDims[pos.index()];
133 VectorType newReadType = VectorType::get(
134 newVectorShape, op.getVectorType().getElementType(), newScalableDims);
135 Operation *newRead = vector::TransferReadOp::create(
136 rewriter, op.getLoc(), newReadType, op.getBase(), op.getIndices(),
137 AffineMapAttr::get(newMap), op.getPadding(), op.getMask(),
140 SmallVector<int64_t> transposePerm(permutation.begin(), permutation.end());
145 Value passthru = maskOp.getPassthru();
148 vector::TransposeOp::create(rewriter, op.getLoc(), passthru,
155 return vector::TransposeOp::create(rewriter, op.getLoc(),
177struct TransferWritePermutationLowering
179 using MaskableOpRewritePattern::MaskableOpRewritePattern;
181 FailureOr<mlir::Value>
182 matchAndRewriteMaskableOp(vector::TransferWriteOp op,
183 MaskingOpInterface maskOp,
184 PatternRewriter &rewriter)
const override {
186 if (op.getTransferRank() == 0)
189 SmallVector<unsigned> permutation;
196 op,
"map is not permutable to minor identity, apply another pattern");
207 [](AffineExpr expr) {
208 return dyn_cast<AffineDimExpr>(expr).getPosition();
216 Value newVec = vector::TransposeOp::create(rewriter, op.getLoc(),
220 auto newWrite = vector::TransferWriteOp::create(
221 rewriter, op.getLoc(), newVec, op.getBase(), op.getIndices(),
222 AffineMapAttr::get(newMap), op.getMask(), newInBoundsAttr);
226 Operation *rewritten = newWrite;
230 if (newWrite.hasPureTensorSemantics())
253struct TransferWriteNonPermutationLowering
255 using MaskableOpRewritePattern::MaskableOpRewritePattern;
257 FailureOr<mlir::Value>
258 matchAndRewriteMaskableOp(vector::TransferWriteOp op,
259 MaskingOpInterface maskOp,
260 PatternRewriter &rewriter)
const override {
262 if (op.getTransferRank() == 0)
268 SmallVector<unsigned> permutation;
273 "map is already permutable to minor identity, apply another pattern");
278 SmallVector<bool> foundDim(map.
getNumDims(),
false);
280 foundDim[cast<AffineDimExpr>(exp).getPosition()] =
true;
281 SmallVector<AffineExpr> exprs;
282 bool foundFirstDim =
false;
283 SmallVector<int64_t> missingInnerDim;
284 for (
size_t i = 0; i < foundDim.size(); i++) {
286 foundFirstDim =
true;
293 missingInnerDim.push_back(i);
298 missingInnerDim.size());
303 missingInnerDim.size());
308 SmallVector<bool> newInBoundsValues(missingInnerDim.size(),
true);
309 for (int64_t i = 0, e = op.getVectorType().getRank(); i < e; ++i) {
310 newInBoundsValues.push_back(op.isDimInBounds(i));
313 auto newWrite = vector::TransferWriteOp::create(
314 rewriter, op.getLoc(), newVec, op.getBase(), op.getIndices(),
315 AffineMapAttr::get(newMap), newMask, newInBoundsAttr);
316 if (newWrite.hasPureTensorSemantics())
317 return newWrite.getResult();
332struct TransferOpReduceRank
334 using MaskableOpRewritePattern::MaskableOpRewritePattern;
336 FailureOr<mlir::Value>
337 matchAndRewriteMaskableOp(vector::TransferReadOp op,
338 MaskingOpInterface maskOp,
339 PatternRewriter &rewriter)
const override {
341 if (op.getTransferRank() == 0)
348 unsigned numLeadingBroadcast = 0;
350 auto dimExpr = dyn_cast<AffineConstantExpr>(expr);
351 if (!dimExpr || dimExpr.getValue() != 0)
353 numLeadingBroadcast++;
356 if (numLeadingBroadcast == 0)
359 VectorType originalVecType = op.getVectorType();
360 unsigned reducedShapeRank = originalVecType.getRank() - numLeadingBroadcast;
369 op,
"map is not a minor identity with broadcasting");
372 SmallVector<int64_t> newShape(
373 originalVecType.getShape().take_back(reducedShapeRank));
374 SmallVector<bool> newScalableDims(
375 originalVecType.getScalableDims().take_back(reducedShapeRank));
377 VectorType newReadType = VectorType::get(
378 newShape, originalVecType.getElementType(), newScalableDims);
382 op.getInBoundsAttr().getValue().take_back(reducedShapeRank))
384 Value newRead = vector::TransferReadOp::create(
385 rewriter, op.getLoc(), newReadType, op.getBase(), op.getIndices(),
386 AffineMapAttr::get(newMap), op.getPadding(), op.getMask(),
388 return vector::BroadcastOp::create(rewriter, op.getLoc(), originalVecType,
399 .
add<TransferReadPermutationLowering, TransferWritePermutationLowering,
400 TransferOpReduceRank, TransferWriteNonPermutationLowering>(
417struct TransferReadToVectorLoadLowering
419 TransferReadToVectorLoadLowering(
MLIRContext *context,
420 std::optional<unsigned> maxRank,
423 maxTransferRank(maxRank) {}
425 FailureOr<mlir::Value>
426 matchAndRewriteMaskableOp(vector::TransferReadOp read,
427 MaskingOpInterface maskOp,
429 if (maxTransferRank && read.getVectorType().getRank() > *maxTransferRank) {
431 read,
"vector type is greater than max transfer rank");
436 SmallVector<unsigned> broadcastedDims;
440 if (!read.getPermutationMap().isMinorIdentityWithBroadcasting(
444 auto memRefType = dyn_cast<MemRefType>(read.getShapedType());
449 if (!memRefType.isLastDimUnitStride())
454 ArrayRef<int64_t>
vectorShape = read.getVectorType().getShape();
455 SmallVector<int64_t> unbroadcastedVectorShape(
vectorShape);
456 for (
unsigned i : broadcastedDims)
457 unbroadcastedVectorShape[i] = 1;
458 VectorType unbroadcastedVectorType = read.getVectorType().cloneWith(
459 unbroadcastedVectorShape, read.getVectorType().getElementType());
463 auto memrefElTy = memRefType.getElementType();
464 if (isa<VectorType>(memrefElTy) && memrefElTy != unbroadcastedVectorType)
468 if (!isa<VectorType>(memrefElTy) &&
469 memrefElTy != read.getVectorType().getElementType())
473 if (read.hasOutOfBoundsDim())
478 if (read.getMask()) {
479 if (read.getVectorType().getRank() != 1)
482 read,
"vector type is not rank 1, can't create masked load, needs "
485 Value fill = vector::BroadcastOp::create(
486 rewriter, read.getLoc(), unbroadcastedVectorType, read.getPadding());
487 res = vector::MaskedLoadOp::create(
488 rewriter, read.getLoc(), unbroadcastedVectorType, read.getBase(),
489 read.getIndices(), read.getMask(), fill);
491 res = vector::LoadOp::create(rewriter, read.getLoc(),
492 unbroadcastedVectorType, read.getBase(),
497 if (!broadcastedDims.empty())
498 res = vector::BroadcastOp::create(
499 rewriter, read.getLoc(), read.getVectorType(), res->
getResult(0));
503 std::optional<unsigned> maxTransferRank;
514struct TransferWriteToVectorStoreLowering
516 TransferWriteToVectorStoreLowering(MLIRContext *context,
517 std::optional<unsigned> maxRank,
518 PatternBenefit benefit = 1)
519 : MaskableOpRewritePattern<vector::TransferWriteOp>(context, benefit),
520 maxTransferRank(maxRank) {}
522 FailureOr<mlir::Value>
523 matchAndRewriteMaskableOp(vector::TransferWriteOp write,
524 MaskingOpInterface maskOp,
525 PatternRewriter &rewriter)
const override {
526 if (maxTransferRank && write.getVectorType().getRank() > *maxTransferRank) {
528 write,
"vector type is greater than max transfer rank");
536 !write.getPermutationMap().isMinorIdentity())
538 diag <<
"permutation map is not minor identity: " << write;
541 auto memRefType = dyn_cast<MemRefType>(write.getShapedType());
544 diag <<
"not a memref type: " << write;
548 if (!memRefType.isLastDimUnitStride())
550 diag <<
"most minor stride is not 1: " << write;
555 auto memrefElTy = memRefType.getElementType();
556 if (isa<VectorType>(memrefElTy) && memrefElTy != write.getVectorType())
558 diag <<
"elemental type mismatch: " << write;
562 if (!isa<VectorType>(memrefElTy) &&
563 memrefElTy != write.getVectorType().getElementType())
565 diag <<
"elemental type mismatch: " << write;
569 if (write.hasOutOfBoundsDim())
571 diag <<
"out of bounds dim: " << write;
573 if (write.getMask()) {
574 if (write.getVectorType().getRank() != 1)
577 write.getLoc(), [=](Diagnostic &
diag) {
578 diag <<
"vector type is not rank 1, can't create masked store, "
579 "needs VectorToSCF: "
583 vector::MaskedStoreOp::create(rewriter, write.getLoc(), write.getBase(),
584 write.getIndices(), write.getMask(),
587 vector::StoreOp::create(rewriter, write.getLoc(), write.getVector(),
588 write.getBase(), write.getIndices());
595 std::optional<unsigned> maxTransferRank;
602 patterns.
add<TransferReadToVectorLoadLowering,
603 TransferWriteToVectorStoreLowering>(patterns.
getContext(),
604 maxTransferRank, benefit);
static ArrayAttr inverseTransposeInBoundsAttr(OpBuilder &builder, ArrayAttr attr, const SmallVector< unsigned > &permutation)
Transpose a vector transfer op's in_bounds attribute by applying reverse permutation based on the giv...
static Value extendMaskRank(OpBuilder &builder, Location loc, Value vec, int64_t addedRank)
Extend the rank of a vector Value by addedRanks by adding inner unit dimensions.
static Value extendVectorRank(OpBuilder &builder, Location loc, Value vec, int64_t addedRank)
Extend the rank of a vector Value by addedRanks by adding outer unit dimensions.
static std::string diag(const llvm::Value &value)
static std::optional< VectorShape > vectorShape(Type type)
static AffineMap getMinorIdentityMap(unsigned dims, unsigned results, MLIRContext *context)
Returns an identity affine map (d0, ..., dn) -> (dp, ..., dn) on the most minor dimensions.
bool isMinorIdentity() const
Returns true if this affine map is a minor identity, i.e.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
bool isMinorIdentityWithBroadcasting(SmallVectorImpl< unsigned > *broadcastedDims=nullptr) const
Returns true if this affine map is a minor identity up to broadcasted dimensions which are indicated ...
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
bool isPermutationOfMinorIdentityWithBroadcasting(SmallVectorImpl< unsigned > &permutedDims) const
Return true if this affine map can be converted to a minor identity with broadcast by doing a permute...
unsigned getNumResults() const
static AffineMap getPermutationMap(ArrayRef< unsigned > permutation, MLIRContext *context)
Returns an AffineMap representing a permutation.
AffineMap compose(AffineMap map) const
Returns the AffineMap resulting from composing this with map.
bool isIdentity() const
Returns true if this affine map is an identity affine map.
AffineExpr getAffineDimExpr(unsigned position)
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
MLIRContext * getContext() const
ArrayAttr getBoolArrayAttr(ArrayRef< bool > values)
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
This class helps build Operations.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
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.
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,...
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.
void populateVectorTransferPermutationMapLoweringPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Collect a set of transfer read/write lowering patterns that simplify the permutation map (e....
Operation * maskOperation(OpBuilder &builder, Operation *maskableOp, Value mask, Value passthru=Value())
Creates a vector.mask operation around a maskable operation.
void populateVectorTransferLoweringPatterns(RewritePatternSet &patterns, std::optional< unsigned > maxTransferRank=std::nullopt, PatternBenefit benefit=1)
Populate the pattern set with the following patterns:
Include the generated interface declarations.
AffineMap inversePermutation(AffineMap map)
Returns a map of codomain to domain dimensions such that the first codomain dimension for a particula...
AffineMap compressUnusedDims(AffineMap map)
Drop the dims that are not used.
SmallVector< int64_t > invertPermutationVector(ArrayRef< int64_t > permutation)
Helper method to apply to inverse a permutation.
A pattern for ops that implement MaskableOpInterface and that might be masked (i.e.