32#include "llvm/ADT/SmallVectorExtras.h"
33#include "llvm/ADT/TypeSwitch.h"
34#include "llvm/Support/Debug.h"
35#include "llvm/Support/DebugLog.h"
36#include "llvm/Support/InterleavedRange.h"
37#include "llvm/Support/raw_ostream.h"
41#define DEBUG_TYPE "linalg-transforms"
60 .Case([&](scf::ForOp forOp) {
61 scf::ForOp partialIteration;
64 return partialIteration->getResults();
65 assert(!partialIteration &&
"expected that loop was not peeled");
66 return forOp->getResults();
75 for (
auto loopOp : loops)
88 if (!e.isFunctionOfDim(dim))
99 return llvm::interleaved(ri,
", ",
"|",
"");
150static FailureOr<SmallVector<std::optional<int64_t>>>
154 int64_t newDim = iteratorTypes.size();
155 iteratorTypes.push_back(iteratorTypes[dim]);
158 indexingMaps.size(), std::nullopt);
160 for (
int64_t operandIdx = 0, e = indexingMaps.size(); operandIdx < e;
162 AffineMap map = indexingMaps[operandIdx];
165 assert(map.
getNumDims() == newDim &&
"num dims invariant violation");
173 "num results invariant violation");
175 if (!maybeOperandDimensionToPack.has_value()) {
176 newMaps.push_back(map);
181 if (!isa<AffineDimExpr>(map.
getResult(maybeOperandDimensionToPack.value())))
187 newMaps.push_back(map);
190 packedDimPerIndexingMap[operandIdx] = maybeOperandDimensionToPack;
192 indexingMaps = newMaps;
194 return packedDimPerIndexingMap;
200struct PackedOperandsDim {
201 OpFoldResult packedSize;
202 SmallVector<std::optional<int64_t>> packedDimForEachOperand;
206struct PackedOperandsDimList {
207 void pushBack(PackedOperandsDim &&packedOperandsDims) {
208 spec.emplace_back(packedOperandsDims);
211 SmallVector<int64_t> extractPackedDimsForOperand(int64_t operandPos);
213 SmallVector<OpFoldResult> extractPackSizesForOperand(int64_t operandPos);
216 SmallVector<PackedOperandsDim> spec;
222 linalg::PackOp packOp,
223 bool lowerPadLikeWithInsertSlice) {
227 if (!packOp.hasPureTensorSemantics())
230 auto packedTensorType =
231 cast<RankedTensorType>(packOp->getResultTypes().front());
239 PackingMetadata packingMetadata;
262 for (
auto [pos, innerSize] :
263 llvm::zip_equal(packOp.getInnerDimsPos(), packOp.getMixedTiles())) {
265 packedToStripMinedShapePerm[packingMetadata.outerPositions[pos]];
275 rewriter, loc, map, {outerSize, origSize, innerSize});
277 RankedTensorType collapsed = tensor::CollapseShapeOp::inferCollapsedType(
279 packingMetadata.reassociations);
280 Value paddingValue = packOp.getPaddingValue();
282 paddingValue = arith::ConstantOp::create(
286 tensor::PadOp::create(rewriter, loc, collapsed, packOp.getSource(), lows,
287 highs, paddingValue,
false);
289 LDBG() <<
"insertPositions: "
290 << llvm::interleaved(packingMetadata.insertPositions);
291 LDBG() <<
"outerPositions: "
292 << llvm::interleaved(packingMetadata.outerPositions);
293 LDBG() <<
"packedShape: " << llvm::interleaved(packedTensorType.getShape());
294 LDBG() <<
"packedToStripMinedShapePerm: "
295 << llvm::interleaved(packedToStripMinedShapePerm);
296 LDBG() <<
"reassociations: "
297 << llvm::interleaved(llvm::map_range(packingMetadata.reassociations,
299 LDBG() <<
"stripMinedShape: " << llvm::interleaved(stripMinedShape);
300 LDBG() <<
"collapsed type: " << collapsed;
302 if (lowerPadLikeWithInsertSlice && packOp.isLikePad()) {
321 auto insertSliceOp = tensor::InsertSliceOp::create(
322 rewriter, loc, padOp, packOp.getDest(),
325 LDBG() <<
"insert_slice op: " << insertSliceOp;
327 rewriter.
replaceOp(packOp, insertSliceOp->getResults());
335 auto expandShapeResultType =
337 auto reshapeOp = tensor::ExpandShapeOp::create(
338 rewriter, loc, expandShapeResultType, padOp.getResult(),
339 packingMetadata.reassociations, stripMinedMixedSizes);
344 auto transposeOp = linalg::TransposeOp::create(
345 rewriter, loc, reshapeOp.getResult(), packOp.getDest(), transpPerm);
347 LDBG() <<
"reshape op: " << reshapeOp;
348 LDBG() <<
"transpPerm: " << llvm::interleaved(transpPerm);
349 LDBG() <<
"transpose op: " << transposeOp;
352 rewriter.
replaceOp(packOp, transposeOp->getResults());
357FailureOr<LowerUnPackOpResult>
359 bool lowerUnpadLikeWithExtractSlice) {
363 if (!unPackOp.hasPureTensorSemantics())
370 auto packedTensorType = cast<RankedTensorType>(unPackOp.getSourceType());
371 int64_t packedRank = packedTensorType.getRank();
374 auto destTensorType = cast<RankedTensorType>(unPackOp.getDest().getType());
375 if (lowerUnpadLikeWithExtractSlice && unPackOp.isLikeUnPad()) {
384 auto extractSliceOp = tensor::ExtractSliceOp::create(
385 rewriter, loc, destTensorType, unPackOp.getSource(),
389 rewriter.
replaceOp(unPackOp, extractSliceOp->getResults());
392 nullptr, extractSliceOp,
398 PackingMetadata packingMetadata;
408 RankedTensorType stripMinedTensorType =
410 RankedTensorType collapsedType = tensor::CollapseShapeOp::inferCollapsedType(
411 stripMinedTensorType, packingMetadata.reassociations);
418 auto emptyOp = tensor::EmptyOp::create(rewriter, loc, dims,
419 stripMinedTensorType.getElementType());
421 linalg::TransposeOp::create(rewriter, loc, unPackOp.getSource(), emptyOp,
422 packedToStripMinedShapePerm);
424 LDBG() <<
"insertPositions: "
425 << llvm::interleaved(packingMetadata.insertPositions);
426 LDBG() <<
"packedShape: " << llvm::interleaved(packedTensorType.getShape());
427 LDBG() <<
"packedToStripMinedShapePerm: "
428 << llvm::interleaved(packedToStripMinedShapePerm);
429 LDBG() <<
"reassociations: "
430 << llvm::interleaved(llvm::map_range(packingMetadata.reassociations,
432 LDBG() <<
"stripMinedShape: " << llvm::interleaved(stripMinedShape);
433 LDBG() <<
"collapsed type: " << collapsedType;
436 auto reshapeOp = tensor::CollapseShapeOp::create(
437 rewriter, loc, collapsedType, transposeOp->getResult(0),
438 packingMetadata.reassociations);
441 int64_t destRank = destTensorType.getRank();
442 auto extractSliceOp = tensor::ExtractSliceOp::create(
443 rewriter, loc, destTensorType, reshapeOp->getResult(0),
449 auto copyOp = linalg::CopyOp::create(
450 rewriter, loc, extractSliceOp->getResult(0), unPackOp.getDest());
453 rewriter.
replaceOp(unPackOp, copyOp->getResults());
460PackedOperandsDimList::extractPackedDimsForOperand(
int64_t operandPos) {
462 for (
auto &i : spec) {
463 if (!i.packedDimForEachOperand[operandPos].has_value())
465 res.push_back(i.packedDimForEachOperand[operandPos].value());
470SmallVector<OpFoldResult>
471PackedOperandsDimList::extractPackSizesForOperand(int64_t operandPos) {
472 SmallVector<OpFoldResult> res;
473 for (
auto &i : spec) {
474 if (!i.packedDimForEachOperand[operandPos].has_value())
476 res.push_back(i.packedSize);
485 linalg::LinalgOp linalgOp,
487 if (packedSizes.size() != linalgOp.getNumLoops()) {
489 "incorrect number of pack sizes");
491 if (!linalgOp.hasPureTensorSemantics()) {
493 linalgOp,
"expects LinalgOp with pure tensor semantics");
499 linalgOp.getIteratorTypesArray();
500 LDBG() <<
"Start packing: " << linalgOp;
501 LDBG() <<
"maps: " << llvm::interleaved(indexingMaps);
502 LDBG() <<
"iterators: " << llvm::interleaved(iteratorTypes);
507 PackedOperandsDimList listOfPackedOperandsDim;
508 for (
int64_t i = 0, e = packedSizes.size(); i < e; ++i) {
511 if (maybeConstant.has_value() && maybeConstant.value() == 0)
514 PackedOperandsDim packedOperandsDims;
515 packedOperandsDims.packedSize = packedSizes[i];
516 FailureOr<SmallVector<std::optional<int64_t>>>
517 maybePackedDimForEachOperand =
519 if (failed(maybePackedDimForEachOperand))
521 packedOperandsDims.packedDimForEachOperand = *maybePackedDimForEachOperand;
523 LDBG() <<
"++++ After pack size #" << i <<
": " << packedSizes[i];
524 LDBG() <<
"maps: " << llvm::interleaved(indexingMaps);
525 LDBG() <<
"iterators: " << llvm::interleaved(iteratorTypes);
526 LDBG() <<
"packedDimForEachOperand: "
527 << llvm::interleaved(packedOperandsDims.packedDimForEachOperand);
529 listOfPackedOperandsDim.pushBack(std::move(packedOperandsDims));
535 llvm::to_vector(llvm::make_pointer_range(linalgOp.getDpsInitsMutable()));
537 for (
const auto &operandsList : {inputOperands, initOperands}) {
538 for (
OpOperand *opOperand : operandsList) {
539 int64_t pos = opOperand->getOperandNumber();
540 Value operand = opOperand->get();
542 listOfPackedOperandsDim.extractPackedDimsForOperand(pos);
544 listOfPackedOperandsDim.extractPackSizesForOperand(pos);
545 LDBG() <<
"operand: " << operand;
546 LDBG() <<
"innerPos: " << llvm::interleaved(innerPos);
547 LDBG() <<
"innerPackSizes: " << llvm::interleaved(innerPackSizes);
548 if (innerPackSizes.empty()) {
549 inputsAndInits.push_back(operand);
552 Value dest = linalg::PackOp::createDestinationTensor(
553 rewriter, loc, operand, innerPackSizes, innerPos,
555 ShapedType operandType = cast<ShapedType>(operand.
getType());
556 bool areConstantTiles =
560 if (areConstantTiles && operandType.hasStaticShape() &&
561 !linalg::PackOp::requirePaddingValue(
562 operandType.getShape(), innerPos,
563 cast<ShapedType>(dest.
getType()).getShape(), {},
565 packOps.push_back(linalg::PackOp::create(rewriter, loc, operand, dest,
566 innerPos, innerPackSizes));
572 Value zero = arith::ConstantOp::create(rewriter, loc, zeroAttr);
573 packOps.push_back(linalg::PackOp::create(
574 rewriter, loc, operand, dest, innerPos, innerPackSizes, zero));
576 inputsAndInits.push_back(packOps.back().getResult());
582 ValueRange{inputsAndInits}.take_front(linalgOp.getNumDpsInputs());
584 ValueRange{inputsAndInits}.take_back(linalgOp.getNumDpsInits());
585 auto packedLinalgOp =
586 linalg::GenericOp::create(rewriter, linalgOp.getLoc(), inits.
getTypes(),
587 inputs, inits, indexingMaps, iteratorTypes);
588 packedLinalgOp.getRegion().takeBody(linalgOp->getRegion(0));
593 linalg::PackOp maybePackedInit =
594 inits[resultNum].getDefiningOp<linalg::PackOp>();
595 if (!maybePackedInit) {
596 results.push_back(
result);
600 unPackOps.push_back(linalg::UnPackOp::create(
601 rewriter, packedLinalgOp->getLoc(),
result, maybePackedInit.getSource(),
602 maybePackedInit.getInnerDimsPos(), maybePackedInit.getMixedTiles()));
603 results.push_back(unPackOps.back().getResult());
611 cast<linalg::LinalgOp>(packedLinalgOp.getOperation()),
640 assert(linalgOp == opOperand.
getOwner() &&
"linalg op must own the operand");
644 cast<RankedTensorType>(opOperand.
get().
getType()), permutation);
646 assert(tensorType == transposedValue.
getType() &&
647 "expected tensor type mismatch");
652 llvm::map_to_vector(permutation, [](
int64_t i) ->
unsigned {
return i; });
656 permutationMap.
compose(linalgOp.getMatchingIndexingMap(&opOperand));
660 indexingMaps[linalgOp.getIndexingMapIndex(&opOperand)] = transposedMap;
666 auto transposedGenericOp = linalg::GenericOp::create(
670 operandsRef.drop_front(linalgOp.getNumDpsInputs()).
getTypes(),
671 operandsRef.take_front(linalgOp.getNumDpsInputs()),
672 operandsRef.drop_front(linalgOp.getNumDpsInputs()),
674 linalgOp.getIteratorTypesArray());
675 transposedGenericOp.getRegion().takeBody(linalgOp->getRegion(0));
676 rewriter.
replaceOp(linalgOp, transposedGenericOp->getResults());
678 return cast<linalg::LinalgOp>(transposedGenericOp.getOperation());
681FailureOr<PackTransposeResult>
683 linalg::LinalgOp linalgOp, linalg::UnPackOp maybeUnPackOp,
690 linalg::PackOp transposedPackOp =
691 packOp.createTransposedClone(rewriter, loc, innerPerm, outerPerm);
693 if (packOp.hasPureBufferSemantics() || !packOp.getResult().hasOneUse())
696 OpOperand &packUse = *packOp->getUses().begin();
697 if (packUse.
getOwner() != linalgOp) {
699 linalgOp,
"not a single use by the LinalgOp target");
702 (!linalgOp.isDpsInit(&packUse) ||
703 maybeUnPackOp.getSource() != linalgOp.getTiedOpResult(&packUse))) {
705 "not produced by the LinalgOp target");
711 int64_t numLeadingDims = packOp.getSourceRank();
712 int64_t numTrailingDims = packOp.getInnerDimsPos().size();
716 if (permutation.empty())
717 llvm::append_range(permutation, llvm::seq<int64_t>(0, numLeadingDims));
719 if (innerPerm.empty()) {
722 llvm::seq<int64_t>(numLeadingDims, numLeadingDims + numTrailingDims));
724 llvm::append_range(permutation,
725 llvm::map_range(innerPerm, [&](
int64_t pos) {
726 return numLeadingDims + pos;
738 rewriter, linalgOp, packUse, permutation, transposedPackOp.getResult());
741 linalg::UnPackOp transposedUnPackOp;
744 transposedLinalgOp->getOpOperand(packUseOperandNumber);
745 OpResult transposedResult = transposedLinalgOp.getTiedOpResult(&opOperand);
747 transposedUnPackOp = maybeUnPackOp.createTransposedClone(
748 rewriter, loc, transposedResult, innerPerm, outerPerm);
750 rewriter.
replaceOp(maybeUnPackOp, transposedUnPackOp->getResults());
754 if (packOp.hasPureTensorSemantics())
755 rewriter.
replaceOp(packOp, transposedPackOp->getResults());
780 assert(mnkPackedSizes.size() == 3 &&
"unexpected num of packing sizes");
781 assert((mnkPaddedSizesNextMultipleOf.empty() ||
782 mnkPaddedSizesNextMultipleOf.size() == 3) &&
783 "num of packing sizes next multiple should be empty or of size 3");
784 assert(mnkOrder.size() == 3 &&
"unexpected mnkOrder size");
787 int64_t numLoops = linalgOp.getNumLoops();
789 LDBG() <<
"need 3+ loops to find a matmul to pack, got " << numLoops
790 <<
" in: " << linalgOp;
792 linalgOp,
"need 3+ loops to find a matmul to pack");
796 int64_t numPackedDims = mnkPackedSizes.size();
798 for (
int64_t i = 0, e = numPackedDims; i < e; ++i)
799 mmnnkkPos[i] = numLoops - numPackedDims + mnkOrder[i];
801 for (
int64_t i = 0, e = numPackedDims; i < e; ++i)
802 packedSizes[mnkOrder[i]] = mnkPackedSizes[i];
804 for (
int64_t i = 0, e = numPackedDims; i < e; ++i) {
805 paddedSizesNextMultipleOf[mnkOrder[i]] =
806 mnkPaddedSizesNextMultipleOf.empty() ? 0
807 : mnkPaddedSizesNextMultipleOf[i];
811 FailureOr<ContractionDimensions> maybeDimensions =
813 if (failed(maybeDimensions)) {
814 LDBG() <<
"couldn't infer matmul iterators in: " << linalgOp;
816 "couldn't infer matmul iterators");
824 int64_t mPos = maybeDimensions->m.back(), nPos = maybeDimensions->n.back(),
825 kPos = maybeDimensions->k.back();
826 LDBG() <<
"Start packing generic op greedily with (m@" << mPos <<
", n@"
827 << nPos <<
", k@" << kPos <<
"): " << linalgOp;
830 auto genericOp = dyn_cast<GenericOp>(linalgOp.getOperation());
832 FailureOr<LinalgOp> generalizeResult =
834 assert(succeeded(generalizeResult) &&
835 isa<GenericOp>(generalizeResult->getOperation()) &&
836 "unexpected failure generalizing op");
837 genericOp = cast<GenericOp>(generalizeResult->getOperation());
845 LDBG() <<
"perm: " << llvm::interleaved(permutation);
848 FailureOr<GenericOp> interchangeResult =
850 assert(succeeded(interchangeResult) &&
"unexpected failure interchanging op");
851 genericOp = *interchangeResult;
852 LDBG() <<
"Generalized Op to pack: " << genericOp;
869 cast<LinalgOp>(genericOp.getOperation())
870 .createLoopRanges(rewriter, genericOp.getLoc());
874 LDBG() <<
"paddedSizesNextMultipleOf: "
875 << llvm::interleaved(paddedSizesNextMultipleOf);
876 LDBG() <<
"loopRanges: "
877 << llvm::interleaved(
878 llvm::map_range(loopRanges, [](
Range r) {
return r.
size; }));
881 for (
int64_t i = 0, e = numPackedDims; i < e; ++i) {
882 if (paddedSizesNextMultipleOf[i] == 0) {
883 adjustedPackedSizes.push_back(packedSizes[i]);
890 rewriter, genericOp->getLoc(), d0.
ceilDiv(s0) * s0,
891 {loopRanges[adjustedPackedSizes.size()].size,
892 rewriter.getIndexAttr(paddedSizesNextMultipleOf[i])}));
894 LDBG() <<
"adjustedPackedSizes: " << llvm::interleaved(adjustedPackedSizes);
900 return pack(rewriter, genericOp, adjustedPackedSizes);
913 b.setInsertionPointToStart(
914 &op->getParentOfType<func::FuncOp>().getBody().front());
915 return llvm::map_to_vector<4>(tileSizes, [&](
int64_t s) {
933 auto padValue = padOp.getConstantPaddingValue();
936 if (padValue.getParentBlock() == &padOp.getRegion().front())
938 return FillOp::create(rewriter, padOp.getLoc(), padValue, dest).result();
942 auto generateOp = tensor::GenerateOp::create(rewriter, padOp.getLoc(),
943 padOp.getResultType(), dynSizes);
946 padOp.getRegion().cloneInto(&generateOp.getRegion(), bvm);
955 if (
auto val = llvm::dyn_cast_if_present<Value>(ofr))
958 rewriter, padOp.getLoc(),
959 cast<IntegerAttr>(cast<Attribute>(ofr)).getInt())
963 auto resultType = padOp.getResultType();
967 for (
unsigned dim = 0; dim < resultType.getRank(); ++dim) {
968 if (resultType.isDynamicDim(dim)) {
970 padOp.getSource(), dim));
973 padOp.getLoc(), srcSize, getIdxValue(padOp.getMixedLowPad()[dim]));
975 padOp.getLoc(), plusLow, getIdxValue(padOp.getMixedHighPad()[dim]));
976 dynSizes.push_back(plusHigh);
978 staticSizes.push_back(resultType.getDimSize(dim));
983 tensor::EmptyOp::create(rewriter, padOp.getLoc(), staticSizes,
984 resultType.getElementType(), dynSizes);
988 auto sourceType = padOp.getSourceType();
996 padOp, padOp.getSource(), fill, padOp.getMixedLowPad(), srcSizes,
1004 if (!sliceOp.hasUnitStride())
1007 auto padOp = sliceOp.getSource().getDefiningOp<tensor::PadOp>();
1011 bool zeroSliceGuard =
true;
1013 if (std::optional<bool> control = controlFn(sliceOp))
1014 zeroSliceGuard = *control;
1019 FailureOr<TilingResult> tilingResult =
1021 sliceOp.getMixedSizes(), zeroSliceGuard);
1022 if (failed(tilingResult))
1025 RankedTensorType sourceType = sliceOp.getSourceType();
1026 RankedTensorType resultType = sliceOp.getResultType();
1030 if (sourceType.getRank() == resultType.getRank()) {
1031 rewriter.
replaceOp(sliceOp, tilingResult->tiledValues);
1037 rewriter, sliceOp.getLoc(), tilingResult->tiledValues[0], resultType);
1039 rewriter.
replaceOp(sliceOp, rankReduced);
1049 linalg::PackOp packOp) {
1050 Value input = packOp.getSource();
1054 if (!packOp.hasPureTensorSemantics())
1057 if (!packOp.getPaddingValue()) {
1061 assert(llvm::all_of(packOp.getAllOuterDims(),
1062 [](
int64_t val) { return val == 1; }) &&
1063 "some outer dims are != 1");
1066 ShapedType inputType = packOp.getSourceType();
1067 int64_t inputRank = inputType.getRank();
1070 packOp.getDimAndTileMapping();
1077 for (
int64_t dimIdx = 0; dimIdx < inputRank; ++dimIdx) {
1080 if (!tileAndPosMapping.count(dimIdx)) {
1081 int64_t inputDimSize = inputType.getDimSize(dimIdx);
1082 assert(inputDimSize == 1 &&
1083 "with all outer dims == 1, this non-tiled input dim should be 1!");
1084 paddedShape.push_back(inputDimSize);
1091 OpFoldResult tileSizeForDim = tileAndPosMapping.lookup(dimIdx);
1095 if (cstTileSize.has_value()) {
1096 paddedShape.push_back(cstTileSize.value());
1101 paddedShape.push_back(ShapedType::kDynamic);
1104 dynamicTileSizes.push_back(llvm::dyn_cast<Value>(tileSizeForDim));
1107 RankedTensorType::get(paddedShape, inputType.getElementType());
1109 false, loc, builder,
1117static SmallVector<int64_t>
1119 constexpr int64_t kNonTiledMarker = -1;
1121 for (
auto [
index, value] : llvm::enumerate(perm))
1124 vec, [&](
int64_t v) {
return v != kNonTiledMarker; });
1131static SmallVector<int64_t>
1140 for (
auto i : llvm::seq<unsigned>(0, unpackedRank)) {
1141 if (llvm::is_contained(innerDimsPos, i)) {
1142 innerDims.push_back(dim++);
1147 outerDims.push_back(dim++);
1148 if (!outerDimsPerm.empty())
1149 rankReducedOuterDimsPerm.push_back(outerDimsPerm[i]);
1160 rankReducedOuterDimsPerm =
1162 if (!rankReducedOuterDimsPerm.empty())
1166 perm.append(innerDims);
1176 if (!packOp.hasPureTensorSemantics())
1179 if (llvm::any_of(packOp.getTiledOuterDims(),
1180 [](
int64_t dim) { return dim != 1; })) {
1182 packOp,
"not all outer dimensions of the result are 1s");
1190 if (packOp.getPaddingValue() &&
1191 llvm::any_of(packOp.getAllOuterDims(),
1192 [](
int64_t dim) { return dim != 1; })) {
1194 packOp,
"cannot decompose padded pack with a non-unit un-tiled outer "
1199 auto outerDimsPerm = packOp.getOuterDimsPerm();
1205 if (!llvm::all_of(outerDimsPerm, [&innerDimsPos, &packOp](
int64_t dim) {
1206 static int prev = 0;
1208 if (llvm::is_contained(innerDimsPos, dim))
1213 if (dim < prev && (packOp.getResult().getType().getShape()[prev] != 1 ||
1214 packOp.getResult().getType().getShape()[dim] != 1))
1221 packOp,
"At least one non-unit and un-tiled outer dim is permuted, "
1222 "this is not supported ATM!");
1227 int64_t srcRank = packOp.getSourceRank();
1246 for (
int64_t i = 0; i < srcRank; i++) {
1254 if (llvm::is_contained(innerDimsPos, i))
1256 srcPermForTranspose.push_back(i);
1258 srcPermForTranspose.append(innerDimsPos.begin(), innerDimsPos.end());
1262 ShapedType inputTy = cast<ShapedType>(input.
getType());
1264 for (
int64_t i = 0; i < srcRank; i++) {
1265 if (llvm::is_contained(innerDimsPos, i)) {
1269 if (inputTy.isStaticDim(i))
1270 shapeForEmptyOp.push_back(rewriter.
getIndexAttr(inputTy.getShape()[i]));
1272 shapeForEmptyOp.emplace_back(
1273 tensor::DimOp::create(rewriter, loc, input, i).getResult());
1275 shapeForEmptyOp.append(packOp.getMixedTiles());
1282 llvm::transform(shapeForEmptyOp, shapeForEmptyOp.begin(),
1284 if (auto val = llvm::dyn_cast<Value>(ofr))
1285 return getAsOpFoldResult(val);
1289 LDBG() <<
"Pack permutation: " << packOp;
1290 LDBG() <<
"perm: " << llvm::interleaved(srcPermForTranspose);
1291 LDBG() <<
"Shape of empty tensor: " << llvm::interleaved(shapeForEmptyOp);
1293 Value empty = tensor::EmptyOp::create(
1294 rewriter, loc, shapeForEmptyOp, packOp.getSourceType().getElementType());
1297 auto transposedOp = linalg::TransposeOp::create(rewriter, loc, input, empty,
1298 srcPermForTranspose);
1310 for (
auto size : packOp.getAllOuterDims()) {
1314 for (
auto tileSize : packOp.getMixedTiles()) {
1315 auto [_, tileSizeOfr] =
1317 writeSizes.push_back(tileSizeOfr);
1320 auto insert = tensor::InsertSliceOp::create(
1321 rewriter, loc, transposedOp.getResult()[0], packOp.getDest(), writeSizes);
1324 rewriter.
replaceOp(packOp, insert.getResult());
1331 if (!unpackOp.hasPureTensorSemantics())
1334 int64_t destRank = unpackOp.getDestRank();
1337 if (llvm::any_of(unpackOp.getTiledOuterDims(),
1338 [](
int64_t dim) { return dim != 1; })) {
1341 "require the tiled outer dimensions of the result are all 1s");
1347 Value source = unpackOp.getSource();
1349 unpackOp.getDimAndTileMapping();
1368 for (
auto i : llvm::seq<unsigned>(0, destRank)) {
1377 if (dimAndTileMapping.count(i)) {
1378 extractSliceSizes.push_back(oneIdxAttr);
1384 if (ShapedType::isDynamic(srcShape[i])) {
1386 tensor::DimOp::create(rewriter, loc, source, i).getResult();
1387 extractSliceSizes.push_back(dynamicDim);
1388 shapeForEmptyOp.push_back(dynamicDim);
1390 extractSliceSizes.push_back(rewriter.
getIndexAttr(srcShape[i]));
1391 if (srcShape[i] != 1)
1392 shapeForEmptyOp.push_back(rewriter.
getIndexAttr(srcShape[i]));
1396 if (srcShape[i] != 1) {
1397 readShapeForExtractSlice.push_back(srcShape[i]);
1402 auto mixedTiles = unpackOp.getMixedTiles();
1403 extractSliceSizes.append(mixedTiles.begin(), mixedTiles.end());
1404 shapeForEmptyOp.append(mixedTiles.begin(), mixedTiles.end());
1408 auto tileShape = srcShape.drop_front(destRank);
1410 readShapeForExtractSlice.append(tileShape.begin(), tileShape.end());
1411 Type elemType = unpackOp.getSourceType().getElementType();
1412 auto readType = RankedTensorType::get(readShapeForExtractSlice, elemType);
1413 Value innerTile = tensor::ExtractSliceOp::create(
1414 rewriter, loc, readType, unpackOp.getSource(), extractSliceSizes);
1418 srcShape.take_front(destRank), innerDimsPos, unpackOp.getOuterDimsPerm());
1424 tensor::EmptyOp::create(rewriter, loc, shapeForEmptyOp, elemType);
1426 linalg::TransposeOp::create(rewriter, loc, innerTile, empty, perm);
1432 for (
auto i : llvm::seq<unsigned>(0, destRank)) {
1433 if (dimAndTileMapping.count(i) || destShape[i] != 1)
1434 tileSizes.push_back(
1439 tensor::ExtractSliceOp::create(rewriter, loc, RankedTensorType(),
1440 transposedOp.getResult()[0], tileSizes);
1444 for (
int i = 0, idx = 0; i < destRank; ++i) {
1445 if (dimAndTileMapping.count(i) || destShape[i] != 1)
1446 writeSizes.push_back(tileSizes[idx++]);
1448 writeSizes.push_back(oneIdxAttr);
1450 auto insert = tensor::InsertSliceOp::create(rewriter, loc, partialTile,
1451 unpackOp.getDest(), writeSizes);
1452 rewriter.
replaceOp(unpackOp, insert.getResult());
1466 for (
unsigned dim : dims) {
1470 resultIndices.push_back(i);
1475 return resultIndices;
1483 auto tensorType = cast<RankedTensorType>(
tensor.getType());
1484 int64_t rank = tensorType.getRank();
1488 for (
int64_t i = 0; i < rank; ++i) {
1489 if (!llvm::is_contained(dimsToRemove, i))
1490 newShape.push_back(tensorType.getDimSize(i));
1493 auto newType = RankedTensorType::get(newShape, tensorType.getElementType());
1501static std::optional<AffineExpr>
1505 bool onlyReferencesDroppedDims =
true;
1506 for (
unsigned d = 0; d < newNumDims + dimsToDrop.size(); ++d) {
1508 onlyReferencesDroppedDims =
false;
1512 if (onlyReferencesDroppedDims && llvm::any_of(dimsToDrop, [&](
unsigned d) {
1515 return std::nullopt;
1520 unsigned newDimIdx = 0;
1521 for (
unsigned d = 0; d < newNumDims + dimsToDrop.size(); ++d) {
1522 if (llvm::is_contained(dimsToDrop, d)) {
1536 if (failed(maybeDims))
1540 if (maybeDims->outputImage.size() != 2 || maybeDims->filterLoop.size() != 2)
1543 if (op.hasPureBufferSemantics())
1547 unsigned outSpatial0 = maybeDims->outputImage[0];
1548 unsigned outSpatial1 = maybeDims->outputImage[1];
1549 unsigned filterSpatial0 = maybeDims->filterLoop[0];
1550 unsigned filterSpatial1 = maybeDims->filterLoop[1];
1554 int64_t outSize0 = loopRanges[outSpatial0];
1555 int64_t outSize1 = loopRanges[outSpatial1];
1556 int64_t filterSize0 = loopRanges[filterSpatial0];
1557 int64_t filterSize1 = loopRanges[filterSpatial1];
1560 bool canRemoveSpatial0 = (filterSize0 == 1 && outSize0 == 1);
1561 bool canRemoveSpatial1 = (filterSize1 == 1 && outSize1 == 1);
1562 if (!canRemoveSpatial0 && !canRemoveSpatial1)
1569 if (canRemoveSpatial0) {
1570 loopDimsToRemove.push_back(outSpatial0);
1571 loopDimsToRemove.push_back(filterSpatial0);
1573 loopDimsToRemove.push_back(outSpatial1);
1574 loopDimsToRemove.push_back(filterSpatial1);
1576 llvm::sort(loopDimsToRemove);
1581 unsigned numDims = op.getNumLoops();
1582 unsigned newNumDims = numDims - loopDimsToRemove.size();
1583 for (
AffineMap map : op.getIndexingMapsArray()) {
1589 newResults.push_back(*newExpr);
1591 newMaps.push_back(
AffineMap::get(newNumDims, 0, newResults, ctx));
1596 auto iterTypes = op.getIteratorTypesArray();
1597 for (
unsigned idx = 0; idx < iterTypes.size(); ++idx) {
1598 if (!llvm::is_contained(loopDimsToRemove, idx))
1599 newIterTypes.push_back(iterTypes[idx]);
1605 for (
OpOperand *input : op.getDpsInputOperands()) {
1606 AffineMap map = op.getMatchingIndexingMap(input);
1610 tensorDimsToRemove);
1611 newInputs.push_back(reduced);
1614 OpOperand &output = *op.getDpsInitsMutable().begin();
1615 AffineMap outputMap = op.getMatchingIndexingMap(&output);
1619 outputDimsToRemove);
1624 newInputs, newOutput, newMaps, newIterTypes);
1626 newOp.getRegion().begin());
1630 LinalgOp resultOp = newOp;
1631 if (!isa<GenericOp>(op)) {
1633 if (succeeded(specializedOp))
1634 resultOp = *specializedOp;
1639 rewriter, loc, resultOp->getResult(0), output.
get());
1647struct DownscaleSizeOneWindowedConvolution final
1649 DownscaleSizeOneWindowedConvolution(
MLIRContext *context,
1653 LogicalResult matchAndRewrite(LinalgOp op,
1662 patterns.
add<DownscaleSizeOneWindowedConvolution>(patterns.
getContext(),
Base type for affine expression.
bool isFunctionOfDim(unsigned position) const
Return true if the affine expression involves AffineDimExpr position.
AffineExpr replaceDims(ArrayRef< AffineExpr > dimReplacements) const
Dim-only version of replaceDimsAndSymbols.
AffineExpr ceilDiv(uint64_t v) const
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
MLIRContext * getContext() const
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
AffineMap shiftDims(unsigned shift, unsigned offset=0) const
Replace dims[offset ... numDims) by dims[offset + shift ... shift + numDims).
AffineMap insertResult(AffineExpr expr, unsigned pos) const
Returns a new AffineMap with the same number of dims and symbols and an extra result inserted at pos.
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
unsigned getNumResults() const
AffineExpr getResult(unsigned idx) 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.
Attributes are known-constant values of operations.
This class is a general helper class for creating context-global objects like types,...
IntegerAttr getIndexAttr(int64_t value)
TypedAttr getZeroAttr(Type type)
AffineExpr getAffineDimExpr(unsigned position)
MLIRContext * getContext() const
This is a utility class for mapping one set of IR entities to another.
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.
RAII guard to reset the insertion point of the builder when destroyed.
This class helps build Operations.
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
This class represents a single result from folding an operation.
This class represents an operand of an operation.
unsigned getOperandNumber() const
Return which operand this is in the OpOperand list of the Operation.
This is a value defined by a result of an operation.
Operation is the basic unit of execution within MLIR.
result_range getResults()
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...
This is a builder type that keeps local references to arguments.
Builder & setShape(ArrayRef< int64_t > newShape)
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 moveOpBefore(Operation *op, Operation *existingOp)
Unlink this operation from its current block and insert it right before existingOp which may be in th...
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 inlineRegionBefore(Region ®ion, Region &parent, Region::iterator before)
Move the blocks that belong to "region" before the given position in another region "parent".
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
This class provides an abstraction over the different types of ranges over Values.
type_range getTypes() const
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.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Operation * getOwner() const
Return the owner of this operand.
OpFoldResult makeComposedFoldedAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Constructs an AffineApplyOp that applies map to operands after composing the map with the maps of any...
SmallVector< int64_t > getUnPackInverseSrcPerm(linalg::UnPackOp, PackingMetadata &metadata)
Compute inverse permutation for the source tensor (i.e.
FailureOr< PackTransposeResult > packTranspose(RewriterBase &rewriter, linalg::PackOp packOp, linalg::LinalgOp linalgOp, linalg::UnPackOp maybeUnPackOp, ArrayRef< int64_t > outerPerm, ArrayRef< int64_t > innerPerm)
Transpose a single PackOp -> LinalgOp -> UnPackOp chain and return the transposed PackOp -> LinalgOp ...
FailureOr< LowerUnPackOpResult > lowerUnPack(RewriterBase &rewriter, linalg::UnPackOp unPackOp, bool lowerUnpadLikeWithExtractSlice=true)
Rewrite pack as empty + transpose + reshape + extract_slice + copy.
void peelLoops(RewriterBase &rewriter, ArrayRef< scf::ForOp > loops)
Peel 'loops' and applies affine_min/max bounds simplification on the fly where relevant.
FailureOr< ConvolutionDimensions > inferConvolutionDims(LinalgOp linalgOp)
Find at least 1 parallel (output_image) and reduction (filter_loop) dimension candidates that form a ...
FailureOr< LinalgOp > generalizeNamedOp(RewriterBase &rewriter, LinalgOp linalgOp, bool emitCategoryOps=false)
Create a GenericOp or CategoryOp from the given named operation linalgOp and replace the given linalg...
FailureOr< LinalgOp > specializeGenericOp(RewriterBase &rewriter, GenericOp genericOp, bool emitCategoryOps=false)
Replace the given GenericOp with a namedOp or categoryOp.
void populateDecomposeConvolutionPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Linalg decompose convolutions patterns.
LogicalResult vectorizeCopy(RewriterBase &builder, memref::CopyOp copyOp)
Emit a suitable vector form for a Copy op with fully static shape.
FailureOr< GenericOp > interchangeGenericOp(RewriterBase &rewriter, GenericOp genericOp, ArrayRef< unsigned > interchangeVector)
Interchange the iterator_types and iterator_maps dimensions and adapts the index accesses of op.
SmallVector< int64_t > getPackInverseDestPerm(linalg::PackOp packOp, PackingMetadata &metadata)
Compute inverse permutation for the destination tensor (i.e.
void populateDecomposePackUnpackPatterns(RewritePatternSet &patterns)
Populates patterns to decompose linalg.pack and linalg.unpack Ops into e.g.
FailureOr< ContractionDimensions > inferContractionDims(LinalgOp linalgOp)
Find at least 2 parallel (m and n) and 1 reduction (k) dimension candidates that form a matmul subcom...
FailureOr< PackResult > packMatmulGreedily(RewriterBase &rewriter, LinalgOp linalgOp, ArrayRef< OpFoldResult > mnkPackedSizes, ArrayRef< int64_t > mnkPaddedSizesNextMultipleOf, ArrayRef< int64_t > mnkOrder)
Pack a LinalgOp by greedily inferring matmul dimensions (m, n, k) where m and n are proper parallel d...
FailureOr< PackResult > pack(RewriterBase &rewriter, linalg::LinalgOp linalgOp, ArrayRef< OpFoldResult > packedSizes)
Implement packing of a single LinalgOp by packedSizes.
SmallVector< Value > peelLoop(RewriterBase &rewriter, Operation *op)
Try to peel and canonicalize loop op and return the new result.
void populateDecomposePadPatterns(RewritePatternSet &patterns)
Populates patterns to decompose tensor.pad into e.g.
FailureOr< LinalgOp > downscaleSizeOneWindowedConvolution(RewriterBase &rewriter, LinalgOp op)
Rewrite convolution/pooling/depthwise ops with size-1 window dimensions into lower-dimensional ops.
FailureOr< LowerPackResult > lowerPack(RewriterBase &rewriter, linalg::PackOp packOp, bool lowerPadLikeWithInsertSlice=true)
Rewrite pack as pad + reshape + transpose.
LogicalResult peelForLoopAndSimplifyBounds(RewriterBase &rewriter, ForOp forOp, scf::ForOp &partialIteration)
Rewrite a for loop with bounds/step that potentially do not divide evenly into a for loop where the s...
FailureOr< TilingResult > bubbleUpPadSlice(OpBuilder &b, tensor::PadOp padOp, ArrayRef< OpFoldResult > offsets, ArrayRef< OpFoldResult > sizes, bool generateZeroSliceGuard=true)
Bubbles up a slice of this pad by taking the slice first and then performing the padding.
PadOp createPadHighOp(RankedTensorType resType, Value source, Value pad, bool nofold, Location loc, OpBuilder &builder, ValueRange dynOutDims={})
Value createCanonicalRankReducingInsertSliceOp(OpBuilder &b, Location loc, Value tensor, Value dest)
Create a rank-reducing InsertSliceOp @[0 .
Value createCanonicalRankReducingExtractSliceOp(OpBuilder &b, Location loc, Value tensor, RankedTensorType targetType)
Create a rank-reducing ExtractSliceOp @[0 .
OpFoldResult getMixedSize(OpBuilder &builder, Location loc, Value value, int64_t dim)
Return the dimension of the given tensor value.
SmallVector< OpFoldResult > getMixedSizes(OpBuilder &builder, Location loc, Value value)
Return the dimensions of the given tensor value.
Include the generated interface declarations.
SliceVerificationResult
Enum that captures information related to verifier error conditions on slice insert/extract type of o...
ArrayRef< int64_t > ReassociationIndicesRef
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
SmallVector< int64_t > computePermutationVector(int64_t permSize, ArrayRef< int64_t > positions, ArrayRef< int64_t > desiredPositions)
Return a permutation vector of size permSize that would result in moving positions into desiredPositi...
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
void bindSymbols(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to SymbolExpr at positions: [0 .
SmallVector< Loops, 8 > tile(ArrayRef< scf::ForOp > forOps, ArrayRef< Value > sizes, ArrayRef< scf::ForOp > targets)
Performs tiling fo imperfectly nested loops (with interchange) by strip-mining the forOps by sizes an...
AffineExpr getAffineConstantExpr(int64_t constant, MLIRContext *context)
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
std::pair< int64_t, OpFoldResult > getSimplifiedOfrAndStaticSizePair(OpFoldResult ofr, Builder &b)
Given OpFoldResult representing dim size value (*), generates a pair of sizes:
void applyPermutationToVector(SmallVector< T, N > &inVec, ArrayRef< int64_t > permutation)
Apply the permutation defined by permutation to inVec.
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
SliceVerificationResult isRankReducedType(ShapedType originalType, ShapedType candidateReducedType)
Check if originalType can be rank reduced to candidateReducedType type by dropping some dimensions wi...
bool isPermutationVector(ArrayRef< int64_t > interchange)
Method to check if an interchange vector is a permutation.
SmallVector< int64_t > invertPermutationVector(ArrayRef< int64_t > permutation)
Helper method to apply to inverse a permutation.
OpInterfaceRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting a...
Represents a range (offset, size, and stride) where each element of the triple may be dynamic or stat...
LogicalResult matchAndRewrite(memref::CopyOp copyOp, PatternRewriter &rewriter) const override
Rewrites a linalg::PackOp into a sequence of:
LogicalResult matchAndRewrite(linalg::PackOp packOp, PatternRewriter &rewriter) const override
Rewrites a linalg::UnPackOp into a sequence of:
LogicalResult matchAndRewrite(linalg::UnPackOp unpackOp, PatternRewriter &rewriter) const override
Rewrite a tensor::PadOp into a sequence of EmptyOp, FillOp and InsertSliceOp.
LogicalResult matchAndRewrite(tensor::PadOp padOp, PatternRewriter &rewriter) const override
Value createFillOrGenerateOp(RewriterBase &rewriter, tensor::PadOp padOp, Value dest, const SmallVector< Value > &dynSizes) const
Filling dest using FillOp constant padding value if possible.
LinalgTilingOptions & setTileSizes(const SmallVector< Value, 4 > &ts)
Set the tileSizeComputationFunction to return the values ts.
TileSizeComputationFunction tileSizeComputationFunction
Computation function that returns the tile sizes for each operation.
Struct to hold the result of a pack call.
Struct to hold the result of a packTranspose call.