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 ")
208struct LinalgOpTilingInterfaceImpl {
210 SmallVector<Range> getIterationDomain(Operation *op, OpBuilder &
b)
const {
211 OpBuilder::InsertionGuard g(
b);
212 b.setInsertionPoint(op);
213 Location loc = op->
getLoc();
214 LinalgOp linalgOp = cast<LinalgOp>(op);
215 SmallVector<OpFoldResult> allShapesSizes =
216 linalgOp.createFlatListOfOperandDims(
b, loc);
217 AffineMap map = linalgOp.getShapesToLoopsMap();
219 return llvm::map_to_vector(map.
getResults(), [&](AffineExpr loopExpr) {
220 OpFoldResult ofr = affine::makeComposedFoldedAffineApply(b, loc, loopExpr,
222 return Range{b.getIndexAttr(0), ofr, b.getIndexAttr(1)};
227 FailureOr<TilingResult>
234 LinalgOp linalgOp = cast<LinalgOp>(op);
245 b, loc, linalgOp, valuesToTile, offsets, sizes, {},
true);
247 llvm::make_filter_range(
249 [](
Value v) ->
bool {
250 return isa_and_nonnull<tensor::ExtractSliceOp, memref::SubViewOp>(
258 Operation *tiledOp =
clone(
b, linalgOp, resultTensorTypes, tiledOperands);
269 getMappedOffsetAndSize(LinalgOp linalgOp,
OpBuilder &
b,
277 for (
auto [indexingMap, offsets, sizes] :
278 llvm::zip_equal(indexingMaps, allOffsets, allSizes)) {
279 for (
auto [resultExpr, offset, size] :
280 llvm::zip_equal(indexingMap.getResults(), offsets, sizes)) {
281 auto dimExpr = dyn_cast<AffineDimExpr>(resultExpr);
284 unsigned position = dimExpr.getPosition();
285 auto it = mappedOffsets.find(position);
286 if (it != mappedOffsets.end()) {
289 if (seenOffset != offset || seenSize != size) {
291 llvm::dbgs() <<
"inconsistent iteration space mapping from "
292 "offsets/sizes of operands/results";
297 mappedOffsets[position] = offset;
298 mappedSizes[position] = size;
306 cast<TilingInterface>(linalgOp.getOperation()).getIterationDomain(
b);
307 mappedOffsetsVec.resize(iterationDomain.size());
308 mappedSizesVec.resize(iterationDomain.size());
309 for (
auto [
index, domain] : llvm::enumerate(iterationDomain)) {
310 auto it = mappedOffsets.find(
index);
311 if (it != mappedOffsets.end()) {
312 mappedOffsetsVec[
index] = it->second;
313 mappedSizesVec[
index] = mappedSizes.lookup(
index);
316 mappedOffsetsVec[
index] = domain.offset;
317 mappedSizesVec[
index] = domain.size;
324 LogicalResult getIterationDomainTileFromOperandTiles(
330 auto linalgOp = cast<LinalgOp>(op);
333 llvm::map_to_vector(operandNumbers, [&](
unsigned operandNumber) {
334 OpOperand &opOperand = linalgOp->getOpOperand(operandNumber);
335 return linalgOp.getMatchingIndexingMap(&opOperand);
337 if (
failed(getMappedOffsetAndSize(linalgOp,
b, indexingMaps, allOffsets,
338 allSizes, iterDomainOffsets,
354 LinalgOp linalgOp = cast<LinalgOp>(op);
363 OpOperand *outOperand = linalgOp.getDpsInitOperand(resultNumber);
365 b, loc, outOperand->get(), sizes,
366 linalgOp.getMatchingIndexingMap(outOperand), offsets,
367 {}, subShapeSizes,
true);
368 resultOffsets = sliceParams.
offsets;
369 resultSizes = sliceParams.
sizes;
373 LogicalResult getIterationDomainTileFromResultTile(
378 auto linalgOp = cast<LinalgOp>(op);
385 linalgOp.getIndexingMapMatchingResult(op->
getResult(resultNumber));
388 "unhandled tiled implementation generation when result is not "
389 "accessed using a permuted projection");
395 getMappedOffsetAndSize(linalgOp,
b, indexingMap, {allOffsets},
396 {allSizes}, iterDomainOffsets, iterDomainSizes);
398 assert(succeeded(status) &&
"unexpected error in offset calculation");
402 FailureOr<TilingResult>
407 if (
failed(getIterationDomainTileFromResultTile(
408 op,
b, resultNumber, offsets, sizes, mappedOffsets, mappedSizes))) {
411 auto tilingInterfaceOp = cast<TilingInterface>(op);
412 FailureOr<TilingResult> tilingResult =
413 tilingInterfaceOp.getTiledImplementation(
b, mappedOffsets, mappedSizes);
418 if (tilingResult->tiledOps.size() != 1)
419 return op->
emitOpError(
"failed to generate tiled implementation");
422 tilingResult->tiledOps,
424 tilingResult->generatedSlices};
429 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
434 if (
failed(getIterationDomainTileFromOperandTiles(
435 op,
b, operandNumbers, allOffsets, allSizes, mappedOffsets,
445 auto linalgOp = cast<LinalgOp>(op);
446 if (!linalgOp.hasPureBufferSemantics())
447 return op->
emitOpError(
"expected operation to have buffer semantics");
450 indexedValues.reserve(linalgOp->getNumOperands());
454 for (
OpOperand &operand : linalgOp->getOpOperands()) {
455 if (!linalgOp.payloadUsesValueFromOperand(&operand)) {
456 indexedValues.push_back(
nullptr);
459 if (linalgOp.isScalar(&operand)) {
460 indexedValues.push_back(operand.get());
464 builder, linalgOpLoc, linalgOp.getMatchingIndexingMap(&operand), ivs);
466 memref::LoadOp::create(builder, linalgOpLoc, operand.get(),
indices);
467 indexedValues.push_back(
load);
474 bool isOpFusableWithConsumerSlice(
Operation *op,
unsigned resultNumber,
481 bool isOpFusableWithProducerSlices(
486 auto linalgOp = cast<LinalgOp>(op);
488 llvm::map_to_vector(operandNumbers, [&](
unsigned operandNumber) {
489 OpOperand &opOperand = linalgOp->getOpOperand(operandNumber);
490 return linalgOp.getMatchingIndexingMap(&opOperand);
495 return succeeded(getMappedOffsetAndSize(linalgOp,
b, indexingMaps,
496 allOffsets, allSizes, mappedOffsets,
501template <
typename LinalgOpTy>
502struct LinalgOpTilingInterfaceModel
503 :
public TilingInterface::ExternalModel<
504 LinalgOpTilingInterfaceModel<LinalgOpTy>, LinalgOpTy>,
505 public LinalgOpTilingInterfaceImpl {
506 using ExternalModel =
507 TilingInterface::ExternalModel<LinalgOpTilingInterfaceModel<LinalgOpTy>,
510 using LinalgOpTilingInterfaceImpl::generateScalarImplementation;
511 using LinalgOpTilingInterfaceImpl::getIterationDomain;
512 using LinalgOpTilingInterfaceImpl::getIterationDomainTileFromResultTile;
513 using LinalgOpTilingInterfaceImpl::getResultTilePosition;
514 using LinalgOpTilingInterfaceImpl::isOpFusableWithConsumerSlice;
515 using LinalgOpTilingInterfaceImpl::isOpFusableWithProducerSlices;
518 SmallVector<utils::IteratorType> getLoopIteratorTypes(Operation *op)
const {
519 return cast<LinalgOpTy>(op).getIteratorTypesArray();
524 using ExternalModel::generateResultTileValue;
525 FailureOr<TilingResult>
526 generateResultTileValue(Operation *op, OpBuilder &
b,
unsigned resultNumber,
527 ArrayRef<OpFoldResult> offsets,
528 ArrayRef<OpFoldResult> sizes)
const {
529 return LinalgOpTilingInterfaceImpl::generateResultTileValue(
530 op,
b, resultNumber, offsets, sizes);
533 using ExternalModel::getIterationDomainTileFromOperandTiles;
534 LogicalResult getIterationDomainTileFromOperandTiles(
535 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
536 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
537 ArrayRef<SmallVector<OpFoldResult>> allSizes,
538 SmallVectorImpl<OpFoldResult> &iterDomainOffsets,
539 SmallVectorImpl<OpFoldResult> &iterDomainSizes)
const {
540 return LinalgOpTilingInterfaceImpl::getIterationDomainTileFromOperandTiles(
541 op,
b, operandNumbers, allOffsets, allSizes, iterDomainOffsets,
545 using ExternalModel::getTiledImplementation;
546 FailureOr<TilingResult>
548 ArrayRef<OpFoldResult> offsets,
549 ArrayRef<OpFoldResult> sizes)
const {
550 return LinalgOpTilingInterfaceImpl::getTiledImplementation(op,
b, offsets,
554 using ExternalModel::getTiledImplementationFromOperandTiles;
555 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
556 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
557 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
558 ArrayRef<SmallVector<OpFoldResult>> allSizes)
const {
559 return LinalgOpTilingInterfaceImpl::getTiledImplementationFromOperandTiles(
560 op,
b, operandNumbers, allOffsets, allSizes);
571 for (
auto [
index, reductionDim] : llvm::enumerate(reductionDims)) {
572 if (reductionDim == value) {
584getPartialResultAffineMaps(LinalgOp linalgOp,
586 auto partialReductionMaps = llvm::map_to_vector(
587 linalgOp.getDpsInitsMutable(), [&](
OpOperand &opOperand) {
588 AffineMap map = linalgOp.getMatchingIndexingMap(&opOperand);
589 for (auto redPos : reductionDims) {
591 map.insertResult(getAffineDimExpr(redPos, linalgOp.getContext()),
592 map.getNumResults());
596 return partialReductionMaps;
599struct InitSliceInfo {
600 SmallVector<int64_t> resultShape;
601 SmallVector<OpFoldResult> offsets;
602 SmallVector<OpFoldResult> sizes;
603 SmallVector<OpFoldResult> strides;
609static InitSliceInfo getInitSliceInfoForOuterReduction(
616 Attribute zero = IntegerAttr::get(IndexType::get(context), 0);
617 Attribute one = IntegerAttr::get(IndexType::get(context), 1);
619 for (
auto [resultIdx, dimExpr] :
620 llvm::enumerate(partialReductionMap.
getResults())) {
621 if (isa<AffineConstantExpr>(dimExpr)) {
624 initOffsets.push_back(zero);
625 initSizes.push_back(initOperandShape[resultIdx]);
628 unsigned dim = cast<AffineDimExpr>(dimExpr).getPosition();
629 if (reductionDims.contains(dim)) {
630 initOffsets.push_back(zero);
632 initOffsets.push_back(offsets[dim]);
634 initSizes.push_back(sizes[dim]);
638 return {resultShape, initOffsets, initSizes, initStrides};
644static InitSliceInfo getInitSliceInfoForOuterParallel(
651 Attribute zero = IntegerAttr::get(IndexType::get(context), 0);
652 Attribute one = IntegerAttr::get(IndexType::get(context), 1);
655 for (
auto [resultIdx, dimExpr] :
656 llvm::enumerate(partialReductionMap.
getResults())) {
657 if (isa<AffineConstantExpr>(dimExpr)) {
660 initOffsets.push_back(zero);
661 initSizes.push_back(initOperandShape[resultIdx]);
662 resultShape.push_back(initOperandShape[resultIdx]);
665 unsigned dim = cast<AffineDimExpr>(dimExpr).getPosition();
666 if (std::optional<unsigned> dimPos = getPositionIn(reductionDims, dim)) {
667 initOffsets.push_back(splitReductionIvs[dimPos.value()]);
668 initSizes.push_back(one);
670 initOffsets.push_back(offsets[dim]);
671 initSizes.push_back(sizes[dim]);
672 resultShape.push_back(sizes[dim]);
677 return {staticShapes, initOffsets, initSizes, initStrides};
682static InitSliceInfo getInitSliceInfo(
MLIRContext *context,
691 return getInitSliceInfoForOuterReduction(
692 context, offsets, sizes, reductionDims, splitReductionIvs,
693 partialReductionMap, initOperandShape);
696 "unexpected ReductionTilingStrategy");
697 return getInitSliceInfoForOuterParallel(
698 context, offsets, sizes, reductionDims, splitReductionIvs,
699 partialReductionMap, initOperandShape);
704struct LinalgOpPartialReductionInterfaceImpl {
705 FailureOr<SmallVector<Value>> generateInitialTensorForPartialReduction(
706 Operation *op, OpBuilder &
b, Location loc, ArrayRef<OpFoldResult> sizes,
708 auto linalgOp = cast<LinalgOp>(op);
710 OpBuilder::InsertionGuard guard(
b);
711 if (linalgOp.hasPureBufferSemantics())
712 return op->
emitOpError(
"expected operation to have tensor semantics");
714 SmallVector<AffineMap> partialResultMaps =
715 getPartialResultAffineMaps(linalgOp, reductionDims);
717 SmallVector<Value> inits;
718 for (
auto [initIdx,
result, partialMap] :
719 llvm::enumerate(linalgOp->getResults(), partialResultMaps)) {
720 SmallVector<Operation *, 4> combinerOps;
723 combinerOps.size() != 1)
724 return op->
emitOpError(
"Failed to anaysis the reduction operation.");
726 Operation *reductionOp = combinerOps[0];
727 std::optional<TypedAttr> identity = arith::getNeutralElement(reductionOp);
728 if (!identity.has_value())
730 "Failed to get an identity value for the reduction operation.");
733 SmallVector<OpFoldResult> partialResultShape;
734 Value initValue = linalgOp.getDpsInits()[initIdx];
735 SmallVector<OpFoldResult> initShape =
737 for (
auto [resultIdx, dimExpr] :
738 llvm::enumerate(partialMap.getResults())) {
739 if (isa<AffineConstantExpr>(dimExpr)) {
742 partialResultShape.push_back(initShape[resultIdx]);
745 auto dim = cast<AffineDimExpr>(dimExpr);
746 partialResultShape.push_back(sizes[dim.getPosition()]);
751 tensor::EmptyOp::create(
b, loc, partialResultShape, elType);
752 Value constantOp = arith::ConstantOp::create(
b, loc, *identity);
753 auto identityTensor =
754 linalg::FillOp::create(
b, loc, constantOp, emptyTensor);
755 inits.push_back(identityTensor.getResult(0));
761 FailureOr<TilingResult>
762 tileToPartialReduction(Operation *op, OpBuilder &
b, Location loc,
764 ValueRange init, ArrayRef<OpFoldResult> offsets,
765 ArrayRef<OpFoldResult> sizes,
767 ArrayRef<OpFoldResult> splitReductionIvs)
const {
768 OpBuilder::InsertionGuard guard(
b);
769 auto linalgOp = cast<LinalgOp>(op);
771 SmallVector<AffineMap> partialReductionMaps =
772 getPartialResultAffineMaps(linalgOp, reductionDims);
776 SmallVector<AffineMap> newInitMaps;
777 if (tilingStrategy ==
778 ReductionTilingStrategy::PartialReductionOuterReduction) {
779 newInitMaps = llvm::to_vector(partialReductionMaps);
781 newInitMaps = llvm::map_to_vector(
782 linalgOp.getDpsInitsMutable(), [&](OpOperand &opOperand) {
783 return linalgOp.getMatchingIndexingMap(&opOperand);
789 b, loc, linalgOp, linalgOp.getDpsInputs(), offsets, sizes, {},
true);
790 SmallVector<Operation *> generatedSlices = llvm::map_to_vector(
791 llvm::make_filter_range(
792 tiledInputs, [](Value v) ->
bool {
return v.
getDefiningOp(); }),
796 SmallVector<Value, 1> tiledInits;
797 for (
auto [partialReductionMap, valueToTile, initOperandValue] :
798 llvm::zip_equal(partialReductionMaps, init, linalgOp.getDpsInits())) {
801 SmallVector<OpFoldResult> initOperandShape =
803 InitSliceInfo sliceInfo = getInitSliceInfo(
804 b.getContext(), tilingStrategy, offsets, sizes, reductionDims,
805 splitReductionIvs, partialReductionMap, initOperandShape);
806 auto valueToTileType = cast<RankedTensorType>(valueToTile.getType());
808 sliceInfo.resultShape, valueToTileType.getElementType(),
809 valueToTileType.getEncoding());
810 auto sliceOp = tensor::ExtractSliceOp::create(
812 sliceInfo.sizes, sliceInfo.strides);
813 tiledInits.push_back(sliceOp.getResult());
814 generatedSlices.push_back(sliceOp);
818 SmallVector<AffineMap> newMaps = linalgOp.getIndexingMapsArray();
819 for (
auto [initOperand, newInitMap] :
820 llvm::zip_equal(linalgOp.getDpsInitsMutable(), newInitMaps)) {
821 int mapIdx = linalgOp.getIndexingMapIndex(&initOperand);
822 newMaps[mapIdx] = newInitMap;
826 SmallVector<utils::IteratorType> newIteratorTypes =
827 linalgOp.getIteratorTypesArray();
828 if (tilingStrategy ==
829 ReductionTilingStrategy::PartialReductionOuterReduction) {
830 for (
int dim : reductionDims)
831 newIteratorTypes[dim] = utils::IteratorType::parallel;
835 Operation *partialReductionOp;
836 auto resultTypes =
ValueRange(tiledInits).getTypes();
837 if (tilingStrategy ==
838 ReductionTilingStrategy::PartialReductionOuterReduction) {
839 auto genericOp = GenericOp::create(
b, loc, resultTypes, tiledInputs,
840 tiledInits, newMaps, newIteratorTypes);
843 genericOp.getRegion().begin(), mapping);
845 partialReductionOp = genericOp.getOperation();
847 SmallVector<Value> operands = std::move(tiledInputs);
848 llvm::append_range(operands, tiledInits);
849 partialReductionOp =
mlir::clone(
b, op, resultTypes, operands);
853 {partialReductionOp},
854 llvm::map_to_vector(partialReductionOp->
getResults(),
855 [](OpResult r) -> Value { return r; }),
859 FailureOr<MergeResult>
860 mergeReductions(Operation *op, OpBuilder &
b, Location loc,
863 auto linalgOp = cast<LinalgOp>(op);
864 SmallVector<AffineMap> partialReductionMaps =
865 getPartialResultAffineMaps(linalgOp, reductionDims);
868 SmallVector<Operation *> mergeOperations;
869 SmallVector<Value> replacements;
870 for (
auto [idx, init, partialResult, partialMap] : llvm::enumerate(
871 linalgOp.getDpsInits(), partialReduce, partialReductionMaps)) {
872 unsigned initIdx = idx;
877 SmallVector<int64_t> partialReductionDims;
878 for (
auto [resultNum, dimExpr] :
879 llvm::enumerate(partialMap.getResults())) {
880 if (isa<AffineConstantExpr>(dimExpr))
882 unsigned dim = cast<AffineDimExpr>(dimExpr).getPosition();
883 if (llvm::is_contained(reductionDims, dim)) {
884 partialReductionDims.push_back(resultNum);
888 auto reduction = linalg::ReduceOp::create(
889 b, loc, partialResult, init, partialReductionDims,
890 [&linalgOp, &initIdx](OpBuilder &
b, Location loc,
ValueRange inputs) {
892 SmallVector<Operation *, 4> combinerOps;
895 Operation *clonedReductionOp =
b.clone(*combinerOps[0]);
899 linalg::YieldOp::create(
b, loc, clonedReductionOp->
getResult(0));
902 mergeOperations.push_back(reduction);
903 replacements.push_back(reduction->getResult(0));
906 return MergeResult{mergeOperations, replacements};
909 LogicalResult getPartialResultTilePosition(
910 Operation *op, OpBuilder &
b,
unsigned resultNumber,
913 ArrayRef<OpFoldResult> splitReductionIvs,
914 SmallVector<OpFoldResult> &resultOffsets,
915 SmallVector<OpFoldResult> &resultSizes)
const {
916 auto linalgOp = cast<LinalgOp>(op);
917 SmallVector<AffineMap> partialReductionMaps =
918 getPartialResultAffineMaps(linalgOp, reductionDims);
921 Value initOperandValue = linalgOp.getDpsInits()[resultNumber];
922 Location loc = op->
getLoc();
923 SmallVector<OpFoldResult> initOperandShape =
925 InitSliceInfo sliceInfo =
926 getInitSliceInfo(
b.getContext(), tilingStrategy, offsets, sizes,
927 reductionDims, splitReductionIvs,
928 partialReductionMaps[resultNumber], initOperandShape);
929 std::swap(resultOffsets, sliceInfo.offsets);
930 std::swap(resultSizes, sliceInfo.sizes);
936template <
typename LinalgOpTy>
937struct LinalgOpPartialReductionInterfaceModel
938 :
public PartialReductionOpInterface::ExternalModel<
939 LinalgOpPartialReductionInterfaceModel<LinalgOpTy>, LinalgOpTy>,
940 public LinalgOpPartialReductionInterfaceImpl {
941 using LinalgOpPartialReductionInterfaceImpl::
942 generateInitialTensorForPartialReduction;
943 using LinalgOpPartialReductionInterfaceImpl::getPartialResultTilePosition;
944 using LinalgOpPartialReductionInterfaceImpl::mergeReductions;
945 using LinalgOpPartialReductionInterfaceImpl::tileToPartialReduction;
948template <
typename OpTy>
951 static_assert(llvm::is_one_of<OpTy, PackOp, UnPackOp>::value,
952 "applies to only pack or unpack operations");
954 int64_t rank = (std::is_same<OpTy, PackOp>::value) ? op.getSourceRank()
959 (
void)op.reifyResultShapes(builder, resultShape);
961 for (
auto dim : llvm::seq<int64_t>(0, rank)) {
962 loopBounds[dim].offset = zero;
963 loopBounds[dim].stride = one;
964 loopBounds[dim].size = resultShape[0][dim];
972 if (permutation.empty())
984 interchangeVector.reserve(dimsPos.size());
993 for (
int64_t dimsIdx = 0, end = dimsPos.size(); dimsIdx < end; dimsIdx++)
994 dimsAndPosMapping[dimsPos[dimsIdx]] = dimsIdx;
998 for (
int64_t dimsIdx = 0; dimsIdx < rank; dimsIdx++) {
999 if (dimsAndPosMapping.count(dimsIdx))
1000 interchangeVector.push_back(dimsAndPosMapping[dimsIdx]);
1002 return interchangeVector;
1017template <
typename T>
1022 for (
auto [idx, val] : llvm::enumerate(interchangeVector))
1023 vec[idx + offset] = elements[val + offset];
1029static void generatePackOpScalarImplementationBody(PackOp packOp,
1044 computeInterchangeFromDimPos(dimsToInnerBlock, packOp.getSourceRank());
1045 interchangedIvs = interchange<Value>(interchangedIvs, interchangeVector,
1046 packOp.getSourceRank());
1047 if (!dimsToOuterBlock.empty()) {
1049 computeInterchangeFromDimPos(dimsToOuterBlock, packOp.getSourceRank());
1051 interchange<Value>(interchangedIvs, interchangeVector, 0);
1054 packOp.getDimAndTileMapping();
1056 size_t pointLoopsOffset = 0;
1057 int64_t sourceRank = packOp.getSourceRank();
1058 for (
auto dim : llvm::seq<int64_t>(0, sourceRank)) {
1059 if (dimAndTileMapping.contains(dim)) {
1064 builder, loc, i *
tile +
j,
1066 interchangedIvs[dim],
1067 interchangedIvs[pointLoopsOffset + packOp.getSourceRank()],
1068 dimAndTileMapping[dim]});
1069 sourceIndices.push_back(sourceIndex);
1072 sourceIndices.push_back(interchangedIvs[dim]);
1076 auto createLoad = [&]() ->
Value {
1077 return memref::LoadOp::create(
1078 builder, loc, packOp.getSource(),
1082 if (
auto paddingValue = packOp.getPaddingValue()) {
1085 for (
auto dim : llvm::seq<int64_t>(0, sourceRank)) {
1088 Value cond = arithBuilder.slt(
1092 scalar = scf::IfOp::create(
1095 scf::YieldOp::create(
b, l, createLoad());
1099 scf::YieldOp::create(
b, l, paddingValue);
1103 scalar = createLoad();
1106 memref::StoreOp::create(builder, loc, scalar, packOp.getDest(), ivs);
1110 :
public TilingInterface::ExternalModel<PackOpTiling, linalg::PackOp> {
1111 using Base = TilingInterface::ExternalModel<PackOpTiling, linalg::PackOp>;
1112 using Base::getTiledImplementation;
1114 SmallVector<utils::IteratorType> getLoopIteratorTypes(Operation *op)
const {
1118 auto packOp = cast<PackOp>(op);
1119 SmallVector<utils::IteratorType> iteratorTypes(
1120 packOp.getSourceRank(), utils::IteratorType::parallel);
1121 return iteratorTypes;
1124 SmallVector<Range> getIterationDomain(Operation *op, OpBuilder &
b)
const {
1125 return getPackUnPackIterationDomain<PackOp>(cast<PackOp>(op),
b);
1128 FailureOr<TilingResult>
1130 ArrayRef<OpFoldResult> offsets,
1131 ArrayRef<OpFoldResult> sizes)
const {
1132 auto packOp = cast<PackOp>(op);
1134 if (!packOp.hasPureTensorSemantics())
1137 Location loc = packOp.getLoc();
1141 int64_t inputRank = packOp.getSourceRank();
1142 SmallVector<OpFoldResult> origOffsets(offsets);
1143 SmallVector<OpFoldResult> origSizes(sizes);
1144 applyPermToRange(origOffsets, origSizes,
1148 packOp.getDimAndTileMapping();
1149 SmallVector<OpFoldResult> srcDimValues =
1151 SmallVector<OpFoldResult> inputIndices, inputSizes;
1152 for (
auto dim : llvm::seq<int64_t>(0, inputRank)) {
1153 using AV = affine::AffineValueExpr;
1154 affine::AffineBuilder ab(
b, loc);
1155 AffineExpr dim0, dim1, sym;
1158 if (dimAndTileMapping.count(dim)) {
1162 auto avOffset = AV(dim0).bind(origOffsets[dim]);
1163 auto avSize = AV(dim0).bind(origSizes[dim]);
1164 auto avTileSize = AV(sym).bind(dimAndTileMapping[dim]);
1165 inputIndices.push_back(ab.mul(avOffset, avTileSize));
1166 inputSizes.push_back(ab.mul(avSize, avTileSize));
1168 inputIndices.push_back(origOffsets[dim]);
1169 inputSizes.push_back(origSizes[dim]);
1173 if (packOp.getPaddingValue()) {
1174 OpFoldResult dimSize = srcDimValues[dim];
1175 auto avDimSize = AV(dim0).bind(dimSize);
1176 auto avInputIdx = AV(dim1).bind(inputIndices.back());
1178 ab.min({inputSizes.back(), ab.sub(avDimSize, avInputIdx)});
1182 auto oneAttr =
b.getI64IntegerAttr(1);
1183 SmallVector<OpFoldResult> strides(inputRank, oneAttr);
1185 SmallVector<Value> tiledOperands;
1186 auto sourceSlice = tensor::ExtractSliceOp::create(
1187 b, loc, packOp.getSource(), inputIndices, inputSizes, strides);
1188 tiledOperands.push_back(sourceSlice);
1190 SmallVector<OpFoldResult> outputOffsets, outputSizes;
1195 strides.append(packOp.getDestRank() - inputRank, oneAttr);
1196 auto outSlice = tensor::ExtractSliceOp::create(
1197 b, loc, packOp.getDest(), outputOffsets, outputSizes, strides);
1198 tiledOperands.push_back(outSlice);
1200 if (
auto val = packOp.getPaddingValue())
1201 tiledOperands.push_back(val);
1202 for (
auto tile : packOp.getInnerTiles())
1203 tiledOperands.push_back(
tile);
1205 PackOp tiledPackOp =
1206 PackOp::create(
b, loc,
TypeRange{outSlice.getType()}, tiledOperands,
1207 packOp.getProperties(),
1208 packOp->getDiscardableAttrDictionary().getValue());
1210 return TilingResult{
1212 SmallVector<Value>(tiledPackOp->getResults()),
1213 llvm::to_vector(ArrayRef<Operation *>{sourceSlice, outSlice})};
1218 ArrayRef<OpFoldResult> offsets,
1219 ArrayRef<OpFoldResult> sizes,
1220 SmallVector<OpFoldResult> &resultOffsets,
1221 SmallVector<OpFoldResult> &resultSizes)
const {
1226 auto packOp = cast<PackOp>(op);
1227 int64_t inputRank = packOp.getSourceRank();
1228 int64_t outputRank = packOp.getDestRank();
1229 auto zeroAttr =
b.getI64IntegerAttr(0);
1230 resultOffsets.assign(offsets.begin(), offsets.end());
1231 resultOffsets.append(outputRank - inputRank, zeroAttr);
1235 resultSizes.assign(sizes.begin(), sizes.end());
1236 for (
auto dataTileDim : llvm::seq<unsigned>(inputRank, outputRank))
1237 resultSizes.push_back(outputShape[0][dataTileDim]);
1242 FailureOr<TilingResult>
1243 generateResultTileValue(Operation *op, OpBuilder &
b,
unsigned resultNumber,
1244 ArrayRef<OpFoldResult> offsets,
1245 ArrayRef<OpFoldResult> sizes)
const {
1246 return generateResultTileValue(op,
b, resultNumber, offsets, sizes,
1250 FailureOr<TilingResult> generateResultTileValue(
1251 Operation *op, OpBuilder &
b,
unsigned resultNumber,
1252 ArrayRef<OpFoldResult> offsets, ArrayRef<OpFoldResult> sizes,
1253 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1254 auto packOp = cast<PackOp>(op);
1255 int64_t numTiles = packOp.getInnerDimsPos().size();
1260 for (
auto offset : offsets.take_back(numTiles))
1267 ArrayRef<int64_t> innerDimsPos = packOp.getInnerDimsPos();
1268 SmallVector<OpFoldResult> mixedTiles = packOp.getMixedTiles();
1269 ArrayRef<OpFoldResult> innerSizes = sizes.take_back(numTiles);
1270 for (
auto [i, pos] : llvm::enumerate(innerDimsPos)) {
1272 pos < static_cast<int64_t>(innerTileAlignments.size())
1273 ? innerTileAlignments[pos]
1274 : InnerTileAlignment::Unknown;
1275 if (alignment != InnerTileAlignment::Equal &&
1281 op,
b, offsets.drop_back(numTiles), sizes.drop_back(numTiles));
1282 if (
failed(tilingResult))
1284 return tilingResult.value();
1287 LogicalResult generateScalarImplementation(Operation *op, OpBuilder &builder,
1290 auto packOp = cast<PackOp>(op);
1291 assert(packOp.hasPureBufferSemantics() &&
1292 "expected operation to have buffer semantics");
1293 OpBuilder::InsertionGuard g(builder);
1296 SmallVector<Value> ivVec(ivs);
1299 SmallVector<OpFoldResult> outputShape;
1300 Value dest = packOp.getDest();
1301 for (
auto dim : llvm::seq<int64_t>(0, packOp.getDestRank()))
1310 for (
auto dataTileDim : llvm::seq<unsigned>(packOp.getSourceRank(),
1311 packOp.getDestRank() - 1)) {
1313 outputShape[dataTileDim]);
1314 scf::ForOp loop = scf::ForOp::create(builder, loc, zero, ub, one);
1316 ivVec.push_back(loop.getInductionVar());
1323 [&](OpBuilder &bodyBuilder, Location bodyLoc, Value iv,
1325 ivVec.push_back(iv);
1326 generatePackOpScalarImplementationBody(packOp, bodyBuilder, bodyLoc,
1328 scf::YieldOp::create(bodyBuilder, bodyLoc);
1333 LogicalResult getIterationDomainTileFromOperandTiles(
1334 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1335 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1336 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1337 SmallVectorImpl<OpFoldResult> &resultOffsets,
1338 SmallVectorImpl<OpFoldResult> &resultSizes)
const {
1339 return getIterationDomainTileFromOperandTiles(
1340 op,
b, operandNumbers, allOffsets, allSizes, resultOffsets, resultSizes,
1347 LogicalResult getIterationDomainTileFromOperandTiles(
1348 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1349 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1350 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1351 SmallVectorImpl<OpFoldResult> &resultOffsets,
1352 SmallVectorImpl<OpFoldResult> &resultSizes,
1353 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1354 if (operandNumbers.size() != 1 || operandNumbers[0] != 0) {
1356 { llvm::dbgs() <<
"unsupported operands for consumer fusion"; });
1360 ArrayRef<OpFoldResult> offsets(allOffsets[0]);
1361 ArrayRef<OpFoldResult> sizes(allSizes[0]);
1362 auto packOp = cast<PackOp>(op);
1363 Location loc = packOp.getLoc();
1364 SmallVector<OpFoldResult> outerDimOffsets, outerDimSizes;
1366 packOp.getDimAndTileMapping();
1367 SmallVector<int64_t> outerShapeWithoutTranspose(
1368 packOp.getDestType().getShape().take_front(packOp.getSourceRank()));
1369 if (!packOp.getOuterDimsPerm().empty()) {
1371 outerShapeWithoutTranspose,
1374 for (
auto dim : llvm::seq<int64_t>(packOp.getSourceRank())) {
1375 if (dimAndTileMapping.count(dim)) {
1376 FailureOr<int64_t> cstTileSize =
1378 presburger::BoundType::UB, sizes[dim],
1380 ValueBoundsOptions{
true});
1381 std::optional<int64_t> cstInnerSize =
1388 dim < static_cast<int64_t>(innerTileAlignments.size())
1389 ? innerTileAlignments[dim]
1390 : InnerTileAlignment::Unknown;
1403 int64_t srcDimSize = packOp.getSourceType().getDimSize(dim);
1404 int64_t destDimSize = outerShapeWithoutTranspose[dim];
1405 bool isTiled = innerTileAlignment != InnerTileAlignment::Unknown ||
1407 ShapedType::isDynamic(srcDimSize) ||
1408 cstTileSize.value() < srcDimSize;
1410 outerDimOffsets.push_back(offsets[dim]);
1411 if (ShapedType::isStatic(destDimSize)) {
1412 outerDimSizes.push_back(
b.getIndexAttr(destDimSize));
1414 outerDimSizes.push_back(
1415 b.createOrFold<tensor::DimOp>(loc, packOp.getDest(), dim));
1442 bool assumeInnerTileSizesMatchTiles =
1443 innerTileAlignment == InnerTileAlignment::Equal;
1444 bool staticallyDecidable =
1445 !
failed(cstTileSize) && cstInnerSize.has_value();
1446 if (innerTileAlignment == InnerTileAlignment::Unknown) {
1447 if (!staticallyDecidable || *cstTileSize % *cstInnerSize != 0)
1449 }
else if (staticallyDecidable) {
1450 assert(*cstTileSize % *cstInnerSize == 0 &&
1451 "InnerTileAlignment hint contradicts statically known tile "
1453 assert((innerTileAlignment != InnerTileAlignment::Equal ||
1454 *cstTileSize == *cstInnerSize) &&
1455 "InnerTileAlignment::Equal contradicts statically known tile "
1459 using AV = affine::AffineValueExpr;
1460 affine::AffineBuilder ab(
b, loc);
1461 AffineExpr dim0, sym;
1464 auto avOffset = AV(dim0).bind(offsets[dim]);
1465 auto avSize = AV(dim0).bind(sizes[dim]);
1466 auto avTileSize = AV(sym).bind(dimAndTileMapping[dim]);
1467 outerDimOffsets.push_back(ab.floor(avOffset, avTileSize));
1470 outerDimSizes.push_back(assumeInnerTileSizesMatchTiles
1472 : ab.ceil(avSize, avTileSize));
1474 outerDimOffsets.push_back(offsets[dim]);
1475 outerDimSizes.push_back(sizes[dim]);
1478 applyPermToRange(outerDimOffsets, outerDimSizes, packOp.getOuterDimsPerm());
1479 resultOffsets = outerDimOffsets;
1480 resultSizes = outerDimSizes;
1484 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
1485 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1486 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1487 ArrayRef<SmallVector<OpFoldResult>> allSizes)
const {
1488 return getTiledImplementationFromOperandTiles(op,
b, operandNumbers,
1489 allOffsets, allSizes,
1494 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
1495 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1496 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1497 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1498 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1499 if (operandNumbers.size() != 1 || operandNumbers[0] != 0) {
1500 LLVM_DEBUG({ llvm::dbgs() <<
"unhandled operands for consumer fusion"; });
1504 ArrayRef<OpFoldResult> offsets(allOffsets[0]);
1505 ArrayRef<OpFoldResult> sizes(allSizes[0]);
1507 auto packOp = cast<PackOp>(op);
1509 if (!packOp.hasPureTensorSemantics())
1512 Location loc = packOp.getLoc();
1514 int64_t inputRank = packOp.getSourceRank();
1515 auto oneAttr =
b.getI64IntegerAttr(1);
1516 SmallVector<OpFoldResult> strides(inputRank, oneAttr);
1518 SmallVector<Value> tiledOperands;
1519 auto sourceSlice = tensor::ExtractSliceOp::create(
1520 b, loc, packOp.getSource(), offsets, sizes, strides);
1521 tiledOperands.push_back(sourceSlice);
1523 SmallVector<OpFoldResult> outerDimOffsets, outerDimSizes;
1524 if (
failed(getIterationDomainTileFromOperandTiles(
1525 op,
b, operandNumbers, allOffsets, allSizes, outerDimOffsets,
1526 outerDimSizes, innerTileAlignments)))
1529 SmallVector<OpFoldResult> outputOffsets, outputSizes;
1531 outputOffsets, outputSizes)))
1534 strides.append(packOp.getDestRank() - inputRank, oneAttr);
1535 auto outSlice = tensor::ExtractSliceOp::create(
1536 b, loc, packOp.getDest(), outputOffsets, outputSizes, strides);
1537 tiledOperands.push_back(outSlice);
1539 if (
auto val = packOp.getPaddingValue())
1540 tiledOperands.push_back(val);
1541 for (
auto tile : packOp.getInnerTiles())
1542 tiledOperands.push_back(
tile);
1544 PackOp tiledPackOp =
1545 PackOp::create(
b, loc,
TypeRange{outSlice.getType()}, tiledOperands,
1546 packOp.getProperties(),
1547 packOp->getDiscardableAttrDictionary().getValue());
1549 return TilingResult{
1551 SmallVector<Value>(tiledPackOp->getResults()),
1552 llvm::to_vector(ArrayRef<Operation *>{sourceSlice, outSlice})};
1556struct UnpackTileDimInfo {
1557 bool isAlignedToInnerTileSize;
1558 OpFoldResult sourceOffset;
1559 OpFoldResult sourceSize;
1560 OpFoldResult resultOffset;
1561 OpFoldResult destExpandedSize;
1567static UnpackTileDimInfo
1571 UnpackTileDimInfo info;
1575 unpackOp.getDimAndTileMapping();
1577 if (!dimAndTileMapping.count(tileDim)) {
1578 info.isAlignedToInnerTileSize =
true;
1579 info.sourceOffset = tileOffset;
1580 info.sourceSize = tileSize;
1581 info.resultOffset = zeroAttr;
1582 info.destExpandedSize = tileSize;
1593 OpFoldResult innerTileSize = dimAndTileMapping[tileDim];
1595 info.isAlignedToInnerTileSize =
false;
1608 bool assumeInnerTileSizesMatchTiles =
1610 bool staticallyDecidable = !
failed(cstSize) && cstInnerSize.has_value();
1612 info.isAlignedToInnerTileSize =
true;
1613 if (staticallyDecidable) {
1614 assert(*cstSize % *cstInnerSize == 0 &&
1615 "InnerTileAlignment hint contradicts statically known tile sizes");
1617 *cstSize == *cstInnerSize) &&
1618 "InnerTileAlignment::Equal contradicts statically known tile "
1622 if (info.isAlignedToInnerTileSize || (!
failed(cstSize) && cstInnerSize)) {
1623 if (!info.isAlignedToInnerTileSize && *cstSize % *cstInnerSize == 0)
1624 info.isAlignedToInnerTileSize =
true;
1628 if (assumeInnerTileSizesMatchTiles ||
1629 (cstInnerSize && !
failed(cstSize) && *cstInnerSize == *cstSize)) {
1630 auto lhs = AV(dim0).bind(tileOffset);
1631 auto rhs = AV(dim1).bind(innerTileSize);
1632 info.sourceOffset = ab.floor(lhs, rhs);
1633 info.sourceSize = oneAttr;
1634 info.resultOffset = zeroAttr;
1635 info.destExpandedSize = tileSize;
1640 if (info.isAlignedToInnerTileSize) {
1642 ab.floor(AV(dim0).bind(tileOffset), AV(dim1).bind(innerTileSize));
1643 info.resultOffset = zeroAttr;
1644 info.destExpandedSize = tileSize;
1653 ab.ceil(AV(dim0).bind(tileSize), AV(dim1).bind(innerTileSize));
1657 affine::DivModValue firstCoord = affine::getDivMod(
1661 ab.add(AV(dim0).bind(tileOffset), AV(dim1).bind(tileSize));
1662 affine::DivModValue lastCoord = affine::getDivMod(
1666 ab.sub(AV(dim0).bind(tileExclusiveBound), AV(dim1).bind(oneAttr))),
1669 OpFoldResult lengthMinusOne = ab.sub(AV(dim0).bind(lastCoord.quotient),
1670 AV(dim1).bind(firstCoord.quotient));
1672 ab.add(AV(dim0).bind(lengthMinusOne), AV(dim1).bind(oneAttr));
1673 info.sourceOffset = firstCoord.quotient;
1674 info.resultOffset = firstCoord.remainder;
1677 info.destExpandedSize =
b.createOrFold<arith::MulIOp>(
1683struct UnPackOpTiling
1684 :
public TilingInterface::ExternalModel<UnPackOpTiling, linalg::UnPackOp> {
1685 using Base = TilingInterface::ExternalModel<UnPackOpTiling, linalg::UnPackOp>;
1686 using Base::getIterationDomainTileFromOperandTiles;
1688 SmallVector<utils::IteratorType> getLoopIteratorTypes(Operation *op)
const {
1689 auto unpackOp = cast<UnPackOp>(op);
1690 SmallVector<utils::IteratorType> iteratorTypes(
1691 unpackOp.getDestRank(), utils::IteratorType::parallel);
1692 return iteratorTypes;
1695 SmallVector<Range> getIterationDomain(Operation *op, OpBuilder &
b)
const {
1696 return getPackUnPackIterationDomain<UnPackOp>(cast<UnPackOp>(op),
b);
1713 FailureOr<TilingResult>
1715 ArrayRef<OpFoldResult> offsets,
1716 ArrayRef<OpFoldResult> sizes)
const {
1722 Operation *op, OpBuilder &
b, ArrayRef<OpFoldResult> offsets,
1723 ArrayRef<OpFoldResult> sizes,
1724 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1725 auto unpackOp = cast<UnPackOp>(op);
1727 if (!unpackOp.hasPureTensorSemantics())
1730 int64_t srcRank = unpackOp.getSourceRank();
1731 int64_t destRank = unpackOp.getDestRank();
1732 int64_t numInnerTiles = srcRank - destRank;
1733 Location loc = unpackOp.getLoc();
1738 bool isPerfectTilingCase =
true;
1739 Attribute oneAttr =
b.getIndexAttr(1);
1740 SmallVector<OpFoldResult> sliceSrcStrides(destRank, oneAttr);
1741 SmallVector<OpFoldResult> sliceSrcIndices, sliceSrcSizes;
1742 SmallVector<OpFoldResult> destExpandedSizes, resultOffsetsFromDest;
1743 for (
auto dim : llvm::seq<int64_t>(0, destRank)) {
1744 UnpackTileDimInfo info = getUnpackTileDimInfo(
1745 b, unpackOp, dim, offsets[dim], sizes[dim],
1746 dim <
static_cast<int64_t
>(innerTileAlignments.size())
1747 ? innerTileAlignments[dim]
1748 : InnerTileAlignment::Unknown);
1749 if (!info.isAlignedToInnerTileSize)
1750 isPerfectTilingCase =
false;
1751 sliceSrcIndices.push_back(info.sourceOffset);
1752 sliceSrcSizes.push_back(info.sourceSize);
1753 destExpandedSizes.push_back(info.destExpandedSize);
1754 resultOffsetsFromDest.push_back(info.resultOffset);
1759 applyPermToRange(sliceSrcIndices, sliceSrcSizes,
1760 unpackOp.getOuterDimsPerm());
1761 Attribute zeroAttr =
b.getIndexAttr(0);
1762 sliceSrcIndices.append(numInnerTiles, zeroAttr);
1763 sliceSrcSizes.append(unpackOp.getMixedTiles());
1764 sliceSrcStrides.append(numInnerTiles, oneAttr);
1765 SmallVector<Operation *> generatedSlices;
1766 tensor::ExtractSliceOp sliceSource = tensor::ExtractSliceOp::create(
1767 b, loc, unpackOp.getSource(), sliceSrcIndices, sliceSrcSizes,
1769 generatedSlices.push_back(sliceSource);
1771 SmallVector<OpFoldResult> destStrides(destRank, oneAttr);
1773 if (isPerfectTilingCase) {
1774 auto destSliceOp = tensor::ExtractSliceOp::create(
1775 b, loc, unpackOp.getDest(), offsets, sizes, destStrides);
1776 sliceDest = destSliceOp;
1777 generatedSlices.push_back(destSliceOp);
1779 sliceDest = tensor::EmptyOp::create(
1780 b, loc, destExpandedSizes, unpackOp.getDestType().getElementType());
1783 SmallVector<Value> tiledOperands = {sliceSource.getResult(), sliceDest};
1784 for (
auto tile : unpackOp.getInnerTiles())
1785 tiledOperands.push_back(
tile);
1787 UnPackOp tiledUnpackOp =
1789 unpackOp.getProperties(),
1790 unpackOp->getDiscardableAttrDictionary().getValue());
1792 if (isPerfectTilingCase)
1793 return TilingResult{{tiledUnpackOp},
1794 SmallVector<Value>(tiledUnpackOp->getResults()),
1797 auto extractSlice = tensor::ExtractSliceOp::create(
1798 b, loc, tiledUnpackOp->getResult(0), resultOffsetsFromDest, sizes,
1800 return TilingResult{
1801 {tiledUnpackOp}, {extractSlice.getResult()}, generatedSlices};
1806 ArrayRef<OpFoldResult> offsets,
1807 ArrayRef<OpFoldResult> sizes,
1808 SmallVector<OpFoldResult> &resultOffsets,
1809 SmallVector<OpFoldResult> &resultSizes)
const {
1810 resultOffsets = llvm::to_vector(offsets);
1811 resultSizes = llvm::to_vector(sizes);
1815 FailureOr<TilingResult>
1816 generateResultTileValue(Operation *op, OpBuilder &
b,
unsigned resultNumber,
1817 ArrayRef<OpFoldResult> offsets,
1818 ArrayRef<OpFoldResult> sizes)
const {
1819 return generateResultTileValue(op,
b, resultNumber, offsets, sizes,
1823 FailureOr<TilingResult> generateResultTileValue(
1824 Operation *op, OpBuilder &
b,
unsigned resultNumber,
1825 ArrayRef<OpFoldResult> offsets, ArrayRef<OpFoldResult> sizes,
1826 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1827 FailureOr<TilingResult> tilingResult =
1829 if (
failed(tilingResult))
1831 return tilingResult.value();
1834 LogicalResult generateScalarImplementation(Operation *op, OpBuilder &builder,
1837 auto unpackOp = cast<UnPackOp>(op);
1838 assert(unpackOp.hasPureBufferSemantics() &&
1839 "expected operation to have buffer semantics");
1840 assert(ivs.size() == unpackOp.getDestRank() &&
1841 "number of ivs must match the rank of the output tensor");
1842 OpBuilder::InsertionGuard g(builder);
1845 unpackOp.getDimAndTileMapping();
1847 SmallVector<Value> inputIvs;
1849 SmallVector<Value> inputIvsPointLoops;
1850 inputIvs.reserve(unpackOp.getDestRank());
1851 inputIvsPointLoops.reserve(dimAndTileMapping.size());
1852 for (
auto dim : llvm::seq<int64_t>(0, unpackOp.getDestRank())) {
1853 if (dimAndTileMapping.count(dim)) {
1854 affine::DivModValue divMod =
1855 affine::getDivMod(builder, loc, ivs[dim],
1857 builder, loc, dimAndTileMapping[dim]));
1858 inputIvsPointLoops.push_back(divMod.remainder);
1859 inputIvs.push_back(divMod.quotient);
1861 inputIvs.push_back(ivs[dim]);
1867 assert(inputIvsPointLoops.size() + inputIvs.size() ==
1868 unpackOp.getSourceRank() &&
1869 "expect same number of induction variables equals to input rank");
1871 ArrayRef<int64_t> innerDims = unpackOp.getInnerDimsPos();
1872 SmallVector<int64_t> interchangeVector =
1873 computeInterchangeFromDimPos(innerDims, unpackOp.getDestRank());
1874 SmallVector<Value> interchangedInputIvsPointLoops = inputIvsPointLoops;
1875 interchangedInputIvsPointLoops = interchange<Value>(
1876 interchangedInputIvsPointLoops, interchangeVector, 0);
1879 ArrayRef<int64_t> outerDims = unpackOp.getOuterDimsPerm();
1880 if (!outerDims.empty())
1881 inputIvs = interchange<Value>(inputIvs, outerDims, 0);
1883 llvm::append_range(inputIvs, interchangedInputIvsPointLoops);
1885 memref::LoadOp::create(builder, loc, unpackOp.getSource(), inputIvs);
1886 memref::StoreOp::create(builder, loc, scalar, unpackOp.getDest(), ivs);
1892 LogicalResult getIterationDomainTileFromOperandTiles(
1893 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1894 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1895 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1896 SmallVectorImpl<OpFoldResult> &resultOffsets,
1897 SmallVectorImpl<OpFoldResult> &resultSizes)
const {
1898 if (operandNumbers.size() != 1) {
1899 LLVM_DEBUG({ llvm::dbgs() <<
"unable to handle multiple operands"; });
1902 auto unPackOp = cast<UnPackOp>(op);
1903 unsigned operandNumber = operandNumbers[0];
1904 ArrayRef<OpFoldResult> offsets(allOffsets[0]);
1905 ArrayRef<OpFoldResult> sizes(allSizes[0]);
1908 if (operandNumber == unPackOp.getDestMutable().getOperandNumber()) {
1909 resultOffsets = llvm::to_vector(offsets);
1910 resultSizes = llvm::to_vector(sizes);
1913 Location loc = unPackOp.getLoc();
1915 int64_t numTiles = unPackOp.getInnerDimsPos().size();
1916 auto destOffsets = offsets.drop_back(numTiles);
1917 auto destSizes = sizes.drop_back(numTiles);
1920 int64_t outputRank = unPackOp.getDestRank();
1924 SmallVector<OpFoldResult> outputMixedSizes = reifiedReturnShapes.front();
1925 SmallVector<OpFoldResult> origOffsets(destOffsets);
1926 SmallVector<OpFoldResult> origSizes(destSizes);
1927 applyPermToRange(origOffsets, origSizes,
1931 unPackOp.getDimAndTileMapping();
1933 for (
auto dim : llvm::seq<int64_t>(0, outputRank)) {
1934 using AV = affine::AffineValueExpr;
1935 affine::AffineBuilder ab(
b, loc);
1936 AffineExpr dim0, dim1, sym0;
1939 if (dimAndTileMapping.count(dim)) {
1943 auto avOffset = AV(dim0).bind(origOffsets[dim]);
1944 auto avSize = AV(dim0).bind(origSizes[dim]);
1945 auto avTileSize = AV(sym0).bind(dimAndTileMapping[dim]);
1946 auto avResultSize = AV(dim0).bind(outputMixedSizes[dim]);
1947 resultOffsets.push_back(ab.mul(avOffset, avTileSize));
1948 auto avResultOffset = AV(dim1).bind(resultOffsets.back());
1949 resultSizes.push_back(ab.min({ab.mul(avSize, avTileSize),
1950 ab.sub(avResultSize, avResultOffset)}));
1952 resultOffsets.push_back(origOffsets[dim]);
1953 resultSizes.push_back(origSizes[dim]);
1959 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
1960 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1961 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1962 ArrayRef<SmallVector<OpFoldResult>> allSizes)
const {
1963 return getTiledImplementationFromOperandTiles(op,
b, operandNumbers,
1964 allOffsets, allSizes,
1969 FailureOr<TilingResult> getTiledImplementationFromOperandTiles(
1970 Operation *op, OpBuilder &
b, ArrayRef<unsigned> operandNumbers,
1971 ArrayRef<SmallVector<OpFoldResult>> allOffsets,
1972 ArrayRef<SmallVector<OpFoldResult>> allSizes,
1973 ArrayRef<InnerTileAlignment> innerTileAlignments)
const {
1974 if (operandNumbers.size() != 1 || operandNumbers[0] != 0) {
1975 LLVM_DEBUG({ llvm::dbgs() <<
"unhandled operands for consumer fusion"; });
1978 auto unPackOp = cast<UnPackOp>(op);
1980 if (!unPackOp.hasPureTensorSemantics())
1983 ArrayRef<OpFoldResult> offsets(allOffsets[0]);
1984 ArrayRef<OpFoldResult> sizes(allSizes[0]);
1990 int64_t numTiles = unPackOp.getInnerDimsPos().size();
1991 ArrayRef<int64_t> innerDimsPos = unPackOp.getInnerDimsPos();
1992 SmallVector<OpFoldResult> mixedTiles = unPackOp.getMixedTiles();
1993 ArrayRef<OpFoldResult> innerSizes = sizes.take_back(numTiles);
1994 for (int64_t i = 0; i < numTiles; ++i) {
1997 int64_t destDim = innerDimsPos[i];
1999 destDim < static_cast<int64_t>(innerTileAlignments.size()) &&
2000 innerTileAlignments[destDim] == InnerTileAlignment::Equal;
2010 "InnerTileAlignment::Equal contradicts statically known tile "
2019 Location loc = unPackOp.getLoc();
2023 SmallVector<OpFoldResult> outputOffsets, outputSizes;
2024 if (
failed(getIterationDomainTileFromOperandTiles(
2025 op,
b, operandNumbers, allOffsets, allSizes, outputOffsets,
2029 auto oneAttr =
b.getI64IntegerAttr(1);
2030 int64_t outputRank = unPackOp.getDestRank();
2031 SmallVector<OpFoldResult> strides(outputRank, oneAttr);
2033 SmallVector<Value> tiledOperands;
2035 auto extractDestSlice = tensor::ExtractSliceOp::create(
2036 b, loc, unPackOp.getDest(), outputOffsets, outputSizes, strides);
2037 tiledOperands.push_back(extractDestSlice);
2039 strides.append(unPackOp.getSourceRank() - outputRank, oneAttr);
2041 auto extractSourceSlice = tensor::ExtractSliceOp::create(
2042 b, loc, unPackOp.getSource(), offsets, sizes, strides);
2043 tiledOperands.insert(tiledOperands.begin(), extractSourceSlice);
2044 for (
auto tile : unPackOp.getInnerTiles())
2045 tiledOperands.push_back(
tile);
2048 UnPackOp tiledUnPackOp =
2049 UnPackOp::create(
b, loc,
TypeRange{extractDestSlice.getType()},
2050 tiledOperands, unPackOp.getProperties(),
2051 unPackOp->getDiscardableAttrDictionary().getValue());
2053 return TilingResult{{tiledUnPackOp},
2054 SmallVector<Value>(tiledUnPackOp->getResults()),
2055 llvm::to_vector(ArrayRef<Operation *>{
2056 extractSourceSlice, extractDestSlice})};
2062template <
typename OpType>
2064 OpType::template attachInterface<LinalgOpTilingInterfaceModel<OpType>>(*ctx);
2065 OpType::template attachInterface<
2066 LinalgOpPartialReductionInterfaceModel<OpType>>(*ctx);
2070template <
typename... OpTypes>
2081 linalg::PackOp::attachInterface<PackOpTiling>(*ctx);
2082 linalg::UnPackOp::attachInterface<UnPackOpTiling>(*ctx);
2084#include "mlir/Dialect/Linalg/IR/LinalgStructuredOps.cpp.inc"
2092 linalg::PackOp::attachInterface<PackOpTiling>(*ctx);
2093 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)
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.