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(),
100 for (
auto [pos, tileSize] : llvm::enumerate(tileSizeBounds)) {
101 if (failed(tileSize)) {
102 tiledDims[pos] =
true;
108 ShapedType::isDynamic(loopRanges[pos]) || *tileSize < loopRanges[pos];
111 for (
AffineMap map : linalgOp.getIndexingMapsArray()) {
114 auto binExpr = dyn_cast<AffineBinaryOpExpr>(expr);
126 auto dim = dyn_cast<AffineDimExpr>(e);
127 if (dim && tiledDims[dim.getPosition()])
132 if (!involvesTiledDim)
139 auto dimExpr = dyn_cast<AffineDimExpr>(binExpr.getLHS());
140 auto stepExpr = dyn_cast<AffineConstantExpr>(binExpr.getRHS());
141 if (!dimExpr || !stepExpr || stepExpr.getValue() <= 0) {
142 linalgOp.emitOpError()
143 <<
"tiling is not supported for the semi-affine indexing map: "
144 "only a single iteration dimension divided by a positive "
145 "constant step can be tiled over a tiled dimension";
156 unsigned dimPos = dimExpr.getPosition();
157 FailureOr<int64_t> tileSize = tileSizeBounds[dimPos];
161 if (failed(tileSize) || *tileSize == 1)
176 int64_t step = stepExpr.getValue();
178 bool safe = *tileSize % step == 0 || (!isCeil && step % *tileSize == 0);
180 linalgOp.emitOpError()
181 <<
"tiling is not supported for the semi-affine indexing map: "
183 << *tileSize <<
" for dimension d" << dimPos
184 << (isCeil ?
" must be a multiple of the step "
185 :
" must divide or be divisible by the step ")
208template <
typename LinalgOpTy>
209struct LinalgOpTilingInterface
210 :
public TilingInterface::ExternalModel<LinalgOpTilingInterface<LinalgOpTy>,
213 TilingInterface::ExternalModel<LinalgOpTilingInterface<LinalgOpTy>,
217 using Base::generateResultTileValue;
218 using Base::getIterationDomainTileFromOperandTiles;
219 using Base::getTiledImplementation;
220 using Base::getTiledImplementationFromOperandTiles;
223 SmallVector<utils::IteratorType> getLoopIteratorTypes(Operation *op)
const {
224 LinalgOpTy concreteOp = cast<LinalgOpTy>(op);
225 return concreteOp.getIteratorTypesArray();
229 SmallVector<Range> getIterationDomain(Operation *op, OpBuilder &
b)
const {
230 OpBuilder::InsertionGuard g(
b);
231 b.setInsertionPoint(op);
232 Location loc = op->
getLoc();
233 LinalgOp linalgOp = cast<LinalgOp>(op);
234 SmallVector<OpFoldResult> allShapesSizes =
235 linalgOp.createFlatListOfOperandDims(
b, loc);
236 AffineMap map = linalgOp.getShapesToLoopsMap();
238 return llvm::map_to_vector(map.
getResults(), [&](AffineExpr loopExpr) {
239 OpFoldResult ofr = affine::makeComposedFoldedAffineApply(b, loc, loopExpr,
241 return Range{b.getIndexAttr(0), ofr, b.getIndexAttr(1)};
246 FailureOr<TilingResult>
253 LinalgOp linalgOp = cast<LinalgOp>(op);
264 b, loc, linalgOp, valuesToTile, offsets, sizes, {},
true);
266 llvm::make_filter_range(
268 [](
Value v) ->
bool {
269 return isa_and_nonnull<tensor::ExtractSliceOp, memref::SubViewOp>(
277 Operation *tiledOp =
clone(
b, linalgOp, resultTensorTypes, tiledOperands);
288 getMappedOffsetAndSize(LinalgOp linalgOp,
OpBuilder &
b,
296 for (
auto [indexingMap, offsets, sizes] :
297 llvm::zip_equal(indexingMaps, allOffsets, allSizes)) {
298 for (
auto [resultExpr, offset, size] :
299 llvm::zip_equal(indexingMap.getResults(), offsets, sizes)) {
300 auto dimExpr = dyn_cast<AffineDimExpr>(resultExpr);
303 unsigned position = dimExpr.getPosition();
304 auto it = mappedOffsets.find(position);
305 if (it != mappedOffsets.end()) {
308 if (seenOffset != offset || seenSize != size) {
310 llvm::dbgs() <<
"inconsistent iteration space mapping from "
311 "offsets/sizes of operands/results";
316 mappedOffsets[position] = offset;
317 mappedSizes[position] = size;
325 cast<TilingInterface>(linalgOp.getOperation()).getIterationDomain(
b);
326 mappedOffsetsVec.resize(iterationDomain.size());
327 mappedSizesVec.resize(iterationDomain.size());
328 for (
auto [
index, domain] : llvm::enumerate(iterationDomain)) {
329 auto it = mappedOffsets.find(
index);
330 if (it != mappedOffsets.end()) {
331 mappedOffsetsVec[
index] = it->second;
332 mappedSizesVec[
index] = mappedSizes.lookup(
index);
335 mappedOffsetsVec[
index] = domain.offset;
336 mappedSizesVec[
index] = domain.size;
343 LogicalResult getIterationDomainTileFromOperandTiles(
349 auto linalgOp = cast<LinalgOp>(op);
352 llvm::map_to_vector(operandNumbers, [&](
unsigned operandNumber) {
353 OpOperand &opOperand = linalgOp->getOpOperand(operandNumber);
354 return linalgOp.getMatchingIndexingMap(&opOperand);
356 if (
failed(getMappedOffsetAndSize(linalgOp,
b, indexingMaps, allOffsets,
357 allSizes, iterDomainOffsets,
373 LinalgOp linalgOp = cast<LinalgOp>(op);
382 OpOperand *outOperand = linalgOp.getDpsInitOperand(resultNumber);
384 b, loc, outOperand->get(), sizes,
385 linalgOp.getMatchingIndexingMap(outOperand), offsets,
386 {}, subShapeSizes,
true);
387 resultOffsets = sliceParams.
offsets;
388 resultSizes = sliceParams.
sizes;
392 LogicalResult getIterationDomainTileFromResultTile(
397 auto linalgOp = cast<LinalgOp>(op);
404 linalgOp.getIndexingMapMatchingResult(op->
getResult(resultNumber));
407 "unhandled tiled implementation generation when result is not "
408 "accessed using a permuted projection");
414 getMappedOffsetAndSize(linalgOp,
b, indexingMap, {allOffsets},
415 {allSizes}, iterDomainOffsets, iterDomainSizes);
417 assert(succeeded(status) &&
"unexpected error in offset calculation");
421 FailureOr<TilingResult>
426 if (
failed(getIterationDomainTileFromResultTile(
427 op,
b, resultNumber, offsets, sizes, mappedOffsets, mappedSizes))) {
430 auto tilingInterfaceOp = cast<TilingInterface>(op);
431 FailureOr<TilingResult> tilingResult =
432 tilingInterfaceOp.getTiledImplementation(
b, mappedOffsets, mappedSizes);
437 if (tilingResult->tiledOps.size() != 1)
438 return op->
emitOpError(
"failed to generate tiled implementation");
441 tilingResult->tiledOps,
443 tilingResult->generatedSlices};
448 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
453 if (
failed(getIterationDomainTileFromOperandTiles(
454 op,
b, operandNumbers, allOffsets, allSizes, mappedOffsets,
464 auto linalgOp = cast<LinalgOp>(op);
465 if (!linalgOp.hasPureBufferSemantics())
466 return op->
emitOpError(
"expected operation to have buffer semantics");
469 indexedValues.reserve(linalgOp->getNumOperands());
473 for (
OpOperand &operand : linalgOp->getOpOperands()) {
474 if (!linalgOp.payloadUsesValueFromOperand(&operand)) {
475 indexedValues.push_back(
nullptr);
478 if (linalgOp.isScalar(&operand)) {
479 indexedValues.push_back(operand.get());
483 builder, linalgOpLoc, linalgOp.getMatchingIndexingMap(&operand), ivs);
485 memref::LoadOp::create(builder, linalgOpLoc, operand.get(),
indices);
486 indexedValues.push_back(
load);
493 bool isOpFusableWithConsumerSlice(
Operation *op,
unsigned resultNumber,
500 bool isOpFusableWithProducerSlices(
505 auto linalgOp = cast<LinalgOp>(op);
507 llvm::map_to_vector(operandNumbers, [&](
unsigned operandNumber) {
508 OpOperand &opOperand = linalgOp->getOpOperand(operandNumber);
509 return linalgOp.getMatchingIndexingMap(&opOperand);
514 return succeeded(getMappedOffsetAndSize(linalgOp,
b, indexingMaps,
515 allOffsets, allSizes, mappedOffsets,
527 for (
auto [
index, reductionDim] : llvm::enumerate(reductionDims)) {
528 if (reductionDim == value) {
540getPartialResultAffineMaps(LinalgOp linalgOp,
542 auto partialReductionMaps = llvm::map_to_vector(
543 linalgOp.getDpsInitsMutable(), [&](
OpOperand &opOperand) {
544 AffineMap map = linalgOp.getMatchingIndexingMap(&opOperand);
545 for (auto redPos : reductionDims) {
547 map.insertResult(getAffineDimExpr(redPos, linalgOp.getContext()),
548 map.getNumResults());
552 return partialReductionMaps;
555struct InitSliceInfo {
556 SmallVector<int64_t> resultShape;
557 SmallVector<OpFoldResult> offsets;
558 SmallVector<OpFoldResult> sizes;
559 SmallVector<OpFoldResult> strides;
565static InitSliceInfo getInitSliceInfoForOuterReduction(
572 Attribute zero = IntegerAttr::get(IndexType::get(context), 0);
573 Attribute one = IntegerAttr::get(IndexType::get(context), 1);
575 for (
auto [resultIdx, dimExpr] :
576 llvm::enumerate(partialReductionMap.
getResults())) {
577 if (isa<AffineConstantExpr>(dimExpr)) {
580 initOffsets.push_back(zero);
581 initSizes.push_back(initOperandShape[resultIdx]);
584 unsigned dim = cast<AffineDimExpr>(dimExpr).getPosition();
585 if (reductionDims.contains(dim)) {
586 initOffsets.push_back(zero);
588 initOffsets.push_back(offsets[dim]);
590 initSizes.push_back(sizes[dim]);
594 return {resultShape, initOffsets, initSizes, initStrides};
600static InitSliceInfo getInitSliceInfoForOuterParallel(
607 Attribute zero = IntegerAttr::get(IndexType::get(context), 0);
608 Attribute one = IntegerAttr::get(IndexType::get(context), 1);
611 for (
auto [resultIdx, dimExpr] :
612 llvm::enumerate(partialReductionMap.
getResults())) {
613 if (isa<AffineConstantExpr>(dimExpr)) {
616 initOffsets.push_back(zero);
617 initSizes.push_back(initOperandShape[resultIdx]);
618 resultShape.push_back(initOperandShape[resultIdx]);
621 unsigned dim = cast<AffineDimExpr>(dimExpr).getPosition();
622 if (std::optional<unsigned> dimPos = getPositionIn(reductionDims, dim)) {
623 initOffsets.push_back(splitReductionIvs[dimPos.value()]);
624 initSizes.push_back(one);
626 initOffsets.push_back(offsets[dim]);
627 initSizes.push_back(sizes[dim]);
628 resultShape.push_back(sizes[dim]);
633 return {staticShapes, initOffsets, initSizes, initStrides};
638static InitSliceInfo getInitSliceInfo(
MLIRContext *context,
647 return getInitSliceInfoForOuterReduction(
648 context, offsets, sizes, reductionDims, splitReductionIvs,
649 partialReductionMap, initOperandShape);
652 "unexpected ReductionTilingStrategy");
653 return getInitSliceInfoForOuterParallel(
654 context, offsets, sizes, reductionDims, splitReductionIvs,
655 partialReductionMap, initOperandShape);
660template <
typename LinalgOpTy>
661struct LinalgOpPartialReductionInterface
662 :
public PartialReductionOpInterface::ExternalModel<
663 LinalgOpPartialReductionInterface<LinalgOpTy>, LinalgOpTy> {
664 FailureOr<SmallVector<Value>> generateInitialTensorForPartialReduction(
665 Operation *op, OpBuilder &
b, Location loc, ArrayRef<OpFoldResult> sizes,
667 auto linalgOp = cast<LinalgOp>(op);
669 OpBuilder::InsertionGuard guard(
b);
670 if (linalgOp.hasPureBufferSemantics())
671 return op->
emitOpError(
"expected operation to have tensor semantics");
673 SmallVector<AffineMap> partialResultMaps =
674 getPartialResultAffineMaps(linalgOp, reductionDims);
676 SmallVector<Value> inits;
677 for (
auto [initIdx,
result, partialMap] :
678 llvm::enumerate(linalgOp->getResults(), partialResultMaps)) {
679 SmallVector<Operation *, 4> combinerOps;
682 combinerOps.size() != 1)
683 return op->
emitOpError(
"Failed to anaysis the reduction operation.");
685 Operation *reductionOp = combinerOps[0];
686 std::optional<TypedAttr> identity = arith::getNeutralElement(reductionOp);
687 if (!identity.has_value())
689 "Failed to get an identity value for the reduction operation.");
692 SmallVector<OpFoldResult> partialResultShape;
693 Value initValue = linalgOp.getDpsInits()[initIdx];
694 SmallVector<OpFoldResult> initShape =
696 for (
auto [resultIdx, dimExpr] :
697 llvm::enumerate(partialMap.getResults())) {
698 if (isa<AffineConstantExpr>(dimExpr)) {
701 partialResultShape.push_back(initShape[resultIdx]);
704 auto dim = cast<AffineDimExpr>(dimExpr);
705 partialResultShape.push_back(sizes[dim.getPosition()]);
710 tensor::EmptyOp::create(
b, loc, partialResultShape, elType);
711 Value constantOp = arith::ConstantOp::create(
b, loc, *identity);
712 auto identityTensor =
713 linalg::FillOp::create(
b, loc, constantOp, emptyTensor);
714 inits.push_back(identityTensor.getResult(0));
720 FailureOr<TilingResult>
721 tileToPartialReduction(Operation *op, OpBuilder &
b, Location loc,
723 ValueRange init, ArrayRef<OpFoldResult> offsets,
724 ArrayRef<OpFoldResult> sizes,
726 ArrayRef<OpFoldResult> splitReductionIvs)
const {
727 OpBuilder::InsertionGuard guard(
b);
728 auto linalgOp = cast<LinalgOp>(op);
730 SmallVector<AffineMap> partialReductionMaps =
731 getPartialResultAffineMaps(linalgOp, reductionDims);
735 SmallVector<AffineMap> newInitMaps;
736 if (tilingStrategy ==
737 ReductionTilingStrategy::PartialReductionOuterReduction) {
738 newInitMaps = llvm::to_vector(partialReductionMaps);
740 newInitMaps = llvm::map_to_vector(
741 linalgOp.getDpsInitsMutable(), [&](OpOperand &opOperand) {
742 return linalgOp.getMatchingIndexingMap(&opOperand);
748 b, loc, linalgOp, linalgOp.getDpsInputs(), offsets, sizes, {},
true);
749 SmallVector<Operation *> generatedSlices = llvm::map_to_vector(
750 llvm::make_filter_range(
751 tiledInputs, [](Value v) ->
bool {
return v.
getDefiningOp(); }),
755 SmallVector<Value, 1> tiledInits;
756 for (
auto [partialReductionMap, valueToTile, initOperandValue] :
757 llvm::zip_equal(partialReductionMaps, init, linalgOp.getDpsInits())) {
760 SmallVector<OpFoldResult> initOperandShape =
762 InitSliceInfo sliceInfo = getInitSliceInfo(
763 b.getContext(), tilingStrategy, offsets, sizes, reductionDims,
764 splitReductionIvs, partialReductionMap, initOperandShape);
765 auto valueToTileType = cast<RankedTensorType>(valueToTile.getType());
767 sliceInfo.resultShape, valueToTileType.getElementType(),
768 valueToTileType.getEncoding());
769 auto sliceOp = tensor::ExtractSliceOp::create(
771 sliceInfo.sizes, sliceInfo.strides);
772 tiledInits.push_back(sliceOp.getResult());
773 generatedSlices.push_back(sliceOp);
777 SmallVector<AffineMap> newMaps = linalgOp.getIndexingMapsArray();
778 for (
auto [initOperand, newInitMap] :
779 llvm::zip_equal(linalgOp.getDpsInitsMutable(), newInitMaps)) {
780 int mapIdx = linalgOp.getIndexingMapIndex(&initOperand);
781 newMaps[mapIdx] = newInitMap;
785 SmallVector<utils::IteratorType> newIteratorTypes =
786 linalgOp.getIteratorTypesArray();
787 if (tilingStrategy ==
788 ReductionTilingStrategy::PartialReductionOuterReduction) {
789 for (
int dim : reductionDims)
790 newIteratorTypes[dim] = utils::IteratorType::parallel;
794 Operation *partialReductionOp;
795 auto resultTypes =
ValueRange(tiledInits).getTypes();
796 if (tilingStrategy ==
797 ReductionTilingStrategy::PartialReductionOuterReduction) {
798 auto genericOp = GenericOp::create(
b, loc, resultTypes, tiledInputs,
799 tiledInits, newMaps, newIteratorTypes);
802 genericOp.getRegion().begin(), mapping);
804 partialReductionOp = genericOp.getOperation();
806 SmallVector<Value> operands = std::move(tiledInputs);
807 llvm::append_range(operands, tiledInits);
808 partialReductionOp =
mlir::clone(
b, op, resultTypes, operands);
812 {partialReductionOp},
813 llvm::map_to_vector(partialReductionOp->
getResults(),
814 [](OpResult r) -> Value { return r; }),
818 FailureOr<MergeResult>
819 mergeReductions(Operation *op, OpBuilder &
b, Location loc,
822 auto linalgOp = cast<LinalgOp>(op);
823 SmallVector<AffineMap> partialReductionMaps =
824 getPartialResultAffineMaps(linalgOp, reductionDims);
827 SmallVector<Operation *> mergeOperations;
828 SmallVector<Value> replacements;
829 for (
auto [idx, init, partialResult, partialMap] : llvm::enumerate(
830 linalgOp.getDpsInits(), partialReduce, partialReductionMaps)) {
831 unsigned initIdx = idx;
836 SmallVector<int64_t> partialReductionDims;
837 for (
auto [resultNum, dimExpr] :
838 llvm::enumerate(partialMap.getResults())) {
839 if (isa<AffineConstantExpr>(dimExpr))
841 unsigned dim = cast<AffineDimExpr>(dimExpr).getPosition();
842 if (llvm::is_contained(reductionDims, dim)) {
843 partialReductionDims.push_back(resultNum);
847 auto reduction = linalg::ReduceOp::create(
848 b, loc, partialResult, init, partialReductionDims,
849 [&linalgOp, &initIdx](OpBuilder &
b, Location loc,
ValueRange inputs) {
851 SmallVector<Operation *, 4> combinerOps;
854 Operation *clonedReductionOp =
b.clone(*combinerOps[0]);
858 linalg::YieldOp::create(
b, loc, clonedReductionOp->
getResult(0));
861 mergeOperations.push_back(reduction);
862 replacements.push_back(reduction->getResult(0));
865 return MergeResult{mergeOperations, replacements};
868 LogicalResult getPartialResultTilePosition(
869 Operation *op, OpBuilder &
b,
unsigned resultNumber,
872 ArrayRef<OpFoldResult> splitReductionIvs,
873 SmallVector<OpFoldResult> &resultOffsets,
874 SmallVector<OpFoldResult> &resultSizes)
const {
875 auto linalgOp = cast<LinalgOp>(op);
876 SmallVector<AffineMap> partialReductionMaps =
877 getPartialResultAffineMaps(linalgOp, reductionDims);
880 Value initOperandValue = linalgOp.getDpsInits()[resultNumber];
881 Location loc = op->
getLoc();
882 SmallVector<OpFoldResult> initOperandShape =
884 InitSliceInfo sliceInfo =
885 getInitSliceInfo(
b.getContext(), tilingStrategy, offsets, sizes,
886 reductionDims, splitReductionIvs,
887 partialReductionMaps[resultNumber], initOperandShape);
888 std::swap(resultOffsets, sliceInfo.offsets);
889 std::swap(resultSizes, sliceInfo.sizes);
895template <
typename OpTy>
898 static_assert(llvm::is_one_of<OpTy, PackOp, UnPackOp>::value,
899 "applies to only pack or unpack operations");
901 int64_t rank = (std::is_same<OpTy, PackOp>::value) ? op.getSourceRank()
906 (
void)op.reifyResultShapes(builder, resultShape);
908 for (
auto dim : llvm::seq<int64_t>(0, rank)) {
909 loopBounds[dim].offset = zero;
910 loopBounds[dim].stride = one;
911 loopBounds[dim].size = resultShape[0][dim];
919 if (permutation.empty())
931 interchangeVector.reserve(dimsPos.size());
940 for (
int64_t dimsIdx = 0, end = dimsPos.size(); dimsIdx < end; dimsIdx++)
941 dimsAndPosMapping[dimsPos[dimsIdx]] = dimsIdx;
945 for (
int64_t dimsIdx = 0; dimsIdx < rank; dimsIdx++) {
946 if (dimsAndPosMapping.count(dimsIdx))
947 interchangeVector.push_back(dimsAndPosMapping[dimsIdx]);
949 return interchangeVector;
969 for (
auto [idx, val] : llvm::enumerate(interchangeVector))
970 vec[idx + offset] = elements[val + offset];
976static void generatePackOpScalarImplementationBody(PackOp packOp,
991 computeInterchangeFromDimPos(dimsToInnerBlock, packOp.getSourceRank());
992 interchangedIvs = interchange<Value>(interchangedIvs, interchangeVector,
993 packOp.getSourceRank());
994 if (!dimsToOuterBlock.empty()) {
996 computeInterchangeFromDimPos(dimsToOuterBlock, packOp.getSourceRank());
998 interchange<Value>(interchangedIvs, interchangeVector, 0);
1001 packOp.getDimAndTileMapping();
1003 size_t pointLoopsOffset = 0;
1004 int64_t sourceRank = packOp.getSourceRank();
1005 for (
auto dim : llvm::seq<int64_t>(0, sourceRank)) {
1006 if (dimAndTileMapping.contains(dim)) {
1011 builder, loc, i *
tile +
j,
1013 interchangedIvs[dim],
1014 interchangedIvs[pointLoopsOffset + packOp.getSourceRank()],
1015 dimAndTileMapping[dim]});
1016 sourceIndices.push_back(sourceIndex);
1019 sourceIndices.push_back(interchangedIvs[dim]);
1023 auto createLoad = [&]() ->
Value {
1024 return memref::LoadOp::create(
1025 builder, loc, packOp.getSource(),
1029 if (
auto paddingValue = packOp.getPaddingValue()) {
1032 for (
auto dim : llvm::seq<int64_t>(0, sourceRank)) {
1035 Value cond = arithBuilder.slt(
1039 scalar = scf::IfOp::create(
1042 scf::YieldOp::create(
b, l, createLoad());
1046 scf::YieldOp::create(
b, l, paddingValue);
1050 scalar = createLoad();
1053 memref::StoreOp::create(builder, loc, scalar, packOp.getDest(), ivs);
1057 :
public TilingInterface::ExternalModel<PackOpTiling, linalg::PackOp> {
1058 using Base = TilingInterface::ExternalModel<PackOpTiling, linalg::PackOp>;
1059 using Base::getTiledImplementation;
1061 SmallVector<utils::IteratorType> getLoopIteratorTypes(Operation *op)
const {
1065 auto packOp = cast<PackOp>(op);
1066 SmallVector<utils::IteratorType> iteratorTypes(
1067 packOp.getSourceRank(), utils::IteratorType::parallel);
1068 return iteratorTypes;
1071 SmallVector<Range> getIterationDomain(Operation *op, OpBuilder &
b)
const {
1072 return getPackUnPackIterationDomain<PackOp>(cast<PackOp>(op),
b);
1075 FailureOr<TilingResult>
1077 ArrayRef<OpFoldResult> offsets,
1078 ArrayRef<OpFoldResult> sizes)
const {
1079 auto packOp = cast<PackOp>(op);
1081 if (!packOp.hasPureTensorSemantics())
1084 Location loc = packOp.getLoc();
1088 int64_t inputRank = packOp.getSourceRank();
1089 SmallVector<OpFoldResult> origOffsets(offsets);
1090 SmallVector<OpFoldResult> origSizes(sizes);
1091 applyPermToRange(origOffsets, origSizes,
1095 packOp.getDimAndTileMapping();
1096 SmallVector<OpFoldResult> srcDimValues =
1098 SmallVector<OpFoldResult> inputIndices, inputSizes;
1099 for (
auto dim : llvm::seq<int64_t>(0, inputRank)) {
1100 using AV = affine::AffineValueExpr;
1101 affine::AffineBuilder ab(
b, loc);
1102 AffineExpr dim0, dim1, sym;
1105 if (dimAndTileMapping.count(dim)) {
1109 auto avOffset = AV(dim0).bind(origOffsets[dim]);
1110 auto avSize = AV(dim0).bind(origSizes[dim]);
1111 auto avTileSize = AV(sym).bind(dimAndTileMapping[dim]);
1112 inputIndices.push_back(ab.mul(avOffset, avTileSize));
1113 inputSizes.push_back(ab.mul(avSize, avTileSize));
1115 inputIndices.push_back(origOffsets[dim]);
1116 inputSizes.push_back(origSizes[dim]);
1120 if (packOp.getPaddingValue()) {
1121 OpFoldResult dimSize = srcDimValues[dim];
1122 auto avDimSize = AV(dim0).bind(dimSize);
1123 auto avInputIdx = AV(dim1).bind(inputIndices.back());
1125 ab.min({inputSizes.back(), ab.sub(avDimSize, avInputIdx)});
1129 auto oneAttr =
b.getI64IntegerAttr(1);
1130 SmallVector<OpFoldResult> strides(inputRank, oneAttr);
1132 SmallVector<Value> tiledOperands;
1133 auto sourceSlice = tensor::ExtractSliceOp::create(
1134 b, loc, packOp.getSource(), inputIndices, inputSizes, strides);
1135 tiledOperands.push_back(sourceSlice);
1137 SmallVector<OpFoldResult> outputOffsets, outputSizes;
1142 strides.append(packOp.getDestRank() - inputRank, oneAttr);
1143 auto outSlice = tensor::ExtractSliceOp::create(
1144 b, loc, packOp.getDest(), outputOffsets, outputSizes, strides);
1145 tiledOperands.push_back(outSlice);
1147 if (
auto val = packOp.getPaddingValue())
1148 tiledOperands.push_back(val);
1149 for (
auto tile : packOp.getInnerTiles())
1150 tiledOperands.push_back(
tile);
1152 Operation *tiledPackOp = PackOp::create(
1155 return TilingResult{
1157 SmallVector<Value>(tiledPackOp->
getResults()),
1158 llvm::to_vector(ArrayRef<Operation *>{sourceSlice, outSlice})};
1163 ArrayRef<OpFoldResult> offsets,
1164 ArrayRef<OpFoldResult> sizes,
1165 SmallVector<OpFoldResult> &resultOffsets,
1166 SmallVector<OpFoldResult> &resultSizes)
const {
1171 auto packOp = cast<PackOp>(op);
1172 int64_t inputRank = packOp.getSourceRank();
1173 int64_t outputRank = packOp.getDestRank();
1174 auto zeroAttr =
b.getI64IntegerAttr(0);
1175 resultOffsets.assign(offsets.begin(), offsets.end());
1176 resultOffsets.append(outputRank - inputRank, zeroAttr);
1180 resultSizes.assign(sizes.begin(), sizes.end());
1181 for (
auto dataTileDim : llvm::seq<unsigned>(inputRank, outputRank))
1182 resultSizes.push_back(outputShape[0][dataTileDim]);
1187 FailureOr<TilingResult>
1188 generateResultTileValue(Operation *op, OpBuilder &
b,
unsigned resultNumber,
1189 ArrayRef<OpFoldResult> offsets,
1190 ArrayRef<OpFoldResult> sizes)
const {
1191 return generateResultTileValue(op,
b, resultNumber, offsets, sizes,
1195 FailureOr<TilingResult> generateResultTileValue(
1196 Operation *op, OpBuilder &
b,
unsigned resultNumber,
1197 ArrayRef<OpFoldResult> offsets, ArrayRef<OpFoldResult> sizes,
1198 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1199 auto packOp = cast<PackOp>(op);
1200 int64_t numTiles = packOp.getInnerDimsPos().size();
1205 for (
auto offset : offsets.take_back(numTiles))
1212 ArrayRef<int64_t> innerDimsPos = packOp.getInnerDimsPos();
1213 SmallVector<OpFoldResult> mixedTiles = packOp.getMixedTiles();
1214 ArrayRef<OpFoldResult> innerSizes = sizes.take_back(numTiles);
1215 for (
auto [i, pos] : llvm::enumerate(innerDimsPos)) {
1217 pos < static_cast<int64_t>(innerTileAlignments.size())
1218 ? innerTileAlignments[pos]
1219 : InnerTileAlignment::Unknown;
1220 if (alignment != InnerTileAlignment::Equal &&
1226 op,
b, offsets.drop_back(numTiles), sizes.drop_back(numTiles));
1227 if (
failed(tilingResult))
1229 return tilingResult.value();
1232 LogicalResult generateScalarImplementation(Operation *op, OpBuilder &builder,
1235 auto packOp = cast<PackOp>(op);
1236 assert(packOp.hasPureBufferSemantics() &&
1237 "expected operation to have buffer semantics");
1238 OpBuilder::InsertionGuard g(builder);
1241 SmallVector<Value> ivVec(ivs);
1244 SmallVector<OpFoldResult> outputShape;
1245 Value dest = packOp.getDest();
1246 for (
auto dim : llvm::seq<int64_t>(0, packOp.getDestRank()))
1255 for (
auto dataTileDim : llvm::seq<unsigned>(packOp.getSourceRank(),
1256 packOp.getDestRank() - 1)) {
1258 outputShape[dataTileDim]);
1259 scf::ForOp loop = scf::ForOp::create(builder, loc, zero, ub, one);
1261 ivVec.push_back(loop.getInductionVar());
1268 [&](OpBuilder &bodyBuilder, Location bodyLoc, Value iv,
1270 ivVec.push_back(iv);
1271 generatePackOpScalarImplementationBody(packOp, bodyBuilder, bodyLoc,
1273 scf::YieldOp::create(bodyBuilder, bodyLoc);
1278 LogicalResult getIterationDomainTileFromOperandTiles(
1279 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1280 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1281 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1282 SmallVectorImpl<OpFoldResult> &resultOffsets,
1283 SmallVectorImpl<OpFoldResult> &resultSizes)
const {
1284 return getIterationDomainTileFromOperandTiles(
1285 op,
b, operandNumbers, allOffsets, allSizes, resultOffsets, resultSizes,
1292 LogicalResult getIterationDomainTileFromOperandTiles(
1293 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1294 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1295 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1296 SmallVectorImpl<OpFoldResult> &resultOffsets,
1297 SmallVectorImpl<OpFoldResult> &resultSizes,
1298 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1299 if (operandNumbers.size() != 1 || operandNumbers[0] != 0) {
1301 { llvm::dbgs() <<
"unsupported operands for consumer fusion"; });
1305 ArrayRef<OpFoldResult> offsets(allOffsets[0]);
1306 ArrayRef<OpFoldResult> sizes(allSizes[0]);
1307 auto packOp = cast<PackOp>(op);
1308 Location loc = packOp.getLoc();
1309 SmallVector<OpFoldResult> outerDimOffsets, outerDimSizes;
1311 packOp.getDimAndTileMapping();
1312 SmallVector<int64_t> outerShapeWithoutTranspose(
1313 packOp.getDestType().getShape().take_front(packOp.getSourceRank()));
1314 if (!packOp.getOuterDimsPerm().empty()) {
1316 outerShapeWithoutTranspose,
1319 for (
auto dim : llvm::seq<int64_t>(packOp.getSourceRank())) {
1320 if (dimAndTileMapping.count(dim)) {
1321 FailureOr<int64_t> cstTileSize =
1323 presburger::BoundType::UB, sizes[dim],
1325 ValueBoundsOptions{
true});
1326 std::optional<int64_t> cstInnerSize =
1333 dim < static_cast<int64_t>(innerTileAlignments.size())
1334 ? innerTileAlignments[dim]
1335 : InnerTileAlignment::Unknown;
1348 int64_t srcDimSize = packOp.getSourceType().getDimSize(dim);
1349 int64_t destDimSize = outerShapeWithoutTranspose[dim];
1350 bool isTiled = innerTileAlignment != InnerTileAlignment::Unknown ||
1352 ShapedType::isDynamic(srcDimSize) ||
1353 cstTileSize.value() < srcDimSize;
1355 outerDimOffsets.push_back(offsets[dim]);
1356 if (ShapedType::isStatic(destDimSize)) {
1357 outerDimSizes.push_back(
b.getIndexAttr(destDimSize));
1359 outerDimSizes.push_back(
1360 b.createOrFold<tensor::DimOp>(loc, packOp.getDest(), dim));
1387 bool assumeInnerTileSizesMatchTiles =
1388 innerTileAlignment == InnerTileAlignment::Equal;
1389 bool staticallyDecidable =
1390 !
failed(cstTileSize) && cstInnerSize.has_value();
1391 if (innerTileAlignment == InnerTileAlignment::Unknown) {
1392 if (!staticallyDecidable || *cstTileSize % *cstInnerSize != 0)
1394 }
else if (staticallyDecidable) {
1395 assert(*cstTileSize % *cstInnerSize == 0 &&
1396 "InnerTileAlignment hint contradicts statically known tile "
1398 assert((innerTileAlignment != InnerTileAlignment::Equal ||
1399 *cstTileSize == *cstInnerSize) &&
1400 "InnerTileAlignment::Equal contradicts statically known tile "
1404 using AV = affine::AffineValueExpr;
1405 affine::AffineBuilder ab(
b, loc);
1406 AffineExpr dim0, sym;
1409 auto avOffset = AV(dim0).bind(offsets[dim]);
1410 auto avSize = AV(dim0).bind(sizes[dim]);
1411 auto avTileSize = AV(sym).bind(dimAndTileMapping[dim]);
1412 outerDimOffsets.push_back(ab.floor(avOffset, avTileSize));
1415 outerDimSizes.push_back(assumeInnerTileSizesMatchTiles
1417 : ab.ceil(avSize, avTileSize));
1419 outerDimOffsets.push_back(offsets[dim]);
1420 outerDimSizes.push_back(sizes[dim]);
1423 applyPermToRange(outerDimOffsets, outerDimSizes, packOp.getOuterDimsPerm());
1424 resultOffsets = outerDimOffsets;
1425 resultSizes = outerDimSizes;
1429 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
1430 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1431 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1432 ArrayRef<SmallVector<OpFoldResult>> allSizes)
const {
1433 return getTiledImplementationFromOperandTiles(op,
b, operandNumbers,
1434 allOffsets, allSizes,
1439 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
1440 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1441 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1442 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1443 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1444 if (operandNumbers.size() != 1 || operandNumbers[0] != 0) {
1445 LLVM_DEBUG({ llvm::dbgs() <<
"unhandled operands for consumer fusion"; });
1449 ArrayRef<OpFoldResult> offsets(allOffsets[0]);
1450 ArrayRef<OpFoldResult> sizes(allSizes[0]);
1452 auto packOp = cast<PackOp>(op);
1454 if (!packOp.hasPureTensorSemantics())
1457 Location loc = packOp.getLoc();
1459 int64_t inputRank = packOp.getSourceRank();
1460 auto oneAttr =
b.getI64IntegerAttr(1);
1461 SmallVector<OpFoldResult> strides(inputRank, oneAttr);
1463 SmallVector<Value> tiledOperands;
1464 auto sourceSlice = tensor::ExtractSliceOp::create(
1465 b, loc, packOp.getSource(), offsets, sizes, strides);
1466 tiledOperands.push_back(sourceSlice);
1468 SmallVector<OpFoldResult> outerDimOffsets, outerDimSizes;
1469 if (
failed(getIterationDomainTileFromOperandTiles(
1470 op,
b, operandNumbers, allOffsets, allSizes, outerDimOffsets,
1471 outerDimSizes, innerTileAlignments)))
1474 SmallVector<OpFoldResult> outputOffsets, outputSizes;
1476 outputOffsets, outputSizes)))
1479 strides.append(packOp.getDestRank() - inputRank, oneAttr);
1480 auto outSlice = tensor::ExtractSliceOp::create(
1481 b, loc, packOp.getDest(), outputOffsets, outputSizes, strides);
1482 tiledOperands.push_back(outSlice);
1484 if (
auto val = packOp.getPaddingValue())
1485 tiledOperands.push_back(val);
1486 for (
auto tile : packOp.getInnerTiles())
1487 tiledOperands.push_back(
tile);
1489 Operation *tiledPackOp = PackOp::create(
1492 return TilingResult{
1494 SmallVector<Value>(tiledPackOp->
getResults()),
1495 llvm::to_vector(ArrayRef<Operation *>{sourceSlice, outSlice})};
1499struct UnpackTileDimInfo {
1500 bool isAlignedToInnerTileSize;
1501 OpFoldResult sourceOffset;
1502 OpFoldResult sourceSize;
1503 OpFoldResult resultOffset;
1504 OpFoldResult destExpandedSize;
1510static UnpackTileDimInfo
1514 UnpackTileDimInfo info;
1518 unpackOp.getDimAndTileMapping();
1520 if (!dimAndTileMapping.count(tileDim)) {
1521 info.isAlignedToInnerTileSize =
true;
1522 info.sourceOffset = tileOffset;
1523 info.sourceSize = tileSize;
1524 info.resultOffset = zeroAttr;
1525 info.destExpandedSize = tileSize;
1536 OpFoldResult innerTileSize = dimAndTileMapping[tileDim];
1538 info.isAlignedToInnerTileSize =
false;
1551 bool assumeInnerTileSizesMatchTiles =
1553 bool staticallyDecidable = !
failed(cstSize) && cstInnerSize.has_value();
1555 info.isAlignedToInnerTileSize =
true;
1556 if (staticallyDecidable) {
1557 assert(*cstSize % *cstInnerSize == 0 &&
1558 "InnerTileAlignment hint contradicts statically known tile sizes");
1560 *cstSize == *cstInnerSize) &&
1561 "InnerTileAlignment::Equal contradicts statically known tile "
1565 if (info.isAlignedToInnerTileSize || (!
failed(cstSize) && cstInnerSize)) {
1566 if (!info.isAlignedToInnerTileSize && *cstSize % *cstInnerSize == 0)
1567 info.isAlignedToInnerTileSize =
true;
1571 if (assumeInnerTileSizesMatchTiles ||
1572 (cstInnerSize && !
failed(cstSize) && *cstInnerSize == *cstSize)) {
1573 auto lhs = AV(dim0).bind(tileOffset);
1574 auto rhs = AV(dim1).bind(innerTileSize);
1575 info.sourceOffset = ab.floor(
lhs,
rhs);
1576 info.sourceSize = oneAttr;
1577 info.resultOffset = zeroAttr;
1578 info.destExpandedSize = tileSize;
1583 if (info.isAlignedToInnerTileSize) {
1585 ab.floor(AV(dim0).bind(tileOffset), AV(dim1).bind(innerTileSize));
1586 info.resultOffset = zeroAttr;
1587 info.destExpandedSize = tileSize;
1596 ab.ceil(AV(dim0).bind(tileSize), AV(dim1).bind(innerTileSize));
1600 affine::DivModValue firstCoord = affine::getDivMod(
1604 ab.add(AV(dim0).bind(tileOffset), AV(dim1).bind(tileSize));
1605 affine::DivModValue lastCoord = affine::getDivMod(
1609 ab.sub(AV(dim0).bind(tileExclusiveBound), AV(dim1).bind(oneAttr))),
1612 OpFoldResult lengthMinusOne = ab.sub(AV(dim0).bind(lastCoord.quotient),
1613 AV(dim1).bind(firstCoord.quotient));
1615 ab.add(AV(dim0).bind(lengthMinusOne), AV(dim1).bind(oneAttr));
1616 info.sourceOffset = firstCoord.quotient;
1617 info.resultOffset = firstCoord.remainder;
1620 info.destExpandedSize =
b.createOrFold<arith::MulIOp>(
1626struct UnPackOpTiling
1627 :
public TilingInterface::ExternalModel<UnPackOpTiling, linalg::UnPackOp> {
1628 using Base = TilingInterface::ExternalModel<UnPackOpTiling, linalg::UnPackOp>;
1629 using Base::getIterationDomainTileFromOperandTiles;
1631 SmallVector<utils::IteratorType> getLoopIteratorTypes(Operation *op)
const {
1632 auto unpackOp = cast<UnPackOp>(op);
1633 SmallVector<utils::IteratorType> iteratorTypes(
1634 unpackOp.getDestRank(), utils::IteratorType::parallel);
1635 return iteratorTypes;
1638 SmallVector<Range> getIterationDomain(Operation *op, OpBuilder &
b)
const {
1639 return getPackUnPackIterationDomain<UnPackOp>(cast<UnPackOp>(op),
b);
1656 FailureOr<TilingResult>
1658 ArrayRef<OpFoldResult> offsets,
1659 ArrayRef<OpFoldResult> sizes)
const {
1665 Operation *op, OpBuilder &
b, ArrayRef<OpFoldResult> offsets,
1666 ArrayRef<OpFoldResult> sizes,
1667 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1668 auto unpackOp = cast<UnPackOp>(op);
1670 if (!unpackOp.hasPureTensorSemantics())
1673 int64_t srcRank = unpackOp.getSourceRank();
1674 int64_t destRank = unpackOp.getDestRank();
1675 int64_t numInnerTiles = srcRank - destRank;
1676 Location loc = unpackOp.getLoc();
1681 bool isPerfectTilingCase =
true;
1682 Attribute oneAttr =
b.getIndexAttr(1);
1683 SmallVector<OpFoldResult> sliceSrcStrides(destRank, oneAttr);
1684 SmallVector<OpFoldResult> sliceSrcIndices, sliceSrcSizes;
1685 SmallVector<OpFoldResult> destExpandedSizes, resultOffsetsFromDest;
1686 for (
auto dim : llvm::seq<int64_t>(0, destRank)) {
1687 UnpackTileDimInfo info = getUnpackTileDimInfo(
1688 b, unpackOp, dim, offsets[dim], sizes[dim],
1689 dim <
static_cast<int64_t
>(innerTileAlignments.size())
1690 ? innerTileAlignments[dim]
1691 : InnerTileAlignment::Unknown);
1692 if (!info.isAlignedToInnerTileSize)
1693 isPerfectTilingCase =
false;
1694 sliceSrcIndices.push_back(info.sourceOffset);
1695 sliceSrcSizes.push_back(info.sourceSize);
1696 destExpandedSizes.push_back(info.destExpandedSize);
1697 resultOffsetsFromDest.push_back(info.resultOffset);
1702 applyPermToRange(sliceSrcIndices, sliceSrcSizes,
1703 unpackOp.getOuterDimsPerm());
1704 Attribute zeroAttr =
b.getIndexAttr(0);
1705 sliceSrcIndices.append(numInnerTiles, zeroAttr);
1706 sliceSrcSizes.append(unpackOp.getMixedTiles());
1707 sliceSrcStrides.append(numInnerTiles, oneAttr);
1708 SmallVector<Operation *> generatedSlices;
1709 tensor::ExtractSliceOp sliceSource = tensor::ExtractSliceOp::create(
1710 b, loc, unpackOp.getSource(), sliceSrcIndices, sliceSrcSizes,
1712 generatedSlices.push_back(sliceSource);
1714 SmallVector<OpFoldResult> destStrides(destRank, oneAttr);
1716 if (isPerfectTilingCase) {
1717 auto destSliceOp = tensor::ExtractSliceOp::create(
1718 b, loc, unpackOp.getDest(), offsets, sizes, destStrides);
1719 sliceDest = destSliceOp;
1720 generatedSlices.push_back(destSliceOp);
1722 sliceDest = tensor::EmptyOp::create(
1723 b, loc, destExpandedSizes, unpackOp.getDestType().getElementType());
1726 SmallVector<Value> tiledOperands = {sliceSource.getResult(), sliceDest};
1727 for (
auto tile : unpackOp.getInnerTiles())
1728 tiledOperands.push_back(
tile);
1730 Operation *tiledUnpackOp = UnPackOp::create(
1733 if (isPerfectTilingCase)
1734 return TilingResult{{tiledUnpackOp},
1735 SmallVector<Value>(tiledUnpackOp->
getResults()),
1738 auto extractSlice = tensor::ExtractSliceOp::create(
1739 b, loc, tiledUnpackOp->
getResult(0), resultOffsetsFromDest, sizes,
1741 return TilingResult{
1742 {tiledUnpackOp}, {extractSlice.getResult()}, generatedSlices};
1747 ArrayRef<OpFoldResult> offsets,
1748 ArrayRef<OpFoldResult> sizes,
1749 SmallVector<OpFoldResult> &resultOffsets,
1750 SmallVector<OpFoldResult> &resultSizes)
const {
1751 resultOffsets = llvm::to_vector(offsets);
1752 resultSizes = llvm::to_vector(sizes);
1756 FailureOr<TilingResult>
1757 generateResultTileValue(Operation *op, OpBuilder &
b,
unsigned resultNumber,
1758 ArrayRef<OpFoldResult> offsets,
1759 ArrayRef<OpFoldResult> sizes)
const {
1760 return generateResultTileValue(op,
b, resultNumber, offsets, sizes,
1764 FailureOr<TilingResult> generateResultTileValue(
1765 Operation *op, OpBuilder &
b,
unsigned resultNumber,
1766 ArrayRef<OpFoldResult> offsets, ArrayRef<OpFoldResult> sizes,
1767 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1768 FailureOr<TilingResult> tilingResult =
1770 if (
failed(tilingResult))
1772 return tilingResult.value();
1775 LogicalResult generateScalarImplementation(Operation *op, OpBuilder &builder,
1778 auto unpackOp = cast<UnPackOp>(op);
1779 assert(unpackOp.hasPureBufferSemantics() &&
1780 "expected operation to have buffer semantics");
1781 assert(ivs.size() == unpackOp.getDestRank() &&
1782 "number of ivs must match the rank of the output tensor");
1783 OpBuilder::InsertionGuard g(builder);
1786 unpackOp.getDimAndTileMapping();
1788 SmallVector<Value> inputIvs;
1790 SmallVector<Value> inputIvsPointLoops;
1791 inputIvs.reserve(unpackOp.getDestRank());
1792 inputIvsPointLoops.reserve(dimAndTileMapping.size());
1793 for (
auto dim : llvm::seq<int64_t>(0, unpackOp.getDestRank())) {
1794 if (dimAndTileMapping.count(dim)) {
1795 affine::DivModValue divMod =
1796 affine::getDivMod(builder, loc, ivs[dim],
1798 builder, loc, dimAndTileMapping[dim]));
1799 inputIvsPointLoops.push_back(divMod.remainder);
1800 inputIvs.push_back(divMod.quotient);
1802 inputIvs.push_back(ivs[dim]);
1808 assert(inputIvsPointLoops.size() + inputIvs.size() ==
1809 unpackOp.getSourceRank() &&
1810 "expect same number of induction variables equals to input rank");
1812 ArrayRef<int64_t> innerDims = unpackOp.getInnerDimsPos();
1813 SmallVector<int64_t> interchangeVector =
1814 computeInterchangeFromDimPos(innerDims, unpackOp.getDestRank());
1815 SmallVector<Value> interchangedInputIvsPointLoops = inputIvsPointLoops;
1816 interchangedInputIvsPointLoops = interchange<Value>(
1817 interchangedInputIvsPointLoops, interchangeVector, 0);
1820 ArrayRef<int64_t> outerDims = unpackOp.getOuterDimsPerm();
1821 if (!outerDims.empty())
1822 inputIvs = interchange<Value>(inputIvs, outerDims, 0);
1824 llvm::append_range(inputIvs, interchangedInputIvsPointLoops);
1826 memref::LoadOp::create(builder, loc, unpackOp.getSource(), inputIvs);
1827 memref::StoreOp::create(builder, loc, scalar, unpackOp.getDest(), ivs);
1833 LogicalResult getIterationDomainTileFromOperandTiles(
1834 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1835 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1836 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1837 SmallVectorImpl<OpFoldResult> &resultOffsets,
1838 SmallVectorImpl<OpFoldResult> &resultSizes)
const {
1839 if (operandNumbers.size() != 1) {
1840 LLVM_DEBUG({ llvm::dbgs() <<
"unable to handle multiple operands"; });
1843 auto unPackOp = cast<UnPackOp>(op);
1844 unsigned operandNumber = operandNumbers[0];
1845 ArrayRef<OpFoldResult> offsets(allOffsets[0]);
1846 ArrayRef<OpFoldResult> sizes(allSizes[0]);
1849 if (operandNumber == unPackOp.getDestMutable().getOperandNumber()) {
1850 resultOffsets = llvm::to_vector(offsets);
1851 resultSizes = llvm::to_vector(sizes);
1854 Location loc = unPackOp.getLoc();
1856 int64_t numTiles = unPackOp.getInnerDimsPos().size();
1857 auto destOffsets = offsets.drop_back(numTiles);
1858 auto destSizes = sizes.drop_back(numTiles);
1861 int64_t outputRank = unPackOp.getDestRank();
1865 SmallVector<OpFoldResult> outputMixedSizes = reifiedReturnShapes.front();
1866 SmallVector<OpFoldResult> origOffsets(destOffsets);
1867 SmallVector<OpFoldResult> origSizes(destSizes);
1868 applyPermToRange(origOffsets, origSizes,
1872 unPackOp.getDimAndTileMapping();
1874 for (
auto dim : llvm::seq<int64_t>(0, outputRank)) {
1875 using AV = affine::AffineValueExpr;
1876 affine::AffineBuilder ab(
b, loc);
1877 AffineExpr dim0, dim1, sym0;
1880 if (dimAndTileMapping.count(dim)) {
1884 auto avOffset = AV(dim0).bind(origOffsets[dim]);
1885 auto avSize = AV(dim0).bind(origSizes[dim]);
1886 auto avTileSize = AV(sym0).bind(dimAndTileMapping[dim]);
1887 auto avResultSize = AV(dim0).bind(outputMixedSizes[dim]);
1888 resultOffsets.push_back(ab.mul(avOffset, avTileSize));
1889 auto avResultOffset = AV(dim1).bind(resultOffsets.back());
1890 resultSizes.push_back(ab.min({ab.mul(avSize, avTileSize),
1891 ab.sub(avResultSize, avResultOffset)}));
1893 resultOffsets.push_back(origOffsets[dim]);
1894 resultSizes.push_back(origSizes[dim]);
1900 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
1901 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1902 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1903 ArrayRef<SmallVector<OpFoldResult>> allSizes)
const {
1904 return getTiledImplementationFromOperandTiles(op,
b, operandNumbers,
1905 allOffsets, allSizes,
1910 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
1911 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1912 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1913 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1914 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1915 if (operandNumbers.size() != 1 || operandNumbers[0] != 0) {
1916 LLVM_DEBUG({ llvm::dbgs() <<
"unhandled operands for consumer fusion"; });
1919 auto unPackOp = cast<UnPackOp>(op);
1921 if (!unPackOp.hasPureTensorSemantics())
1924 ArrayRef<OpFoldResult> offsets(allOffsets[0]);
1925 ArrayRef<OpFoldResult> sizes(allSizes[0]);
1931 int64_t numTiles = unPackOp.getInnerDimsPos().size();
1932 ArrayRef<int64_t> innerDimsPos = unPackOp.getInnerDimsPos();
1933 SmallVector<OpFoldResult> mixedTiles = unPackOp.getMixedTiles();
1934 ArrayRef<OpFoldResult> innerSizes = sizes.take_back(numTiles);
1935 for (int64_t i = 0; i < numTiles; ++i) {
1938 int64_t destDim = innerDimsPos[i];
1940 destDim < static_cast<int64_t>(innerTileAlignments.size()) &&
1941 innerTileAlignments[destDim] == InnerTileAlignment::Equal;
1951 "InnerTileAlignment::Equal contradicts statically known tile "
1960 Location loc = unPackOp.getLoc();
1964 SmallVector<OpFoldResult> outputOffsets, outputSizes;
1965 if (
failed(getIterationDomainTileFromOperandTiles(
1966 op,
b, operandNumbers, allOffsets, allSizes, outputOffsets,
1970 auto oneAttr =
b.getI64IntegerAttr(1);
1971 int64_t outputRank = unPackOp.getDestRank();
1972 SmallVector<OpFoldResult> strides(outputRank, oneAttr);
1974 SmallVector<Value> tiledOperands;
1976 auto extractDestSlice = tensor::ExtractSliceOp::create(
1977 b, loc, unPackOp.getDest(), outputOffsets, outputSizes, strides);
1978 tiledOperands.push_back(extractDestSlice);
1980 strides.append(unPackOp.getSourceRank() - outputRank, oneAttr);
1982 auto extractSourceSlice = tensor::ExtractSliceOp::create(
1983 b, loc, unPackOp.getSource(), offsets, sizes, strides);
1984 tiledOperands.insert(tiledOperands.begin(), extractSourceSlice);
1985 for (
auto tile : unPackOp.getInnerTiles())
1986 tiledOperands.push_back(
tile);
1989 Operation *tiledUnPackOp =
1990 UnPackOp::create(
b, loc,
TypeRange{extractDestSlice.getType()},
1993 return TilingResult{{tiledUnPackOp},
1994 SmallVector<Value>(tiledUnPackOp->
getResults()),
1995 llvm::to_vector(ArrayRef<Operation *>{
1996 extractSourceSlice, extractDestSlice})};
2002template <
typename OpType>
2004 OpType::template attachInterface<LinalgOpTilingInterface<OpType>>(*ctx);
2005 OpType::template attachInterface<LinalgOpPartialReductionInterface<OpType>>(
2010template <
typename... OpTypes>
2021 linalg::PackOp::attachInterface<PackOpTiling>(*ctx);
2022 linalg::UnPackOp::attachInterface<UnPackOpTiling>(*ctx);
2024#include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc"
2032 linalg::PackOp::attachInterface<PackOpTiling>(*ctx);
2033 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 LogicalResult validateTilingSemiAffineMaps(LinalgOp linalgOp, ArrayRef< OpFoldResult > sizes)
Verify that tiling can be applied in presence of semi-affine maps.
static bool isInBounds(TransferOp op, int64_t resultIdx, int64_t indicesIdx)
Base type for affine expression.
RetT walk(FnT &&callback) const
Walk all of the AffineExpr's in this expression in postorder.
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.
A utility result that is used to signal how to proceed with an ongoing walk:
static WalkResult advance()
bool wasInterrupted() const
Returns true if the walk was interrupted.
static WalkResult interrupt()
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...
@ CeilDiv
RHS of ceildiv is always a constant or a symbolic expression.
@ Mod
RHS of mod is always a constant or a symbolic expression with a positive value.
@ FloorDiv
RHS of floordiv is always a constant or a symbolic expression.
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.