27#include "llvm/ADT/SmallVectorExtras.h"
28#include "llvm/Support/Debug.h"
31#define DEBUG_TYPE "linalg-tiling-interface-impl"
50 Value v = affine::AffineApplyOp::create(
b, loc, m, ivs);
60 Block *body = linalgOp.getBlock();
64 if (
auto indexOp = dyn_cast<IndexOp>(&op)) {
65 map.
map(indexOp.getResult(), ivs[indexOp.getDim()]);
73 for (
const auto &operand : llvm::enumerate(terminator->
getOperands())) {
75 OpOperand *storeInto = linalgOp.getDpsInitOperand(operand.index());
77 b, loc, linalgOp.getMatchingIndexingMap(storeInto), ivs);
78 memref::StoreOp::create(
b, loc, toStore,
79 linalgOp.getDpsInitOperand(operand.index())->get(),
95template <
typename LinalgOpTy>
96struct LinalgOpTilingInterface
97 :
public TilingInterface::ExternalModel<LinalgOpTilingInterface<LinalgOpTy>,
100 TilingInterface::ExternalModel<LinalgOpTilingInterface<LinalgOpTy>,
104 using Base::generateResultTileValue;
105 using Base::getIterationDomainTileFromOperandTiles;
106 using Base::getTiledImplementation;
107 using Base::getTiledImplementationFromOperandTiles;
110 SmallVector<utils::IteratorType> getLoopIteratorTypes(Operation *op)
const {
111 LinalgOpTy concreteOp = cast<LinalgOpTy>(op);
112 return concreteOp.getIteratorTypesArray();
116 SmallVector<Range> getIterationDomain(Operation *op, OpBuilder &
b)
const {
117 OpBuilder::InsertionGuard g(
b);
118 b.setInsertionPoint(op);
119 Location loc = op->
getLoc();
120 LinalgOp linalgOp = cast<LinalgOp>(op);
121 SmallVector<OpFoldResult> allShapesSizes =
122 linalgOp.createFlatListOfOperandDims(
b, loc);
123 AffineMap map = linalgOp.getShapesToLoopsMap();
125 return llvm::map_to_vector(map.
getResults(), [&](AffineExpr loopExpr) {
126 OpFoldResult ofr = affine::makeComposedFoldedAffineApply(b, loc, loopExpr,
128 return Range{b.getIndexAttr(0), ofr, b.getIndexAttr(1)};
133 FailureOr<TilingResult>
140 LinalgOp linalgOp = cast<LinalgOp>(op);
143 b, loc, linalgOp, valuesToTile, offsets, sizes, {},
true);
145 llvm::make_filter_range(
147 [](
Value v) ->
bool {
148 return isa_and_nonnull<tensor::ExtractSliceOp, memref::SubViewOp>(
156 Operation *tiledOp =
clone(
b, linalgOp, resultTensorTypes, tiledOperands);
167 getMappedOffsetAndSize(LinalgOp linalgOp,
OpBuilder &
b,
175 for (
auto [indexingMap, offsets, sizes] :
176 llvm::zip_equal(indexingMaps, allOffsets, allSizes)) {
177 for (
auto [resultExpr, offset, size] :
178 llvm::zip_equal(indexingMap.getResults(), offsets, sizes)) {
179 auto dimExpr = dyn_cast<AffineDimExpr>(resultExpr);
182 unsigned position = dimExpr.getPosition();
183 auto it = mappedOffsets.find(position);
184 if (it != mappedOffsets.end()) {
187 if (seenOffset != offset || seenSize != size) {
189 llvm::dbgs() <<
"inconsistent iteration space mapping from "
190 "offsets/sizes of operands/results";
195 mappedOffsets[position] = offset;
196 mappedSizes[position] = size;
204 cast<TilingInterface>(linalgOp.getOperation()).getIterationDomain(
b);
205 mappedOffsetsVec.resize(iterationDomain.size());
206 mappedSizesVec.resize(iterationDomain.size());
207 for (
auto [
index, domain] : llvm::enumerate(iterationDomain)) {
208 auto it = mappedOffsets.find(
index);
209 if (it != mappedOffsets.end()) {
210 mappedOffsetsVec[
index] = it->second;
211 mappedSizesVec[
index] = mappedSizes.lookup(
index);
214 mappedOffsetsVec[
index] = domain.offset;
215 mappedSizesVec[
index] = domain.size;
222 LogicalResult getIterationDomainTileFromOperandTiles(
228 auto linalgOp = cast<LinalgOp>(op);
231 llvm::map_to_vector(operandNumbers, [&](
unsigned operandNumber) {
232 OpOperand &opOperand = linalgOp->getOpOperand(operandNumber);
233 return linalgOp.getMatchingIndexingMap(&opOperand);
235 if (
failed(getMappedOffsetAndSize(linalgOp,
b, indexingMaps, allOffsets,
236 allSizes, iterDomainOffsets,
252 LinalgOp linalgOp = cast<LinalgOp>(op);
261 OpOperand *outOperand = linalgOp.getDpsInitOperand(resultNumber);
263 b, loc, outOperand->get(), sizes,
264 linalgOp.getMatchingIndexingMap(outOperand), offsets,
265 {}, subShapeSizes,
true);
266 resultOffsets = sliceParams.
offsets;
267 resultSizes = sliceParams.
sizes;
271 LogicalResult getIterationDomainTileFromResultTile(
276 auto linalgOp = cast<LinalgOp>(op);
283 linalgOp.getIndexingMapMatchingResult(op->
getResult(resultNumber));
286 "unhandled tiled implementation generation when result is not "
287 "accessed using a permuted projection");
293 getMappedOffsetAndSize(linalgOp,
b, indexingMap, {allOffsets},
294 {allSizes}, iterDomainOffsets, iterDomainSizes);
296 assert(succeeded(status) &&
"unexpected error in offset calculation");
300 FailureOr<TilingResult>
305 if (
failed(getIterationDomainTileFromResultTile(
306 op,
b, resultNumber, offsets, sizes, mappedOffsets, mappedSizes))) {
309 auto tilingInterfaceOp = cast<TilingInterface>(op);
310 FailureOr<TilingResult> tilingResult =
311 tilingInterfaceOp.getTiledImplementation(
b, mappedOffsets, mappedSizes);
316 if (tilingResult->tiledOps.size() != 1)
317 return op->
emitOpError(
"failed to generate tiled implementation");
320 tilingResult->tiledOps,
322 tilingResult->generatedSlices};
327 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
332 if (
failed(getIterationDomainTileFromOperandTiles(
333 op,
b, operandNumbers, allOffsets, allSizes, mappedOffsets,
343 auto linalgOp = cast<LinalgOp>(op);
344 if (!linalgOp.hasPureBufferSemantics())
345 return op->
emitOpError(
"expected operation to have buffer semantics");
348 indexedValues.reserve(linalgOp->getNumOperands());
352 for (
OpOperand &operand : linalgOp->getOpOperands()) {
353 if (!linalgOp.payloadUsesValueFromOperand(&operand)) {
354 indexedValues.push_back(
nullptr);
357 if (linalgOp.isScalar(&operand)) {
358 indexedValues.push_back(operand.get());
362 builder, linalgOpLoc, linalgOp.getMatchingIndexingMap(&operand), ivs);
364 memref::LoadOp::create(builder, linalgOpLoc, operand.get(),
indices);
365 indexedValues.push_back(
load);
372 bool isOpFusableWithConsumerSlice(
Operation *op,
unsigned resultNumber,
379 bool isOpFusableWithProducerSlices(
384 auto linalgOp = cast<LinalgOp>(op);
386 llvm::map_to_vector(operandNumbers, [&](
unsigned operandNumber) {
387 OpOperand &opOperand = linalgOp->getOpOperand(operandNumber);
388 return linalgOp.getMatchingIndexingMap(&opOperand);
393 return succeeded(getMappedOffsetAndSize(linalgOp,
b, indexingMaps,
394 allOffsets, allSizes, mappedOffsets,
406 for (
auto [
index, reductionDim] : llvm::enumerate(reductionDims)) {
407 if (reductionDim == value) {
419getPartialResultAffineMaps(LinalgOp linalgOp,
421 auto partialReductionMaps = llvm::map_to_vector(
422 linalgOp.getDpsInitsMutable(), [&](
OpOperand &opOperand) {
423 AffineMap map = linalgOp.getMatchingIndexingMap(&opOperand);
424 for (auto redPos : reductionDims) {
426 map.insertResult(getAffineDimExpr(redPos, linalgOp.getContext()),
427 map.getNumResults());
431 return partialReductionMaps;
434struct InitSliceInfo {
435 SmallVector<int64_t> resultShape;
436 SmallVector<OpFoldResult> offsets;
437 SmallVector<OpFoldResult> sizes;
438 SmallVector<OpFoldResult> strides;
444static InitSliceInfo getInitSliceInfoForOuterReduction(
451 Attribute zero = IntegerAttr::get(IndexType::get(context), 0);
452 Attribute one = IntegerAttr::get(IndexType::get(context), 1);
454 for (
auto [resultIdx, dimExpr] :
455 llvm::enumerate(partialReductionMap.
getResults())) {
456 if (isa<AffineConstantExpr>(dimExpr)) {
459 initOffsets.push_back(zero);
460 initSizes.push_back(initOperandShape[resultIdx]);
463 unsigned dim = cast<AffineDimExpr>(dimExpr).getPosition();
464 if (reductionDims.contains(dim)) {
465 initOffsets.push_back(zero);
467 initOffsets.push_back(offsets[dim]);
469 initSizes.push_back(sizes[dim]);
473 return {resultShape, initOffsets, initSizes, initStrides};
479static InitSliceInfo getInitSliceInfoForOuterParallel(
486 Attribute zero = IntegerAttr::get(IndexType::get(context), 0);
487 Attribute one = IntegerAttr::get(IndexType::get(context), 1);
490 for (
auto [resultIdx, dimExpr] :
491 llvm::enumerate(partialReductionMap.
getResults())) {
492 if (isa<AffineConstantExpr>(dimExpr)) {
495 initOffsets.push_back(zero);
496 initSizes.push_back(initOperandShape[resultIdx]);
497 resultShape.push_back(initOperandShape[resultIdx]);
500 unsigned dim = cast<AffineDimExpr>(dimExpr).getPosition();
501 if (std::optional<unsigned> dimPos = getPositionIn(reductionDims, dim)) {
502 initOffsets.push_back(splitReductionIvs[dimPos.value()]);
503 initSizes.push_back(one);
505 initOffsets.push_back(offsets[dim]);
506 initSizes.push_back(sizes[dim]);
507 resultShape.push_back(sizes[dim]);
512 return {staticShapes, initOffsets, initSizes, initStrides};
517static InitSliceInfo getInitSliceInfo(
MLIRContext *context,
526 return getInitSliceInfoForOuterReduction(
527 context, offsets, sizes, reductionDims, splitReductionIvs,
528 partialReductionMap, initOperandShape);
531 "unexpected ReductionTilingStrategy");
532 return getInitSliceInfoForOuterParallel(
533 context, offsets, sizes, reductionDims, splitReductionIvs,
534 partialReductionMap, initOperandShape);
539template <
typename LinalgOpTy>
540struct LinalgOpPartialReductionInterface
541 :
public PartialReductionOpInterface::ExternalModel<
542 LinalgOpPartialReductionInterface<LinalgOpTy>, LinalgOpTy> {
543 FailureOr<SmallVector<Value>> generateInitialTensorForPartialReduction(
544 Operation *op, OpBuilder &
b, Location loc, ArrayRef<OpFoldResult> sizes,
546 auto linalgOp = cast<LinalgOp>(op);
548 OpBuilder::InsertionGuard guard(
b);
549 if (linalgOp.hasPureBufferSemantics())
550 return op->
emitOpError(
"expected operation to have tensor semantics");
552 SmallVector<AffineMap> partialResultMaps =
553 getPartialResultAffineMaps(linalgOp, reductionDims);
555 SmallVector<Value> inits;
556 for (
auto [initIdx,
result, partialMap] :
557 llvm::enumerate(linalgOp->getResults(), partialResultMaps)) {
558 SmallVector<Operation *, 4> combinerOps;
561 combinerOps.size() != 1)
562 return op->
emitOpError(
"Failed to anaysis the reduction operation.");
564 Operation *reductionOp = combinerOps[0];
565 std::optional<TypedAttr> identity = arith::getNeutralElement(reductionOp);
566 if (!identity.has_value())
568 "Failed to get an identity value for the reduction operation.");
571 SmallVector<OpFoldResult> partialResultShape;
572 Value initValue = linalgOp.getDpsInits()[initIdx];
573 SmallVector<OpFoldResult> initShape =
575 for (
auto [resultIdx, dimExpr] :
576 llvm::enumerate(partialMap.getResults())) {
577 if (isa<AffineConstantExpr>(dimExpr)) {
580 partialResultShape.push_back(initShape[resultIdx]);
583 auto dim = cast<AffineDimExpr>(dimExpr);
584 partialResultShape.push_back(sizes[dim.getPosition()]);
589 tensor::EmptyOp::create(
b, loc, partialResultShape, elType);
590 Value constantOp = arith::ConstantOp::create(
b, loc, *identity);
591 auto identityTensor =
592 linalg::FillOp::create(
b, loc, constantOp, emptyTensor);
593 inits.push_back(identityTensor.getResult(0));
599 FailureOr<TilingResult>
600 tileToPartialReduction(Operation *op, OpBuilder &
b, Location loc,
602 ValueRange init, ArrayRef<OpFoldResult> offsets,
603 ArrayRef<OpFoldResult> sizes,
605 ArrayRef<OpFoldResult> splitReductionIvs)
const {
606 OpBuilder::InsertionGuard guard(
b);
607 auto linalgOp = cast<LinalgOp>(op);
609 SmallVector<AffineMap> partialReductionMaps =
610 getPartialResultAffineMaps(linalgOp, reductionDims);
614 SmallVector<AffineMap> newInitMaps;
615 if (tilingStrategy ==
616 ReductionTilingStrategy::PartialReductionOuterReduction) {
617 newInitMaps = llvm::to_vector(partialReductionMaps);
619 newInitMaps = llvm::map_to_vector(
620 linalgOp.getDpsInitsMutable(), [&](OpOperand &opOperand) {
621 return linalgOp.getMatchingIndexingMap(&opOperand);
627 b, loc, linalgOp, linalgOp.getDpsInputs(), offsets, sizes, {},
true);
628 SmallVector<Operation *> generatedSlices = llvm::map_to_vector(
629 llvm::make_filter_range(
630 tiledInputs, [](Value v) ->
bool {
return v.
getDefiningOp(); }),
634 SmallVector<Value, 1> tiledInits;
635 for (
auto [partialReductionMap, valueToTile, initOperandValue] :
636 llvm::zip_equal(partialReductionMaps, init, linalgOp.getDpsInits())) {
639 SmallVector<OpFoldResult> initOperandShape =
641 InitSliceInfo sliceInfo = getInitSliceInfo(
642 b.getContext(), tilingStrategy, offsets, sizes, reductionDims,
643 splitReductionIvs, partialReductionMap, initOperandShape);
644 auto valueToTileType = cast<RankedTensorType>(valueToTile.getType());
646 sliceInfo.resultShape, valueToTileType.getElementType(),
647 valueToTileType.getEncoding());
648 auto sliceOp = tensor::ExtractSliceOp::create(
650 sliceInfo.sizes, sliceInfo.strides);
651 tiledInits.push_back(sliceOp.getResult());
652 generatedSlices.push_back(sliceOp);
656 SmallVector<AffineMap> newMaps = linalgOp.getIndexingMapsArray();
657 for (
auto [initOperand, newInitMap] :
658 llvm::zip_equal(linalgOp.getDpsInitsMutable(), newInitMaps)) {
659 int mapIdx = linalgOp.getIndexingMapIndex(&initOperand);
660 newMaps[mapIdx] = newInitMap;
664 SmallVector<utils::IteratorType> newIteratorTypes =
665 linalgOp.getIteratorTypesArray();
666 if (tilingStrategy ==
667 ReductionTilingStrategy::PartialReductionOuterReduction) {
668 for (
int dim : reductionDims)
669 newIteratorTypes[dim] = utils::IteratorType::parallel;
673 Operation *partialReductionOp;
674 auto resultTypes =
ValueRange(tiledInits).getTypes();
675 if (tilingStrategy ==
676 ReductionTilingStrategy::PartialReductionOuterReduction) {
677 auto genericOp = GenericOp::create(
b, loc, resultTypes, tiledInputs,
678 tiledInits, newMaps, newIteratorTypes);
681 genericOp.getRegion().begin(), mapping);
683 partialReductionOp = genericOp.getOperation();
685 SmallVector<Value> operands = std::move(tiledInputs);
686 llvm::append_range(operands, tiledInits);
687 partialReductionOp =
mlir::clone(
b, op, resultTypes, operands);
691 {partialReductionOp},
692 llvm::map_to_vector(partialReductionOp->
getResults(),
693 [](OpResult r) -> Value { return r; }),
697 FailureOr<MergeResult>
698 mergeReductions(Operation *op, OpBuilder &
b, Location loc,
701 auto linalgOp = cast<LinalgOp>(op);
702 SmallVector<AffineMap> partialReductionMaps =
703 getPartialResultAffineMaps(linalgOp, reductionDims);
706 SmallVector<Operation *> mergeOperations;
707 SmallVector<Value> replacements;
708 for (
auto [idx, init, partialResult, partialMap] : llvm::enumerate(
709 linalgOp.getDpsInits(), partialReduce, partialReductionMaps)) {
710 unsigned initIdx = idx;
715 SmallVector<int64_t> partialReductionDims;
716 for (
auto [resultNum, dimExpr] :
717 llvm::enumerate(partialMap.getResults())) {
718 if (isa<AffineConstantExpr>(dimExpr))
720 unsigned dim = cast<AffineDimExpr>(dimExpr).getPosition();
721 if (llvm::is_contained(reductionDims, dim)) {
722 partialReductionDims.push_back(resultNum);
726 auto reduction = linalg::ReduceOp::create(
727 b, loc, partialResult, init, partialReductionDims,
728 [&linalgOp, &initIdx](OpBuilder &
b, Location loc,
ValueRange inputs) {
730 SmallVector<Operation *, 4> combinerOps;
733 Operation *clonedReductionOp =
b.clone(*combinerOps[0]);
737 linalg::YieldOp::create(
b, loc, clonedReductionOp->
getResult(0));
740 mergeOperations.push_back(reduction);
741 replacements.push_back(reduction->getResult(0));
744 return MergeResult{mergeOperations, replacements};
747 LogicalResult getPartialResultTilePosition(
748 Operation *op, OpBuilder &
b,
unsigned resultNumber,
751 ArrayRef<OpFoldResult> splitReductionIvs,
752 SmallVector<OpFoldResult> &resultOffsets,
753 SmallVector<OpFoldResult> &resultSizes)
const {
754 auto linalgOp = cast<LinalgOp>(op);
755 SmallVector<AffineMap> partialReductionMaps =
756 getPartialResultAffineMaps(linalgOp, reductionDims);
759 Value initOperandValue = linalgOp.getDpsInits()[resultNumber];
760 Location loc = op->
getLoc();
761 SmallVector<OpFoldResult> initOperandShape =
763 InitSliceInfo sliceInfo =
764 getInitSliceInfo(
b.getContext(), tilingStrategy, offsets, sizes,
765 reductionDims, splitReductionIvs,
766 partialReductionMaps[resultNumber], initOperandShape);
767 std::swap(resultOffsets, sliceInfo.offsets);
768 std::swap(resultSizes, sliceInfo.sizes);
774template <
typename OpTy>
777 static_assert(llvm::is_one_of<OpTy, PackOp, UnPackOp>::value,
778 "applies to only pack or unpack operations");
780 int64_t rank = (std::is_same<OpTy, PackOp>::value) ? op.getSourceRank()
785 (
void)op.reifyResultShapes(builder, resultShape);
787 for (
auto dim : llvm::seq<int64_t>(0, rank)) {
788 loopBounds[dim].offset = zero;
789 loopBounds[dim].stride = one;
790 loopBounds[dim].size = resultShape[0][dim];
798 if (permutation.empty())
810 interchangeVector.reserve(dimsPos.size());
819 for (
int64_t dimsIdx = 0, end = dimsPos.size(); dimsIdx < end; dimsIdx++)
820 dimsAndPosMapping[dimsPos[dimsIdx]] = dimsIdx;
824 for (
int64_t dimsIdx = 0; dimsIdx < rank; dimsIdx++) {
825 if (dimsAndPosMapping.count(dimsIdx))
826 interchangeVector.push_back(dimsAndPosMapping[dimsIdx]);
828 return interchangeVector;
848 for (
auto [idx, val] : llvm::enumerate(interchangeVector))
849 vec[idx + offset] = elements[val + offset];
855static void generatePackOpScalarImplementationBody(PackOp packOp,
870 computeInterchangeFromDimPos(dimsToInnerBlock, packOp.getSourceRank());
871 interchangedIvs = interchange<Value>(interchangedIvs, interchangeVector,
872 packOp.getSourceRank());
873 if (!dimsToOuterBlock.empty()) {
875 computeInterchangeFromDimPos(dimsToOuterBlock, packOp.getSourceRank());
877 interchange<Value>(interchangedIvs, interchangeVector, 0);
880 packOp.getDimAndTileMapping();
882 size_t pointLoopsOffset = 0;
883 int64_t sourceRank = packOp.getSourceRank();
884 for (
auto dim : llvm::seq<int64_t>(0, sourceRank)) {
885 if (dimAndTileMapping.contains(dim)) {
890 builder, loc, i *
tile +
j,
892 interchangedIvs[dim],
893 interchangedIvs[pointLoopsOffset + packOp.getSourceRank()],
894 dimAndTileMapping[dim]});
895 sourceIndices.push_back(sourceIndex);
898 sourceIndices.push_back(interchangedIvs[dim]);
902 auto createLoad = [&]() ->
Value {
903 return memref::LoadOp::create(
904 builder, loc, packOp.getSource(),
908 if (
auto paddingValue = packOp.getPaddingValue()) {
911 for (
auto dim : llvm::seq<int64_t>(0, sourceRank)) {
914 Value cond = arithBuilder.slt(
918 scalar = scf::IfOp::create(
921 scf::YieldOp::create(
b, l, createLoad());
925 scf::YieldOp::create(
b, l, paddingValue);
929 scalar = createLoad();
932 memref::StoreOp::create(builder, loc, scalar, packOp.getDest(), ivs);
936 :
public TilingInterface::ExternalModel<PackOpTiling, linalg::PackOp> {
937 using Base = TilingInterface::ExternalModel<PackOpTiling, linalg::PackOp>;
938 using Base::getTiledImplementation;
940 SmallVector<utils::IteratorType> getLoopIteratorTypes(Operation *op)
const {
944 auto packOp = cast<PackOp>(op);
945 SmallVector<utils::IteratorType> iteratorTypes(
946 packOp.getSourceRank(), utils::IteratorType::parallel);
947 return iteratorTypes;
950 SmallVector<Range> getIterationDomain(Operation *op, OpBuilder &
b)
const {
951 return getPackUnPackIterationDomain<PackOp>(cast<PackOp>(op),
b);
954 FailureOr<TilingResult>
956 ArrayRef<OpFoldResult> offsets,
957 ArrayRef<OpFoldResult> sizes)
const {
958 auto packOp = cast<PackOp>(op);
960 if (!packOp.hasPureTensorSemantics())
963 Location loc = packOp.getLoc();
967 int64_t inputRank = packOp.getSourceRank();
968 SmallVector<OpFoldResult> origOffsets(offsets);
969 SmallVector<OpFoldResult> origSizes(sizes);
970 applyPermToRange(origOffsets, origSizes,
974 packOp.getDimAndTileMapping();
975 SmallVector<OpFoldResult> srcDimValues =
977 SmallVector<OpFoldResult> inputIndices, inputSizes;
978 for (
auto dim : llvm::seq<int64_t>(0, inputRank)) {
979 using AV = affine::AffineValueExpr;
980 affine::AffineBuilder ab(
b, loc);
981 AffineExpr dim0, dim1, sym;
984 if (dimAndTileMapping.count(dim)) {
988 auto avOffset = AV(dim0).bind(origOffsets[dim]);
989 auto avSize = AV(dim0).bind(origSizes[dim]);
990 auto avTileSize = AV(sym).bind(dimAndTileMapping[dim]);
991 inputIndices.push_back(ab.mul(avOffset, avTileSize));
992 inputSizes.push_back(ab.mul(avSize, avTileSize));
994 inputIndices.push_back(origOffsets[dim]);
995 inputSizes.push_back(origSizes[dim]);
999 if (packOp.getPaddingValue()) {
1000 OpFoldResult dimSize = srcDimValues[dim];
1001 auto avDimSize = AV(dim0).bind(dimSize);
1002 auto avInputIdx = AV(dim1).bind(inputIndices.back());
1004 ab.min({inputSizes.back(), ab.sub(avDimSize, avInputIdx)});
1008 auto oneAttr =
b.getI64IntegerAttr(1);
1009 SmallVector<OpFoldResult> strides(inputRank, oneAttr);
1011 SmallVector<Value> tiledOperands;
1012 auto sourceSlice = tensor::ExtractSliceOp::create(
1013 b, loc, packOp.getSource(), inputIndices, inputSizes, strides);
1014 tiledOperands.push_back(sourceSlice);
1016 SmallVector<OpFoldResult> outputOffsets, outputSizes;
1021 strides.append(packOp.getDestRank() - inputRank, oneAttr);
1022 auto outSlice = tensor::ExtractSliceOp::create(
1023 b, loc, packOp.getDest(), outputOffsets, outputSizes, strides);
1024 tiledOperands.push_back(outSlice);
1026 if (
auto val = packOp.getPaddingValue())
1027 tiledOperands.push_back(val);
1028 for (
auto tile : packOp.getInnerTiles())
1029 tiledOperands.push_back(
tile);
1031 Operation *tiledPackOp = PackOp::create(
1034 return TilingResult{
1036 SmallVector<Value>(tiledPackOp->
getResults()),
1037 llvm::to_vector(ArrayRef<Operation *>{sourceSlice, outSlice})};
1042 ArrayRef<OpFoldResult> offsets,
1043 ArrayRef<OpFoldResult> sizes,
1044 SmallVector<OpFoldResult> &resultOffsets,
1045 SmallVector<OpFoldResult> &resultSizes)
const {
1050 auto packOp = cast<PackOp>(op);
1051 int64_t inputRank = packOp.getSourceRank();
1052 int64_t outputRank = packOp.getDestRank();
1053 auto zeroAttr =
b.getI64IntegerAttr(0);
1054 resultOffsets.assign(offsets.begin(), offsets.end());
1055 resultOffsets.append(outputRank - inputRank, zeroAttr);
1059 resultSizes.assign(sizes.begin(), sizes.end());
1060 for (
auto dataTileDim : llvm::seq<unsigned>(inputRank, outputRank))
1061 resultSizes.push_back(outputShape[0][dataTileDim]);
1066 FailureOr<TilingResult>
1067 generateResultTileValue(Operation *op, OpBuilder &
b,
unsigned resultNumber,
1068 ArrayRef<OpFoldResult> offsets,
1069 ArrayRef<OpFoldResult> sizes)
const {
1070 return generateResultTileValue(op,
b, resultNumber, offsets, sizes,
1074 FailureOr<TilingResult> generateResultTileValue(
1075 Operation *op, OpBuilder &
b,
unsigned resultNumber,
1076 ArrayRef<OpFoldResult> offsets, ArrayRef<OpFoldResult> sizes,
1077 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1078 auto packOp = cast<PackOp>(op);
1079 int64_t numTiles = packOp.getInnerDimsPos().size();
1084 for (
auto offset : offsets.take_back(numTiles))
1091 ArrayRef<int64_t> innerDimsPos = packOp.getInnerDimsPos();
1092 SmallVector<OpFoldResult> mixedTiles = packOp.getMixedTiles();
1093 ArrayRef<OpFoldResult> innerSizes = sizes.take_back(numTiles);
1094 for (
auto [i, pos] : llvm::enumerate(innerDimsPos)) {
1096 pos < static_cast<int64_t>(innerTileAlignments.size())
1097 ? innerTileAlignments[pos]
1098 : InnerTileAlignment::Unknown;
1099 if (alignment != InnerTileAlignment::Equal &&
1105 op,
b, offsets.drop_back(numTiles), sizes.drop_back(numTiles));
1106 if (
failed(tilingResult))
1108 return tilingResult.value();
1111 LogicalResult generateScalarImplementation(Operation *op, OpBuilder &builder,
1114 auto packOp = cast<PackOp>(op);
1115 assert(packOp.hasPureBufferSemantics() &&
1116 "expected operation to have buffer semantics");
1117 OpBuilder::InsertionGuard g(builder);
1120 SmallVector<Value> ivVec(ivs);
1123 SmallVector<OpFoldResult> outputShape;
1124 Value dest = packOp.getDest();
1125 for (
auto dim : llvm::seq<int64_t>(0, packOp.getDestRank()))
1134 for (
auto dataTileDim : llvm::seq<unsigned>(packOp.getSourceRank(),
1135 packOp.getDestRank() - 1)) {
1137 outputShape[dataTileDim]);
1138 scf::ForOp loop = scf::ForOp::create(builder, loc, zero, ub, one);
1140 ivVec.push_back(loop.getInductionVar());
1147 [&](OpBuilder &bodyBuilder, Location bodyLoc, Value iv,
1149 ivVec.push_back(iv);
1150 generatePackOpScalarImplementationBody(packOp, bodyBuilder, bodyLoc,
1152 scf::YieldOp::create(bodyBuilder, bodyLoc);
1157 LogicalResult getIterationDomainTileFromOperandTiles(
1158 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1159 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1160 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1161 SmallVectorImpl<OpFoldResult> &resultOffsets,
1162 SmallVectorImpl<OpFoldResult> &resultSizes)
const {
1163 return getIterationDomainTileFromOperandTiles(
1164 op,
b, operandNumbers, allOffsets, allSizes, resultOffsets, resultSizes,
1171 LogicalResult getIterationDomainTileFromOperandTiles(
1172 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1173 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1174 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1175 SmallVectorImpl<OpFoldResult> &resultOffsets,
1176 SmallVectorImpl<OpFoldResult> &resultSizes,
1177 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1178 if (operandNumbers.size() != 1 || operandNumbers[0] != 0) {
1180 { llvm::dbgs() <<
"unsupported operands for consumer fusion"; });
1184 ArrayRef<OpFoldResult> offsets(allOffsets[0]);
1185 ArrayRef<OpFoldResult> sizes(allSizes[0]);
1186 auto packOp = cast<PackOp>(op);
1187 Location loc = packOp.getLoc();
1188 SmallVector<OpFoldResult> outerDimOffsets, outerDimSizes;
1190 packOp.getDimAndTileMapping();
1191 SmallVector<int64_t> outerShapeWithoutTranspose(
1192 packOp.getDestType().getShape().take_front(packOp.getSourceRank()));
1193 if (!packOp.getOuterDimsPerm().empty()) {
1195 outerShapeWithoutTranspose,
1198 for (
auto dim : llvm::seq<int64_t>(packOp.getSourceRank())) {
1199 if (dimAndTileMapping.count(dim)) {
1200 FailureOr<int64_t> cstTileSize =
1202 presburger::BoundType::UB, sizes[dim],
1204 ValueBoundsOptions{
true});
1205 std::optional<int64_t> cstInnerSize =
1212 dim < static_cast<int64_t>(innerTileAlignments.size())
1213 ? innerTileAlignments[dim]
1214 : InnerTileAlignment::Unknown;
1227 int64_t srcDimSize = packOp.getSourceType().getDimSize(dim);
1228 int64_t destDimSize = outerShapeWithoutTranspose[dim];
1229 bool isTiled = innerTileAlignment != InnerTileAlignment::Unknown ||
1231 ShapedType::isDynamic(srcDimSize) ||
1232 cstTileSize.value() < srcDimSize;
1234 outerDimOffsets.push_back(offsets[dim]);
1235 if (ShapedType::isStatic(destDimSize)) {
1236 outerDimSizes.push_back(
b.getIndexAttr(destDimSize));
1238 outerDimSizes.push_back(
1239 b.createOrFold<tensor::DimOp>(loc, packOp.getDest(), dim));
1266 bool assumeInnerTileSizesMatchTiles =
1267 innerTileAlignment == InnerTileAlignment::Equal;
1268 bool staticallyDecidable =
1269 !
failed(cstTileSize) && cstInnerSize.has_value();
1270 if (innerTileAlignment == InnerTileAlignment::Unknown) {
1271 if (!staticallyDecidable || *cstTileSize % *cstInnerSize != 0)
1273 }
else if (staticallyDecidable) {
1274 assert(*cstTileSize % *cstInnerSize == 0 &&
1275 "InnerTileAlignment hint contradicts statically known tile "
1277 assert((innerTileAlignment != InnerTileAlignment::Equal ||
1278 *cstTileSize == *cstInnerSize) &&
1279 "InnerTileAlignment::Equal contradicts statically known tile "
1283 using AV = affine::AffineValueExpr;
1284 affine::AffineBuilder ab(
b, loc);
1285 AffineExpr dim0, sym;
1288 auto avOffset = AV(dim0).bind(offsets[dim]);
1289 auto avSize = AV(dim0).bind(sizes[dim]);
1290 auto avTileSize = AV(sym).bind(dimAndTileMapping[dim]);
1291 outerDimOffsets.push_back(ab.floor(avOffset, avTileSize));
1294 outerDimSizes.push_back(assumeInnerTileSizesMatchTiles
1296 : ab.ceil(avSize, avTileSize));
1298 outerDimOffsets.push_back(offsets[dim]);
1299 outerDimSizes.push_back(sizes[dim]);
1302 applyPermToRange(outerDimOffsets, outerDimSizes, packOp.getOuterDimsPerm());
1303 resultOffsets = outerDimOffsets;
1304 resultSizes = outerDimSizes;
1308 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
1309 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1310 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1311 ArrayRef<SmallVector<OpFoldResult>> allSizes)
const {
1312 return getTiledImplementationFromOperandTiles(op,
b, operandNumbers,
1313 allOffsets, allSizes,
1318 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
1319 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1320 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1321 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1322 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1323 if (operandNumbers.size() != 1 || operandNumbers[0] != 0) {
1324 LLVM_DEBUG({ llvm::dbgs() <<
"unhandled operands for consumer fusion"; });
1328 ArrayRef<OpFoldResult> offsets(allOffsets[0]);
1329 ArrayRef<OpFoldResult> sizes(allSizes[0]);
1331 auto packOp = cast<PackOp>(op);
1333 if (!packOp.hasPureTensorSemantics())
1336 Location loc = packOp.getLoc();
1338 int64_t inputRank = packOp.getSourceRank();
1339 auto oneAttr =
b.getI64IntegerAttr(1);
1340 SmallVector<OpFoldResult> strides(inputRank, oneAttr);
1342 SmallVector<Value> tiledOperands;
1343 auto sourceSlice = tensor::ExtractSliceOp::create(
1344 b, loc, packOp.getSource(), offsets, sizes, strides);
1345 tiledOperands.push_back(sourceSlice);
1347 SmallVector<OpFoldResult> outerDimOffsets, outerDimSizes;
1348 if (
failed(getIterationDomainTileFromOperandTiles(
1349 op,
b, operandNumbers, allOffsets, allSizes, outerDimOffsets,
1350 outerDimSizes, innerTileAlignments)))
1353 SmallVector<OpFoldResult> outputOffsets, outputSizes;
1355 outputOffsets, outputSizes)))
1358 strides.append(packOp.getDestRank() - inputRank, oneAttr);
1359 auto outSlice = tensor::ExtractSliceOp::create(
1360 b, loc, packOp.getDest(), outputOffsets, outputSizes, strides);
1361 tiledOperands.push_back(outSlice);
1363 if (
auto val = packOp.getPaddingValue())
1364 tiledOperands.push_back(val);
1365 for (
auto tile : packOp.getInnerTiles())
1366 tiledOperands.push_back(
tile);
1368 Operation *tiledPackOp = PackOp::create(
1371 return TilingResult{
1373 SmallVector<Value>(tiledPackOp->
getResults()),
1374 llvm::to_vector(ArrayRef<Operation *>{sourceSlice, outSlice})};
1378struct UnpackTileDimInfo {
1379 bool isAlignedToInnerTileSize;
1380 OpFoldResult sourceOffset;
1381 OpFoldResult sourceSize;
1382 OpFoldResult resultOffset;
1383 OpFoldResult destExpandedSize;
1389static UnpackTileDimInfo
1393 UnpackTileDimInfo info;
1397 unpackOp.getDimAndTileMapping();
1399 if (!dimAndTileMapping.count(tileDim)) {
1400 info.isAlignedToInnerTileSize =
true;
1401 info.sourceOffset = tileOffset;
1402 info.sourceSize = tileSize;
1403 info.resultOffset = zeroAttr;
1404 info.destExpandedSize = tileSize;
1415 OpFoldResult innerTileSize = dimAndTileMapping[tileDim];
1417 info.isAlignedToInnerTileSize =
false;
1430 bool assumeInnerTileSizesMatchTiles =
1432 bool staticallyDecidable = !
failed(cstSize) && cstInnerSize.has_value();
1434 info.isAlignedToInnerTileSize =
true;
1435 if (staticallyDecidable) {
1436 assert(*cstSize % *cstInnerSize == 0 &&
1437 "InnerTileAlignment hint contradicts statically known tile sizes");
1439 *cstSize == *cstInnerSize) &&
1440 "InnerTileAlignment::Equal contradicts statically known tile "
1444 if (info.isAlignedToInnerTileSize || (!
failed(cstSize) && cstInnerSize)) {
1445 if (!info.isAlignedToInnerTileSize && *cstSize % *cstInnerSize == 0)
1446 info.isAlignedToInnerTileSize =
true;
1450 if (assumeInnerTileSizesMatchTiles ||
1451 (cstInnerSize && !
failed(cstSize) && *cstInnerSize == *cstSize)) {
1452 auto lhs = AV(dim0).bind(tileOffset);
1453 auto rhs = AV(dim1).bind(innerTileSize);
1454 info.sourceOffset = ab.floor(
lhs,
rhs);
1455 info.sourceSize = oneAttr;
1456 info.resultOffset = zeroAttr;
1457 info.destExpandedSize = tileSize;
1462 if (info.isAlignedToInnerTileSize) {
1464 ab.floor(AV(dim0).bind(tileOffset), AV(dim1).bind(innerTileSize));
1465 info.resultOffset = zeroAttr;
1466 info.destExpandedSize = tileSize;
1475 ab.ceil(AV(dim0).bind(tileSize), AV(dim1).bind(innerTileSize));
1479 affine::DivModValue firstCoord = affine::getDivMod(
1483 ab.add(AV(dim0).bind(tileOffset), AV(dim1).bind(tileSize));
1484 affine::DivModValue lastCoord = affine::getDivMod(
1488 ab.sub(AV(dim0).bind(tileExclusiveBound), AV(dim1).bind(oneAttr))),
1491 OpFoldResult lengthMinusOne = ab.sub(AV(dim0).bind(lastCoord.quotient),
1492 AV(dim1).bind(firstCoord.quotient));
1494 ab.add(AV(dim0).bind(lengthMinusOne), AV(dim1).bind(oneAttr));
1495 info.sourceOffset = firstCoord.quotient;
1496 info.resultOffset = firstCoord.remainder;
1499 info.destExpandedSize =
b.createOrFold<arith::MulIOp>(
1505struct UnPackOpTiling
1506 :
public TilingInterface::ExternalModel<UnPackOpTiling, linalg::UnPackOp> {
1507 using Base = TilingInterface::ExternalModel<UnPackOpTiling, linalg::UnPackOp>;
1508 using Base::getIterationDomainTileFromOperandTiles;
1510 SmallVector<utils::IteratorType> getLoopIteratorTypes(Operation *op)
const {
1511 auto unpackOp = cast<UnPackOp>(op);
1512 SmallVector<utils::IteratorType> iteratorTypes(
1513 unpackOp.getDestRank(), utils::IteratorType::parallel);
1514 return iteratorTypes;
1517 SmallVector<Range> getIterationDomain(Operation *op, OpBuilder &
b)
const {
1518 return getPackUnPackIterationDomain<UnPackOp>(cast<UnPackOp>(op),
b);
1535 FailureOr<TilingResult>
1537 ArrayRef<OpFoldResult> offsets,
1538 ArrayRef<OpFoldResult> sizes)
const {
1544 Operation *op, OpBuilder &
b, ArrayRef<OpFoldResult> offsets,
1545 ArrayRef<OpFoldResult> sizes,
1546 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1547 auto unpackOp = cast<UnPackOp>(op);
1549 if (!unpackOp.hasPureTensorSemantics())
1552 int64_t srcRank = unpackOp.getSourceRank();
1553 int64_t destRank = unpackOp.getDestRank();
1554 int64_t numInnerTiles = srcRank - destRank;
1555 Location loc = unpackOp.getLoc();
1560 bool isPerfectTilingCase =
true;
1561 Attribute oneAttr =
b.getIndexAttr(1);
1562 SmallVector<OpFoldResult> sliceSrcStrides(destRank, oneAttr);
1563 SmallVector<OpFoldResult> sliceSrcIndices, sliceSrcSizes;
1564 SmallVector<OpFoldResult> destExpandedSizes, resultOffsetsFromDest;
1565 for (
auto dim : llvm::seq<int64_t>(0, destRank)) {
1566 UnpackTileDimInfo info = getUnpackTileDimInfo(
1567 b, unpackOp, dim, offsets[dim], sizes[dim],
1568 dim <
static_cast<int64_t
>(innerTileAlignments.size())
1569 ? innerTileAlignments[dim]
1570 : InnerTileAlignment::Unknown);
1571 if (!info.isAlignedToInnerTileSize)
1572 isPerfectTilingCase =
false;
1573 sliceSrcIndices.push_back(info.sourceOffset);
1574 sliceSrcSizes.push_back(info.sourceSize);
1575 destExpandedSizes.push_back(info.destExpandedSize);
1576 resultOffsetsFromDest.push_back(info.resultOffset);
1581 applyPermToRange(sliceSrcIndices, sliceSrcSizes,
1582 unpackOp.getOuterDimsPerm());
1583 Attribute zeroAttr =
b.getIndexAttr(0);
1584 sliceSrcIndices.append(numInnerTiles, zeroAttr);
1585 sliceSrcSizes.append(unpackOp.getMixedTiles());
1586 sliceSrcStrides.append(numInnerTiles, oneAttr);
1587 SmallVector<Operation *> generatedSlices;
1588 tensor::ExtractSliceOp sliceSource = tensor::ExtractSliceOp::create(
1589 b, loc, unpackOp.getSource(), sliceSrcIndices, sliceSrcSizes,
1591 generatedSlices.push_back(sliceSource);
1593 SmallVector<OpFoldResult> destStrides(destRank, oneAttr);
1595 if (isPerfectTilingCase) {
1596 auto destSliceOp = tensor::ExtractSliceOp::create(
1597 b, loc, unpackOp.getDest(), offsets, sizes, destStrides);
1598 sliceDest = destSliceOp;
1599 generatedSlices.push_back(destSliceOp);
1601 sliceDest = tensor::EmptyOp::create(
1602 b, loc, destExpandedSizes, unpackOp.getDestType().getElementType());
1605 SmallVector<Value> tiledOperands = {sliceSource.getResult(), sliceDest};
1606 for (
auto tile : unpackOp.getInnerTiles())
1607 tiledOperands.push_back(
tile);
1609 Operation *tiledUnpackOp = UnPackOp::create(
1612 if (isPerfectTilingCase)
1613 return TilingResult{{tiledUnpackOp},
1614 SmallVector<Value>(tiledUnpackOp->
getResults()),
1617 auto extractSlice = tensor::ExtractSliceOp::create(
1618 b, loc, tiledUnpackOp->
getResult(0), resultOffsetsFromDest, sizes,
1620 return TilingResult{
1621 {tiledUnpackOp}, {extractSlice.getResult()}, generatedSlices};
1626 ArrayRef<OpFoldResult> offsets,
1627 ArrayRef<OpFoldResult> sizes,
1628 SmallVector<OpFoldResult> &resultOffsets,
1629 SmallVector<OpFoldResult> &resultSizes)
const {
1630 resultOffsets = llvm::to_vector(offsets);
1631 resultSizes = llvm::to_vector(sizes);
1635 FailureOr<TilingResult>
1636 generateResultTileValue(Operation *op, OpBuilder &
b,
unsigned resultNumber,
1637 ArrayRef<OpFoldResult> offsets,
1638 ArrayRef<OpFoldResult> sizes)
const {
1639 return generateResultTileValue(op,
b, resultNumber, offsets, sizes,
1643 FailureOr<TilingResult> generateResultTileValue(
1644 Operation *op, OpBuilder &
b,
unsigned resultNumber,
1645 ArrayRef<OpFoldResult> offsets, ArrayRef<OpFoldResult> sizes,
1646 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1647 FailureOr<TilingResult> tilingResult =
1649 if (
failed(tilingResult))
1651 return tilingResult.value();
1654 LogicalResult generateScalarImplementation(Operation *op, OpBuilder &builder,
1657 auto unpackOp = cast<UnPackOp>(op);
1658 assert(unpackOp.hasPureBufferSemantics() &&
1659 "expected operation to have buffer semantics");
1660 assert(ivs.size() == unpackOp.getDestRank() &&
1661 "number of ivs must match the rank of the output tensor");
1662 OpBuilder::InsertionGuard g(builder);
1665 unpackOp.getDimAndTileMapping();
1667 SmallVector<Value> inputIvs;
1669 SmallVector<Value> inputIvsPointLoops;
1670 inputIvs.reserve(unpackOp.getDestRank());
1671 inputIvsPointLoops.reserve(dimAndTileMapping.size());
1672 for (
auto dim : llvm::seq<int64_t>(0, unpackOp.getDestRank())) {
1673 if (dimAndTileMapping.count(dim)) {
1674 affine::DivModValue divMod =
1675 affine::getDivMod(builder, loc, ivs[dim],
1677 builder, loc, dimAndTileMapping[dim]));
1678 inputIvsPointLoops.push_back(divMod.remainder);
1679 inputIvs.push_back(divMod.quotient);
1681 inputIvs.push_back(ivs[dim]);
1687 assert(inputIvsPointLoops.size() + inputIvs.size() ==
1688 unpackOp.getSourceRank() &&
1689 "expect same number of induction variables equals to input rank");
1691 ArrayRef<int64_t> innerDims = unpackOp.getInnerDimsPos();
1692 SmallVector<int64_t> interchangeVector =
1693 computeInterchangeFromDimPos(innerDims, unpackOp.getDestRank());
1694 SmallVector<Value> interchangedInputIvsPointLoops = inputIvsPointLoops;
1695 interchangedInputIvsPointLoops = interchange<Value>(
1696 interchangedInputIvsPointLoops, interchangeVector, 0);
1699 ArrayRef<int64_t> outerDims = unpackOp.getOuterDimsPerm();
1700 if (!outerDims.empty())
1701 inputIvs = interchange<Value>(inputIvs, outerDims, 0);
1703 llvm::append_range(inputIvs, interchangedInputIvsPointLoops);
1705 memref::LoadOp::create(builder, loc, unpackOp.getSource(), inputIvs);
1706 memref::StoreOp::create(builder, loc, scalar, unpackOp.getDest(), ivs);
1712 LogicalResult getIterationDomainTileFromOperandTiles(
1713 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1714 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1715 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1716 SmallVectorImpl<OpFoldResult> &resultOffsets,
1717 SmallVectorImpl<OpFoldResult> &resultSizes)
const {
1718 if (operandNumbers.size() != 1) {
1719 LLVM_DEBUG({ llvm::dbgs() <<
"unable to handle multiple operands"; });
1722 auto unPackOp = cast<UnPackOp>(op);
1723 unsigned operandNumber = operandNumbers[0];
1724 ArrayRef<OpFoldResult> offsets(allOffsets[0]);
1725 ArrayRef<OpFoldResult> sizes(allSizes[0]);
1728 if (operandNumber == unPackOp.getDestMutable().getOperandNumber()) {
1729 resultOffsets = llvm::to_vector(offsets);
1730 resultSizes = llvm::to_vector(sizes);
1733 Location loc = unPackOp.getLoc();
1735 int64_t numTiles = unPackOp.getInnerDimsPos().size();
1736 auto destOffsets = offsets.drop_back(numTiles);
1737 auto destSizes = sizes.drop_back(numTiles);
1740 int64_t outputRank = unPackOp.getDestRank();
1744 SmallVector<OpFoldResult> outputMixedSizes = reifiedReturnShapes.front();
1745 SmallVector<OpFoldResult> origOffsets(destOffsets);
1746 SmallVector<OpFoldResult> origSizes(destSizes);
1747 applyPermToRange(origOffsets, origSizes,
1751 unPackOp.getDimAndTileMapping();
1753 for (
auto dim : llvm::seq<int64_t>(0, outputRank)) {
1754 using AV = affine::AffineValueExpr;
1755 affine::AffineBuilder ab(
b, loc);
1756 AffineExpr dim0, dim1, sym0;
1759 if (dimAndTileMapping.count(dim)) {
1763 auto avOffset = AV(dim0).bind(origOffsets[dim]);
1764 auto avSize = AV(dim0).bind(origSizes[dim]);
1765 auto avTileSize = AV(sym0).bind(dimAndTileMapping[dim]);
1766 auto avResultSize = AV(dim0).bind(outputMixedSizes[dim]);
1767 resultOffsets.push_back(ab.mul(avOffset, avTileSize));
1768 auto avResultOffset = AV(dim1).bind(resultOffsets.back());
1769 resultSizes.push_back(ab.min({ab.mul(avSize, avTileSize),
1770 ab.sub(avResultSize, avResultOffset)}));
1772 resultOffsets.push_back(origOffsets[dim]);
1773 resultSizes.push_back(origSizes[dim]);
1779 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
1780 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1781 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1782 ArrayRef<SmallVector<OpFoldResult>> allSizes)
const {
1783 return getTiledImplementationFromOperandTiles(op,
b, operandNumbers,
1784 allOffsets, allSizes,
1789 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
1790 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1791 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1792 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1793 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1794 if (operandNumbers.size() != 1 || operandNumbers[0] != 0) {
1795 LLVM_DEBUG({ llvm::dbgs() <<
"unhandled operands for consumer fusion"; });
1798 auto unPackOp = cast<UnPackOp>(op);
1800 if (!unPackOp.hasPureTensorSemantics())
1803 ArrayRef<OpFoldResult> offsets(allOffsets[0]);
1804 ArrayRef<OpFoldResult> sizes(allSizes[0]);
1810 int64_t numTiles = unPackOp.getInnerDimsPos().size();
1811 ArrayRef<int64_t> innerDimsPos = unPackOp.getInnerDimsPos();
1812 SmallVector<OpFoldResult> mixedTiles = unPackOp.getMixedTiles();
1813 ArrayRef<OpFoldResult> innerSizes = sizes.take_back(numTiles);
1814 for (int64_t i = 0; i < numTiles; ++i) {
1817 int64_t destDim = innerDimsPos[i];
1819 destDim < static_cast<int64_t>(innerTileAlignments.size()) &&
1820 innerTileAlignments[destDim] == InnerTileAlignment::Equal;
1830 "InnerTileAlignment::Equal contradicts statically known tile "
1839 Location loc = unPackOp.getLoc();
1843 SmallVector<OpFoldResult> outputOffsets, outputSizes;
1844 if (
failed(getIterationDomainTileFromOperandTiles(
1845 op,
b, operandNumbers, allOffsets, allSizes, outputOffsets,
1849 auto oneAttr =
b.getI64IntegerAttr(1);
1850 int64_t outputRank = unPackOp.getDestRank();
1851 SmallVector<OpFoldResult> strides(outputRank, oneAttr);
1853 SmallVector<Value> tiledOperands;
1855 auto extractDestSlice = tensor::ExtractSliceOp::create(
1856 b, loc, unPackOp.getDest(), outputOffsets, outputSizes, strides);
1857 tiledOperands.push_back(extractDestSlice);
1859 strides.append(unPackOp.getSourceRank() - outputRank, oneAttr);
1861 auto extractSourceSlice = tensor::ExtractSliceOp::create(
1862 b, loc, unPackOp.getSource(), offsets, sizes, strides);
1863 tiledOperands.insert(tiledOperands.begin(), extractSourceSlice);
1864 for (
auto tile : unPackOp.getInnerTiles())
1865 tiledOperands.push_back(
tile);
1868 Operation *tiledUnPackOp =
1869 UnPackOp::create(
b, loc,
TypeRange{extractDestSlice.getType()},
1872 return TilingResult{{tiledUnPackOp},
1873 SmallVector<Value>(tiledUnPackOp->
getResults()),
1874 llvm::to_vector(ArrayRef<Operation *>{
1875 extractSourceSlice, extractDestSlice})};
1881template <
typename OpType>
1883 OpType::template attachInterface<LinalgOpTilingInterface<OpType>>(*ctx);
1884 OpType::template attachInterface<LinalgOpPartialReductionInterface<OpType>>(
1889template <
typename... OpTypes>
1900 linalg::PackOp::attachInterface<PackOpTiling>(*ctx);
1901 linalg::UnPackOp::attachInterface<UnPackOpTiling>(*ctx);
1903#include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc"
1911 linalg::PackOp::attachInterface<PackOpTiling>(*ctx);
1912 linalg::UnPackOp::attachInterface<UnPackOpTiling>(*ctx);
static bool isTiled(AffineExpr expr, ArrayRef< OpFoldResult > tileSizes)
static RankedTensorType sliceResultType(Type operandType, GridOp grid, ArrayRef< GridAxis > gridAxes, int64_t sliceAxis)
static LogicalResult getResultTilePosition(RewriterBase &rewriter, ReductionTilingStrategy reductionStrategy, int64_t index, Value tiledResult, TilingInterface op, ArrayRef< OpFoldResult > offsets, ArrayRef< OpFoldResult > sizes, ValueRange ivs, ArrayRef< OpFoldResult > numThreads, ArrayRef< OpFoldResult > givenTileSizes, const SetVector< unsigned > &reductionDims, SmallVector< OpFoldResult > &resultOffset, SmallVector< OpFoldResult > &resultSize)
static FailureOr< TilingResult > getTiledImplementation(RewriterBase &rewriter, TilingInterface op, ReductionTilingStrategy reductionStrategy, ValueRange regionIterArg, ArrayRef< OpFoldResult > offsets, ArrayRef< OpFoldResult > sizes, ValueRange ivs, ArrayRef< OpFoldResult > numThreads, ArrayRef< OpFoldResult > givenTileSizes, ArrayRef< InnerTileAlignment > innerTileAlignments, const SetVector< unsigned > &reductionDims)
static LogicalResult inlinePayload(OpBuilder &b, LinalgOp linalgOp, ValueRange ivs, ValueRange argValues)
Method to inline the payload of a linalgOp given the iteration space point and values for the argumen...
static SmallVector< Value > getIndicesForAccess(OpBuilder &b, Location loc, AffineMap indexingMap, ValueRange ivs)
Return the SSA values that represent the data point accessed using a given indexingMap for a given po...
static bool isInBounds(TransferOp op, int64_t resultIdx, int64_t indicesIdx)
Base type for affine expression.
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
bool isProjectedPermutation(bool allowZeroInResults=false) const
Returns true if the AffineMap represents a subset (i.e.
unsigned getNumSymbols() const
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
unsigned getNumResults() const
Attributes are known-constant values of operations.
Block represents an ordered list of Operations.
Operation * getTerminator()
Get the terminator operation of this block.
BlockArgListType getArguments()
iterator_range< iterator > without_terminator()
Return an iterator range over the operation within this block excluding the terminator operation at t...
IntegerAttr getIndexAttr(int64_t value)
MLIRContext * getContext() const
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
This is a utility class for mapping one set of IR entities to another.
auto lookupOrDefault(T from) const
Lookup a mapped value within the map.
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
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 setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
This class represents a single result from folding an operation.
This class represents an operand of an operation.
Operation is the basic unit of execution within MLIR.
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
void setOperand(unsigned idx, Value value)
ArrayRef< NamedAttribute > getAttrs()
Return all of the attributes on 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.
operand_range getOperands()
Returns an iterator on the underlying Value's.
result_range getResults()
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
void cloneInto(Region *dest, IRMapping &mapper)
Clone the internal blocks from this region into dest.
static FailureOr< int64_t > computeConstantBound(presburger::BoundType type, const Variable &var, const StopConditionFn &stopCondition=nullptr, ValueBoundsOptions options={})
Compute a constant bound for the given variable.
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
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)
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< Value > makeTiledShapes(OpBuilder &builder, Location loc, LinalgOp linalgOp, ValueRange valuesToTile, ArrayRef< OpFoldResult > ivs, ArrayRef< OpFoldResult > tileSizes, ArrayRef< OpFoldResult > sizeBounds, bool omitPartialTileCheck)
Creates extract_slice/subview ops for all valuesToTile of the given linalgOp with builder,...
void registerTilingInterfaceExternalModelsForPackUnPackOps(DialectRegistry ®istry)
Similar to the above registeration, but it is only for tensor.pack and tensor.unpack ops.
static void registerOne(MLIRContext *ctx)
static void registerAll(MLIRContext *ctx)
Variadic helper function.
void offsetIndices(OpBuilder &b, LinalgOp linalgOp, ArrayRef< OpFoldResult > offests)
Add the specified offsets to any linalg.index ops contained in the given linalgOp.
Value createOrFoldDimOp(OpBuilder &b, Location loc, Value val, int64_t dim)
Create one memref::DimOp or tensor::DimOp depending on the type of val.
void registerTilingInterfaceExternalModels(DialectRegistry ®istry)
SmallVector< Type > getTensorOutputTypes(LinalgOp op, ValueRange operands)
Returns the list of tensor output types produced when the given structured operation op is applied to...
SliceParameters computeSliceParameters(OpBuilder &builder, Location loc, Value valueToTile, ArrayRef< OpFoldResult > tileSizes, AffineMap map, ArrayRef< OpFoldResult > lbs, ArrayRef< OpFoldResult > ubs, ArrayRef< OpFoldResult > subShapeSizes, bool omitPartialTileCheck)
Computes SliceParameters for a single valueToTile assuming that its user is being tiled with the give...
SmallVector< OpFoldResult > getMixedSizes(OpBuilder &builder, Location loc, Value value)
Return the dimensions of the given tensor value.
Include the generated interface declarations.
ReductionTilingStrategy
Tiling can be thought of as splitting a dimension into 2 and materializing the outer dimension as a l...
@ PartialReductionOuterReduction
@ PartialReductionOuterParallel
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
LogicalResult reifyResultShapes(OpBuilder &b, Operation *op, ReifiedRankedShapedTypeDims &reifiedReturnShapes)
Reify the shape of the result of an operation (typically in terms of the shape of its operands).
bool isEqualConstantIntOrValue(OpFoldResult ofr1, OpFoldResult ofr2)
Return true if ofr1 and ofr2 are the same integer constant attribute values or the same SSA value.
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
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.
bool isZeroInteger(OpFoldResult v)
Return "true" if v is an integer value/attribute with constant value 0.
void bindSymbols(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to SymbolExpr at positions: [0 .
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
Operation * clone(OpBuilder &b, Operation *op, TypeRange newResultTypes, ValueRange newOperands)
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...
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
void applyPermutationToVector(SmallVector< T, N > &inVec, ArrayRef< int64_t > permutation)
Apply the permutation defined by permutation to inVec.
InnerTileAlignment
Per-dimension alignment of a loop tile size to a linalg.pack / linalg.unpack inner tile size,...
std::pair< SmallVector< int64_t >, SmallVector< Value > > decomposeMixedValues(ArrayRef< OpFoldResult > mixedValues)
Decompose a vector of mixed static or dynamic values into the corresponding pair of arrays.
SmallVector< int64_t > invertPermutationVector(ArrayRef< int64_t > permutation)
Helper method to apply to inverse a permutation.
Helper struct to build simple arithmetic quantities with minimal type inference support.
Container for result values of tiling.
Options that control value bound computation.
Helper struct to build simple AffineValueExprs with minimal type inference support.
A struct containg offsets-sizes-strides arguments of the tiled shape.
SmallVector< OpFoldResult > sizes
SmallVector< OpFoldResult > offsets
Eliminates variable at the specified position using Fourier-Motzkin variable elimination.