38#include "llvm/ADT/STLExtras.h"
39#include "llvm/ADT/Sequence.h"
40#include "llvm/ADT/SmallPtrSet.h"
41#include "llvm/ADT/SmallVector.h"
42#include "llvm/ADT/SmallVectorExtras.h"
43#include "llvm/ADT/TypeSwitch.h"
44#include "llvm/Support/DebugLog.h"
45#include "llvm/Support/InterleavedRange.h"
46#include "llvm/Support/MathExtras.h"
47#include "llvm/Support/raw_ostream.h"
53#define DEBUG_TYPE "linalg-vectorization"
56static FailureOr<Operation *>
60 bool flatten1DDepthwiseConv =
false);
98 int64_t kwSize,
int strideW,
int dilationW,
99 int64_t wSizeStep,
bool isSingleChanneled) {
101 if (isSingleChanneled) {
106 for (
int64_t kw = 0; kw < kwSize; ++kw) {
107 for (
int64_t w = 0; w < wSize; w += wSizeStep) {
108 result.push_back(vector::ExtractStridedSliceOp::create(
118 for (
int64_t kw = 0; kw < kwSize; ++kw) {
119 for (
int64_t w = 0; w < wSize; w += wSizeStep) {
120 result.push_back(vector::ExtractStridedSliceOp::create(
121 rewriter, loc, input,
138 for (
int64_t kw = 0; kw < kwSize; ++kw) {
139 result.push_back(vector::ExtractOp::create(
150 int64_t wSizeStep,
bool isSingleChanneled) {
152 if (isSingleChanneled) {
156 for (
int64_t w = 0; w < wSize; w += wSizeStep) {
157 result.push_back(vector::ExtractStridedSliceOp::create(
166 for (
int64_t w = 0; w < wSize; w += wSizeStep) {
167 result.push_back(vector::ExtractStridedSliceOp::create(
179 bool isSingleChanneled) {
181 if (isSingleChanneled) {
185 for (
int64_t w = 0; w < wSize; w += wSizeStep) {
186 res = vector::InsertStridedSliceOp::create(
194 for (
int64_t w = 0; w < wSize; w += wSizeStep) {
195 res = vector::InsertStridedSliceOp::create(
196 rewriter, loc, resVals[w], res,
210 LogicalResult initState(
RewriterBase &rewriter, LinalgOp linalgOp,
213 bool assumeDynamicDimsMatchVecSizes =
false);
228 std::optional<AffineMap> dimPermutation = std::nullopt)
const {
231 if (dimPermutation.has_value()) {
237 vectorShape.append(canonicalVecShape.begin(), canonicalVecShape.end());
238 scalableDims.append(scalableVecDims.begin(), scalableVecDims.end());
241 return VectorType::get(
vectorShape, elementType, scalableDims);
250 std::optional<AffineMap> maybeIndexingMap = std::nullopt);
255 void initIterSpaceStaticSizes(LinalgOp linalgOp) {
256 iterSpaceStaticSizes.append(linalgOp.getStaticLoopRanges());
262 LogicalResult precomputeIterSpaceValueSizes(RewriterBase &rewriter,
269 Value getOrCreateMaskFor(RewriterBase &rewriter, Operation *opToMask,
271 std::optional<AffineMap> maybeMaskingMap);
276 bool isValidMaskingMap(AffineMap maskingMap) {
295 AffineMap getMaskingMapFromIndexingMap(AffineMap &indexingMap) {
301 SmallVector<int64_t> iterSpaceStaticSizes;
306 SmallVector<Value> iterSpaceValueSizes;
309 SmallVector<int64_t> canonicalVecShape;
313 SmallVector<bool> scalableVecDims;
321 OpBuilder::InsertionGuard rewriterGuard;
329 bool assumeDynamicDimsMatchVecSizes =
false;
333VectorizationState::precomputeIterSpaceValueSizes(
RewriterBase &rewriter,
336 for (
int vecDim = 0, end = canonicalVecShape.size(); vecDim < end; ++vecDim) {
337 if (ShapedType::isStatic(iterSpaceStaticSizes[vecDim])) {
340 rewriter, linalgOp.getLoc(), iterSpaceStaticSizes[vecDim]));
347 unsigned operandDimPos;
348 if (
failed(linalgOp.mapIterationSpaceDimToOperandDim(vecDim, operand,
353 linalgOp.hasPureTensorSemantics()
354 ? (Value)tensor::DimOp::create(rewriter, linalgOp.getLoc(), operand,
356 : (Value)memref::DimOp::create(rewriter, linalgOp.getLoc(), operand,
358 iterSpaceValueSizes.push_back(dynamicDim);
371 bool assumeDimsMatchVec) {
372 assumeDynamicDimsMatchVecSizes = assumeDimsMatchVec;
376 if (!inputVectorSizes.empty()) {
380 canonicalVecShape.append(inputVectorSizes.begin(), inputVectorSizes.end());
381 scalableVecDims.append(inputScalableVecDims.begin(),
382 inputScalableVecDims.end());
387 canonicalVecShape = linalgOp.getStaticLoopRanges();
388 scalableVecDims.append(linalgOp.getNumLoops(),
false);
391 LDBG() <<
"Canonical vector shape: " << llvm::interleaved(canonicalVecShape);
392 LDBG() <<
"Scalable vector dims: " << llvm::interleaved(scalableVecDims);
394 if (ShapedType::isDynamicShape(canonicalVecShape))
398 initIterSpaceStaticSizes(linalgOp);
403 if (failed(precomputeIterSpaceValueSizes(rewriter, linalgOp)))
413Value VectorizationState::getOrCreateMaskFor(
415 std::optional<AffineMap> maybeMaskingMap) {
417 assert((!maybeMaskingMap || isValidMaskingMap(*maybeMaskingMap)) &&
418 "Ill-formed masking map.");
421 auto maskableOp = dyn_cast<vector::MaskableOpInterface>(opToMask);
425 assert(!maskableOp.isMasked() &&
426 "Masking an operation that is already masked");
429 assert((!maybeMaskingMap || *maybeMaskingMap) &&
430 "Unexpected null mask permutation map");
432 maybeMaskingMap ? *maybeMaskingMap
434 linalgOp.getNumLoops(), rewriter.
getContext());
436 LDBG() <<
"Masking map: " << maskingMap;
440 auto activeMaskIt = activeMaskCache.find(maskingMap);
441 if (activeMaskIt != activeMaskCache.end()) {
442 Value mask = activeMaskIt->second;
443 LDBG() <<
"Reusing mask: " << mask;
453 SmallVector<int64_t> permutedStaticSizes =
455 auto maskType = getCanonicalVecType(rewriter.
getI1Type(), maskingMap);
456 auto maskShape = maskType.getShape();
458 LDBG() <<
"Mask shape: " << llvm::interleaved(maskShape);
460 if (permutedStaticSizes == maskShape) {
461 LDBG() <<
"Masking is not needed for masking map: " << maskingMap;
462 activeMaskCache[maskingMap] = Value();
466 if (assumeDynamicDimsMatchVecSizes) {
470 if (llvm::all_of(llvm::zip(permutedStaticSizes, maskType.getShape()),
472 return std::get<0>(it) == ShapedType::kDynamic
474 : std::get<0>(it) == std::get<1>(it);
477 <<
"Dynamic + static dimensions match vector sizes, masking is not "
479 activeMaskCache[maskingMap] = Value();
485 SmallVector<Value> upperBounds =
487 assert(!maskShape.empty() && !upperBounds.empty() &&
488 "Masked 0-d vectors are not supported yet");
491 Value mask = vector::CreateMaskOp::create(rewriter, linalgOp.getLoc(),
492 maskType, upperBounds);
493 LDBG() <<
"Creating new mask: " << mask;
494 activeMaskCache[maskingMap] = mask;
501 std::optional<AffineMap> maybeIndexingMap) {
502 LDBG() <<
"Trying to mask: " << *opToMask;
504 std::optional<AffineMap> maybeMaskingMap = std::nullopt;
505 if (maybeIndexingMap)
506 maybeMaskingMap = getMaskingMapFromIndexingMap(*maybeIndexingMap);
510 getOrCreateMaskFor(rewriter, opToMask, linalgOp, maybeMaskingMap);
513 LDBG() <<
"No mask required";
514 if (assumeDynamicDimsMatchVecSizes) {
516 .Case<vector::TransferReadOp, vector::TransferWriteOp>(
522 LDBG() <<
"Assuming dynamic dimensions match vector sizes and "
523 "setting their in-bounds to true!";
525 ShapedType xferType = xferOp.getShapedType();
530 for (
unsigned i = 0; i < xferOp.getTransferRank(); i++) {
531 auto dimExpr = dyn_cast<AffineDimExpr>(permMap.
getResult(i));
535 unsigned pos = dimExpr.getPosition();
536 if (xferType.isDynamicDim(pos))
537 inBoundsMap[i] =
true;
540 xferOp.setInBoundsAttr(
552 assert(opToMask &&
"Expected a valid operation to mask");
553 auto maskOp = cast<vector::MaskOp>(
555 Operation *maskOpTerminator = &maskOp.getMaskRegion().front().back();
557 for (
auto [resIdx, resVal] : llvm::enumerate(opToMask->
getResults()))
561 LDBG() <<
"Masked operation: " << *maskOp;
584 "expected projected permutation");
586 assert(res.getNumDims() ==
587 (res.getNumResults() - res.getNumOfZeroResults()) &&
588 "expected reindexed map with same number of dims and results");
624std::optional<vector::CombiningKind>
626 using ::mlir::vector::CombiningKind;
631 .Case<arith::AddIOp, arith::AddFOp>(
632 [&](
auto op) {
return CombiningKind::ADD; })
633 .Case([&](arith::AndIOp op) {
return CombiningKind::AND; })
634 .Case([&](arith::MaxSIOp op) {
return CombiningKind::MAXSI; })
635 .Case([&](arith::MaxUIOp op) {
return CombiningKind::MAXUI; })
636 .Case([&](arith::MaximumFOp op) {
return CombiningKind::MAXIMUMF; })
637 .Case([&](arith::MaxNumFOp op) {
return CombiningKind::MAXNUMF; })
638 .Case([&](arith::MaximumNumFOp op) {
return CombiningKind::MAXIMUMNUMF; })
639 .Case([&](arith::MinSIOp op) {
return CombiningKind::MINSI; })
640 .Case([&](arith::MinUIOp op) {
return CombiningKind::MINUI; })
641 .Case([&](arith::MinimumFOp op) {
return CombiningKind::MINIMUMF; })
642 .Case([&](arith::MinNumFOp op) {
return CombiningKind::MINNUMF; })
643 .Case([&](arith::MinimumNumFOp op) {
return CombiningKind::MINIMUMNUMF; })
644 .Case<arith::MulIOp, arith::MulFOp>(
645 [&](
auto op) {
return CombiningKind::MUL; })
646 .Case([&](arith::OrIOp op) {
return CombiningKind::OR; })
647 .Case([&](arith::XOrIOp op) {
return CombiningKind::XOR; })
648 .Default(std::nullopt);
659 auto linalgOp = cast<LinalgOp>(outputOperand->
getOwner());
664 if (!
matchReduction(linalgOp.getRegionOutputArgs(), outputPos, combinerOps) ||
665 combinerOps.size() != 1)
669 return combinerOps[0];
675 auto dstVecType = dyn_cast<VectorType>(dstType);
677 if (dstVecType.getRank() == 0)
682 Location loc =
b.getInsertionPoint()->getLoc();
683 return b.createOrFold<vector::BroadcastOp>(loc, dstVecType, value);
695 assert(maybeKind &&
"Failed precondition: could not get reduction kind");
696 return vector::MultiDimReductionOp::create(
697 b, reduceOp->
getLoc(), valueToReduce,
acc, dimsToMask, *maybeKind);
701 return llvm::map_to_vector(linalgOp.getIteratorTypesArray(),
708 return isa<linalg::ReduceOp>(op) ||
709 (isa<linalg::GenericOp>(op) &&
721 VectorizationState &state) {
723 auto linalgOp = cast<LinalgOp>(outputOperand->
getOwner());
724 AffineMap opOperandMap = linalgOp.getMatchingIndexingMap(outputOperand);
733 return llvm::is_contained(opOperandMap.getResults(), dimExpr);
735 auto vectorType = state.getCanonicalVecType(
742 if (vectorType.getRank() > 0) {
745 assert(value.
getType() == vectorType &&
"Incorrect type");
746 write = vector::TransferWriteOp::create(
747 rewriter, loc, value, outputOperand->
get(),
indices, writeMap);
750 if (!isa<VectorType>(value.
getType()))
751 value = vector::BroadcastOp::create(rewriter, loc, vectorType, value);
752 assert(value.
getType() == vectorType &&
"Incorrect type");
753 write = vector::TransferWriteOp::create(rewriter, loc, value,
757 write = state.maskOperation(rewriter, write, linalgOp, opOperandMap);
761 if (
auto maskOp = dyn_cast<vector::MaskingOpInterface>(write)) {
762 auto maskedWriteOp = cast<vector::TransferWriteOp>(maskOp.getMaskableOp());
767 LDBG() <<
"vectorized op: " << *write;
777 std::function<LogicalResult(
Operation *,
bool)>;
794 const IRMapping &bvm, VectorizationState &state,
796 auto yieldOp = dyn_cast<linalg::YieldOp>(op);
799 for (
const auto &output : llvm::enumerate(yieldOp.getValues())) {
805 linalgOp.getDpsInitOperand(output.index()), state);
807 newResults.push_back(newResult);
818 VectorizationState &state,
821 IndexOp indexOp = dyn_cast<linalg::IndexOp>(op);
824 auto loc = indexOp.getLoc();
827 auto dim = indexOp.getDim();
829 auto indexVectorType =
830 VectorType::get({targetShape[dim]}, rewriter.
getIndexType(),
831 state.getScalableVecDims()[dim]);
832 auto indexSteps = vector::StepOp::create(rewriter, loc, indexVectorType);
836 if (dim == targetShape.size() - 1)
842 llvm::to_vector(llvm::seq<unsigned>(0, targetShape.size()));
843 std::swap(permPattern[dim], permPattern.back());
847 auto broadCastOp = vector::BroadcastOp::create(
849 state.getCanonicalVecType(rewriter.
getIndexType(), permMap), indexSteps);
851 llvm::to_vector<16>(llvm::seq<int64_t>(0, linalgOp.getNumLoops()));
852 std::swap(transposition.back(), transposition[dim]);
854 vector::TransposeOp::create(rewriter, loc, broadCastOp, transposition);
862 tensor::ExtractOp extractOp = dyn_cast<tensor::ExtractOp>(op);
866 if (extractOp.getIndices().size() != 1 && !vectorizeNDExtract)
871 if (not extractOp.getIndices().empty()) {
872 if (!VectorType::isValidElementType(extractOp.getIndices()[0].getType()))
876 if (!llvm::all_of(extractOp->getResultTypes(),
877 VectorType::isValidElementType)) {
895 VectorizationState &state,
896 tensor::ExtractOp extractOp,
899 auto indexVecType = state.getCanonicalVecType(rewriter.
getIndexType());
900 auto loc = extractOp.getLoc();
903 rewriter, bvm.
lookup(extractOp.getIndices()[0]), indexVecType);
905 const size_t numIndices = extractOp.getIndices().size();
906 for (
size_t i = 1; i < numIndices; i++) {
911 tensor::DimOp::create(rewriter, loc, extractOp.getTensor(), dimIdx),
914 offset = arith::MulIOp::create(rewriter, loc, offset, dimSize);
917 rewriter, bvm.
lookup(extractOp.getIndices()[i]), indexVecType);
919 offset = arith::AddIOp::create(rewriter, loc, extractOpIndex, offset);
945 (linalgOp.hasDynamicShape() ||
946 llvm::count_if(loopRanges, [](
int64_t dim) { return dim != 1; }) == 1) &&
947 "For statically shaped Linalg Ops, only one "
948 "non-unit loop dim is expected");
949 assert(!loopRanges.empty() &&
"Empty loops, nothing to analyse.");
951 size_t idx = loopRanges.size() - 1;
952 for (; idx != 0; idx--)
953 if (loopRanges[idx] != 1)
961 VectorType resType) {
963 assert(((llvm::count_if(resType.getShape(),
964 [](
int64_t dimSize) { return dimSize > 1; }) == 1)) &&
965 "n-D vectors are not yet supported");
967 auto *block = linalgOp.getBlock();
973 while (!worklist.empty()) {
974 Value v = worklist.pop_back_val();
980 if (isa<BlockArgument>(v)) {
981 if (llvm::is_contained(block->getArguments(), v))
987 assert(defOp &&
"This is neither a block argument nor an operation result");
992 if (
auto indexOp = dyn_cast<linalg::IndexOp>(defOp)) {
993 if (linalgOp.getStaticLoopRanges()[indexOp.getDim()] != 1)
998 auto *ancestor = block->findAncestorOpInBlock(*defOp);
1005 if (isa<arith::ConstantOp>(ancestor))
1008 if (visited.insert(ancestor).second)
1009 llvm::append_range(worklist, ancestor->getOperands());
1033 bool &foundIndexOp, VectorType resType) {
1035 assert(((llvm::count_if(resType.getShape(),
1036 [](
int64_t dimSize) { return dimSize > 1; }) == 1)) &&
1037 "n-D vectors are not yet supported");
1043 auto *block = linalgOp.getBlock();
1044 if (isa<BlockArgument>(val))
1045 return !llvm::is_contained(block->getArguments(), val);
1048 assert(defOp &&
"This is neither a block argument nor an operation result");
1050 if (
auto indexOp = dyn_cast<linalg::IndexOp>(defOp)) {
1053 foundIndexOp = (indexOp.getDim() == loopDimThatIncrementsByOne);
1057 auto *ancestor = block->findAncestorOpInBlock(*defOp);
1064 if (!isa<arith::AddIOp, arith::ConstantOp, linalg::IndexOp>(ancestor))
1068 for (
auto op : ancestor->getOperands())
1088 LinalgOp &linalgOp, VectorType resType) {
1090 auto inputShape = cast<ShapedType>(extractOp.getTensor().getType());
1093 if (inputShape.getShape().empty())
1098 if (resType.getRank() == 0)
1103 bool isOutput1DVector =
1104 (llvm::count_if(resType.getShape(),
1105 [](
int64_t dimSize) { return dimSize > 1; }) == 1);
1107 if (!isOutput1DVector)
1110 bool leadingIdxsLoopInvariant =
true;
1116 auto indices = extractOp.getIndices();
1117 auto leadIndices =
indices.drop_back(1);
1119 for (
auto [i, indexVal] : llvm::enumerate(leadIndices)) {
1120 if (inputShape.getShape()[i] == 1)
1126 if (!leadingIdxsLoopInvariant) {
1127 LDBG() <<
"Found gather load: " << extractOp;
1135 auto extractOpTrailingIdx =
indices.back();
1139 if (leadingIdxsLoopInvariant &&
1141 LDBG() <<
"Found scalar broadcast load: " << extractOp;
1150 bool foundIndexOp =
false;
1152 foundIndexOp, resType);
1155 bool isRowVector = resType.getShape().back() != 1;
1156 isContiguousLoad &= (foundIndexOp && isRowVector);
1158 if (isContiguousLoad) {
1159 LDBG() <<
"Found contigous load: " << extractOp;
1164 LDBG() <<
"Found gather load: " << extractOp;
1172static VectorizationHookResult
1175 tensor::ExtractOp extractOp = dyn_cast<tensor::ExtractOp>(op);
1178 auto loc = extractOp.getLoc();
1181 auto resultType = state.getCanonicalVecType(extractOp.getResult().getType());
1182 auto maskConstantOp = arith::ConstantOp::create(
1186 auto passThruConstantOp = arith::ConstantOp::create(
1192 extractOp.getIndices().size(),
1203 Operation *gatherOp = vector::GatherOp::create(
1204 rewriter, loc, resultType, extractOp.getTensor(), baseIndices, offset,
1205 maskConstantOp, passThruConstantOp);
1206 gatherOp = state.maskOperation(rewriter, gatherOp, linalgOp);
1208 LDBG() <<
"Vectorised as gather load: " << extractOp;
1231 for (
size_t i = 0; i < extractOp.getIndices().size(); i++) {
1232 Value idx = bvm.
lookup(extractOp.getIndices()[i]);
1234 transferReadIdxs.push_back(idx);
1238 auto indexAs1dVector = vector::ShapeCastOp::create(
1240 VectorType::get(resultType.getShape().back(), rewriter.
getIndexType(),
1241 resultType.getScalableDims().back()),
1243 transferReadIdxs.push_back(
1244 vector::ExtractOp::create(rewriter, loc, indexAs1dVector, 0));
1248 auto dstRank = resultType.getRank();
1249 auto srcRank = extractOp.getTensor().getType().getRank();
1258 auto transferReadOp = vector::TransferReadOp::create(
1259 rewriter, loc, resultType, extractOp.getTensor(), transferReadIdxs,
1260 std::nullopt, permutationMap, inBounds);
1262 Operation *readOrMaskedReadOp = transferReadOp;
1268 auto readMaskType = VectorType::get(readMaskShape, rewriter.
getI1Type());
1269 auto allTrue = vector::ConstantMaskOp::create(
1271 readOrMaskedReadOp =
1275 LDBG() <<
"Vectorised as scalar broadcast load: " << extractOp;
1277 readOrMaskedReadOp};
1282 srcRank, std::min(dstRank, srcRank), rewriter.
getContext());
1284 int32_t rankDiff = dstRank - srcRank;
1292 while (rankDiff > 0) {
1293 permutationMap = permutationMap.insertResult(
1298 auto transferReadOp = vector::TransferReadOp::create(
1299 rewriter, loc, resultType, extractOp.getTensor(), transferReadIdxs,
1300 std::nullopt, permutationMap, inBounds);
1309 int64_t numReadDims = std::min(dstRank, srcRank);
1311 linalgOp.getNumLoops(), numReadDims, rewriter.
getContext());
1313 state.maskOperation(rewriter, transferReadOp, linalgOp, maskingMap);
1315 LDBG() <<
"Vectorised as contiguous load: " << extractOp;
1328 auto reduceType = dyn_cast<VectorType>(reduceVec.
getType());
1329 auto outputType = dyn_cast<VectorType>(outputVec.
getType());
1333 (outputType && reduceType.getShape() == outputType.getShape()))
1358static VectorizationHookResult
1362 LDBG() <<
"vectorize op " << *op;
1365 if (!customVectorizationHooks.empty()) {
1366 for (
auto &customFunc : customVectorizationHooks) {
1376 if (isa<arith::ConstantOp, func::ConstantOp>(op))
1378 rewriter.
clone(*op)};
1387 auto blockArg = dyn_cast<BlockArgument>(operand);
1388 if (!blockArg || blockArg.getOwner() != linalgOp.getBlock() ||
1389 blockArg.getArgNumber() < linalgOp.getNumDpsInputs())
1393 linalgOp.getRegionOutputArgs(),
1394 blockArg.getArgNumber() - linalgOp.getNumDpsInputs(), reductionOps);
1397 reductionOperands.push_back(std::make_pair(reduceValue, operand));
1399 if (!reductionOperands.empty()) {
1400 assert(reductionOperands.size() == 1);
1402 reduceIfNeeded(rewriter, linalgOp, op, reductionOperands[0].first,
1403 reductionOperands[0].second, bvm);
1410 VectorType firstMaxRankedType;
1412 auto vecOperand = bvm.
lookup(operand);
1413 assert(vecOperand &&
"Vector operand couldn't be found");
1415 auto vecType = dyn_cast<VectorType>(vecOperand.getType());
1416 if (vecType && (!firstMaxRankedType ||
1417 firstMaxRankedType.getRank() < vecType.getRank()))
1418 firstMaxRankedType = vecType;
1424 assert(vecOperand &&
"Vector operand couldn't be found");
1426 if (firstMaxRankedType) {
1427 auto vecType = VectorType::get(firstMaxRankedType.getShape(),
1429 firstMaxRankedType.getScalableDims());
1432 vecOperands.push_back(vecOperand);
1438 resultTypes.push_back(
1440 ? VectorType::get(firstMaxRankedType.getShape(), resultType,
1441 firstMaxRankedType.getScalableDims())
1479 LDBG() <<
"Vectorizing operation as linalg generic/n";
1480 Block *block = linalgOp.getBlock();
1487 bvm.
map(valuesSet.getArrayRef(), valuesSet.getArrayRef());
1489 if (linalgOp.getNumDpsInits() == 0)
1495 for (
OpOperand *opOperand : linalgOp.getOpOperandsMatchingBBargs()) {
1496 BlockArgument bbarg = linalgOp.getMatchingBlockArgument(opOperand);
1497 if (linalgOp.isScalar(opOperand)) {
1498 bvm.
map(bbarg, opOperand->get());
1504 AffineMap indexingMap = linalgOp.getMatchingIndexingMap(opOperand);
1507 VectorType readType;
1509 if (linalgOp.isDpsInput(opOperand)) {
1512 readType = state.getCanonicalVecType(elemType);
1519 state.getCanonicalVecType(elemType, readMap.
compose(indexingMap));
1524 Operation *read = vector::TransferReadOp::create(
1525 rewriter, loc, readType, opOperand->get(),
indices,
1526 std::nullopt, readMap);
1527 read = state.maskOperation(rewriter, read, linalgOp, indexingMap);
1532 if (
auto maskOp = dyn_cast<vector::MaskingOpInterface>(read)) {
1534 cast<vector::TransferReadOp>(maskOp.getMaskableOp())
1540 if (readType.getRank() == 0)
1541 readValue = vector::ExtractOp::create(rewriter, loc, readValue,
1544 LDBG() <<
"New vectorized bbarg(" << bbarg.
getArgNumber()
1545 <<
"): " << readValue;
1546 bvm.
map(bbarg, readValue);
1547 bvm.
map(opOperand->get(), readValue);
1556 hooks.push_back(vectorizeYield);
1563 hooks.push_back(vectorizeIndex);
1570 hooks.push_back(vectorizeExtract);
1577 LDBG() <<
"failed to vectorize: " << op;
1582 state.maskOperation(rewriter,
result.newOp, linalgOp);
1583 LDBG() <<
"New vector op: " << *maybeMaskedOp;
1609 assert(type.getNumScalableDims() < 2 &&
1610 "Collapsing more than 1 scalable dim is not supported ATM");
1616 auto shape = type.getShape();
1617 auto scalableFlags = type.getScalableDims();
1621 unsigned currentDim = 0;
1623 unsigned dim = m.getNumResults();
1626 for (
unsigned d = 0; d < dim; ++d) {
1627 size *=
shape[currentDim + d];
1628 flag |= scalableFlags[currentDim + d];
1630 newShape.push_back(size);
1631 newScalableFlags.push_back(flag);
1635 return VectorType::get(newShape, type.getElementType(), newScalableFlags);
1668vectorizeAsTensorPackOp(RewriterBase &rewriter, linalg::PackOp packOp,
1669 ArrayRef<int64_t> inputVectorSizes,
1670 SmallVectorImpl<Value> &newResults) {
1671 if (!inputVectorSizes.empty()) {
1672 assert(inputVectorSizes.size() == packOp.getDestRank() &&
1673 "Invalid number of input vector sizes!");
1677 OpBuilder::InsertionGuard g(rewriter);
1680 Location loc = packOp.getLoc();
1681 std::optional<Value> padValue = packOp.getPaddingValue()
1682 ? std::optional(packOp.getPaddingValue())
1685 SmallVector<int64_t> destShape =
1686 SmallVector<int64_t>(packOp.getDestType().getShape());
1690 ArrayRef<int64_t> &writeVectorSizes = inputVectorSizes;
1694 bool useInBoundsInsteadOfMasking =
false;
1695 if (writeVectorSizes.empty()) {
1696 if (ShapedType::isDynamicShape(destShape))
1698 "unable to infer vector sizes");
1700 writeVectorSizes = destShape;
1701 useInBoundsInsteadOfMasking =
true;
1710 PackingMetadata packMetadata;
1711 SmallVector<int64_t> preTransposeWriteVecSizses(writeVectorSizes);
1714 auto preTransposeWriteVecType =
1715 VectorType::get(preTransposeWriteVecSizses,
1716 packOp.getResult().getType().getElementType());
1722 preTransposeWriteVecType,
1724 rewriter.
getContext(), packMetadata.reassociations)));
1728 rewriter, loc, packOp.getSource(), readVecType, padValue,
1729 useInBoundsInsteadOfMasking);
1732 auto shapeCastOp = vector::ShapeCastOp::create(
1733 rewriter, loc, preTransposeWriteVecType, maskedRead);
1737 auto transposeOp = vector::TransposeOp::create(
1738 rewriter, loc, shapeCastOp.getResult(), destPermutation);
1742 rewriter, loc, transposeOp.getResult(), packOp.getDest());
1743 newResults.push_back(write->
getResult(0));
1777vectorizeAsTensorUnpackOp(RewriterBase &rewriter, linalg::UnPackOp unpackOp,
1778 ArrayRef<int64_t> inputVectorSizes,
1779 ArrayRef<bool> inputScalableVecDims,
1780 SmallVectorImpl<Value> &newResults) {
1781 if (!inputVectorSizes.empty()) {
1782 assert(inputVectorSizes.size() == unpackOp.getSourceRank() &&
1783 "Invalid number of input vector sizes!");
1784 assert(inputVectorSizes.size() == inputScalableVecDims.size() &&
1785 "Incompatible number of vector sizes and vector scalable flags!");
1789 OpBuilder::InsertionGuard g(rewriter);
1792 ShapedType unpackTensorType = unpackOp.getSourceType();
1794 ArrayRef<int64_t> sourceShape = unpackTensorType.getShape();
1795 bool useInBoundsInsteadOfMasking =
false;
1797 Location loc = unpackOp->getLoc();
1800 SmallVector<int64_t> readVectorSizes(inputVectorSizes);
1801 SmallVector<bool> readScalableVectorFlags(inputScalableVecDims);
1804 if (inputVectorSizes.empty()) {
1805 if (ShapedType::isDynamicShape(sourceShape))
1807 "Unable to infer vector sizes!");
1809 readVectorSizes.assign(sourceShape.begin(), sourceShape.end());
1810 useInBoundsInsteadOfMasking =
true;
1814 VectorType readVecType =
1815 VectorType::get(readVectorSizes, unpackTensorType.getElementType(),
1816 readScalableVectorFlags);
1818 rewriter, loc, unpackOp.getSource(), readVecType, std::nullopt,
1819 useInBoundsInsteadOfMasking);
1822 PackingMetadata packMetadata;
1823 SmallVector<int64_t> lastDimToInsertPosPerm =
1825 vector::TransposeOp transposeOp = vector::TransposeOp::create(
1826 rewriter, loc, readResult, lastDimToInsertPosPerm);
1830 transposeOp.getType(),
1832 rewriter.
getContext(), packMetadata.reassociations)));
1833 vector::ShapeCastOp shapeCastOp = vector::ShapeCastOp::create(
1834 rewriter, loc, collapsedVecType, transposeOp->getResult(0));
1838 rewriter, loc, shapeCastOp.getResult(), unpackOp.getDest(),
1839 {}, useInBoundsInsteadOfMasking);
1841 newResults.push_back(write->
getResult(0));
1849vectorizeAsTensorPadOp(RewriterBase &rewriter, tensor::PadOp padOp,
1850 ArrayRef<int64_t> inputVectorSizes,
1851 SmallVectorImpl<Value> &newResults) {
1852 auto padValue = padOp.getConstantPaddingValue();
1853 Location loc = padOp.getLoc();
1856 OpBuilder::InsertionGuard g(rewriter);
1860 LogicalResult status =
1861 cast<ReifyRankedShapedTypeOpInterface>(padOp.getOperation())
1862 .reifyResultShapes(rewriter, reifiedReturnShapes);
1864 assert(succeeded(status) &&
"failed to reify result shapes");
1865 auto readType = VectorType::get(inputVectorSizes, padValue.getType());
1867 rewriter, loc, padOp.getSource(), readType, padValue,
1871 Value dest = tensor::EmptyOp::create(rewriter, loc, reifiedReturnShapes[0],
1872 padOp.getResultType().getElementType());
1875 newResults.push_back(write->
getResult(0));
1881static LogicalResult reductionPreconditions(LinalgOp op) {
1883 LDBG() <<
"reduction precondition failed: no reduction iterator";
1886 for (OpOperand &opOperand : op.getDpsInitsMutable()) {
1887 AffineMap indexingMap = op.getMatchingIndexingMap(&opOperand);
1893 LDBG() <<
"reduction precondition failed: reduction detection failed";
1901vectorizeDynamicConvOpPrecondition(linalg::LinalgOp conv,
1902 bool flatten1DDepthwiseConv) {
1903 if (flatten1DDepthwiseConv) {
1904 LDBG() <<
"Vectorization of flattened convs with dynamic shapes is not "
1910 LDBG() <<
"Not a 1D depth-wise WC conv, dynamic shapes are not supported";
1916 Value
lhs = conv.getDpsInputOperand(0)->get();
1917 ArrayRef<int64_t> lhsShape = cast<ShapedType>(
lhs.getType()).getShape();
1918 auto shapeWithoutCh = lhsShape.drop_back(1);
1919 if (ShapedType::isDynamicShape(shapeWithoutCh)) {
1920 LDBG() <<
"Dynamically-shaped op vectorization precondition failed: only "
1921 "channel dim can be dynamic";
1929vectorizeDynamicLinalgOpPrecondition(linalg::LinalgOp op,
1930 bool flatten1DDepthwiseConv) {
1932 return vectorizeDynamicConvOpPrecondition(op, flatten1DDepthwiseConv);
1935 return reductionPreconditions(op);
1940 !isa<linalg::GenericOp, linalg::CopyOp, linalg::ContractionOpInterface>(
1944 LDBG() <<
"Dynamically-shaped op meets vectorization pre-conditions";
1954vectorizeUnPackOpPrecondition(linalg::UnPackOp unpackOp,
1955 ArrayRef<int64_t> inputVectorSizes) {
1959 if (!unpackOp.hasPureTensorSemantics())
1964 if (inputVectorSizes.empty() && unpackOp.getDestType().hasStaticShape() &&
1965 unpackOp.getSourceType().hasStaticShape())
1970 if (!inputVectorSizes.empty() &&
1971 (inputVectorSizes.size() != unpackOp.getSourceRank())) {
1972 LDBG() <<
"Incorrect number of input vector sizes";
1978 unpackOp.getSourceType().getShape(), inputVectorSizes))) {
1979 LDBG() <<
"Invalid vector sizes for the read operation";
1987vectorizeInsertSliceOpPrecondition(tensor::InsertSliceOp sliceOp,
1988 ArrayRef<int64_t> inputVectorSizes) {
1991 auto sourceType = source.getType();
1992 if (!VectorType::isValidElementType(sourceType.getElementType()))
2008 bool isOutOfBoundsRead =
2009 !sourceType.hasStaticShape() && inputVectorSizes.empty();
2011 if (!padValue && isOutOfBoundsRead) {
2012 LDBG() <<
"Failed to get a pad value for out-of-bounds read access";
2026vectorizeAsLinalgContraction(RewriterBase &rewriter, VectorizationState &state,
2028 SmallVectorImpl<Value> &newResults) {
2029 Location loc = linalgOp.getLoc();
2030 MLIRContext *ctx = linalgOp.getContext();
2035 if (!isa<ContractionOpInterface>(linalgOp.getOperation()))
2038 OpOperand *outOperand = linalgOp.getDpsInitOperand(0);
2042 LDBG() <<
"Failed to determine contraction combining kind.";
2049 AffineMap lhsMap = linalgOp.getIndexingMapsArray()[0];
2050 AffineMap rhsMap = linalgOp.getIndexingMapsArray()[1];
2052 LDBG() <<
"Contractions with broadcasts are not supported.";
2057 SmallVector<Value> vecOperands;
2058 for (OpOperand &opOperand : linalgOp->getOpOperands()) {
2062 AffineMap indexingMap = linalgOp.getMatchingIndexingMap(&opOperand);
2066 VectorType readType =
2067 state.getCanonicalVecType(elemType, readMap.
compose(indexingMap));
2070 rewriter, loc, opOperand.get(), readType,
2071 arith::getZeroConstant(rewriter, loc, elemType),
2073 vecOperands.push_back(read);
2079 auto castAttr = linalgOp->getAttrOfType<TypeFnAttr>(
"cast");
2080 bool hasUnsignedCast =
2081 castAttr && castAttr.getValue() == TypeFn::cast_unsigned;
2082 auto accType = dyn_cast<VectorType>(vecOperands[2].
getType());
2083 auto accElementType =
2084 accType ? dyn_cast<IntegerType>(accType.getElementType()) :
nullptr;
2085 if (accElementType && accElementType.isSignless()) {
2086 for (Value &operand : MutableArrayRef(vecOperands).take_front(2)) {
2087 auto operandType = cast<VectorType>(operand.
getType());
2088 Type operandElementType = operandType.getElementType();
2089 VectorType castType = operandType.clone(accElementType);
2091 if (isa<FloatType>(operandElementType)) {
2094 ? arith::FPToUIOp::create(rewriter, loc, castType, operand)
2096 : arith::FPToSIOp::create(rewriter, loc, castType, operand)
2101 auto operandIntegerType = dyn_cast<IntegerType>(operandElementType);
2102 if (!operandIntegerType || !operandIntegerType.isSignless())
2104 if (operandIntegerType.getWidth() >= accElementType.getWidth())
2106 if (!hasUnsignedCast)
2111 operand = arith::ExtUIOp::create(rewriter, loc, castType, operand);
2116 SmallVector<Attribute> iterAttrs;
2117 auto iterators = linalgOp.getIteratorTypesArray();
2118 for (utils::IteratorType iter : iterators) {
2119 auto vecIter = iter == utils::IteratorType::parallel
2120 ? vector::IteratorType::parallel
2121 : vector::IteratorType::reduction;
2122 iterAttrs.push_back(vector::IteratorTypeAttr::get(ctx, vecIter));
2126 Operation *contractOp = vector::ContractionOp::create(
2127 rewriter, loc, vecOperands[0],
2128 vecOperands[1], vecOperands[2],
2129 linalgOp.getIndexingMaps(), rewriter.
getArrayAttr(iterAttrs), *maybeKind);
2130 contractOp = state.maskOperation(rewriter, contractOp, linalgOp);
2134 rewriter, loc, contractOp->
getResult(0), outOperand->
get());
2138 newResults.push_back(write->
getResult(0));
2144enum class ConvOperationKind { Conv, Pool };
2147static bool isCastOfBlockArgument(Operation *op) {
2162static std::optional<ConvOperationKind>
2163getConvOperationKind(Operation *reduceOp) {
2164 int numBlockArguments =
2165 llvm::count_if(reduceOp->
getOperands(), llvm::IsaPred<BlockArgument>);
2167 switch (numBlockArguments) {
2173 auto feedValIt = llvm::find_if_not(reduceOp->
getOperands(),
2174 llvm::IsaPred<BlockArgument>);
2176 "Expected a non-block argument operand");
2177 Operation *feedOp = (*feedValIt).getDefiningOp();
2178 if (isCastOfBlockArgument(feedOp)) {
2179 return ConvOperationKind::Pool;
2182 if (!((isa<arith::MulIOp, arith::MulFOp>(feedOp) ||
2183 (isa<arith::AndIOp>(feedOp) &&
2186 if (isa<BlockArgument>(v))
2188 if (Operation *op = v.getDefiningOp())
2189 return isCastOfBlockArgument(op);
2192 return std::nullopt;
2195 return ConvOperationKind::Conv;
2199 return ConvOperationKind::Pool;
2201 return std::nullopt;
2205static bool isSupportedPoolKind(vector::CombiningKind kind) {
2207 case vector::CombiningKind::ADD:
2208 case vector::CombiningKind::MAXNUMF:
2209 case vector::CombiningKind::MAXIMUMF:
2210 case vector::CombiningKind::MAXIMUMNUMF:
2211 case vector::CombiningKind::MAXSI:
2212 case vector::CombiningKind::MAXUI:
2213 case vector::CombiningKind::MINNUMF:
2214 case vector::CombiningKind::MINIMUMF:
2215 case vector::CombiningKind::MINIMUMNUMF:
2216 case vector::CombiningKind::MINSI:
2217 case vector::CombiningKind::MINUI:
2224static LogicalResult vectorizeConvOpPrecondition(linalg::LinalgOp convOp) {
2225 auto getOperandType = [&](
auto operand) {
2226 return dyn_cast<ShapedType>((operand->get()).getType());
2228 ShapedType lhsShapedType = getOperandType(convOp.getDpsInputOperand(0));
2229 ShapedType rhsShapedType = getOperandType(convOp.getDpsInputOperand(1));
2230 ShapedType resShapedType = getOperandType(convOp.getDpsInitOperand(0));
2234 if ((lhsShapedType.getRank() != 3 || resShapedType.getRank() != 3) &&
2235 (lhsShapedType.getRank() != 1 || resShapedType.getRank() != 1))
2242 auto maybeOper = getConvOperationKind(reduceOp);
2243 if (!maybeOper.has_value())
2250 if (!maybeKind || ((*maybeKind != vector::CombiningKind::ADD &&
2251 *maybeKind != vector::CombiningKind::OR) &&
2252 (*maybeOper != ConvOperationKind::Pool ||
2253 !isSupportedPoolKind(*maybeKind)))) {
2257 auto rhsRank = rhsShapedType.getRank();
2258 if (*maybeOper == ConvOperationKind::Pool) {
2262 if (rhsRank != 1 && rhsRank != 2 && rhsRank != 3)
2269static LogicalResult vectorizeLinalgOpPrecondition(
2270 LinalgOp linalgOp, ArrayRef<int64_t> inputVectorSizes,
2271 bool vectorizeNDExtract,
bool flatten1DDepthwiseConv) {
2273 if (llvm::any_of(linalgOp->getOpOperands(), [&](OpOperand &operand) {
2274 return llvm::is_contained(linalgOp.getShape(&operand), 0);
2278 if (!inputVectorSizes.empty() &&
2283 if (linalgOp.hasDynamicShape() &&
failed(vectorizeDynamicLinalgOpPrecondition(
2284 linalgOp, flatten1DDepthwiseConv))) {
2285 LDBG() <<
"Dynamically-shaped op failed vectorization pre-conditions";
2289 SmallVector<CustomVectorizationPrecondition> customPreconditions;
2295 for (Operation &innerOp : linalgOp->getRegion(0).front()) {
2298 customPreconditions,
2301 customPrecondition(&innerOp, vectorizeNDExtract));
2305 if (!llvm::all_of(innerOp.getOperandTypes(),
2306 VectorType::isValidElementType)) {
2309 if (!llvm::all_of(innerOp.getResultTypes(),
2310 VectorType::isValidElementType)) {
2319 return vectorizeConvOpPrecondition(linalgOp);
2325 LDBG() <<
"precondition failed: not projected permutations";
2328 if (
failed(reductionPreconditions(linalgOp))) {
2329 LDBG() <<
"precondition failed: reduction preconditions";
2336vectorizePackOpPrecondition(linalg::PackOp packOp,
2337 ArrayRef<int64_t> inputVectorSizes) {
2341 if (!packOp.hasPureTensorSemantics())
2344 auto padValue = packOp.getPaddingValue();
2348 LDBG() <<
"pad value is not constant: " << packOp;
2352 ArrayRef<int64_t> resultTensorShape = packOp.getDestType().getShape();
2353 bool satisfyEmptyCond =
true;
2354 if (inputVectorSizes.empty()) {
2355 if (!packOp.getDestType().hasStaticShape() ||
2356 !packOp.getSourceType().hasStaticShape())
2357 satisfyEmptyCond =
false;
2360 if (!satisfyEmptyCond &&
2362 resultTensorShape.take_front(packOp.getSourceRank()),
2366 if (llvm::any_of(packOp.getInnerTiles(), [](OpFoldResult v) {
2367 return !getConstantIntValue(v).has_value();
2369 LDBG() <<
"inner_tiles must be constant: " << packOp;
2377vectorizePadOpPrecondition(tensor::PadOp padOp,
2378 ArrayRef<int64_t> inputVectorSizes) {
2379 auto padValue = padOp.getConstantPaddingValue();
2381 LDBG() <<
"pad value is not constant: " << padOp;
2385 ArrayRef<int64_t> resultTensorShape = padOp.getResultType().getShape();
2401 if (llvm::any_of(llvm::enumerate(padOp.getMixedLowPad()),
2402 [&](
const auto &en) {
2403 OpFoldResult padValue = en.value();
2404 unsigned pos = en.index();
2405 std::optional<int64_t> pad = getConstantIntValue(padValue);
2406 return (!pad.has_value() || pad.value() != 0) &&
2407 resultTensorShape[pos] != 1;
2409 LDBG() <<
"low pad must all be zero for all non unit dims: " << padOp;
2423vectorizeScalableVectorPrecondition(Operation *op,
2424 ArrayRef<int64_t> inputVectorSizes,
2425 ArrayRef<bool> inputScalableVecDims) {
2426 assert(inputVectorSizes.size() == inputScalableVecDims.size() &&
2427 "Number of input vector sizes and scalable dims doesn't match");
2429 size_t numOfScalableDims =
2430 llvm::count_if(inputScalableVecDims, [](
bool flag) {
return flag; });
2432 if (numOfScalableDims == 0)
2435 auto linalgOp = dyn_cast<LinalgOp>(op);
2440 return success(isa<linalg::UnPackOp>(op));
2444 if (numOfScalableDims > 2)
2464 bool seenNonUnitParallel =
false;
2465 auto iterators = linalgOp.getIteratorTypesArray();
2466 SmallVector<bool> scalableFlags(inputScalableVecDims);
2467 int64_t idx = scalableFlags.size() - 1;
2468 while (!scalableFlags[idx]) {
2469 bool isNonUnitDim = (inputVectorSizes[idx] != 1);
2470 seenNonUnitParallel |=
2471 (iterators[idx] == utils::IteratorType::parallel && isNonUnitDim);
2473 iterators.pop_back();
2474 scalableFlags.pop_back();
2479 switch (iterators.back()) {
2480 case utils::IteratorType::reduction: {
2482 if (iterators.size() != inputVectorSizes.size()) {
2483 LDBG() <<
"Non-trailing reduction dim requested for scalable "
2487 if (isa<linalg::MatmulOp>(op)) {
2489 <<
"Scalable vectorization of the reduction dim in Matmul-like ops "
2495 case utils::IteratorType::parallel: {
2497 if (seenNonUnitParallel) {
2498 LDBG() <<
"Inner parallel dim not requested for scalable "
2510 if (numOfScalableDims == 2) {
2514 if (iterators.back() == utils::IteratorType::reduction) {
2515 LDBG() <<
"Higher dim than the trailing reduction dim requested for "
2520 scalableFlags.pop_back();
2521 iterators.pop_back();
2523 if (!scalableFlags.back() ||
2524 (iterators.back() != utils::IteratorType::parallel))
2532 isa<linalg::BatchMatmulOp>(op) ||
2534 isa<linalg::MatvecOp>(op) || isa<linalg::Mmt4DOp>(op) ||
2539 Operation *op, ArrayRef<int64_t> inputVectorSizes,
2540 ArrayRef<bool> inputScalableVecDims,
bool vectorizeNDExtract,
2541 bool flatten1DDepthwiseConv) {
2546 if (
failed(vectorizeScalableVectorPrecondition(op, inputVectorSizes,
2547 inputScalableVecDims)))
2551 .Case([&](linalg::LinalgOp linalgOp) {
2552 return vectorizeLinalgOpPrecondition(linalgOp, inputVectorSizes,
2554 flatten1DDepthwiseConv);
2556 .Case([&](tensor::PadOp padOp) {
2557 return vectorizePadOpPrecondition(padOp, inputVectorSizes);
2559 .Case([&](linalg::PackOp packOp) {
2560 return vectorizePackOpPrecondition(packOp, inputVectorSizes);
2562 .Case([&](linalg::UnPackOp unpackOp) {
2563 return vectorizeUnPackOpPrecondition(unpackOp, inputVectorSizes);
2565 .Case([&](tensor::InsertSliceOp sliceOp) {
2566 return vectorizeInsertSliceOpPrecondition(sliceOp, inputVectorSizes);
2568 .Default(failure());
2572static void convertAffineApply(RewriterBase &rewriter, LinalgOp linalgOp) {
2573 OpBuilder::InsertionGuard g(rewriter);
2574 auto toReplace = linalgOp.getBlock()->getOps<affine::AffineApplyOp>();
2576 for (
auto op : make_early_inc_range(toReplace)) {
2578 auto expanded = affine::expandAffineExpr(
2580 op.
getOperands().take_front(op.getAffineMap().getNumDims()),
2581 op.
getOperands().take_back(op.getAffineMap().getNumSymbols()));
2587 return isa<linalg::LinalgOp, tensor::PadOp, linalg::PackOp, linalg::UnPackOp,
2588 tensor::InsertSliceOp>(op);
2592 RewriterBase &rewriter, Operation *op, ArrayRef<int64_t> inputVectorSizes,
2593 ArrayRef<bool> inputScalableVecDims,
bool vectorizeNDExtract,
2594 bool flatten1DDepthwiseConv,
bool assumeDynamicDimsMatchVecSizes,
2595 bool createNamedContraction) {
2596 LDBG() <<
"Attempting to vectorize: " << *op;
2597 LDBG() <<
"Input vector sizes: " << llvm::interleaved(inputVectorSizes);
2598 LDBG() <<
"Input scalable vector dims: "
2599 << llvm::interleaved(inputScalableVecDims);
2603 flatten1DDepthwiseConv))) {
2604 LDBG() <<
"Vectorization pre-conditions failed";
2609 VectorizationState state(rewriter);
2610 if (
auto linalgOp = dyn_cast<linalg::LinalgOp>(op)) {
2611 if (
failed(state.initState(rewriter, linalgOp, inputVectorSizes,
2612 inputScalableVecDims,
2613 assumeDynamicDimsMatchVecSizes))) {
2614 LDBG() <<
"Vectorization state couldn't be initialized";
2619 SmallVector<Value> results;
2620 auto vectorizeResult =
2622 .Case([&](linalg::LinalgOp linalgOp) {
2626 rewriter, linalgOp, inputVectorSizes, inputScalableVecDims,
2627 flatten1DDepthwiseConv);
2628 if (succeeded(convOr)) {
2629 llvm::append_range(results, (*convOr)->getResults());
2633 LDBG() <<
"Unsupported convolution can't be vectorized.";
2637 if (createNamedContraction &&
2638 isa<ContractionOpInterface>(linalgOp.getOperation()))
2639 return vectorizeAsLinalgContraction(rewriter, state, linalgOp,
2643 <<
"Vectorize generic by broadcasting to the canonical vector "
2647 convertAffineApply(rewriter, linalgOp);
2656 .Case([&](tensor::PadOp padOp) {
2657 return vectorizeAsTensorPadOp(rewriter, padOp, inputVectorSizes,
2660 .Case([&](linalg::PackOp packOp) {
2661 return vectorizeAsTensorPackOp(rewriter, packOp, inputVectorSizes,
2664 .Case([&](linalg::UnPackOp unpackOp) {
2665 return vectorizeAsTensorUnpackOp(rewriter, unpackOp,
2667 inputScalableVecDims, results);
2669 .Case([&](tensor::InsertSliceOp sliceOp) {
2673 .Default(failure());
2675 if (
failed(vectorizeResult)) {
2676 LDBG() <<
"Vectorization failed";
2680 return VectorizationResult{results};
2684 memref::CopyOp copyOp) {
2685 auto srcType = cast<MemRefType>(copyOp.getSource().getType());
2686 auto dstType = cast<MemRefType>(copyOp.getTarget().getType());
2687 if (!srcType.hasStaticShape() || !dstType.hasStaticShape())
2692 if (!VectorType::isValidElementType(srcElementType) ||
2693 !VectorType::isValidElementType(dstElementType))
2696 auto readType = VectorType::get(srcType.getShape(), srcElementType);
2697 auto writeType = VectorType::get(dstType.getShape(), dstElementType);
2699 Location loc = copyOp->getLoc();
2701 SmallVector<Value>
indices(srcType.getRank(), zero);
2703 Value
readValue = vector::TransferReadOp::create(
2704 rewriter, loc, readType, copyOp.getSource(),
indices,
2707 if (cast<VectorType>(
readValue.getType()).getRank() == 0) {
2708 readValue = vector::ExtractOp::create(rewriter, loc, readValue,
2709 ArrayRef<int64_t>());
2711 vector::BroadcastOp::create(rewriter, loc, writeType, readValue);
2713 Operation *writeValue = vector::TransferWriteOp::create(
2714 rewriter, loc, readValue, copyOp.getTarget(),
indices,
2725template <
typename OpTy>
2726struct VectorizePadOpUserPattern :
public OpRewritePattern<tensor::PadOp> {
2727 using OpRewritePattern<tensor::PadOp>::OpRewritePattern;
2729 LogicalResult matchAndRewrite(tensor::PadOp padOp,
2730 PatternRewriter &rewriter)
const final {
2731 bool changed =
false;
2733 for (
auto *user : llvm::to_vector<4>(padOp->getUsers()))
2734 if (
auto op = dyn_cast<OpTy>(user))
2735 changed |= rewriteUser(rewriter, padOp, op).succeeded();
2740 virtual LogicalResult rewriteUser(PatternRewriter &rewriter,
2741 tensor::PadOp padOp, OpTy op)
const = 0;
2763struct PadOpVectorizationWithTransferReadPattern
2764 :
public VectorizePadOpUserPattern<vector::TransferReadOp> {
2765 using VectorizePadOpUserPattern<
2766 vector::TransferReadOp>::VectorizePadOpUserPattern;
2768 LogicalResult rewriteUser(PatternRewriter &rewriter, tensor::PadOp padOp,
2769 vector::TransferReadOp xferOp)
const override {
2771 if (!padOp.hasZeroLowPad())
2774 auto padValue = padOp.getConstantPaddingValue();
2778 if (xferOp.hasOutOfBoundsDim() || xferOp.getMask())
2782 SmallVector<bool> inBounds(xferOp.getVectorType().getRank(),
false);
2783 xferOp->setInherentAttr(xferOp.getInBoundsAttrName(),
2785 xferOp.getBaseMutable().assign(padOp.getSource());
2786 xferOp.getPaddingMutable().assign(padValue);
2825struct PadOpVectorizationWithTransferWritePattern
2826 :
public VectorizePadOpUserPattern<vector::TransferWriteOp> {
2827 using VectorizePadOpUserPattern<
2828 vector::TransferWriteOp>::VectorizePadOpUserPattern;
2830 LogicalResult rewriteUser(PatternRewriter &rewriter, tensor::PadOp padOp,
2831 vector::TransferWriteOp xferOp)
const override {
2833 if (xferOp.getTransferRank() == 0)
2837 if (!padOp.hasZeroLowPad())
2840 auto padValue = padOp.getConstantPaddingValue();
2844 if (!xferOp->hasOneUse())
2846 auto trimPadding = dyn_cast<tensor::ExtractSliceOp>(*xferOp->user_begin());
2850 if (!trimPadding.hasZeroOffset())
2853 if (!hasSameTensorSize(padOp.getSource(), trimPadding))
2859 SmallVector<bool> inBounds(xferOp.getVectorType().getRank(),
false);
2861 xferOp, padOp.getSource().
getType(), xferOp.getVector(),
2862 padOp.getSource(), xferOp.getIndices(), xferOp.getPermutationMapAttr(),
2864 rewriter.
replaceOp(trimPadding, newXferOp->getResult(0));
2879 bool hasSameTensorSize(Value beforePadding,
2880 tensor::ExtractSliceOp afterTrimming)
const {
2883 if (
auto castOp = beforePadding.
getDefiningOp<tensor::CastOp>())
2884 if (hasSameTensorSize(castOp.getSource(), afterTrimming))
2887 auto t1 = dyn_cast<RankedTensorType>(beforePadding.
getType());
2888 auto t2 = dyn_cast<RankedTensorType>(afterTrimming.getType());
2893 if (t1.getRank() != t2.getRank())
2898 for (
unsigned i = 0; i < t1.getRank(); ++i) {
2899 if (t1.isDynamicDim(i) != t2.isDynamicDim(i))
2901 if (!t1.isDynamicDim(i) && t1.getDimSize(i) != t2.getDimSize(i))
2906 if (t1.getNumDynamicDims() == 0)
2914 auto beforeSlice = beforePadding.
getDefiningOp<tensor::ExtractSliceOp>();
2918 assert(
static_cast<size_t>(t1.getRank()) ==
2919 beforeSlice.getMixedSizes().size());
2920 assert(
static_cast<size_t>(t2.getRank()) ==
2921 afterTrimming.getMixedSizes().size());
2923 for (
unsigned i = 0; i < t1.getRank(); ++i) {
2925 if (!t1.isDynamicDim(i))
2927 auto size1 = beforeSlice.getMixedSizes()[i];
2928 auto size2 = afterTrimming.getMixedSizes()[i];
2935 auto v1 = llvm::dyn_cast_if_present<Value>(size1);
2936 auto v2 = llvm::dyn_cast_if_present<Value>(size2);
2942 auto minOp1 = v1.getDefiningOp<affine::AffineMinOp>();
2943 auto minOp2 = v2.getDefiningOp<affine::AffineMinOp>();
2944 if (minOp1 && minOp2 && minOp1.getAffineMap() == minOp2.getAffineMap() &&
2945 minOp1.getOperands() == minOp2.getOperands())
2971 if (
auto bcast = llvm::dyn_cast<vector::BroadcastOp>(op)) {
2972 auto source = bcast.getSource();
2973 if (llvm::dyn_cast<VectorType>(source.getType()))
2981 if (
auto fill = llvm::dyn_cast<linalg::FillOp>(op)) {
2982 return fill.getInputs()[0];
2987 if (
auto generate = llvm::dyn_cast<tensor::GenerateOp>(op)) {
2994 if (
auto xferWrite = llvm::dyn_cast<vector::TransferWriteOp>(op))
3002 if (
auto slice = llvm::dyn_cast<tensor::InsertSliceOp>(op))
3010 ArrayRef<int64_t> inputVectorSizes,
3011 SmallVectorImpl<Value> &newResults) {
3013 OpBuilder::InsertionGuard g(rewriter);
3017 auto sourceType = source.getType();
3018 auto resultType = sliceOp.getResultType();
3023 auto elemType = sourceType.getElementType();
3024 padValue = arith::ConstantOp::create(rewriter, sliceOp.getLoc(), elemType,
3031 llvm::SmallBitVector droppedDims = sliceOp.getDroppedDims();
3032 SmallVector<int64_t> resultDimsForSourceDims;
3033 resultDimsForSourceDims.reserve(sourceType.getRank());
3034 for (int64_t resultDim = 0, end = resultType.getRank(); resultDim < end;
3036 if (!droppedDims[resultDim])
3037 resultDimsForSourceDims.push_back(resultDim);
3038 assert(resultDimsForSourceDims.size() ==
3039 static_cast<size_t>(sourceType.getRank()) &&
3040 "expected one non-dropped result dim per source dim");
3042 SmallVector<int64_t> vecShape;
3043 for (int64_t i = 0, end = sourceType.getRank(); i < end; ++i) {
3044 if (!inputVectorSizes.empty()) {
3045 vecShape.push_back(inputVectorSizes[i]);
3046 }
else if (!sourceType.isDynamicDim(i)) {
3047 vecShape.push_back(sourceType.getDimSize(i));
3048 }
else if (!resultType.isDynamicDim(resultDimsForSourceDims[i])) {
3052 vecShape.push_back(resultType.getDimSize(resultDimsForSourceDims[i]));
3059 auto vecType = VectorType::get(vecShape, sourceType.getElementType());
3062 auto loc = sliceOp.getLoc();
3065 SmallVector<Value> readIndices(
3068 rewriter, loc, source, vecType, padValue,
3069 inputVectorSizes.empty());
3076 writeIndices, inputVectorSizes.empty());
3079 newResults.push_back(write->
getResult(0));
3107struct PadOpVectorizationWithInsertSlicePattern
3108 :
public VectorizePadOpUserPattern<tensor::InsertSliceOp> {
3109 using VectorizePadOpUserPattern<
3110 tensor::InsertSliceOp>::VectorizePadOpUserPattern;
3112 LogicalResult rewriteUser(PatternRewriter &rewriter, tensor::PadOp padOp,
3113 tensor::InsertSliceOp insertOp)
const override {
3115 if (!padOp.hasZeroLowPad())
3118 if (!insertOp.hasUnitStride())
3121 auto padValue = padOp.getConstantPaddingValue();
3125 if (!cast<ShapedType>(padOp.getResult().getType()).hasStaticShape())
3128 if (insertOp.getDest() == padOp.getResult())
3131 auto vecType = VectorType::get(padOp.getType().getShape(),
3132 padOp.getType().getElementType());
3133 unsigned vecRank = vecType.getRank();
3134 unsigned tensorRank = insertOp.getType().getRank();
3138 SmallVector<int64_t> expectedSizes(tensorRank - vecRank, 1);
3139 expectedSizes.append(vecType.getShape().begin(), vecType.getShape().end());
3141 llvm::zip(insertOp.getMixedSizes(), expectedSizes), [](
auto it) {
3142 return getConstantIntValue(std::get<0>(it)) == std::get<1>(it);
3152 SmallVector<Value> readIndices(
3154 auto read = vector::TransferReadOp::create(rewriter, padOp.getLoc(),
3155 vecType, padOp.getSource(),
3156 readIndices, padValue);
3162 rewriter, padOp.getLoc(), insertOp.getMixedOffsets());
3163 SmallVector<bool> inBounds(vecRank,
true);
3165 insertOp, read, insertOp.getDest(), writeIndices,
3166 ArrayRef<bool>{inBounds});
3173 RewritePatternSet &patterns, PatternBenefit baseBenefit) {
3174 patterns.
add<PadOpVectorizationWithTransferReadPattern,
3175 PadOpVectorizationWithTransferWritePattern,
3176 PadOpVectorizationWithInsertSlicePattern>(
3187static bool mayExistInterleavedUses(Operation *firstOp, Operation *secondOp,
3191 LDBG() <<
"interleavedUses precondition failed, firstOp: " << *firstOp
3192 <<
", second op: " << *secondOp;
3195 for (
auto v : values) {
3196 for (
auto &u : v.getUses()) {
3197 Operation *owner = u.getOwner();
3198 if (owner == firstOp || owner == secondOp)
3204 LDBG() <<
" found interleaved op " << *owner <<
", firstOp: " << *firstOp
3205 <<
", second op: " << *secondOp;
3214static memref::SubViewOp getSubViewUseIfUnique(Value v) {
3215 memref::SubViewOp subViewOp;
3217 if (
auto newSubViewOp = dyn_cast<memref::SubViewOp>(u.getOwner())) {
3219 return memref::SubViewOp();
3220 subViewOp = newSubViewOp;
3229 vector::TransferReadOp xferOp, PatternRewriter &rewriter)
const {
3232 if (xferOp.getMask())
3236 Value viewOrAlloc = xferOp.getBase();
3242 memref::SubViewOp subViewOp = getSubViewUseIfUnique(viewOrAlloc);
3245 Value subView = subViewOp.getResult();
3248 memref::CopyOp copyOp;
3249 for (
auto &u : subView.
getUses()) {
3250 if (
auto newCopyOp = dyn_cast<memref::CopyOp>(u.getOwner())) {
3251 assert(isa<MemRefType>(newCopyOp.getTarget().getType()));
3252 if (newCopyOp.getTarget() != subView)
3254 if (mayExistInterleavedUses(newCopyOp, xferOp, {viewOrAlloc, subView}))
3266 for (
auto &u : viewOrAlloc.
getUses()) {
3267 if (
auto newFillOp = dyn_cast<FillOp>(u.getOwner())) {
3268 assert(isa<MemRefType>(newFillOp.output().getType()));
3269 if (newFillOp.output() != viewOrAlloc)
3271 if (mayExistInterleavedUses(newFillOp, copyOp, {viewOrAlloc, subView}))
3273 maybeFillOp = newFillOp;
3278 if (maybeFillOp && xferOp.getPadding() != maybeFillOp.value())
3280 "padding value does not match fill");
3283 Value in = copyOp.getSource();
3289 auto vectorType = xferOp.getVectorType();
3290 Value res = vector::TransferReadOp::create(
3291 rewriter, xferOp.getLoc(), vectorType, in, xferOp.getIndices(),
3292 xferOp.getPermutationMapAttr(), xferOp.getPadding(), xferOp.getMask(),
3294 SmallVector<bool>(vectorType.getRank(),
false)));
3297 rewriter.
eraseOp(maybeFillOp);
3307 vector::TransferWriteOp xferOp, PatternRewriter &rewriter)
const {
3309 if (xferOp.getMask())
3313 Value viewOrAlloc = xferOp.getBase();
3319 memref::SubViewOp subViewOp = getSubViewUseIfUnique(viewOrAlloc);
3322 Value subView = subViewOp.getResult();
3325 memref::CopyOp copyOp;
3326 for (
auto &u : subViewOp.getResult().getUses()) {
3327 if (
auto newCopyOp = dyn_cast<memref::CopyOp>(u.getOwner())) {
3328 if (newCopyOp.getSource() != subView)
3330 if (mayExistInterleavedUses(xferOp, newCopyOp, {viewOrAlloc, subView}))
3340 assert(isa<MemRefType>(copyOp.getTarget().getType()));
3341 Value out = copyOp.getTarget();
3348 auto vector = xferOp.getVector();
3349 vector::TransferWriteOp::create(
3350 rewriter, xferOp.getLoc(), vector, out, xferOp.getIndices(),
3351 xferOp.getPermutationMapAttr(), xferOp.getMask(),
3353 dyn_cast<VectorType>(vector.getType()).getRank(),
false)));
3366static void bindShapeDims(ShapedType shapedType) {}
3368template <
int N,
typename IntTy,
typename... IntTy2>
3369static void bindShapeDims(ShapedType shapedType, IntTy &val, IntTy2 &...vals) {
3370 val = shapedType.getShape()[N];
3371 bindShapeDims<N + 1, IntTy2 &...>(shapedType, vals...);
3375template <
typename... IntTy>
3376static void bindShapeDims(ShapedType shapedType, IntTy &...vals) {
3377 bindShapeDims<0>(shapedType, vals...);
3382static std::optional<DilationsAndStrides> match1DConvPoolOp(LinalgOp op) {
3383#define MATCH_1D_CONV_POOL_OP(ConvOpTy) \
3384 if (auto convParams = matchConvolutionOpOfType<ConvOpTy>(op)) \
3406#undef MATCH_1D_CONV_POOL_OP
3408 return std::nullopt;
3446struct Conv1DGenerator
3447 :
public StructuredGenerator<LinalgOp, utils::IteratorType> {
3450 static FailureOr<Conv1DGenerator> create(RewriterBase &rewriter,
3451 LinalgOp linalgOp) {
3454 std::optional<DilationsAndStrides> convParams = match1DConvPoolOp(linalgOp);
3458 int strideW =
static_cast<int>(convParams->strides.front());
3459 int dilationW =
static_cast<int>(convParams->dilations.front());
3460 return Conv1DGenerator(rewriter, linalgOp, strideW, dilationW);
3464 Conv1DGenerator(RewriterBase &rewriter, LinalgOp linalgOp,
int strideW,
3466 : StructuredGenerator<LinalgOp, utils::IteratorType>(rewriter, linalgOp),
3467 strideW(strideW), dilationW(dilationW) {
3469 lhsShaped = linalgOp.getDpsInputOperand(0)->
get();
3470 rhsShaped = linalgOp.getDpsInputOperand(1)->
get();
3471 resShaped = linalgOp.getDpsInitOperand(0)->
get();
3472 lhsShapedType = dyn_cast<ShapedType>(lhsShaped.getType());
3473 rhsShapedType = dyn_cast<ShapedType>(rhsShaped.getType());
3474 resShapedType = dyn_cast<ShapedType>(resShaped.getType());
3479 setConvOperationKind(reduceOp);
3482 reductionKind = maybeKind.value();
3505 int64_t nSize, wSize, cSize, kwSize, fSize;
3506 SmallVector<int64_t, 3> lhsShape, rhsShape, resShape;
3508 switch (conv1DOpOrder) {
3511 nSize = fSize = cSize = 0;
3513 bindShapeDims(resShapedType, wSize);
3515 bindShapeDims(rhsShapedType, kwSize);
3518 (wSize + kwSize - 1)};
3519 rhsShape = {kwSize};
3524 bindShapeDims(resShapedType, nSize, wSize, fSize);
3526 case ConvOperationKind::Conv:
3528 bindShapeDims(rhsShapedType, kwSize, cSize);
3530 case ConvOperationKind::Pool:
3532 bindShapeDims(rhsShapedType, kwSize);
3540 ((wSize - 1) * strideW + 1) + ((kwSize - 1) * dilationW + 1) -
3544 case ConvOperationKind::Conv:
3545 rhsShape = {kwSize, cSize, fSize};
3547 case ConvOperationKind::Pool:
3548 rhsShape = {kwSize};
3551 resShape = {nSize, wSize, fSize};
3555 bindShapeDims(resShapedType, nSize, fSize, wSize);
3557 case ConvOperationKind::Conv:
3559 bindShapeDims(rhsShapedType, fSize, cSize, kwSize);
3561 case ConvOperationKind::Pool:
3563 bindShapeDims(rhsShapedType, kwSize);
3567 lhsShape = {nSize, cSize,
3571 ((wSize - 1) * strideW + 1) + ((kwSize - 1) * dilationW + 1) -
3574 case ConvOperationKind::Conv:
3575 rhsShape = {fSize, cSize, kwSize};
3577 case ConvOperationKind::Pool:
3578 rhsShape = {kwSize};
3581 resShape = {nSize, fSize, wSize};
3585 vector::TransferWriteOp write;
3591 int64_t wSizeStep = strideW == 1 ? wSize : 1;
3593 Type lhsEltType = lhsShapedType.getElementType();
3594 Type rhsEltType = rhsShapedType.getElementType();
3595 Type resEltType = resShapedType.getElementType();
3596 auto lhsType = VectorType::get(lhsShape, lhsEltType);
3597 auto rhsType = VectorType::get(rhsShape, rhsEltType);
3598 auto resType = VectorType::get(resShape, resEltType);
3600 SmallVector<Value> lhsPadding(lhsShape.size(), zero);
3601 SmallVector<Value> rhsPadding(rhsShape.size(), zero);
3602 SmallVector<Value> resPadding(resShape.size(), zero);
3605 Value
lhs = vector::TransferReadOp::create(
3606 rewriter, loc, lhsType, lhsShaped, lhsPadding,
3607 arith::getZeroConstant(rewriter, loc, lhsEltType));
3609 Value
rhs =
nullptr;
3610 if (oper == ConvOperationKind::Conv)
3611 rhs = vector::TransferReadOp::create(
3612 rewriter, loc, rhsType, rhsShaped, rhsPadding,
3613 arith::getZeroConstant(rewriter, loc, rhsEltType));
3614 Value res = vector::TransferReadOp::create(
3615 rewriter, loc, resType, resShaped, resPadding,
3616 arith::getZeroConstant(rewriter, loc, resEltType));
3621 switch (conv1DOpOrder) {
3629 static constexpr std::array<int64_t, 3> permLhs = {0, 2, 1};
3630 lhs = vector::TransposeOp::create(rewriter, loc,
lhs, permLhs);
3632 static constexpr std::array<int64_t, 3> permRhs = {2, 1, 0};
3635 if (oper == ConvOperationKind::Conv)
3636 rhs = vector::TransposeOp::create(rewriter, loc,
rhs, permRhs);
3638 static constexpr std::array<int64_t, 3> permRes = {0, 2, 1};
3639 res = vector::TransposeOp::create(rewriter, loc, res, permRes);
3648 SmallVector<Value> lhsVals, rhsVals, resVals;
3650 kwSize, strideW, dilationW, wSizeStep,
3653 if (oper == ConvOperationKind::Conv)
3656 wSizeStep, isSingleChanneled);
3658 auto linearIndex = [&](int64_t kw, int64_t w) {
3659 return kw * (wSize / wSizeStep) + w;
3665 for (int64_t kw = 0; kw < kwSize; ++kw) {
3666 for (int64_t w = 0; w < wSize; w += wSizeStep) {
3668 case ConvOperationKind::Conv:
3669 if (isSingleChanneled) {
3670 resVals[w] = conv1dSliceAsOuterProduct(rewriter, loc,
3671 lhsVals[linearIndex(kw, w)],
3672 rhsVals[kw], resVals[w]);
3674 resVals[w] = conv1dSliceAsContraction(rewriter, loc,
3675 lhsVals[linearIndex(kw, w)],
3676 rhsVals[kw], resVals[w]);
3679 case ConvOperationKind::Pool:
3680 resVals[w] = pool1dSlice(rewriter, loc, lhsVals[linearIndex(kw, w)],
3696 switch (conv1DOpOrder) {
3703 static constexpr std::array<int64_t, 3> perm = {0, 2, 1};
3704 res = vector::TransposeOp::create(rewriter, loc, res, perm);
3709 return vector::TransferWriteOp::create(rewriter, loc, res, resShaped,
3715 Value
promote(RewriterBase &rewriter, Location loc, Value val, Type ty,
3716 Operation *castOp) {
3721 assert(castOp &&
"expected a payload cast for promoted operand");
3725 if (
auto shapedType = dyn_cast<ShapedType>(val.
getType()))
3726 dstType = shapedType.cloneWith(std::nullopt, dstElementType);
3728 dstType = dstElementType;
3737 Value conv1dSliceAsContraction(RewriterBase &rewriter, Location loc,
3738 Value
lhs, Value
rhs, Value res) {
3739 vector::IteratorType par = vector::IteratorType::parallel;
3740 vector::IteratorType red = vector::IteratorType::reduction;
3741 AffineExpr n, w, f, c;
3745 auto contrationOp = vector::ContractionOp::create(
3746 rewriter, loc,
lhs,
rhs, res,
3747 MapList{{n, w, c}, {c, f}, {n, w, f}},
3748 ArrayRef<vector::IteratorType>{par, par, par, red});
3749 contrationOp.setKind(reductionKind);
3750 return contrationOp;
3755 Value conv1dSliceAsOuterProduct(RewriterBase &rewriter, Location loc,
3756 Value
lhs, Value
rhs, Value res) {
3759 return vector::OuterProductOp::create(rewriter, loc, res.
getType(),
lhs,
3760 rhs, res, vector::CombiningKind::ADD);
3764 Value pool1dSlice(RewriterBase &rewriter, Location loc, Value
lhs,
3782 FailureOr<Operation *> depthwiseConv(uint64_t channelDimVecSize,
3783 bool channelDimScalableFlag,
3785 bool scalableChDim =
false;
3786 bool useMasking =
false;
3787 int64_t nSize, wSize, cSize, kwSize;
3789 bindShapeDims(rhsShapedType, kwSize, cSize);
3790 if (ShapedType::isDynamic(cSize)) {
3791 assert(channelDimVecSize != 0 &&
"Channel dim vec size must be > 0");
3792 cSize = channelDimVecSize;
3796 scalableChDim = channelDimScalableFlag;
3800 assert(!(useMasking && flatten) &&
3801 "Unsupported flattened conv with dynamic shapes");
3804 bindShapeDims(resShapedType, nSize, wSize);
3806 vector::TransferWriteOp write;
3812 int64_t wSizeStep = strideW == 1 ? wSize : 1;
3814 Type lhsEltType = lhsShapedType.getElementType();
3815 Type rhsEltType = rhsShapedType.getElementType();
3816 Type resEltType = resShapedType.getElementType();
3817 VectorType lhsType = VectorType::get(
3821 ((wSize - 1) * strideW + 1) + ((kwSize - 1) * dilationW + 1) - 1,
3823 lhsEltType, {
false,
false, scalableChDim});
3824 VectorType rhsType =
3825 VectorType::get({kwSize, cSize}, rhsEltType,
3826 {
false, scalableChDim});
3827 VectorType resType =
3828 VectorType::get({nSize, wSize, cSize}, resEltType,
3829 {
false,
false, scalableChDim});
3833 auto maybeMaskXferOp = [&](ArrayRef<int64_t> maskShape,
3834 ArrayRef<bool> scalableDims,
3835 Operation *opToMask) {
3839 VectorType::get(maskShape, rewriter.
getI1Type(), scalableDims);
3841 SmallVector<bool> inBounds(maskShape.size(),
true);
3842 auto xferOp = cast<VectorTransferOpInterface>(opToMask);
3843 xferOp->setInherentAttr(
3848 cast<LinalgOp>(op).hasPureTensorSemantics(), opToMask, rewriter);
3851 vector::CreateMaskOp::create(rewriter, loc, maskType, mixedDims);
3858 Value
lhs = vector::TransferReadOp::create(
3859 rewriter, loc, lhsType, lhsShaped,
ValueRange{zero, zero, zero},
3860 arith::getZeroConstant(rewriter, loc, lhsEltType));
3861 auto *maybeMaskedLhs = maybeMaskXferOp(
3862 lhsType.getShape(), lhsType.getScalableDims(),
lhs.getDefiningOp());
3865 Value
rhs = vector::TransferReadOp::create(
3866 rewriter, loc, rhsType, rhsShaped,
ValueRange{zero, zero},
3867 arith::getZeroConstant(rewriter, loc, rhsEltType));
3868 auto *maybeMaskedRhs = maybeMaskXferOp(
3869 rhsType.getShape(), rhsType.getScalableDims(),
rhs.getDefiningOp());
3872 Value res = vector::TransferReadOp::create(
3873 rewriter, loc, resType, resShaped,
ValueRange{zero, zero, zero},
3874 arith::getZeroConstant(rewriter, loc, resEltType));
3875 auto *maybeMaskedRes = maybeMaskXferOp(
3876 resType.getShape(), resType.getScalableDims(), res.
getDefiningOp());
3882 SmallVector<Value> lhsVals, rhsVals, resVals;
3883 SmallVector<int64_t> inOutSliceSizes = {nSize, wSizeStep, cSize};
3884 SmallVector<int64_t> inOutStrides = {1, 1, 1};
3888 for (int64_t kw = 0; kw < kwSize; ++kw) {
3889 for (int64_t w = 0; w < wSize; w += wSizeStep) {
3890 lhsVals.push_back(vector::ExtractStridedSliceOp::create(
3891 rewriter, loc, maybeMaskedLhs->getResult(0),
3892 ArrayRef<int64_t>{0, w * strideW + kw * dilationW, 0},
3893 inOutSliceSizes, inOutStrides));
3897 for (int64_t kw = 0; kw < kwSize; ++kw) {
3899 vector::ExtractOp::create(rewriter, loc, maybeMaskedRhs->getResult(0),
3900 ArrayRef<int64_t>{kw}));
3903 for (int64_t w = 0; w < wSize; w += wSizeStep) {
3904 resVals.push_back(vector::ExtractStridedSliceOp::create(
3905 rewriter, loc, maybeMaskedRes->getResult(0),
3906 ArrayRef<int64_t>{0, w, 0}, inOutSliceSizes,
3910 auto linearIndex = [&](int64_t kw, int64_t w) {
3911 return kw * (wSize / wSizeStep) + w;
3916 SmallVector<int64_t> inOutFlattenSliceSizes = {nSize, wSizeStep * cSize};
3917 auto lhsTypeAfterFlattening =
3918 VectorType::get(inOutFlattenSliceSizes, lhsEltType);
3919 auto resTypeAfterFlattening =
3920 VectorType::get(inOutFlattenSliceSizes, resEltType);
3923 for (int64_t kw = 0; kw < kwSize; ++kw) {
3924 for (int64_t w = 0; w < wSize; w += wSizeStep) {
3925 Value lhsVal = lhsVals[linearIndex(kw, w)];
3926 Value resVal = resVals[w];
3931 vector::ShapeCastOp::create(rewriter, loc, lhsTypeAfterFlattening,
3932 lhsVals[linearIndex(kw, w)]);
3933 resVal = vector::ShapeCastOp::create(
3934 rewriter, loc, resTypeAfterFlattening, resVals[w]);
3936 resVals[w] = depthwiseConv1dSliceAsMulAcc(rewriter, loc, lhsVal,
3937 rhsVals[kw], resVal, flatten);
3940 resVals[w] = vector::ShapeCastOp::create(
3941 rewriter, loc, VectorType::get(inOutSliceSizes, resEltType),
3948 if (!llvm::all_of(resVals, [](Value v) {
return v; })) {
3950 for (
auto &collection :
3951 {resVals, rhsVals, lhsVals, {res,
rhs,
lhs, zero}})
3952 for (Value v : collection)
3959 for (int64_t w = 0; w < wSize; w += wSizeStep) {
3960 maybeMaskedRes = vector::InsertStridedSliceOp::create(
3961 rewriter, loc, resVals[w], maybeMaskedRes->getResult(0),
3962 ArrayRef<int64_t>{0, w, 0},
3963 ArrayRef<int64_t>{1, 1, 1});
3970 Operation *resOut = vector::TransferWriteOp::create(
3971 rewriter, loc, maybeMaskedRes->getResult(0), resShaped,
3973 return maybeMaskXferOp(resType.getShape(), resType.getScalableDims(),
3981 Value depthwiseConv1dSliceAsMulAcc(RewriterBase &rewriter, Location loc,
3982 Value
lhs, Value
rhs, Value res,
3984 auto rhsTy = cast<ShapedType>(
rhs.getType());
3985 auto resTy = cast<ShapedType>(res.
getType());
3999 auto rhsSize = cast<VectorType>(
rhs.getType()).getShape()[0];
4000 auto resSize = cast<VectorType>(res.
getType()).getShape()[1];
4002 SmallVector<int64_t, 16>
indices;
4003 for (
int i = 0; i < resSize / rhsSize; ++i) {
4004 for (
int j = 0; j < rhsSize; ++j)
4011 rhs = vector::BroadcastOp::create(rewriter, loc,
4012 resTy.clone(rhsTy.getElementType()),
rhs);
4019 if (isa<FloatType>(resTy.getElementType()))
4020 return vector::FMAOp::create(rewriter, loc,
lhs,
rhs, res);
4022 auto mul = arith::MulIOp::create(rewriter, loc,
lhs,
rhs);
4023 return arith::AddIOp::create(rewriter, loc,
mul, res);
4028 FailureOr<Operation *> generateNonChanneledConv() {
4031 if (!iters({Par(), Red()}))
4033 "failed to match conv::W 1-par 1-red");
4036 if (layout({ {w + kw},
4046 FailureOr<Operation *> generateNwcConv() {
4047 AffineExpr n, w, f, kw, c;
4049 if (!iters({Par(), Par(), Par(), Red(), Red()}))
4051 op,
"failed to match conv::Nwc 3-par 2-red");
4054 if (layout({ {n, strideW * w + dilationW * kw, c},
4064 FailureOr<Operation *> generateNcwConv() {
4065 AffineExpr n, w, f, kw, c;
4067 if (!iters({Par(), Par(), Par(), Red(), Red()}))
4069 op,
"failed to match conv::Ncw 3-par 2-red");
4071 if (layout({ {n, c, strideW * w + dilationW * kw},
4081 FailureOr<Operation *> generateNwcPooling() {
4082 AffineExpr n, w, c, kw;
4084 if (!iters({Par(), Par(), Par(), Red()}))
4086 "failed to match pooling 3-par 1-red");
4089 if (layout({ {n, strideW * w + dilationW * kw, c},
4099 FailureOr<Operation *> generateNcwPooling() {
4100 AffineExpr n, w, c, kw;
4102 if (!iters({Par(), Par(), Par(), Red()}))
4104 "failed to match pooling 3-par 1-red");
4106 if (layout({ {n, c, strideW * w + dilationW * kw},
4116 FailureOr<Operation *> generateDilatedConv(uint64_t vecChDimSize = 0,
4117 bool vecChDimScalableFlag =
false,
4118 bool flatten =
false) {
4119 AffineExpr n, w, c, kw;
4121 if (!iters({Par(), Par(), Par(), Red()}))
4123 op,
"failed to match depthwise::Nwc conv 3-par 1-red");
4126 if (layout({ {n, strideW * w + dilationW * kw, c},
4129 return depthwiseConv(vecChDimSize, vecChDimScalableFlag, flatten);
4135 ConvOperationKind oper = ConvOperationKind::Conv;
4137 StringAttr poolExtOp;
4138 bool isPoolExt =
false;
4141 Operation *lhsCastOp =
nullptr;
4142 Operation *rhsCastOp =
nullptr;
4143 int strideW, dilationW;
4144 Value lhsShaped, rhsShaped, resShaped;
4145 ShapedType lhsShapedType, rhsShapedType, resShapedType;
4146 vector::CombiningKind reductionKind;
4150 void setConvOperationKind(Operation *reduceOp) {
4151 int numBlockArguments =
4152 llvm::count_if(reduceOp->
getOperands(), llvm::IsaPred<BlockArgument>);
4153 if (numBlockArguments == 1) {
4158 auto feedValIt = llvm::find_if_not(reduceOp->
getOperands(),
4159 llvm::IsaPred<BlockArgument>);
4160 Operation *feedOp = (*feedValIt).getDefiningOp();
4161 if (isCastOfBlockArgument(feedOp)) {
4162 oper = ConvOperationKind::Pool;
4167 oper = ConvOperationKind::Conv;
4168 setConvCastOps(feedOp);
4172 oper = ConvOperationKind::Pool;
4177 void setConvCastOps(Operation *feedOp) {
4187 RewriterBase &rewriter, LinalgOp op, ArrayRef<int64_t> inputVecSizes,
4188 ArrayRef<bool> inputScalableVecDims,
bool flatten1DDepthwiseConv) {
4189 FailureOr<Conv1DGenerator> conv1dGen = Conv1DGenerator::create(rewriter, op);
4192 auto res = conv1dGen->generateNonChanneledConv();
4195 res = conv1dGen->generateNwcConv();
4198 res = conv1dGen->generateNcwConv();
4201 res = conv1dGen->generateNwcPooling();
4204 res = conv1dGen->generateNcwPooling();
4211 uint64_t vecChDimSize = ShapedType::kDynamic;
4212 bool vecChDimScalableFlag =
false;
4213 if (!inputVecSizes.empty()) {
4218 "Not a 1D depthwise conv!");
4219 size_t chDimIdx = 0;
4225 vecChDimSize = inputVecSizes[chDimIdx];
4226 vecChDimScalableFlag = inputScalableVecDims[chDimIdx];
4228 return conv1dGen->generateDilatedConv(vecChDimSize, vecChDimScalableFlag,
4229 flatten1DDepthwiseConv);
4232struct VectorizeConvolution :
public OpInterfaceRewritePattern<LinalgOp> {
4235 LogicalResult matchAndRewrite(LinalgOp op,
4236 PatternRewriter &rewriter)
const override {
4238 if (
failed(resultOrFail))
4240 Operation *newOp = *resultOrFail;
4242 rewriter.
eraseOp(op.getOperation());
4245 assert(newOp->
getNumResults() == 1 &&
"expected single result");
4252 RewritePatternSet &patterns, PatternBenefit benefit) {
4253 patterns.
add<VectorizeConvolution>(patterns.
getContext(), benefit);
static std::optional< VectorShape > vectorShape(Type type)
static bool isLoopInvariantIdx(LinalgOp &linalgOp, Value &val, VectorType resType)
Checks whether val can be used for calculating a loop invariant index.
static Value insertConvResultSlices(RewriterBase &rewriter, Location loc, Value res, int64_t wSize, int64_t wSizeStep, SmallVectorImpl< Value > &resVals, bool isSingleChanneled)
Helper function to insert the computed result slices.
static SmallVector< bool > getDimsToReduce(LinalgOp linalgOp)
static VectorMemoryAccessKind getTensorExtractMemoryAccessPattern(tensor::ExtractOp extractOp, LinalgOp &linalgOp, VectorType resType)
Infer the memory access pattern for the input ExtractOp.
static SmallVector< Value > extractConvInputSlices(RewriterBase &rewriter, Location loc, Value input, int64_t nSize, int64_t wSize, int64_t cSize, int64_t kwSize, int strideW, int dilationW, int64_t wSizeStep, bool isSingleChanneled)
Helper function to extract the input slices after filter is unrolled along kw.
static VectorizationHookResult vectorizeTensorExtract(RewriterBase &rewriter, VectorizationState &state, Operation *op, LinalgOp linalgOp, const IRMapping &bvm)
Helper function to vectorize the tensor.extract operations.
static VectorizationHookResult vectorizeLinalgIndex(RewriterBase &rewriter, VectorizationState &state, Operation *op, LinalgOp linalgOp)
Helper function to vectorize the index operations of a linalgOp.
static LogicalResult vectorizeAsInsertSliceOp(RewriterBase &rewriter, tensor::InsertSliceOp sliceOp, ArrayRef< int64_t > inputVectorSizes, SmallVectorImpl< Value > &newResults)
Vectorize tensor::InsertSliceOp with:
static FailureOr< Operation * > vectorizeConvolution(RewriterBase &rewriter, LinalgOp convOp, ArrayRef< int64_t > inputVecSizes={}, ArrayRef< bool > inputVecScalableFlags={}, bool flatten1DDepthwiseConv=false)
Try to vectorize convOp as a convolution.
static LogicalResult vectorizeAsLinalgGeneric(RewriterBase &rewriter, VectorizationState &state, LinalgOp linalgOp, SmallVectorImpl< Value > &newResults)
Generic vectorization function that rewrites the body of a linalgOp into vector form.
#define MATCH_1D_CONV_POOL_OP(ConvOpTy)
static VectorizationHookResult vectorizeOneOp(RewriterBase &rewriter, VectorizationState &state, LinalgOp linalgOp, Operation *op, const IRMapping &bvm, ArrayRef< CustomVectorizationHook > customVectorizationHooks)
Generic vectorization for a single operation op, given already vectorized operands carried by bvm.
static Operation * matchLinalgReduction(OpOperand *outputOperand)
Check whether outputOperand is a reduction with a single combiner operation.
static Value buildVectorWrite(RewriterBase &rewriter, Value value, OpOperand *outputOperand, VectorizationState &state)
Build a vector.transfer_write of value into outputOperand at indices set to all 0; where outputOperan...
static Value getStaticPadVal(Operation *op)
Returns the effective Pad value for the input op, provided it's a scalar.
static SmallVector< Value > extractConvFilterSlices(RewriterBase &rewriter, Location loc, Value filter, int64_t kwSize)
Helper function to extract the filter slices after filter is unrolled along kw.
static bool hasReductionIterator(LinalgOp &op)
Check if op is a linalg.reduce or a linalg.generic that has at least one reduction iterator.
std::function< LogicalResult(Operation *, bool)> CustomVectorizationPrecondition
static uint64_t getTrailingNonUnitLoopDimIdx(LinalgOp linalgOp)
Find the index of the trailing non-unit dim in linalgOp.
static VectorType getCollapsedVecType(VectorType type, ArrayRef< AffineMap > reassociation)
Given the re-associations, "collapses" the input Vector type.
Conv1DOpOrder
Helper enum to represent conv1d input traversal order.
VectorizationHookStatus
Helper data structure to represent the result of vectorization for a single operation.
@ Failure
Op failed to vectorize.
@ NewOp
Op vectorized into a new Op whose results will replace original Op's results.
@ NoReplace
Op vectorized and custom function took care of replacement logic.
static Operation * reduceIfNeeded(OpBuilder &b, LinalgOp linalgOp, Operation *op, Value reduceValue, Value initialValue, const IRMapping &bvm)
Emit reduction operations if the shapes of the value to reduce is different that the result shape.
std::function< VectorizationHookResult(Operation *, const IRMapping &)> CustomVectorizationHook
static AffineMap reindexIndexingMap(AffineMap map)
Given an indexing map coming from a LinalgOp indexing, restricted to a projectedPermutation,...
static LogicalResult tensorExtractVectorizationPrecondition(Operation *op, bool vectorizeNDExtract)
Helper function to check if the tensor.extract can be vectorized by the custom hook vectorizeTensorEx...
static Value broadcastIfNeeded(OpBuilder &b, Value value, Type dstType)
Broadcast value to a vector of shape if possible.
static Value calculateGatherOffset(RewriterBase &rewriter, VectorizationState &state, tensor::ExtractOp extractOp, const IRMapping &bvm)
Calculates the offsets ($index_vec) for vector.gather operations generated from tensor....
static SmallVector< Value > extractConvResultSlices(RewriterBase &rewriter, Location loc, Value res, int64_t nSize, int64_t wSize, int64_t fSize, int64_t wSizeStep, bool isSingleChanneled)
Helper function to extract the result slices after filter is unrolled along kw.
static bool isContiguousLoadIdx(LinalgOp &linalgOp, Value &val, bool &foundIndexOp, VectorType resType)
Check whether val could be used for calculating the trailing index for a contiguous load operation.
static VectorizationHookResult vectorizeLinalgYield(RewriterBase &rewriter, Operation *op, const IRMapping &bvm, VectorizationState &state, LinalgOp linalgOp, SmallVectorImpl< Value > &newResults)
Helper function to vectorize the terminator of a linalgOp.
static Operation * buildMultiDimReduce(OpBuilder &b, Operation *reduceOp, Value valueToReduce, Value acc, ArrayRef< bool > dimsToMask)
Create MultiDimReductionOp to compute the reduction for reductionOp.
A dimensional identifier appearing in an affine expression.
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
static AffineMap getMinorIdentityMap(unsigned dims, unsigned results, MLIRContext *context)
Returns an identity affine map (d0, ..., dn) -> (dp, ..., dn) on the most minor dimensions.
MLIRContext * getContext() const
static AffineMap getMultiDimIdentityMap(unsigned numDims, MLIRContext *context)
Returns an AffineMap with 'numDims' identity result dim exprs.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
bool isProjectedPermutation(bool allowZeroInResults=false) const
Returns true if the AffineMap represents a subset (i.e.
unsigned getNumResults() const
unsigned getNumInputs() const
AffineExpr getResult(unsigned idx) const
static AffineMap getFilteredIdentityMap(MLIRContext *ctx, unsigned numDims, llvm::function_ref< bool(AffineDimExpr)> keepDimFilter)
Returns an identity affine map with numDims input dimensions and filtered results using keepDimFilter...
AffineMap dropZeroResults()
Returns the AffineMap resulting from removing "zero" results (constant values == 0) from this map.
static AffineMap getPermutationMap(ArrayRef< unsigned > permutation, MLIRContext *context)
Returns an AffineMap representing a permutation.
SmallVector< unsigned > getBroadcastDims() const
Returns the list of broadcast dimensions (i.e.
AffineMap compose(AffineMap map) const
Returns the AffineMap resulting from composing this with map.
bool isPermutation() const
Returns true if the AffineMap represents a symbol-less permutation map.
This class represents an argument of a Block.
unsigned getArgNumber() const
Returns the number of this argument.
Block represents an ordered list of Operations.
OpListType & getOperations()
AffineMap getMultiDimIdentityMap(unsigned rank)
StringAttr getStringAttr(const Twine &bytes)
TypedAttr getZeroAttr(Type type)
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
MLIRContext * getContext() const
ArrayAttr getBoolArrayAttr(ArrayRef< bool > values)
static DenseIntElementsAttr get(const ShapedType &type, Arg &&arg)
Get an instance of a DenseIntElementsAttr with the given arguments.
This is a utility class for mapping one set of IR entities to another.
auto lookup(T from) const
Lookup a mapped value within the map.
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
IRValueT get() const
Return the current value being used by this operand.
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.
Operation * clone(Operation &op, IRMapping &mapper)
Creates a deep copy of the specified operation, remapping any operands that use values outside of the...
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Operation * create(const OperationState &state)
Creates an operation given the fields represented as an OperationState.
Operation * insert(Operation *op)
Insert the given operation at the current insertion point and return it.
This class represents an operand of an operation.
unsigned getOperandNumber() const
Return which operand this is in the OpOperand list of the Operation.
StringAttr getIdentifier() const
Return the name of this operation as a StringAttr.
Operation is the basic unit of execution within MLIR.
PropertyRef getPropertiesStorage()
Return a generic (but typed) reference to the property type storage.
Value getOperand(unsigned idx)
bool isBeforeInBlock(Operation *other)
Given an operation 'other' that is within the same parent block, return whether the current operation...
Block * getBlock()
Returns the operation block that contains this operation.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Location getLoc()
The source location the operation was defined or derived from.
unsigned getNumOperands()
Attribute getPropertiesAsAttribute()
Return the properties converted to an attribute.
operand_iterator operand_end()
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.
static Operation * create(Location location, OperationName name, TypeRange resultTypes, ValueRange operands, NamedAttrList &&attributes, PropertyRef properties, BlockRange successors, unsigned numRegions)
Create a new Operation with the specific fields.
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.
unsigned short getBenefit() const
If the corresponding pattern can match, return its benefit. If the.
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.
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...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
void replaceAllUsesExcept(Value from, Value to, Operation *exceptedUser)
Find uses of from and replace them with to except if the user is exceptedUser.
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,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
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.
use_range getUses() const
Returns a range of all uses, which is useful for iterating over all uses.
Location getLoc() const
Return the location of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Operation * getOwner() const
Return the owner of this operand.
bool hasElementwiseMappableTraits(Operation *op)
Together, Elementwise, Scalarizable, Vectorizable, and Tensorizable provide an easy way for scalar op...
bool hasVectorizationImpl(Operation *)
Return true if there's dedicated logic in the Linalg Vectorizer to vectorize this Op,...
SmallVector< int64_t > getUnPackInverseSrcPerm(linalg::UnPackOp, PackingMetadata &metadata)
Compute inverse permutation for the source tensor (i.e.
bool allIndexingsAreProjectedPermutation(LinalgOp op)
Check if all indexing maps are projected permutations.
FailureOr< VectorizationResult > vectorize(RewriterBase &rewriter, Operation *op, ArrayRef< int64_t > inputVectorSizes={}, ArrayRef< bool > inputScalableVecDims={}, bool vectorizeNDExtract=false, bool flatten1DDepthwiseConv=false, bool assumeDynamicDimsMatchVecSizes=false, bool createNamedContraction=false)
Returns a VectorizationResult containing the results of the vectorized op, or failure if the transfor...
void populatePadOpVectorizationPatterns(RewritePatternSet &patterns, PatternBenefit baseBenefit=1)
Populates patterns with patterns that vectorize tensor.pad.
bool isReductionIterator(utils::IteratorType iteratorType)
Check if iterator type has "reduction" semantics.
bool isaConvolutionOpInterface(LinalgOp linalgOp, bool allowEmptyConvolvedDims=false)
Checks whether linalgOp conforms to ConvolutionOpInterface.
void populateConvolutionVectorizationPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Populate patterns for vectorizing low-D convolution ops.
bool isElementwise(LinalgOp op)
Check if a LinalgOp is an element-wise operation.
LogicalResult vectorizeCopy(RewriterBase &builder, memref::CopyOp copyOp)
Emit a suitable vector form for a Copy op with fully static shape.
LogicalResult vectorizeOpPrecondition(Operation *op, ArrayRef< int64_t > inputVectorSizes={}, ArrayRef< bool > inputScalableVecDims={}, bool vectorizeNDExtract=false, bool flatten1DDepthwiseConv=false)
Return success if the operation can be vectorized.
SmallVector< int64_t > getPackInverseDestPerm(linalg::PackOp packOp, PackingMetadata &metadata)
Compute inverse permutation for the destination tensor (i.e.
bool isaConvolutionOpOfType(LinalgOp op)
Returns true if the linalg op is a convolution op of type ConvOpTy.
std::optional< vector::CombiningKind > getCombinerOpKind(Operation *combinerOp)
Return vector::CombiningKind for the given op.
void promote(RewriterBase &rewriter, scf::ForallOp forallOp)
Promotes the loop body of a scf::ForallOp to its containing block.
std::enable_if_t<!is_complex< V >::value, V > readValue(char **linePtr)
Returns an element-value of non-complex type.
Operation * maskOperation(OpBuilder &builder, Operation *maskableOp, Value mask, Value passthru=Value())
Creates a vector.mask operation around a maskable operation.
LogicalResult isValidMaskedInputVector(ArrayRef< int64_t > shape, ArrayRef< int64_t > inputVectorSizes)
Returns success if inputVectorSizes is a valid masking configuraion for given shape,...
BroadcastableToResult isBroadcastableTo(Type srcType, VectorType dstVectorType, std::pair< VectorDim, VectorDim > *mismatchingDims=nullptr)
Return whether srcType can be broadcast to dstVectorType under the semantics of the vector....
Operation * createWriteOrMaskedWrite(OpBuilder &builder, Location loc, Value vecToStore, Value dest, SmallVector< Value > writeIndices={}, bool useInBoundsInsteadOfMasking=false, AffineMap permutationMap=AffineMap())
Create a TransferWriteOp of vecToStore into dest.
Value createReadOrMaskedRead(OpBuilder &builder, Location loc, Value source, const VectorType &vecToReadTy, std::optional< Value > padValue=std::nullopt, bool useInBoundsInsteadOfMasking=false, ArrayRef< Value > indices={}, AffineMap permutationMap=AffineMap())
Creates a TransferReadOp from source.
SmallVector< OpFoldResult > getMixedSizesXfer(bool hasTensorSemantics, Operation *xfer, RewriterBase &rewriter)
A wrapper for getMixedSizes for vector.transfer_read and vector.transfer_write Ops (for source and de...
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
bool isEqualConstantIntOrValue(OpFoldResult ofr1, OpFoldResult ofr2)
Return true if ofr1 and ofr2 are the same integer constant attribute values or the same SSA value.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
AffineMap inverseAndBroadcastProjectedPermutation(AffineMap map)
Return the reverse map of a projected permutation where the projected dimensions are transformed into...
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
AffineMap inversePermutation(AffineMap map)
Returns a map of codomain to domain dimensions such that the first codomain dimension for a particula...
SmallVector< AffineMap, 4 > getSymbolLessAffineMaps(ArrayRef< ReassociationExprs > reassociation)
Constructs affine maps out of Array<Array<AffineExpr>>.
SmallVector< SmallVector< OpFoldResult > > ReifiedRankedShapedTypeDims
Value matchReduction(ArrayRef< BlockArgument > iterCarriedArgs, unsigned redPos, SmallVectorImpl< Operation * > &combinerOps)
Utility to match a generic reduction given a list of iteration-carried arguments, iterCarriedArgs and...
llvm::SetVector< T, Vector, Set, N > SetVector
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
std::conditional_t< std::is_same_v< Ty, mlir::Type >, mlir::Value, detail::TypedValue< Ty > > TypedValue
If Ty is mlir::Type this will select Value instead of having a wrapper around it.
void getUsedValuesDefinedAbove(Region ®ion, Region &limit, SetVector< Value > &values)
Fill values with a list of values defined at the ancestors of the limit region and used within region...
AffineMap compressUnusedDims(AffineMap map)
Drop the dims that are not used.
SmallVector< SmallVector< AffineExpr, 2 >, 2 > convertReassociationIndicesToExprs(MLIRContext *context, ArrayRef< ReassociationIndices > reassociationIndices)
Convert reassociation indices to affine expressions.
bool isReassociationValid(ArrayRef< AffineMap > reassociation, int *invalidIndex=nullptr)
Return true if the reassociation specification is valid, false otherwise.
llvm::TypeSwitch< T, ResultT > TypeSwitch
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
AffineExpr getAffineConstantExpr(int64_t constant, MLIRContext *context)
auto get(MLIRContext *context, Ts &&...params)
Helper method that injects context only if needed, this helps unify some of the attribute constructio...
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
SmallVector< T > applyPermutationMap(AffineMap map, llvm::ArrayRef< T > source)
Apply a permutation from map to source and return the result.
llvm::SmallBitVector getUnusedDimsBitVector(ArrayRef< AffineMap > maps)
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
void applyPermutationToVector(SmallVector< T, N > &inVec, ArrayRef< int64_t > permutation)
Apply the permutation defined by permutation to inVec.
SmallVector< int64_t > invertPermutationVector(ArrayRef< int64_t > permutation)
Helper method to apply to inverse a permutation.
VectorizationHookResult contains the vectorized op returned from a CustomVectorizationHook.
enum VectorizationHookStatus status
Return status from vectorizing the current op.
Operation * newOp
New vectorized operation to replace the current op.
ArrayRef< int64_t > getCanonicalVecShape() const
Returns the canonical vector shape used to vectorize the iteration space.
LogicalResult initState(RewriterBase &rewriter, LinalgOp linalgOp, ArrayRef< int64_t > inputVectorSizes, ArrayRef< bool > inputScalableVecDims, bool assumeDynamicDimsMatchVecSizes=false)
Initializes the vectorization state, including the computation of the canonical vector shape for vect...
Operation * maskOperation(RewriterBase &rewriter, Operation *opToMask, LinalgOp linalgOp, std::optional< AffineMap > maybeIndexingMap=std::nullopt)
Masks an operation with the canonical vector mask if the operation needs masking.
VectorType getCanonicalVecType(Type elementType, std::optional< AffineMap > dimPermutation=std::nullopt) const
Returns a vector type of the provided elementType with the canonical vector shape and the correspondi...
ArrayRef< bool > getScalableVecDims() const
Returns the vector dimensions that are scalable in the canonical vector shape.
VectorizationState(RewriterBase &rewriter)
OpInterfaceRewritePattern(MLIRContext *context, PatternBenefit benefit=1)
LogicalResult matchAndRewrite(vector::TransferReadOp xferOp, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(vector::TransferWriteOp xferOp, PatternRewriter &rewriter) const override