2589 lines · cpp
1//===- Tiling.cpp - Implementation of tiling using TilingInterface -------===//2//3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.4// See https://llvm.org/LICENSE.txt for license information.5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception6//7//===----------------------------------------------------------------------===//8//9// This file implements the tiling using TilingInterface.10//11//===----------------------------------------------------------------------===//12 13#include "mlir/Dialect/SCF/Transforms/TileUsingInterface.h"14 15#include "mlir/Analysis/SliceAnalysis.h"16#include "mlir/Analysis/TopologicalSortUtils.h"17#include "mlir/Dialect/Affine/IR/AffineOps.h"18#include "mlir/Dialect/Arith/Utils/Utils.h"19#include "mlir/Dialect/SCF/Utils/Utils.h"20#include "mlir/Dialect/Tensor/IR/Tensor.h"21#include "mlir/Dialect/Utils/IndexingUtils.h"22#include "mlir/IR/Dominance.h"23#include "mlir/IR/PatternMatch.h"24#include "mlir/Interfaces/DestinationStyleOpInterface.h"25#include "mlir/Interfaces/TilingInterface.h"26#include "mlir/Rewrite/FrozenRewritePatternSet.h"27#include "mlir/Transforms/GreedyPatternRewriteDriver.h"28#include "llvm/ADT/ScopeExit.h"29#include "llvm/ADT/TypeSwitch.h"30#include "llvm/Support/Debug.h"31#include <optional>32 33#define DEBUG_TYPE "tile-using-interface"34 35using namespace mlir;36 37scf::SCFTilingOptions &38scf::SCFTilingOptions::setTileSizes(ArrayRef<OpFoldResult> ts) {39 assert(!tileSizeComputationFunction && "tile sizes already set");40 auto tileSizes = llvm::to_vector(ts);41 tileSizeComputationFunction = [tileSizes](OpBuilder &b, Operation *op) {42 return tileSizes;43 };44 return *this;45}46 47scf::SCFTilingOptions &48scf::SCFTilingOptions::setNumThreads(ArrayRef<OpFoldResult> nt) {49 assert(!numThreadsComputationFunction && "num tiles already set");50 auto numThreads = llvm::to_vector(nt);51 numThreadsComputationFunction = [numThreads](OpBuilder &b, Operation *op) {52 return numThreads;53 };54 return *this;55}56 57/// Helper method to adjust the interchange vector to match the iteration58/// domain.59static SmallVector<int64_t>60fillInterchangeVector(ArrayRef<int64_t> interchangeVector,61 size_t iterationDomainSize) {62 SmallVector<int64_t> filledVector = llvm::to_vector(interchangeVector);63 if (filledVector.size() < iterationDomainSize) {64 auto range = llvm::seq<int64_t>(filledVector.size(), iterationDomainSize);65 filledVector.append(range.begin(), range.end());66 }67 if (filledVector.size() > iterationDomainSize)68 filledVector.resize(iterationDomainSize);69 return filledVector;70}71 72//===----------------------------------------------------------------------===//73// tileUsingSCF implementation.74//===----------------------------------------------------------------------===//75 76/// Verify the tile size options are set in a consistent manner.77static LogicalResult verifyOptions(RewriterBase &rewriter, Location loc,78 const scf::SCFTilingOptions &options) {79 // Specifying number of threads is only supported on `scf.forall` op.80 if (options.numThreadsComputationFunction &&81 options.loopType != scf::SCFTilingOptions::LoopType::ForallOp) {82 return rewriter.notifyMatchFailure(83 loc, "number of threads can only by specified when loop type is "84 "set to use `scf.forall`");85 }86 87 // If specified, check that the interchange vector is a permutation.88 if (!options.interchangeVector.empty()) {89 if (!isPermutationVector(options.interchangeVector)) {90 return rewriter.notifyMatchFailure(91 loc, "invalid interchange vector, not a permutation of the entire "92 "iteration space");93 }94 }95 return success();96}97 98/// Method to instantiate the tile sizes and/or number of threads specified99/// by the user.100static std::tuple<SmallVector<OpFoldResult>, SmallVector<OpFoldResult>>101getUserTileSizesAndNumThreads(RewriterBase &rewriter, TilingInterface op,102 ArrayRef<Range> iterationDomain,103 const scf::SCFTilingOptions &options) {104 OpFoldResult zero = rewriter.getIndexAttr(0);105 SmallVector<OpFoldResult> tileSizes, numThreads;106 size_t numLoops = iterationDomain.size();107 108 // Check whether the number of tiles to use is specified.109 if (options.numThreadsComputationFunction) {110 numThreads = options.numThreadsComputationFunction(rewriter, op);111 numThreads.resize(numLoops, zero);112 113 // If the number of tiles is also specified, use that.114 if (options.tileSizeComputationFunction) {115 tileSizes = options.tileSizeComputationFunction(rewriter, op);116 tileSizes.resize(numLoops, zero);117 return {tileSizes, numThreads};118 }119 120 // Compute the tile sizes from the iteration domain and number121 // of tiles as follows122 // - niters = ceilDiv(ub - lb, step)123 // - tileSize = ceilDiv(niters, numThreads)124 AffineExpr s0, s1, s2;125 bindSymbols(rewriter.getContext(), s0, s1, s2);126 // TODO: The step here is assumed to be 1.127 AffineExpr numItersExpr = (s1 - s0);128 AffineExpr tileSizeExpr = numItersExpr.ceilDiv(s2);129 tileSizes.resize(numLoops, zero);130 for (auto [index, range, nt] :131 llvm::enumerate(iterationDomain, numThreads)) {132 if (isZeroInteger(nt))133 continue;134 135 tileSizes[index] = affine::makeComposedFoldedAffineApply(136 rewriter, op.getLoc(), tileSizeExpr, {range.offset, range.size, nt});137 }138 tileSizes.resize(numLoops, zero);139 return {tileSizes, numThreads};140 }141 142 // Enforce the convention that "tiling by zero"143 // skips tiling a particular dimension. This convention is significantly144 // simpler to handle instead of adjusting affine maps to account for missing145 // dimensions.146 assert(options.tileSizeComputationFunction &&147 "expected tile sizes to be specified");148 tileSizes = options.tileSizeComputationFunction(rewriter, op);149 tileSizes.resize(numLoops, zero);150 151 return {tileSizes, numThreads};152}153 154/// Checks if any of the tiled loops are not parallel.155static LogicalResult checkTileSizes(TilingInterface op,156 scf::SCFTilingOptions::LoopType loopType,157 ReductionTilingStrategy reductionStrategy,158 ArrayRef<OpFoldResult> givenTileSizes,159 ArrayRef<OpFoldResult> numThreads) {160 auto iterators = op.getLoopIteratorTypes();161 assert(iterators.size() == givenTileSizes.size() &&162 "expected as many tile size values as number of loops");163 assert((numThreads.empty() || (numThreads.size() == iterators.size())) &&164 "when specified, expected number of threads to use for each loop");165 166 bool isParallelTiling = false;167 for (auto [index, iterator, givenTileSize] :168 llvm::enumerate(iterators, givenTileSizes)) {169 if (!isConstantIntValue(givenTileSize, 0)) {170 isParallelTiling |= iterator == utils::IteratorType::parallel;171 }172 173 if (loopType == scf::SCFTilingOptions::LoopType::ForallOp &&174 reductionStrategy == ReductionTilingStrategy::FullReduction) {175 // If num threads is specified, check that it is greater than one only for176 // parallel dimensions.177 if (!numThreads.empty()) {178 if (std::optional<int64_t> constNumThreads =179 getConstantIntValue(numThreads[index])) {180 if (constNumThreads.value() > 1 &&181 iterator != utils::IteratorType::parallel) {182 op.emitWarning() << "tiling is not thread safe at axis #" << index;183 }184 }185 continue;186 }187 188 if (std::optional<int64_t> constTileSize =189 getConstantIntValue(givenTileSize)) {190 if (constTileSize.value() > 0 &&191 iterator != utils::IteratorType::parallel) {192 op.emitWarning() << "tiling is not thread safe at axis #" << index;193 }194 }195 }196 }197 198 if (reductionStrategy != ReductionTilingStrategy::FullReduction) {199 if (isParallelTiling) {200 return op->emitOpError("tiling parallel dimensions is not supported with "201 "partial reduction tiling strategies");202 }203 }204 return success();205}206 207/// Get the reduction dims that are tiled. This accounts for reduction dims208/// that are specified as tiled, but the tile size is 0.209static SetVector<unsigned>210getSanitizedReductionDims(ArrayRef<OpFoldResult> givenTileSizes,211 const scf::SCFTilingOptions &options) {212 SetVector<unsigned> reductionDims;213 for (auto dim : options.reductionDims) {214 if (isConstantIntValue(givenTileSizes[dim], 0))215 continue;216 reductionDims.insert(dim);217 }218 return reductionDims;219}220 221/// Check if `stride` evenly divides the trip count `size - offset`.222static bool tileDividesIterationDomain(Range loopRange) {223 std::optional<int64_t> offsetAsInt = getConstantIntValue(loopRange.offset);224 if (!offsetAsInt)225 return false;226 std::optional<int64_t> sizeAsInt = getConstantIntValue(loopRange.size);227 if (!sizeAsInt)228 return false;229 std::optional<int64_t> strideAsInt = getConstantIntValue(loopRange.stride);230 if (!strideAsInt)231 return false;232 return ((sizeAsInt.value() - offsetAsInt.value()) % strideAsInt.value() == 0);233}234 235/// Returns the bounded tile size given the current `offset`, `loopRange` and236/// `tileSize`, i.e., `min(tileSize, range.end() - offset)`.237static OpFoldResult getBoundedTileSize(OpBuilder &b, Location loc,238 Range loopRange, OpFoldResult offset,239 OpFoldResult givenTileSize) {240 std::optional<int64_t> ts = getConstantIntValue(givenTileSize);241 if (ts && ts.value() == 1)242 return givenTileSize;243 244 if (tileDividesIterationDomain(245 Range{loopRange.offset, loopRange.size, givenTileSize}))246 return givenTileSize;247 248 // The tile size to use (to avoid out of bounds access) is minimum of249 // `tileSize` and `ub - iv`, where `iv` is the induction variable of the tiled250 // loop.251 AffineExpr s0, s1, d0;252 bindDims(b.getContext(), d0);253 bindSymbols(b.getContext(), s0, s1);254 AffineMap minMap = AffineMap::get(1, 2, {s0 - d0, s1}, b.getContext());255 Value size = getValueOrCreateConstantIndexOp(b, loc, loopRange.size);256 return affine::makeComposedFoldedAffineMin(257 b, loc, minMap, SmallVector<OpFoldResult>{offset, size, givenTileSize});258}259 260/// Returns true if the maximum tile offset `tileSize * numThreads-1` is less261/// than `iterationSize`.262static bool canOmitTileOffsetInBoundsCheck(OpFoldResult givenTileSize,263 OpFoldResult numThreads,264 OpFoldResult iterationSize) {265 std::optional<int64_t> tileSizeConst = getConstantIntValue(givenTileSize);266 std::optional<int64_t> numThreadsConst = getConstantIntValue(numThreads);267 std::optional<int64_t> iterSizeConst = getConstantIntValue(iterationSize);268 if (!tileSizeConst || !numThreadsConst || !iterSizeConst)269 return false;270 return *tileSizeConst * (*numThreadsConst - 1) < *iterSizeConst;271}272 273/// Compute the `OpFoldResult`s that represents the multi-dimensional274/// `offset`s and `size`s of the tile of the iteration space that the275/// innermost loop body of the generated tiled loops corresponds to.276static std::tuple<SmallVector<OpFoldResult>, SmallVector<OpFoldResult>>277getTileOffsetAndSizes(RewriterBase &rewriter, Location loc, ValueRange ivs,278 ArrayRef<Range> iterationDomain,279 ArrayRef<OpFoldResult> givenTileSizes) {280 SmallVector<OpFoldResult> offsets, sizes;281 int materializedLoopNum = 0;282 for (auto [givenTileSize, loopRange] :283 llvm::zip_equal(givenTileSizes, iterationDomain)) {284 285 // Non-tiled cases, set the offset and size to the286 // `loopRange.offset/size`.287 if (isZeroInteger(givenTileSize)) {288 offsets.push_back(loopRange.offset);289 sizes.push_back(loopRange.size);290 continue;291 }292 293 Value iv = ivs[materializedLoopNum++];294 OpFoldResult offset = getAsOpFoldResult(iv);295 offsets.push_back(offset);296 OpFoldResult size =297 getBoundedTileSize(rewriter, loc, loopRange, offset, givenTileSize);298 sizes.push_back(size);299 }300 return {offsets, sizes};301}302 303/// Function to return the bounds of the loops to be generated.304static std::tuple<SmallVector<OpFoldResult>, SmallVector<OpFoldResult>,305 SmallVector<OpFoldResult>>306getLoopBounds(RewriterBase &rewriter, Location loc, ArrayRef<Range> loopRanges,307 ArrayRef<OpFoldResult> givenTileSizes) {308 SmallVector<OpFoldResult> lbs, ubs, steps;309 for (auto [loopRange, givenTileSize] :310 llvm::zip_equal(loopRanges, givenTileSizes)) {311 // No loop if the tile size is 0.312 if (isZeroInteger(givenTileSize))313 continue;314 lbs.push_back(loopRange.offset);315 ubs.push_back(loopRange.size);316 steps.push_back(givenTileSize);317 }318 return {lbs, ubs, steps};319}320 321/// Typedef for function that allows returning additional yielded values during322/// `yieldTiledValuesAndReplace`.323/// - `ivs` induction variable for the loop.324/// - `newBbArgs` basic block arguments corresponding to newly added iter_args.325/// - `tiledValues` the tiled values to return. Must be of same size as326/// `newbbArgs`, each element of this array is inserted into the corresponding327/// element in `newbbArgs`.328/// - `resultOffsets` is of the same size as `tiledValues` and represents329/// the offsets to use when inserting corresponding element from `tiledValues`330/// into the element from `newBbArgs`.331/// - `resultSizes` is of the same size as `tiledValues` and represents332/// the size of the corresponding element from `tiledValues` inserted into333/// the element from `newBbArgs`.334/// In case the method needs to return `failure()` the method is expected335/// to clean up any inserted operations.336using YieldTiledValuesFn = std::function<LogicalResult(337 RewriterBase &rewriter, Location loc, ValueRange ivs, ValueRange newBbArgs,338 SmallVector<Value> &tiledValues,339 SmallVector<SmallVector<OpFoldResult>> &resultOffsets,340 SmallVector<SmallVector<OpFoldResult>> &resultSizes)>;341 342/// Typedef for function that implements the body of a tiled loop.343/// - `ivs` induction variable for the loop.344/// - `tileOffsets` represents offsets for the tiled iteration space.345/// - `tileSizes` represents the sizes for the tiled iteraiton space.346/// - `outerDestinationTensors` tensor that holds the result. Is same size347/// as the destination operands of the original operations.348/// - `tiledResults` results of the tiled computation, corresponds to349/// tiles of the original operation computed by the loop body.350/// Should be same size as the `destinationTensors`351/// - `resultOffsets` is of the same size as `tiledResults` and represents352/// the offset to use when writing the corresponding element from353/// `tiledResults` into `destinationTensors`.354/// - `resultOffsets` is of the same size as `tiledResults` and represents355/// the size to use when writing the corresponding element from356/// `tiledResults` into `destinationTensors`.357/// In case the method needs to return `failure()` the method is expected358/// to clean up any inserted operations.359using GenerateTiledBodyFn = std::function<LogicalResult(360 RewriterBase &rewriter, Location Loc, ValueRange ivs,361 ArrayRef<OpFoldResult> tileOffsets, ArrayRef<OpFoldResult> tileSizes,362 ValueRange outerDestinationTensors, SmallVector<Value> &tiledResults,363 SmallVector<SmallVector<OpFoldResult>> &resultOffsets,364 SmallVector<SmallVector<OpFoldResult>> &resultSizes)>;365 366/// Clones the operation and updates the destination if the operation367/// implements the `DestinationStyleOpInterface`.368static Operation *cloneOpAndUpdateDestinationArgs(RewriterBase &rewriter,369 Operation *op,370 ValueRange newDestArgs) {371 Operation *clonedOp = rewriter.clone(*op);372 if (newDestArgs.empty())373 return clonedOp;374 if (auto destinationStyleOp = dyn_cast<DestinationStyleOpInterface>(clonedOp))375 destinationStyleOp.getDpsInitsMutable().assign(newDestArgs);376 return clonedOp;377}378 379/// Generate the tile-loop nest using `scf.for` operation.380/// - `loopRanges` specifies the lb, ub and step of the untiled iteration space.381/// - `givenTileSizes` is the tile sizes to use. Zero represent untiled loops.382/// - `outerDestinationTensors` are the init values to use for the outer most383/// loop.384/// - `tiledBodyFn` is called to generated the loop body of the inner385/// most386/// loop.387/// Returns the generated `scf.for` loops on success.388static FailureOr<SmallVector<LoopLikeOpInterface>> generateLoopNestUsingForOp(389 RewriterBase &rewriter, Location loc, ArrayRef<Range> loopRanges,390 ArrayRef<OpFoldResult> givenTileSizes, ValueRange outerDestinationTensors,391 GenerateTiledBodyFn tiledBodyFn) {392 assert(!loopRanges.empty() && "unexpected empty loop ranges");393 assert(loopRanges.size() == givenTileSizes.size() &&394 "expected as many tile sizes as loop ranges");395 OpBuilder::InsertionGuard guard(rewriter);396 397 SmallVector<OpFoldResult> lbs, ubs, steps;398 std::tie(lbs, ubs, steps) =399 getLoopBounds(rewriter, loc, loopRanges, givenTileSizes);400 SmallVector<Value> lbVals =401 getValueOrCreateConstantIndexOp(rewriter, loc, lbs);402 SmallVector<Value> ubVals =403 getValueOrCreateConstantIndexOp(rewriter, loc, ubs);404 SmallVector<Value> stepVals =405 getValueOrCreateConstantIndexOp(rewriter, loc, steps);406 407 SmallVector<Value> ivs;408 SmallVector<LoopLikeOpInterface> loops;409 ValueRange innerDestinationTensors(outerDestinationTensors);410 for (auto [lb, ub, step] : llvm::zip_equal(lbVals, ubVals, stepVals)) {411 auto loop =412 scf::ForOp::create(rewriter, loc, lb, ub, step, innerDestinationTensors,413 [](OpBuilder &bodyBuilder, Location bodyLoc,414 Value iv, ValueRange /*iterArgs*/) {});415 loops.push_back(loop);416 ivs.push_back(loop.getInductionVar());417 rewriter.setInsertionPointToEnd(loop.getBody());418 innerDestinationTensors = loop.getRegionIterArgs();419 }420 if (loops.empty())421 return success();422 423 // Compute the `offsets` and `sizes` to use for tiling.424 SmallVector<OpFoldResult> offsets, sizes;425 std::tie(offsets, sizes) =426 getTileOffsetAndSizes(rewriter, loc, ivs, loopRanges, givenTileSizes);427 428 SmallVector<Value> tiledResults;429 SmallVector<SmallVector<OpFoldResult>> resultOffsets, resultSizes;430 if (failed(tiledBodyFn(rewriter, loc, ivs, offsets, sizes,431 innerDestinationTensors, tiledResults, resultOffsets,432 resultSizes))) {433 return rewriter.notifyMatchFailure(434 loc, "failed to generate inner tile loop body");435 }436 if (loops.empty())437 return loops;438 439 assert(tiledResults.size() == innerDestinationTensors.size() &&440 "Number of results of body should be equal to number of iter args");441 442 // 6. Yield all the results of the tiled operation.443 SmallVector<Value> yieldedValues;444 for (auto [tiledValue, destinationTensor, resultOffset, resultSize] :445 llvm::zip_equal(tiledResults, innerDestinationTensors, resultOffsets,446 resultSizes)) {447 SmallVector<OpFoldResult> resultStride(resultOffset.size(),448 rewriter.getIndexAttr(1));449 auto insertSlice = tensor::InsertSliceOp::create(450 rewriter, loc, tiledValue, destinationTensor, resultOffset, resultSize,451 resultStride);452 yieldedValues.push_back(insertSlice);453 }454 scf::YieldOp::create(rewriter, loc, yieldedValues);455 456 // Add the scf.yield operations for all the outer loops.457 for (auto [outerLoop, innerLoop] :458 llvm::zip_equal(MutableArrayRef(loops).drop_back(),459 MutableArrayRef(loops).drop_front())) {460 rewriter.setInsertionPointToEnd(461 cast<scf::ForOp>(outerLoop.getOperation()).getBody());462 scf::YieldOp::create(rewriter, outerLoop.getLoc(), innerLoop->getResults());463 }464 return loops;465}466 467/// Compute the `OpFoldResult`s that represents the multi-dimensional468/// `offset`s and `size`s of the tile of the iteration space that the469/// innermost loop body of the generated tiled loops corresponds to470/// when tiling using `forall` op. This is handle separately due to471/// the special case handling needed for when the tiling is done by472/// specifying number of threads.473static std::tuple<SmallVector<OpFoldResult>, SmallVector<OpFoldResult>>474getTileOffsetAndSizesWithForAllOp(RewriterBase &rewriter, Location loc,475 ValueRange ivs,476 ArrayRef<Range> iterationDomain,477 ArrayRef<OpFoldResult> givenTileSizes,478 ArrayRef<OpFoldResult> numThreads) {479 if (numThreads.empty()) {480 return getTileOffsetAndSizes(rewriter, loc, ivs, iterationDomain,481 givenTileSizes);482 }483 484 SmallVector<OpFoldResult> offsets, sizes;485 int materializedLoopNum = 0;486 487 AffineExpr d0, d1, s0, s1;488 AffineExpr offsetExpr, residualTileSizeExpr;489 bindDims(rewriter.getContext(), d0, d1);490 bindSymbols(rewriter.getContext(), s0, s1);491 offsetExpr = d0 + d1 * s0;492 residualTileSizeExpr = s1 - (d0 + d1 * s0);493 494 for (auto [index, nt, givenTileSize, loopRange] :495 llvm::enumerate(numThreads, givenTileSizes, iterationDomain)) {496 497 // Non-tiled cases, set the offset and size to the498 // `loopRange.offset/size`.499 if (isZeroInteger(nt)) {500 offsets.push_back(loopRange.offset);501 sizes.push_back(loopRange.size);502 continue;503 }504 505 Value iv = ivs[materializedLoopNum++];506 OpFoldResult offset = affine::makeComposedFoldedAffineApply(507 rewriter, loc, offsetExpr,508 ArrayRef<OpFoldResult>{loopRange.offset, iv, givenTileSize});509 OpFoldResult residualTileSize = affine::makeComposedFoldedAffineApply(510 rewriter, loc, residualTileSizeExpr,511 {loopRange.offset, nt, givenTileSize, loopRange.size});512 513 OpFoldResult size = givenTileSize;514 if (!isZeroInteger(residualTileSize)) {515 OpFoldResult sizeMinusOffsetPerThread =516 affine::makeComposedFoldedAffineApply(rewriter, loc, s0 - d0,517 {offset, loopRange.size});518 size = affine::makeComposedFoldedAffineMin(519 rewriter, loc,520 AffineMap::getMultiDimIdentityMap(2, rewriter.getContext()),521 {sizeMinusOffsetPerThread, givenTileSize});522 }523 524 // Consider the case where the original loop was `[0, 100)`.525 // If number of threads are `7`, the tile size would be computed as526 // `ceilDiv(100, 7) = 15`. For the last thread (thread_id = 6)527 // - `offset = 0 + 6 * 15 = 105`528 // - `tileSize = min(15, 100 - 105) = -5`529 // To avoid negative tile sizes, we need to do a further530 // `nonNegativeTileSize = affine.max(0, tileSize)`.531 // This `max` can be avoided if532 // `offset + tileSize * (numThreads - 1) < (ub - lb)`533 if (!canOmitTileOffsetInBoundsCheck(givenTileSize, nt, loopRange.size)) {534 AffineMap maxMap =535 AffineMap::getMultiDimIdentityMap(2, rewriter.getContext());536 size = affine::makeComposedFoldedAffineMax(537 rewriter, loc, maxMap, {rewriter.getIndexAttr(0), size});538 }539 540 offsets.push_back(offset);541 sizes.push_back(size);542 }543 return {offsets, sizes};544}545 546/// Generate the tile-loop nest using `scf.forall` operation.547/// - `loopRanges` specifies the lb, ub and step of the untiled iteration space.548/// - `giventileSizes` is the tile sizes to use. Zero represent untiled loops.549/// - `outerDestinationTensors` are the init values to use for the loop.550/// - `mappingVector` is the mapping attributes to use for loop construction.551/// Can be empty.552/// - `tiledBodyFn` is called to generated the loop body of the inner553/// most554/// loop.555/// Returns the generated `scf.forall` loop on success.556static FailureOr<SmallVector<LoopLikeOpInterface>>557generateLoopNestUsingForallOp(RewriterBase &rewriter, Location loc,558 ArrayRef<Range> loopRanges,559 ArrayRef<OpFoldResult> givenTileSizes,560 ArrayRef<OpFoldResult> numThreads,561 ArrayRef<Attribute> mappingVector,562 ValueRange outerDestinationTensors,563 GenerateTiledBodyFn tiledBodyFn) {564 assert(!loopRanges.empty() && "unexpected empty loop ranges");565 assert(loopRanges.size() == givenTileSizes.size() &&566 "expected as many tile sizes as loop ranges");567 OpBuilder::InsertionGuard guard(rewriter);568 569 std::optional<ArrayAttr> mappingAttr;570 if (!mappingVector.empty())571 mappingAttr = rewriter.getArrayAttr(mappingVector);572 573 scf::ForallOp forallOp;574 bool useNumThreads = !numThreads.empty();575 576 SmallVector<LoopLikeOpInterface> loops;577 if (useNumThreads) {578 // Prune the zero numthreads.579 SmallVector<OpFoldResult> nonZeroNumThreads;580 for (auto nt : numThreads) {581 if (isZeroInteger(nt))582 continue;583 nonZeroNumThreads.push_back(nt);584 }585 forallOp = scf::ForallOp::create(rewriter, loc, nonZeroNumThreads,586 outerDestinationTensors, mappingAttr);587 } else {588 SmallVector<OpFoldResult> lbs, ubs, steps;589 std::tie(lbs, ubs, steps) =590 getLoopBounds(rewriter, loc, loopRanges, givenTileSizes);591 forallOp = scf::ForallOp::create(rewriter, loc, lbs, ubs, steps,592 outerDestinationTensors, mappingAttr);593 }594 loops.push_back(forallOp);595 596 rewriter.setInsertionPoint(forallOp.getTerminator());597 ValueRange innerDestinationTensors = forallOp.getRegionOutArgs();598 SmallVector<Value> ivs = forallOp.getInductionVars();599 600 // Compute the `offsets` and `sizes` to use for tiling.601 SmallVector<OpFoldResult> offsets, sizes;602 std::tie(offsets, sizes) = getTileOffsetAndSizesWithForAllOp(603 rewriter, loc, ivs, loopRanges, givenTileSizes, numThreads);604 605 SmallVector<Value> tiledResults;606 SmallVector<SmallVector<OpFoldResult>> resultOffsets, resultSizes;607 if (failed(tiledBodyFn(rewriter, loc, ivs, offsets, sizes,608 innerDestinationTensors, tiledResults, resultOffsets,609 resultSizes)))610 return rewriter.notifyMatchFailure(loc, "failed to generate loop body");611 612 rewriter.setInsertionPointToEnd(forallOp.getTerminator().getBody());613 for (auto [tiledValue, destinationTensor, resultOffset, resultSize] :614 llvm::zip_equal(tiledResults, innerDestinationTensors, resultOffsets,615 resultSizes)) {616 SmallVector<OpFoldResult> resultStride(resultOffset.size(),617 rewriter.getIndexAttr(1));618 619 tensor::ParallelInsertSliceOp::create(rewriter, loc, tiledValue,620 destinationTensor, resultOffset,621 resultSize, resultStride);622 }623 return loops;624}625 626/// Generate the tile-loop nest using custom loop operation.627/// - `loopRanges` specifies the lb, ub and step of the untiled iteration space.628/// - `tileSizes` is the tile sizes to use. Zero represent untiled loops.629/// - `destinationTensors` are the init values to use for the outer most loop.630/// - `mappingVector` is the mapping attributes to use for loop construction.631/// Can be empty.632/// - `tiledBodyFn` is called to generated the loop body of the inner633/// most634/// loop.635/// Returns the generated `scf.forall` loop on success.636static FailureOr<SmallVector<LoopLikeOpInterface>>637generateLoopNestUsingCustomOp(638 RewriterBase &rewriter, Location loc, ArrayRef<Range> loopRanges,639 ArrayRef<OpFoldResult> givenTileSizes, ValueRange outerDestinationTensors,640 const scf::SCFTilingOptions::GenerateLoopHeaderFn &generateLoopHeaderFn,641 const scf::SCFTilingOptions::GenerateLoopTerminatorFn642 &generateLoopTerminatorFn,643 GenerateTiledBodyFn tiledBodyFn) {644 assert(!loopRanges.empty() && "unexpected empty loop ranges");645 assert(loopRanges.size() == givenTileSizes.size() &&646 "expected as many tile sizes as loop ranges");647 assert(generateLoopHeaderFn && generateLoopTerminatorFn &&648 "expected loop header/terminator generation function");649 OpBuilder::InsertionGuard guard(rewriter);650 651 FailureOr<scf::SCFTilingOptions::CustomLoopHeaderInfo> loopHeaderInfo =652 generateLoopHeaderFn(rewriter, loc, loopRanges, givenTileSizes,653 outerDestinationTensors);654 if (failed(loopHeaderInfo)) {655 return failure();656 }657 658 SmallVector<Value> ivs;659 SmallVector<Value> tiledResults;660 SmallVector<SmallVector<OpFoldResult>> resultOffsets, resultSizes;661 if (failed(tiledBodyFn(rewriter, loc, ivs, loopHeaderInfo->tileOffset,662 loopHeaderInfo->tileSizes,663 loopHeaderInfo->destinationTensors, tiledResults,664 resultOffsets, resultSizes))) {665 return failure();666 }667 668 if (failed(generateLoopTerminatorFn(rewriter, loc, loopHeaderInfo->loops,669 tiledResults, resultOffsets, resultSizes,670 loopHeaderInfo->destinationTensors))) {671 return failure();672 }673 674 return loopHeaderInfo->loops;675}676 677/// Generate the tile-loop nest using the loop construct specifed in `options`.678/// - `options`: Tiling options specified.679/// - `loopRanges` specifies the lb, ub and step of the untiled iteration space.680/// - `tileSizes` is the tile sizes to use. Zero represent untiled loops.681/// - `outerDestinationTensors` are the init values to use for the outer most682/// loop.683/// - `yieldTiledValuesFn` is called to generated the loop body of the inner684/// most685/// loop.686/// Returns the generated loops on success.687static FailureOr<SmallVector<LoopLikeOpInterface>> generateLoopNest(688 RewriterBase &rewriter, Location loc, const scf::SCFTilingOptions &options,689 ArrayRef<Range> loopRanges, ArrayRef<OpFoldResult> givenTileSizes,690 ArrayRef<OpFoldResult> numThreads, ValueRange destinationTensors,691 GenerateTiledBodyFn tiledBodyFn) {692 // If the tile sizes are all zero, no loops are generated. Just call the693 // callback function to handle untiled case.694 if (llvm::all_of(givenTileSizes, isZeroInteger)) {695 SmallVector<Value> tiledResults;696 SmallVector<SmallVector<OpFoldResult>> resultOffsets, resultSizes;697 auto tileOffsets =698 llvm::map_to_vector(loopRanges, [](Range r) { return r.offset; });699 auto tileSizes =700 llvm::map_to_vector(loopRanges, [](Range r) { return r.size; });701 if (failed(tiledBodyFn(rewriter, loc, ValueRange{}, tileOffsets, tileSizes,702 destinationTensors, tiledResults, resultOffsets,703 resultSizes))) {704 return failure();705 }706 return SmallVector<LoopLikeOpInterface>{};707 }708 if (options.loopType == scf::SCFTilingOptions::LoopType::ForOp) {709 return generateLoopNestUsingForOp(rewriter, loc, loopRanges, givenTileSizes,710 destinationTensors, tiledBodyFn);711 }712 if (options.loopType == scf::SCFTilingOptions::LoopType::ForallOp) {713 return generateLoopNestUsingForallOp(714 rewriter, loc, loopRanges, givenTileSizes, numThreads,715 options.mappingVector, destinationTensors, tiledBodyFn);716 }717 if (options.loopType == scf::SCFTilingOptions::LoopType::CustomOp) {718 return generateLoopNestUsingCustomOp(719 rewriter, loc, loopRanges, givenTileSizes, destinationTensors,720 options.generateLoopHeaderFn, options.generateLoopTerminatorFn,721 tiledBodyFn);722 }723 return rewriter.notifyMatchFailure(loc, "unhandled loop type");724}725 726static FailureOr<SmallVector<Value>> createInitialTensorsForTiling(727 RewriterBase &rewriter, TilingInterface op,728 ReductionTilingStrategy reductionStrategy, ArrayRef<Range> iterationDomain,729 ArrayRef<OpFoldResult> numThreads, ArrayRef<OpFoldResult> givenTileSizes,730 const SetVector<unsigned> &reductionDims) {731 SmallVector<Value> initTensors;732 Location loc = op->getLoc();733 if (reductionStrategy == ReductionTilingStrategy::FullReduction) {734 if (failed(tensor::getOrCreateDestinations(rewriter, loc, op, initTensors)))735 return failure();736 return initTensors;737 }738 739 auto redOp = dyn_cast<PartialReductionOpInterface>(op.getOperation());740 if (!redOp) {741 return op->emitOpError(742 "PartialReductionOuterReduction tiling strategy is only supported for "743 "operations implementing PartialReductionOpInterface");744 }745 SmallVector<OpFoldResult> sizes(iterationDomain.size());746 AffineExpr s0, s1, s2;747 bindSymbols(rewriter.getContext(), s0, s1, s2);748 AffineExpr sizeExpr = ((s0 - s1).ceilDiv(s2));749 AffineExpr divExpr = s0.ceilDiv(s1);750 for (auto [index, domain, tileSize] :751 llvm::enumerate(iterationDomain, givenTileSizes)) {752 if (!numThreads.empty()) {753 // Untiled case.754 if (isConstantIntValue(numThreads[index], 0)) {755 sizes[index] = affine::makeComposedFoldedAffineApply(756 rewriter, op.getLoc(), sizeExpr,757 {domain.size, domain.offset, domain.stride});758 continue;759 }760 sizes[index] = numThreads[index];761 continue;762 }763 764 // Non reduction dimensions/non-tiled dimensions.765 if (!reductionDims.contains(index) || isConstantIntValue(tileSize, 0)) {766 sizes[index] = affine::makeComposedFoldedAffineApply(767 rewriter, op.getLoc(), sizeExpr,768 {domain.size, domain.offset, domain.stride});769 continue;770 }771 772 if (reductionStrategy ==773 ReductionTilingStrategy::PartialReductionOuterReduction) {774 sizes[index] = tileSize;775 continue;776 }777 778 assert(reductionStrategy ==779 ReductionTilingStrategy::PartialReductionOuterParallel);780 OpFoldResult normalizedRange = affine::makeComposedFoldedAffineApply(781 rewriter, op.getLoc(), sizeExpr,782 {domain.size, domain.offset, domain.stride});783 sizes[index] = affine::makeComposedFoldedAffineApply(784 rewriter, op.getLoc(), divExpr, {normalizedRange, tileSize});785 }786 return redOp.generateInitialTensorForPartialReduction(rewriter, loc, sizes,787 reductionDims);788}789 790/// For the case of `ReductionTilingStrategy::PartialReductionOuterParallel`791/// the `PartialReductionOpInterface` methods need the index of the parallel792/// split reduction being executed.793static SmallVector<OpFoldResult>794getSplitReductionIvs(RewriterBase &rewriter, Location loc,795 ReductionTilingStrategy reductionStrategy, ValueRange ivs,796 ArrayRef<OpFoldResult> numThreads,797 ArrayRef<OpFoldResult> givenTileSizes,798 const SetVector<unsigned> &reductionDims) {799 SmallVector<OpFoldResult> splitReductionIvs;800 splitReductionIvs.resize(reductionDims.size(), rewriter.getIndexAttr(0));801 AffineExpr s0, s1;802 bindSymbols(rewriter.getContext(), s0, s1);803 AffineExpr divExpr = s0.floorDiv(s1);804 int ivIndex = 0;805 if (reductionStrategy ==806 ReductionTilingStrategy::PartialReductionOuterParallel) {807 for (auto [index, reductionDim] : llvm::enumerate(reductionDims)) {808 if (!numThreads.empty()) {809 splitReductionIvs[index] = ivs[ivIndex++];810 continue;811 }812 splitReductionIvs[index] = affine::makeComposedFoldedAffineApply(813 rewriter, loc, divExpr,814 ArrayRef<OpFoldResult>{ivs[ivIndex++], givenTileSizes[reductionDim]});815 }816 }817 return splitReductionIvs;818}819 820static FailureOr<TilingResult>821getTiledImplementation(RewriterBase &rewriter, TilingInterface op,822 ReductionTilingStrategy reductionStrategy,823 ValueRange regionIterArg, ArrayRef<OpFoldResult> offsets,824 ArrayRef<OpFoldResult> sizes, ValueRange ivs,825 ArrayRef<OpFoldResult> numThreads,826 ArrayRef<OpFoldResult> givenTileSizes,827 const SetVector<unsigned> &reductionDims) {828 if (reductionStrategy == ReductionTilingStrategy::FullReduction) {829 return op.getTiledImplementation(rewriter, offsets, sizes);830 }831 832 auto redOp = dyn_cast<PartialReductionOpInterface>(op.getOperation());833 if (!redOp) {834 return rewriter.notifyMatchFailure(835 op, "PartialReductionOuterReduction tiling strategy is only "836 "supported for operations "837 "implementing PartialReductionOpInterface");838 }839 840 SmallVector<OpFoldResult> splitReductionIvs =841 getSplitReductionIvs(rewriter, op.getLoc(), reductionStrategy, ivs,842 numThreads, givenTileSizes, reductionDims);843 return redOp.tileToPartialReduction(rewriter, op.getLoc(), reductionStrategy,844 regionIterArg, offsets, sizes,845 reductionDims, splitReductionIvs);846}847 848static LogicalResult getResultTilePosition(849 RewriterBase &rewriter, ReductionTilingStrategy reductionStrategy,850 int64_t index, Value tiledResult, TilingInterface op,851 ArrayRef<OpFoldResult> offsets, ArrayRef<OpFoldResult> sizes,852 ValueRange ivs, ArrayRef<OpFoldResult> numThreads,853 ArrayRef<OpFoldResult> givenTileSizes,854 const SetVector<unsigned> &reductionDims,855 SmallVector<OpFoldResult> &resultOffset,856 SmallVector<OpFoldResult> &resultSize) {857 858 if (reductionStrategy == ReductionTilingStrategy::FullReduction) {859 return op.getResultTilePosition(rewriter, index, offsets, sizes,860 resultOffset, resultSize);861 }862 auto redOp = dyn_cast<PartialReductionOpInterface>(op.getOperation());863 if (!redOp) {864 return rewriter.notifyMatchFailure(865 op, "PartialReductionOuterReduction tiling strategy is only supported"866 "for operations implementing PartialReductionOpInterface");867 }868 SmallVector<OpFoldResult> splitReductionIvs =869 getSplitReductionIvs(rewriter, op.getLoc(), reductionStrategy, ivs,870 numThreads, givenTileSizes, reductionDims);871 return redOp.getPartialResultTilePosition(872 rewriter, index, reductionStrategy, offsets, sizes, reductionDims,873 splitReductionIvs, resultOffset, resultSize);874}875 876static FailureOr<MergeResult>877mergeTilingResults(RewriterBase &rewriter, TilingInterface op,878 ReductionTilingStrategy reductionStrategy,879 const SetVector<unsigned> &reductionDims,880 ValueRange partialResults) {881 assert(reductionStrategy != ReductionTilingStrategy::FullReduction &&882 "expected merge to be called for only partial reduction cases");883 884 auto redOp = dyn_cast<PartialReductionOpInterface>(op.getOperation());885 if (!redOp) {886 return rewriter.notifyMatchFailure(887 op, "PartialReductionOuterReduction tiling strategy is only "888 "supported for operations "889 "implementing PartialReductionOpInterface");890 }891 return redOp.mergeReductions(rewriter, op.getLoc(), partialResults,892 reductionDims);893}894 895/// Append the specified additional `newInitOperands` operands to the896/// loops existing `init` operands (or similar), and replace `loopOp` with897/// the new loop that has the additional init operands. The loop body of898/// this loop is moved over to the new loop. `yieldTiledValuesFn`899/// is called to get the new tiled values returned, and the offset900/// and sizes at which the tiled value is inserted into the901/// new region iter_args that correspond to the newly added init operands.902template <typename LoopType>903FailureOr<LoopLikeOpInterface>904yieldTiledValuesAndReplaceLoop(LoopType loopOp, RewriterBase &rewriter,905 ValueRange newInitOperands,906 YieldTiledValuesFn yieldTiledValuesFn) {907 return rewriter.notifyMatchFailure(loopOp, "unhandled loop type");908}909 910/// Implementation of `yieldTiledValuesAndReplaceLoop` for `scf.for`.911template <>912FailureOr<LoopLikeOpInterface> yieldTiledValuesAndReplaceLoop<scf::ForOp>(913 scf::ForOp loopOp, RewriterBase &rewriter, ValueRange newInitOperands,914 YieldTiledValuesFn yieldTiledValuesFn) {915 OpBuilder::InsertionGuard g(rewriter);916 Location loc = loopOp.getLoc();917 rewriter.setInsertionPoint(loopOp);918 919 auto inits = llvm::to_vector(loopOp.getInitArgs());920 inits.append(newInitOperands.begin(), newInitOperands.end());921 auto newLoop = scf::ForOp::create(922 rewriter, loc, loopOp.getLowerBound(), loopOp.getUpperBound(),923 loopOp.getStep(), inits, [](OpBuilder &, Location, Value, ValueRange) {},924 loopOp.getUnsignedCmp());925 926 // Move the loop body to the new op.927 Block *loopBody = loopOp.getBody();928 Block *newLoopBody = newLoop.getBody();929 rewriter.mergeBlocks(930 loopBody, newLoopBody,931 newLoopBody->getArguments().take_front(loopBody->getNumArguments()));932 933 auto yieldOp = cast<scf::YieldOp>(newLoopBody->getTerminator());934 rewriter.setInsertionPoint(yieldOp);935 936 SmallVector<Value> tiledValues;937 SmallVector<SmallVector<OpFoldResult>> resultOffsets, resultSizes;938 ValueRange newRegionIterArgs =939 newLoop.getRegionIterArgs().take_back(newInitOperands.size());940 if (failed(yieldTiledValuesFn(rewriter, loc, newLoop.getInductionVar(),941 newRegionIterArgs, tiledValues, resultOffsets,942 resultSizes))) {943 rewriter.eraseOp(newLoop);944 return rewriter.notifyMatchFailure(loopOp, "failed to get tiled values");945 }946 947 SmallVector<Value> newYieldValues = llvm::to_vector(yieldOp.getOperands());948 for (auto [tiledValue, regionIterArg, resultOffset, resultSize] :949 llvm::zip_equal(tiledValues, newRegionIterArgs, resultOffsets,950 resultSizes)) {951 SmallVector<OpFoldResult> resultStride(resultOffset.size(),952 rewriter.getIndexAttr(1));953 Value insert = tensor::InsertSliceOp::create(954 rewriter, yieldOp->getLoc(), tiledValue, regionIterArg, resultOffset,955 resultSize, resultStride);956 newYieldValues.push_back(insert);957 }958 959 rewriter.replaceOpWithNewOp<scf::YieldOp>(yieldOp, newYieldValues);960 rewriter.replaceOp(loopOp,961 newLoop->getResults().take_front(loopOp.getNumResults()));962 return cast<LoopLikeOpInterface>(newLoop.getOperation());963}964 965/// Implementation of `yieldTiledValuesAndReplaceLoop` for `scf.forall`966template <>967FailureOr<LoopLikeOpInterface> yieldTiledValuesAndReplaceLoop<scf::ForallOp>(968 scf::ForallOp loopOp, RewriterBase &rewriter, ValueRange newInitOperands,969 YieldTiledValuesFn yieldTiledValuesFn) {970 OpBuilder::InsertionGuard g(rewriter);971 Location loc = loopOp.getLoc();972 rewriter.setInsertionPoint(loopOp);973 auto inits = llvm::to_vector(loopOp.getOutputs());974 inits.append(newInitOperands.begin(), newInitOperands.end());975 auto newLoop = scf::ForallOp::create(976 rewriter, loc, loopOp.getMixedLowerBound(), loopOp.getMixedUpperBound(),977 loopOp.getMixedStep(), inits, loopOp.getMapping(),978 [](OpBuilder &, Location, ValueRange) {});979 980 // Move the region of the current block to the newly created op.981 Block *loopBody = loopOp.getBody();982 Block *newLoopBody = newLoop.getBody();983 rewriter.mergeBlocks(984 loopBody, newLoopBody,985 newLoopBody->getArguments().take_front(loopBody->getNumArguments()));986 987 auto terminator = cast<scf::InParallelOp>(newLoopBody->getTerminator());988 rewriter.setInsertionPoint(terminator);989 SmallVector<Value> tiledValues;990 SmallVector<SmallVector<OpFoldResult>> resultOffsets, resultSizes;991 ValueRange regionIterArgs =992 newLoop.getRegionIterArgs().take_back(newInitOperands.size());993 if (failed(yieldTiledValuesFn(rewriter, loc, newLoop.getInductionVars(),994 regionIterArgs, tiledValues, resultOffsets,995 resultSizes))) {996 rewriter.eraseOp(newLoop);997 return rewriter.notifyMatchFailure(loopOp,998 "failed to get yielded tiled values");999 }1000 1001 // Update the terminator.1002 rewriter.setInsertionPointToEnd(terminator.getBody());1003 1004 for (auto [tiledValue, iterArg, resultOffset, resultSize] : llvm::zip_equal(1005 tiledValues, regionIterArgs, resultOffsets, resultSizes)) {1006 SmallVector<OpFoldResult> resultStride(resultOffset.size(),1007 rewriter.getIndexAttr(1));1008 tensor::ParallelInsertSliceOp::create(rewriter, terminator.getLoc(),1009 tiledValue, iterArg, resultOffset,1010 resultSize, resultStride);1011 }1012 1013 rewriter.replaceOp(loopOp,1014 newLoop->getResults().take_front(loopOp.getNumResults()));1015 return cast<LoopLikeOpInterface>(newLoop.getOperation());1016}1017 1018/// Implementation of `yieldTiledValuesAndReplaceLoop` for1019/// `LoopLikeOpInterface`, that just dispatches to the implementation for each1020/// supported loop type.1021FailureOr<LoopLikeOpInterface> yieldTiledValuesAndReplaceLoop(1022 LoopLikeOpInterface loopLikeOp, RewriterBase &rewriter,1023 ValueRange newInitOperands, YieldTiledValuesFn yieldTiledValuesFn) {1024 return TypeSwitch<Operation *, FailureOr<LoopLikeOpInterface>>(1025 loopLikeOp.getOperation())1026 .Case<scf::ForOp, scf::ForallOp>(1027 [&](auto loopOp) -> FailureOr<LoopLikeOpInterface> {1028 return yieldTiledValuesAndReplaceLoop(1029 loopOp, rewriter, newInitOperands, yieldTiledValuesFn);1030 })1031 .Default([&](auto loopOp) -> FailureOr<LoopLikeOpInterface> {1032 return rewriter.notifyMatchFailure(loopOp, "unhandled loop type");1033 });1034}1035 1036/// Method to add new init values to a loop nest. Updates `loops` in-place1037/// with new loops that use the `newInitValues`. The outer-loops are updated1038/// to yield the new result values of the inner loop. For the innermost loop,1039/// the call back `getNewYields` is invoked to get the additional values to1040/// yield form the innermost loop.1041static LogicalResult addInitOperandsToLoopNest(1042 RewriterBase &rewriter, MutableArrayRef<LoopLikeOpInterface> loops,1043 ValueRange newInitValues, YieldTiledValuesFn getNewTiledYieldsFn) {1044 if (loops.empty())1045 return success();1046 OpBuilder::InsertionGuard g(rewriter);1047 rewriter.setInsertionPoint(loops.front());1048 1049 SmallVector<Value> ivs;1050 for (auto &loop : loops.drop_back()) {1051 rewriter.setInsertionPoint(loop);1052 1053 // if loops.size() > 1 we assume that scf.for is used for the loops.1054 auto forLoop = cast<scf::ForOp>(loop.getOperation());1055 1056 // Create a new loop with the new init values for this loop.1057 SmallVector<Value> newInits = llvm::to_vector(forLoop.getInitArgs());1058 newInits.append(newInitValues.begin(), newInitValues.end());1059 auto newLoop = scf::ForOp::create(1060 rewriter, forLoop.getLoc(), forLoop.getLowerBound(),1061 forLoop.getUpperBound(), forLoop.getStep(), newInits,1062 [&](OpBuilder &b, Location loc, Value iv, ValueRange iterArgs) {},1063 forLoop.getUnsignedCmp());1064 1065 // Merge the body of the new loop with the body of the old loops.1066 SmallVector<Value> sourceBlockArgs;1067 sourceBlockArgs.push_back(newLoop.getInductionVar());1068 auto newRegionIterArgs = newLoop.getRegionIterArgs();1069 sourceBlockArgs.append(1070 newRegionIterArgs.begin(),1071 std::next(newRegionIterArgs.begin(), forLoop.getNumResults()));1072 rewriter.mergeBlocks(forLoop.getBody(), newLoop.getBody(), sourceBlockArgs);1073 rewriter.replaceOp(1074 forLoop, newLoop.getResults().take_front(forLoop.getNumResults()));1075 loop = newLoop;1076 ivs.push_back(newLoop.getInductionVar());1077 newInitValues = newLoop.getRegionIterArgs().take_back(newInitValues.size());1078 }1079 1080 // Update the loop body of the innermost loop to get new yield values.1081 LoopLikeOpInterface innerMostLoop = loops.back();1082 FailureOr<LoopLikeOpInterface> newInnerMostLoop =1083 yieldTiledValuesAndReplaceLoop(innerMostLoop, rewriter, newInitValues,1084 getNewTiledYieldsFn);1085 1086 if (failed(newInnerMostLoop))1087 return innerMostLoop.emitOpError("failed to return additional yields");1088 loops.back() = newInnerMostLoop.value();1089 1090 // Make all other loops except the innermost loops yield the values returned1091 // by the inner loop.1092 for (auto [outerLoop, innerLoop] :1093 llvm::zip_equal(loops.drop_back(), loops.drop_front())) {1094 // Again assume that all the outer loops are scf.for operations.1095 auto outerForLoop = cast<scf::ForOp>(outerLoop.getOperation());1096 auto outerLoopYield =1097 cast<scf::YieldOp>(outerForLoop.getBody()->getTerminator());1098 SmallVector<Value> newYields =1099 llvm::to_vector(outerLoopYield.getOperands());1100 ValueRange additionalYields =1101 innerLoop->getResults().take_back(newInitValues.size());1102 newYields.append(additionalYields.begin(), additionalYields.end());1103 rewriter.setInsertionPoint(outerLoopYield);1104 rewriter.replaceOpWithNewOp<scf::YieldOp>(outerLoopYield, newYields);1105 }1106 return success();1107}1108 1109/// Implementation of tiling transformation of `op` that implements the1110/// `TilingInterface` using `scf.for` to iterate over the tiles.1111FailureOr<scf::SCFTilingResult>1112mlir::scf::tileUsingSCF(RewriterBase &rewriter, TilingInterface op,1113 const scf::SCFTilingOptions &options) {1114 if (failed(verifyOptions(rewriter, op.getLoc(), options))) {1115 return failure();1116 }1117 1118 OpBuilder::InsertionGuard guard(rewriter);1119 rewriter.setInsertionPointAfter(op);1120 1121 // 1. Get the range of the loops that are represented by the operation.1122 SmallVector<Range> iterationDomain = op.getIterationDomain(rewriter);1123 1124 // 2. Materialize the tile sizes and/or number of threads;1125 SmallVector<OpFoldResult> givenTileSizes, numThreads;1126 std::tie(givenTileSizes, numThreads) =1127 getUserTileSizesAndNumThreads(rewriter, op, iterationDomain, options);1128 1129 // Check if it is safe to tile. This is hold over from previous iterations1130 // of tile to for-all. Consider dropping it.1131 if (failed(checkTileSizes(op, options.loopType, options.reductionStrategy,1132 givenTileSizes, numThreads))) {1133 return failure();1134 }1135 1136 // Get the reduction dims1137 SetVector<unsigned> reductionDims =1138 getSanitizedReductionDims(givenTileSizes, options);1139 1140 // 3. If there is an interchange specified, permute the iteration domain and1141 // the tile sizes.1142 SmallVector<int64_t> interchangeVector;1143 if (!options.interchangeVector.empty()) {1144 interchangeVector = fillInterchangeVector(options.interchangeVector,1145 iterationDomain.size());1146 assert(isPermutationVector(interchangeVector) &&1147 "expected interchange vector to be a permutation");1148 1149 applyPermutationToVector(iterationDomain, interchangeVector);1150 applyPermutationToVector(givenTileSizes, interchangeVector);1151 if (!numThreads.empty())1152 applyPermutationToVector(numThreads, interchangeVector);1153 }1154 1155 FailureOr<TilingResult> tilingResult;1156 // 4. Define the lambda function used later to generate the body of the1157 // innermost tiled loop.1158 GenerateTiledBodyFn innerYieldTiledValuesFn =1159 [&](RewriterBase &rewriter, Location loc, ValueRange ivs,1160 ArrayRef<OpFoldResult> tileOffsets, ArrayRef<OpFoldResult> tileSizes,1161 ValueRange regionIterArgs, SmallVector<Value> &tiledResults,1162 SmallVector<SmallVector<OpFoldResult>> &resultOffsets,1163 SmallVector<SmallVector<OpFoldResult>> &resultSizes)1164 -> LogicalResult {1165 // 4b. If interchange was provided, apply inverse of the interchange1166 // to get back the offsets/sizes in the order to be specified.1167 SmallVector<OpFoldResult> tileOffsetsVec = llvm::to_vector(tileOffsets);1168 SmallVector<OpFoldResult> tileSizesVec = llvm::to_vector(tileSizes);1169 if (!interchangeVector.empty()) {1170 auto inversePermutation = invertPermutationVector(interchangeVector);1171 applyPermutationToVector(tileOffsetsVec, inversePermutation);1172 applyPermutationToVector(tileSizesVec, inversePermutation);1173 }1174 1175 // 5. Generate the tiled implementation within the inner most loop.1176 1177 // 5a. Clone the operation within the loop body.1178 auto clonedOp = cast<TilingInterface>(1179 cloneOpAndUpdateDestinationArgs(rewriter, op, regionIterArgs));1180 1181 // 5b. Early return cloned op if tiling is not happening. We can not1182 // return the original op because it could lead to `rewriter.replaceOp(op,1183 // op->getResults())` and users would get crash.1184 if (llvm::all_of(givenTileSizes, isZeroInteger)) {1185 tiledResults.append(clonedOp->result_begin(), clonedOp->result_end());1186 tilingResult =1187 TilingResult{/*tiledOps=*/{clonedOp}, clonedOp->getResults(),1188 /*generatedSlices=*/{}};1189 return success();1190 }1191 1192 // 5c. Tile the cloned operation.1193 tilingResult =1194 getTiledImplementation(rewriter, clonedOp, options.reductionStrategy,1195 regionIterArgs, tileOffsetsVec, tileSizesVec,1196 ivs, numThreads, givenTileSizes, reductionDims);1197 if (failed(tilingResult)) {1198 rewriter.eraseOp(clonedOp);1199 return op.emitOpError("faild to tile operation");1200 }1201 1202 // 5d. Delete the cloned operation.1203 rewriter.eraseOp(clonedOp);1204 1205 // 5e. Compute the offsets at which the result values are to be inserted1206 // back into its destinations.1207 for (auto [index, tiledValue] :1208 llvm::enumerate(tilingResult->tiledValues)) {1209 tiledResults.push_back(tiledValue);1210 SmallVector<OpFoldResult> resultOffset, resultSize;1211 if (failed(getResultTilePosition(1212 rewriter, options.reductionStrategy, index, tiledValue, op,1213 tileOffsetsVec, tileSizesVec, ivs, numThreads, givenTileSizes,1214 reductionDims, resultOffset, resultSize))) {1215 for (auto op : tilingResult->tiledOps) {1216 rewriter.eraseOp(op);1217 }1218 return rewriter.notifyMatchFailure(1219 op, "failed to get slice of result produced");1220 }1221 resultOffsets.emplace_back(std::move(resultOffset));1222 resultSizes.emplace_back(std::move(resultSize));1223 }1224 1225 return success();1226 };1227 1228 // 6. Find the destination tensors to use for the operation.1229 FailureOr<SmallVector<Value>> maybeInits = createInitialTensorsForTiling(1230 rewriter, op, options.reductionStrategy, iterationDomain, numThreads,1231 givenTileSizes, reductionDims);1232 if (failed(maybeInits)) {1233 return rewriter.notifyMatchFailure(1234 op, "unable to create initial tensors for tiling");1235 }1236 SmallVector<Value> &initTensors = maybeInits.value();1237 1238 // 7. Generate the tiled loops nest using the callback defined above.1239 SmallVector<LoopLikeOpInterface> loops;1240 {1241 FailureOr<SmallVector<LoopLikeOpInterface>> loopsOr = generateLoopNest(1242 rewriter, op.getLoc(), options, iterationDomain, givenTileSizes,1243 numThreads, initTensors, innerYieldTiledValuesFn);1244 if (failed(loopsOr))1245 return op.emitOpError("failed to generate tiling loops");1246 assert(succeeded(tilingResult) &&1247 "expected tiling result to be computed after loop generation");1248 std::swap(loops, loopsOr.value());1249 }1250 1251 if (loops.empty()) {1252 // If loops are empty, the tiled op is used as the replacement for the1253 // untiled op.1254 return scf::SCFTilingResult{tilingResult->tiledOps,1255 initTensors,1256 loops,1257 tilingResult->tiledValues,1258 tilingResult->generatedSlices,1259 {}};1260 }1261 1262 auto loopResults = llvm::map_to_vector(loops.front()->getResults(),1263 [](OpResult r) -> Value { return r; });1264 1265 // For the full reduction case, there is nothing more to do.1266 if (options.reductionStrategy == ReductionTilingStrategy::FullReduction) {1267 return scf::SCFTilingResult{1268 tilingResult->tiledOps, initTensors, loops, loopResults,1269 tilingResult->generatedSlices, {}};1270 }1271 1272 // The results of the loop needs to be merged.1273 FailureOr<MergeResult> mergeResult = mergeTilingResults(1274 rewriter, op, options.reductionStrategy, reductionDims, loopResults);1275 if (failed(mergeResult)) {1276 return rewriter.notifyMatchFailure(1277 op, "Failed to merge partial results from tiling");1278 }1279 return scf::SCFTilingResult{tilingResult->tiledOps,1280 initTensors,1281 loops,1282 mergeResult->replacements,1283 tilingResult->generatedSlices,1284 mergeResult->mergeOps};1285}1286 1287FailureOr<scf::SCFTilingResult>1288mlir::scf::tileReductionUsingScf(RewriterBase &b,1289 PartialReductionOpInterface op,1290 ArrayRef<OpFoldResult> tileSize) {1291 scf::SCFTilingOptions options;1292 options.setLoopType(scf::SCFTilingOptions::LoopType::ForOp);1293 options.setReductionTilingStrategy(1294 ReductionTilingStrategy::PartialReductionOuterReduction);1295 options.setTileSizes(tileSize);1296 SmallVector<unsigned> reductionDims;1297 for (auto [index, iteratorType] : llvm::enumerate(op.getLoopIteratorTypes()))1298 if (iteratorType == utils::IteratorType::reduction)1299 reductionDims.push_back(index);1300 options.setReductionDims(reductionDims);1301 return tileUsingSCF(b, op, options);1302}1303 1304//===----------------------------------------------------------------------===//1305// tileConsumerAndFuseProducersUsingSCF implementation.1306//===----------------------------------------------------------------------===//1307 1308/// Return the untiled producer whose slice is used in a tiled consumer. The1309/// method traverses the tile loop nest (`loops`) if needed, and returns the1310/// `iter_args` of the outer most that is encountered. Traversing the1311/// iter_args indicates that this is a destination operand of the consumer. If1312/// there was no loop traversal needed, the second value of the returned tuple1313/// is empty.1314static std::tuple<OpResult, std::optional<OpOperand *>>1315getUntiledProducerFromSliceSource(OpOperand *source,1316 ArrayRef<LoopLikeOpInterface> loops) {1317 std::optional<OpOperand *> destinationIterArg;1318 assert(!loops.empty() && "expected non empty loops container");1319 auto loopIt = loops.rbegin();1320 while (loopIt != loops.rend() && isa<BlockArgument>(source->get())) {1321 auto iterArg = cast<BlockArgument>(source->get());1322 auto loop = *loopIt;1323 if (iterArg.getOwner()->getParentOp() != loop)1324 break;1325 source = loop.getTiedLoopInit(iterArg);1326 loopIt++;1327 }1328 if (loopIt == loops.rend())1329 destinationIterArg = source;1330 return {dyn_cast<OpResult>(source->get()), destinationIterArg};1331}1332 1333/// Implementation of fusing producer of a single slice by computing the1334/// slice of the producer in-place.1335std::optional<scf::SCFFuseProducerOfSliceResult>1336mlir::scf::tileAndFuseProducerOfSlice(1337 RewriterBase &rewriter, tensor::ExtractSliceOp candidateSliceOp,1338 MutableArrayRef<LoopLikeOpInterface> loops) {1339 // 1. Get the producer of the source (potentially walking through1340 // `iter_args` of nested `scf.for`)1341 auto [fusableProducer, destinationInitArg] =1342 getUntiledProducerFromSliceSource(&candidateSliceOp.getSourceMutable(),1343 loops);1344 if (!fusableProducer)1345 return std::nullopt;1346 unsigned resultNumber = fusableProducer.getResultNumber();1347 1348 OpBuilder::InsertionGuard g(rewriter);1349 rewriter.setInsertionPoint(candidateSliceOp);1350 1351 // 2. Clone the fused producer1352 // 2a. Compute the destination operands to use for the cloned operation.1353 SmallVector<Value> origDestinationTensors, clonedOpDestinationTensors;1354 Operation *fusableProducerOp = fusableProducer.getOwner();1355 if (isa<DestinationStyleOpInterface>(fusableProducerOp) &&1356 failed(tensor::getOrCreateDestinations(1357 rewriter, fusableProducerOp->getLoc(), fusableProducerOp,1358 origDestinationTensors)))1359 return std::nullopt;1360 1361 clonedOpDestinationTensors = origDestinationTensors;1362 if (destinationInitArg &&1363 isa<DestinationStyleOpInterface>(fusableProducerOp)) {1364 // 2b. If the producer is also destination style, then to maintain the1365 // destination passing style, update the destination of the producer to be1366 // the source of the slice.1367 clonedOpDestinationTensors[resultNumber] = candidateSliceOp.getSource();1368 }1369 // 2c. Clone the fused producer.1370 Operation *clonedProducerOp = cloneOpAndUpdateDestinationArgs(1371 rewriter, fusableProducerOp, clonedOpDestinationTensors);1372 // 2d. Update the source of the candidateSlice to be the cloned producer.1373 // Easier to just clone the slice with different source since1374 // replacements and DCE of cloned ops becomes easier1375 SmallVector<Value> candidateSliceOpOperands =1376 llvm::to_vector(candidateSliceOp->getOperands());1377 candidateSliceOpOperands[0] = clonedProducerOp->getResult(resultNumber);1378 tensor::ExtractSliceOp clonedCandidateSliceOp =1379 mlir::clone(rewriter, candidateSliceOp,1380 candidateSliceOp->getResultTypes(), candidateSliceOpOperands);1381 1382 // 3. Generate the tiled implementation of the producer of the source1383 FailureOr<TilingResult> tileAndFuseResult =1384 tensor::replaceExtractSliceWithTiledProducer(1385 rewriter, clonedCandidateSliceOp,1386 clonedProducerOp->getResult(resultNumber));1387 if (failed(tileAndFuseResult))1388 return std::nullopt;1389 // Note: Do not delete the candidateSliceOp, since its passed in from the1390 // caller.1391 rewriter.replaceAllUsesWith(candidateSliceOp,1392 tileAndFuseResult->tiledValues[0]);1393 rewriter.eraseOp(clonedCandidateSliceOp);1394 rewriter.eraseOp(clonedProducerOp);1395 1396 // 3. If the slice is for a destination operand, for example,1397 //1398 // ```mlir1399 // %0 = linalg.init1400 // %1 = linalg.fill .. outs(%0 : )1401 // %2 = scf.for .. iter_args(%arg0 = %1) {1402 // %3 = scf.for .. iter_args(%arg1 = %arg0) {1403 // %4 = tensor.extract_slice %arg1 [..]1404 // .. = linalg.matmul .. outs(%4 : )1405 // }1406 // }1407 // ```1408 //1409 // the IR is currently1410 //1411 // ```1412 // %0 = linalg.init1413 // %1 = linalg.fill1414 // %2 = scf.for .. iter_args(%arg0 = %1 /* incorrect value */ ) {1415 // %3 = scf.for .. iter_args(%arg1 = %arg0) {1416 // %4 = tensor.extract_slice %arg1[..]1417 // %5 = linalg.fill .. outs(%4 : )1418 // .. = linalg.matmul .. outs(%5 : )1419 // }1420 // }1421 // ```1422 //1423 // The untiled `linalg.fill` is still used as the `init_value` since it1424 // was originally a destination operand of the untiled `linalg.matmul`.1425 // When fusing an operand that is a destination operand, the iter_arg of1426 // the outer most loop should be changed to use the destination of the1427 // fused operation. With this the IR will be.1428 //1429 // ```1430 // %0 = linalg.init1431 // %1 = scf.for .. iter_args(%arg0 = %0 /* corrected value */ ) {1432 // %2 = scf.for .. iter_args(%arg1 = %arg0) {1433 // %3 = tensor.extract_slice %arg1[..]1434 // %4 = linalg.fill .. outs(%3 : )1435 // .. = linalg.matmul .. outs(%4 : )1436 // }1437 // }1438 // ```1439 if (destinationInitArg &&1440 isa<DestinationStyleOpInterface>(fusableProducerOp) && !loops.empty()) {1441 loops.front()1442 ->getOpOperands()[destinationInitArg.value()->getOperandNumber()]1443 .set(origDestinationTensors[resultNumber]);1444 }1445 return scf::SCFFuseProducerOfSliceResult{1446 fusableProducer, tileAndFuseResult->tiledValues[0],1447 tileAndFuseResult->tiledOps, tileAndFuseResult->generatedSlices};1448}1449 1450/// Reconstruct the fused producer from within the tiled-and-fused code.1451FailureOr<SmallVector<Operation *>> mlir::scf::yieldReplacementForFusedProducer(1452 RewriterBase &rewriter, tensor::ExtractSliceOp sliceOp,1453 scf::SCFFuseProducerOfSliceResult fusedProducerInfo,1454 MutableArrayRef<LoopLikeOpInterface> loops,1455 ArrayRef<unsigned> yieldResultNumber) {1456 if (loops.empty())1457 return success();1458 1459 Operation *originalOwner = fusedProducerInfo.origProducer.getOwner(),1460 *tiledOwner = fusedProducerInfo.tiledOps[0];1461 1462 Location loc = originalOwner->getLoc();1463 // a. collect all init Value to be appended1464 SmallVector<unsigned> initNumberList =1465 yieldResultNumber.empty() ? llvm::to_vector(llvm::seq<unsigned>(1466 0, originalOwner->getNumResults()))1467 : llvm::to_vector(yieldResultNumber);1468 SmallVector<Value> initValueList;1469 for (const auto &resultNumber : initNumberList) {1470 FailureOr<Value> initValue = tensor::getOrCreateDestination(1471 rewriter, loc, originalOwner->getResult(resultNumber));1472 if (succeeded(initValue)) {1473 initValueList.push_back(initValue.value());1474 } else {1475 return failure();1476 }1477 }1478 1479 SmallVector<Operation *> generatedSlices;1480 YieldTiledValuesFn newYieldValuesFn =1481 [&](RewriterBase &innerRewriter, Location loc, ValueRange /*ivs*/,1482 ValueRange newRegionIterArgs, SmallVector<Value> &tiledResult,1483 SmallVector<SmallVector<OpFoldResult>> &tiledOffset,1484 SmallVector<SmallVector<OpFoldResult>> &tiledSizes) -> LogicalResult {1485 OpBuilder::InsertionGuard g(innerRewriter);1486 1487 // get sliceOp tile information1488 SmallVector<OpFoldResult> sliceOffset = sliceOp.getMixedOffsets(),1489 sliceSizes = sliceOp.getMixedSizes();1490 1491 // expect all strides of sliceOp being 11492 if (!llvm::all_of(sliceOp.getMixedStrides(), isOneInteger))1493 return failure();1494 1495 unsigned sliceResultNumber =1496 fusedProducerInfo.origProducer.getResultNumber();1497 1498 auto tilableOp = cast<TilingInterface>(originalOwner);1499 // b. get iterDomain Offset and Sizes based on sliceOp tile1500 SmallVector<OpFoldResult> iterDomainOffset, iterDomainSizes;1501 // skip tensor.pack/unpack/pad, which expects single opResult1502 if (tilableOp->getNumResults() > 1 &&1503 failed(tilableOp.getIterationDomainTileFromResultTile(1504 rewriter, sliceResultNumber, sliceOffset, sliceSizes,1505 iterDomainOffset, iterDomainSizes))) {1506 // In theory, it is unnecessary to raise an error here. Actually1507 // although it fails to reconstruct the result tensor, it should not1508 // broke current fusion anyway. The reason why we must return failure1509 // currently is that the callback function `newYieldValuesFn` will be1510 // called after new init operand(s) has already been appended. It will1511 // take more refactoring to make sure the init operands are added1512 // consistently in the future. For more details, please refer to:1513 // https://github.com/llvm/llvm-project/pull/93144#discussion_r16437608141514 return failure();1515 }1516 1517 // c. calculate offsets and sizes info of all OpResults respectively based1518 // on iteration Domain Tile1519 SmallVector<SmallVector<OpFoldResult>> offsetList, sizesList;1520 for (const auto &resultNumber : initNumberList) {1521 if (resultNumber == sliceResultNumber) {1522 offsetList.push_back(sliceOffset);1523 sizesList.push_back(sliceSizes);1524 } else {1525 assert(!iterDomainOffset.empty() && !iterDomainSizes.empty());1526 // infer result tile according to the iteration domain tile1527 SmallVector<OpFoldResult> offset, sizes;1528 if (failed(tilableOp.getResultTilePosition(1529 rewriter, resultNumber, iterDomainOffset, iterDomainSizes,1530 offset, sizes))) {1531 return failure();1532 }1533 offsetList.push_back(offset);1534 sizesList.push_back(sizes);1535 }1536 }1537 1538 // d. create `extract_slice` for `iter_args` for DPS operation if1539 // necessary1540 if (auto tiledDestStyleOp =1541 dyn_cast<DestinationStyleOpInterface>(tiledOwner)) {1542 rewriter.setInsertionPoint(tiledDestStyleOp);1543 for (const auto &&[index, newRegionArg] :1544 llvm::enumerate(newRegionIterArgs)) {1545 auto destSlice = tensor::ExtractSliceOp::create(1546 rewriter, loc, newRegionArg, offsetList[index], sizesList[index],1547 SmallVector<OpFoldResult>(offsetList[index].size(),1548 rewriter.getIndexAttr(1)));1549 generatedSlices.push_back(destSlice);1550 unsigned resultNumber = initNumberList[index];1551 rewriter.modifyOpInPlace(tiledDestStyleOp, [&]() {1552 tiledDestStyleOp.getDpsInitsMutable()[resultNumber].set(destSlice);1553 });1554 }1555 }1556 1557 // e. prepare tiled offset and sizes for later `insert_slice` creation by1558 // caller1559 Block *block = rewriter.getInsertionPoint()->getBlock();1560 rewriter.setInsertionPoint(block->getTerminator());1561 for (const auto &&[index, resultNumber] : llvm::enumerate(initNumberList)) {1562 tiledResult.push_back(tiledOwner->getResult(resultNumber));1563 tiledOffset.emplace_back(offsetList[index]);1564 tiledSizes.emplace_back(sizesList[index]);1565 }1566 return success();1567 };1568 1569 if (failed(addInitOperandsToLoopNest(rewriter, loops, initValueList,1570 newYieldValuesFn))) {1571 return failure();1572 }1573 return generatedSlices;1574}1575 1576namespace {1577 1578//===----------------------------------------------------------------------===//1579// SliceTrackingListener1580//===----------------------------------------------------------------------===//1581 1582/// This class is a listener for tracking the insertion and removal of1583/// `tensor.extract_slice` ops in a worklist. This can be used in a greedy1584/// fusion algorithm to apply cleanup patterns in between fusion steps.1585class SliceTrackingListener : public RewriterBase::Listener {1586public:1587 explicit SliceTrackingListener(1588 std::optional<FrozenRewritePatternSet> patterns);1589 SliceTrackingListener() = default;1590 1591 /// Adds the given list of operations to the worklist, and if present,1592 /// applies the list of `patterns` to the newly added operations. This only1593 /// processes the given operations and any newly inserted ones by the1594 /// pattern set.1595 LogicalResult insertAndApplyPatterns(ArrayRef<Operation *> newOps);1596 1597 /// Add to the new operation worklist if it is an extract_slice.1598 void notifyOperationInserted(Operation *op,1599 OpBuilder::InsertPoint previous) override;1600 1601 /// Shared helper for operation removal from the worklist.1602 void removeOp(Operation *op);1603 1604 /// Remove the operation from the worklist.1605 void notifyOperationErased(Operation *op) override;1606 1607 /// Remove the operation from the worklist.1608 void notifyOperationReplaced(Operation *op, ValueRange replacement) override;1609 1610 /// The worklist for this transformation keeps track of the slices to visit1611 /// next for fusion.1612 std::deque<tensor::ExtractSliceOp> worklist;1613 1614private:1615 /// Optional pattern set to apply when adding new operations to the1616 /// worklist.1617 std::optional<FrozenRewritePatternSet> patterns = std::nullopt;1618};1619 1620SliceTrackingListener::SliceTrackingListener(1621 std::optional<FrozenRewritePatternSet> p) {1622 patterns = std::move(p);1623}1624 1625LogicalResult1626SliceTrackingListener::insertAndApplyPatterns(ArrayRef<Operation *> ops) {1627 for (Operation *op : ops) {1628 if (auto slice = dyn_cast<tensor::ExtractSliceOp>(op))1629 worklist.push_back(slice);1630 }1631 1632 if (!patterns)1633 return success();1634 1635 return applyOpPatternsGreedily(1636 ops, patterns.value(),1637 GreedyRewriteConfig().setListener(this).setStrictness(1638 GreedyRewriteStrictness::ExistingAndNewOps));1639}1640 1641void SliceTrackingListener::notifyOperationInserted(1642 Operation *op, OpBuilder::InsertPoint previous) {1643 auto slice = dyn_cast<tensor::ExtractSliceOp>(op);1644 if (!slice)1645 return;1646 worklist.push_back(slice);1647}1648 1649// Scan the worklist for the given op and remove it if present. The1650// expectation is for the worklist to be small and for removal to be1651// relatively rare.1652void SliceTrackingListener::removeOp(Operation *op) {1653 if (!isa<tensor::ExtractSliceOp>(op))1654 return;1655 auto iter = worklist.begin();1656 while (iter != worklist.end()) {1657 if (*iter == op)1658 break;1659 iter++;1660 }1661 if (iter == worklist.end())1662 return;1663 1664 worklist.erase(iter);1665}1666 1667void SliceTrackingListener::notifyOperationErased(Operation *op) {1668 removeOp(op);1669}1670 1671void SliceTrackingListener::notifyOperationReplaced(Operation *op,1672 ValueRange replacement) {1673 removeOp(op);1674}1675 1676//===----------------------------------------------------------------------===//1677// ReplacementListener1678//===----------------------------------------------------------------------===//1679 1680/// Listener that tracks updates replacements for values which can be mutated.1681/// This listener runs on top of the existing listener for the rewriter,1682/// to make sure external users can still run listeners.1683class ReplacementListener : public RewriterBase::ForwardingListener {1684public:1685 ReplacementListener(DenseMap<Value, Value> &replacements,1686 OpBuilder::Listener *listener)1687 : ForwardingListener(listener), replacements(replacements) {}1688 1689 void updateReplacementValues(ValueRange origValues,1690 ValueRange replaceValues) {1691 // This can probably be written better, but just iterates over the map1692 // and the new replacements for now.1693 for (auto &[key, val] : replacements) {1694 for (auto [orig, replace] : llvm::zip_equal(origValues, replaceValues)) {1695 if (val == orig) {1696 val = replace;1697 }1698 }1699 }1700 }1701 1702 void notifyOperationReplaced(Operation *op, Operation *newOp) override {1703 ForwardingListener::notifyOperationReplaced(op, newOp);1704 updateReplacementValues(op->getResults(), newOp->getResults());1705 }1706 1707 void notifyOperationReplaced(Operation *op, ValueRange values) override {1708 ForwardingListener::notifyOperationReplaced(op, values);1709 updateReplacementValues(op->getResults(), values);1710 }1711 1712private:1713 DenseMap<Value, Value> &replacements;1714};1715 1716} // namespace1717 1718/// Implementation of tile consumer and fuse producer greedily.1719FailureOr<scf::SCFTileAndFuseResult>1720mlir::scf::tileConsumerAndFuseProducersUsingSCF(1721 RewriterBase &rewriter, TilingInterface consumer,1722 const scf::SCFTileAndFuseOptions &options) {1723 // This transformation is only valid for ops that return values (i.e. not1724 // valid to use with operations that have memref operands).1725 if (!consumer->getNumResults()) {1726 return rewriter.notifyMatchFailure(1727 consumer, "invalid pattern for op with no results");1728 }1729 1730 // 1. First tile the consumer.1731 SetVector<Operation *> fusedProducers, tiledAndFusedOps;1732 1733 FailureOr<scf::SCFTilingResult> tilingResult =1734 tileUsingSCF(rewriter, consumer, options.tilingOptions);1735 1736 if (failed(tilingResult))1737 return rewriter.notifyMatchFailure(consumer, "failed to tile consumer");1738 tiledAndFusedOps.insert_range(tilingResult->tiledOps);1739 1740 DenseMap<Value, Value> replacements;1741 for (auto [origVal, replacement] :1742 llvm::zip_equal(consumer->getResults(), tilingResult->replacements)) {1743 replacements[origVal] = replacement;1744 }1745 1746 // If there are no loops generated, fusion is immaterial.1747 auto &loops = tilingResult->loops;1748 if (loops.empty()) {1749 return scf::SCFTileAndFuseResult{fusedProducers, tiledAndFusedOps, loops,1750 replacements};1751 }1752 1753 // Since the loop gets potentially replaced during fusion, we need to track1754 // the mutation of replacement values. To do this, we attach a listener to1755 // update the replacements as they happen.1756 OpBuilder::Listener *previousListener = rewriter.getListener();1757 auto resetListener =1758 llvm::make_scope_exit([&]() { rewriter.setListener(previousListener); });1759 ReplacementListener replaceListener(replacements, previousListener);1760 rewriter.setListener(&replaceListener);1761 1762 // 2. Typically, the operands of the tiled operation are slices of the1763 // operands of the untiled operation. These are expressed in IR using1764 // `tensor.extract_slice` operations with source being the operands of1765 // the untiled operation. Create a worklist of these1766 // `tensor.extract_slice` operations. If the producers of the source of1767 // the `tensor.extract_slice` can be tiled such that the tiled value is1768 // generated in-place, that effectively tiles + fuses the operations.1769 struct WorklistItem {1770 tensor::ExtractSliceOp candidateSlice;1771 SCFTileAndFuseOptions::ControlFnResult controlFnResult;1772 };1773 1774 SliceTrackingListener sliceTracker =1775 SliceTrackingListener(options.cleanupPatterns);1776 1777 if (failed(1778 sliceTracker.insertAndApplyPatterns(tilingResult->generatedSlices))) {1779 return rewriter.notifyMatchFailure(consumer, "cleanup patterns failed");1780 }1781 OpBuilder::InsertionGuard g(rewriter);1782 while (!sliceTracker.worklist.empty()) {1783 auto candidateSlice = sliceTracker.worklist.front();1784 sliceTracker.worklist.pop_front();1785 1786 auto [fusableProducer, destinationInitArg] =1787 getUntiledProducerFromSliceSource(&candidateSlice.getSourceMutable(),1788 loops);1789 if (!fusableProducer)1790 continue;1791 1792 std::optional<SCFTileAndFuseOptions::ControlFnResult> controlFnResult =1793 options.fusionControlFn(candidateSlice, fusableProducer,1794 destinationInitArg.has_value());1795 if (!controlFnResult)1796 continue;1797 1798 WorklistItem worklistItem = {candidateSlice, controlFnResult.value()};1799 1800 // The operands of the fused producer might themselved be slices of1801 // values produced by operations that implement the `TilingInterface`.1802 // Add these operations to the worklist.1803 std::optional<scf::SCFFuseProducerOfSliceResult> fusedResult =1804 tileAndFuseProducerOfSlice(rewriter, worklistItem.candidateSlice,1805 loops);1806 if (!fusedResult)1807 continue;1808 1809 SmallVector<Operation *> worklistCandidates = fusedResult->generatedSlices;1810 1811 if (worklistItem.controlFnResult.yieldProducerReplacement) {1812 // Reconstruct and yield all opResult of fusableProducerOp by default.1813 // The caller can specific which one to yield by designating optional1814 // argument named `yieldResultNumber` of1815 // `yieldReplacementForFusedProducer`.1816 Operation *fusableProducerOp = fusedResult->origProducer.getOwner();1817 FailureOr<SmallVector<Operation *>> newSlices =1818 yieldReplacementForFusedProducer(rewriter,1819 worklistItem.candidateSlice,1820 fusedResult.value(), loops);1821 if (failed(newSlices)) {1822 return rewriter.notifyMatchFailure(1823 fusableProducerOp, "failed to replacement value for this "1824 "operation from within the tiled loop");1825 }1826 worklistCandidates.append(newSlices.value());1827 for (auto [index, result] :1828 llvm::enumerate(fusableProducerOp->getResults())) {1829 replacements[result] = loops.front()->getResult(1830 loops.front()->getNumResults() -1831 fusableProducerOp->getNumResults() + index);1832 }1833 }1834 if (Operation *tiledAndFusedOp =1835 fusedResult->tiledAndFusedProducer.getDefiningOp()) {1836 fusedProducers.insert(fusedResult->origProducer.getDefiningOp());1837 tiledAndFusedOps.insert(tiledAndFusedOp);1838 }1839 1840 if (failed(sliceTracker.insertAndApplyPatterns(worklistCandidates))) {1841 return rewriter.notifyMatchFailure(consumer, "cleanup patterns failed");1842 }1843 }1844 1845 return scf::SCFTileAndFuseResult{fusedProducers, tiledAndFusedOps, loops,1846 replacements};1847}1848 1849//===----------------------------------------------------------------------===//1850// tileAndFuseConsumerUsingSCF implementation.1851//===----------------------------------------------------------------------===//1852 1853/// A utility function that checks whether the only use of the result of a1854/// tensor.insert_slice op is in a scf.yield op.1855static LogicalResult1856checkAssumptionForFusingConsumer(tensor::InsertSliceOp candidateSliceOp) {1857 Value result = candidateSliceOp.getResult();1858 Value::use_range uses = result.getUses();1859 if (!llvm::hasSingleElement(uses)) {1860 LLVM_DEBUG(llvm::dbgs() << "Too many uses of the candidate slice op\n");1861 return failure();1862 }1863 OpOperand &operandUse = (*uses.begin());1864 Operation *userOp = operandUse.getOwner();1865 if (!isa<scf::YieldOp>(userOp)) {1866 LLVM_DEBUG(llvm::dbgs()1867 << "Expected scf.yield to be the only user, but got -> "1868 << (*userOp));1869 return failure();1870 }1871 if (result.getDefiningOp()->getBlock() != userOp->getBlock()) {1872 LLVM_DEBUG(llvm::dbgs() << "Expected tensor.insert_slice and scf.yield to "1873 "be in the same block\n");1874 return failure();1875 }1876 return success();1877}1878 1879/// An utility to get the first user of the given loopOp. If any of user stay1880/// in different block of loopOp, return failure.1881static FailureOr<Operation *> getFirstUserOfLoop(Operation *loopOp) {1882 if (!isa<LoopLikeOpInterface>(loopOp))1883 return failure();1884 Operation *firstUserOfLoop = nullptr;1885 for (Operation *userOp : loopOp->getUsers()) {1886 // `ParallelInsertSlice` located inside `InParallelOp` has no same parent1887 // block with any other types of operation. Thus, just redirecting to its1888 // parent `InParallelOp`. E.g.1889 //1890 // ```1891 // %1 = scf.for {1892 // ...1893 // }1894 // %2 = consumerOp ins(%1, ...)1895 // scf.forall.in_parallel {1896 // tensor.parallel_insert_slice %11897 // }1898 // ```1899 // where `InParallelOp` but not `ParallelInsertSlice` stays in the same1900 // same block with `consumerOp`.1901 if (isa<tensor::ParallelInsertSliceOp>(userOp))1902 userOp = userOp->getParentOfType<scf::InParallelOp>();1903 1904 if (loopOp->getBlock() != userOp->getBlock())1905 return failure();1906 1907 if (!firstUserOfLoop || userOp->isBeforeInBlock(firstUserOfLoop))1908 firstUserOfLoop = userOp;1909 }1910 return firstUserOfLoop;1911}1912 1913/// This utility currently checks whether the first userOp of loop is NOT1914/// before the last defineOp of consumer operand. Because that we need to move1915/// the whole loop structure right before the `firstUserOfLoop`. This utility1916/// thus helps ensuring that no invalid IR is formed, i.e. no backward slice1917/// of consumerOp is dominated by the `firstUserOfLoop`. Saying that:1918///1919/// ```1920/// %0 = scf.for() {1921/// ...1922/// }1923/// ...1924/// %1 = firstUserOfLoop(%0)1925/// ...1926/// %2 = lastDefOfConsumerOperand1927/// ...1928/// %3 = consumerOp(%2)1929/// ```1930///1931/// If the `firstUserOfLoop` is before `lastDefOfConsumerOperand`, then it1932/// would be invalid to move the `loopOp` right before the `firstUserOfLoop`,1933/// a.k.a. use-def chain violation:1934///1935/// ```1936/// %0:2 = scf.for() {1937/// // use before define error1938/// %3 = tiledConsumerOp(%2)1939/// }1940/// %1 = firstUserOfLoop(%0)1941/// ...1942/// %2 = lastDefOfConsumerOperand1943/// ```1944///1945/// @param loopOp: loop operation1946/// @param consumerOp: consumer operation1947/// @param reorderOperations: the flag controls whether to reorder the1948/// backward slice w.r.t. the defineOp of `consumerOp` operands.1949/// @return: computed backward slice of consumerOp, but excluding those1950/// already dominates `firstUserOfLoop`.1951static FailureOr<llvm::SetVector<Operation *>>1952checkAssumptionForLoop(Operation *loopOp, Operation *consumerOp,1953 bool reorderOperations) {1954 FailureOr<Operation *> firstUserOfLoop = getFirstUserOfLoop(loopOp);1955 if (failed(firstUserOfLoop))1956 return failure();1957 1958 BackwardSliceOptions options;1959 DominanceInfo dominanceInfo;1960 options.inclusive = true;1961 options.omitBlockArguments = true;1962 bool includeLoopOp = false;1963 options.filter = [&](Operation *op) {1964 if (op == loopOp) {1965 includeLoopOp = true;1966 return false;1967 }1968 // Cut off the slice to not include any operation that already dominates1969 // firstUserOfLoop.1970 return !dominanceInfo.properlyDominates(op, *firstUserOfLoop);1971 };1972 llvm::SetVector<Operation *> slice;1973 for (auto operand : consumerOp->getOperands()) {1974 LogicalResult result = getBackwardSlice(operand, &slice, options);1975 assert(result.succeeded() && "expected a backward slice");1976 (void)result;1977 }1978 1979 if (!slice.empty()) {1980 // If consumerOp has one producer, which is also the user of loopOp.1981 // E.g.1982 // ```1983 // %0 = %loopOp1984 // %1 = consumerOp1 ins(%0)1985 // %2 = consumerOp2 ins(%0, %1)1986 // ```1987 // We can not fuse consumerOp2 into loopOp due to UD chain, unless1988 // consumerOp1 has already been fused into loopOp before.1989 if (includeLoopOp || !reorderOperations)1990 return failure();1991 }1992 1993 return slice;1994}1995 1996/// Fetches the OpOperand of the first valid user (and use) of the value `val`1997/// which implements `TilingInterface` and `DestinationStyleOpInterface`.1998/// Returns failure otherwise.1999static FailureOr<OpOperand *> getConsumerFromLoopUses(RewriterBase &rewriter,2000 Operation *loopOp,2001 unsigned resultNumber) {2002 if (!isa<LoopLikeOpInterface>(loopOp))2003 return failure();2004 Value val = loopOp->getResult(resultNumber);2005 Block *loopBlock = loopOp->getBlock();2006 for (OpOperand &opOperand : val.getUses()) {2007 Operation *consumerOp = opOperand.getOwner();2008 // Step 1. Check if the user is tilable.2009 if (!isa<TilingInterface>(consumerOp) ||2010 !isa<DestinationStyleOpInterface>(consumerOp)) {2011 // TODO: We have to init result of consumer before scf.for, use2012 // DestinationStyleOpInterface to get result shape from init for now.2013 // Add support for other op such as op has InferTypeOpInterface.2014 continue;2015 }2016 // Step 2. Check if user stay in the same block.2017 if (loopBlock != consumerOp->getBlock())2018 continue;2019 // Step 3. Check if user has succeeding user. Otherwise, it usually2020 // represents already tiled.2021 if (consumerOp->use_empty())2022 continue;2023 // Step 4. Check assumption for loop with `reorderOperations` enabled.2024 FailureOr<llvm::SetVector<Operation *>> slice =2025 checkAssumptionForLoop(loopOp, consumerOp, true);2026 if (failed(slice))2027 continue;2028 // Step 5. If backward sice is not empty, move them before2029 // firstUserOfLoop.2030 if (!slice->empty()) {2031 mlir::topologicalSort(*slice);2032 FailureOr<Operation *> firstUserOfLoop = getFirstUserOfLoop(loopOp);2033 assert(succeeded(firstUserOfLoop) && "First user of loop is not found");2034 for (auto op : *slice) {2035 rewriter.moveOpBefore(op, *firstUserOfLoop);2036 }2037 }2038 return &opOperand;2039 }2040 return failure();2041}2042 2043/// Fetch the untiled consumer of the outermost scf.for's result which is2044/// yielded by a tensor.insert_slice from the innermost scf.for. This function2045/// makes the following assumptions :2046/// 1. tensor.insert_slice has scf.yield as its only user.2047/// 2. scf.for's corresponding result has only one use.2048/// 3. The `loops` passed in are perfectly nested `scf.for` operations.2049static FailureOr<OpOperand *>2050getUntiledConsumerFromSlice(RewriterBase &rewriter,2051 tensor::InsertSliceOp candidateSliceOp,2052 MutableArrayRef<LoopLikeOpInterface> loops) {2053 assert(!loops.empty() && "unexpected loops to be empty");2054 // 1. Expect slice to be part of the body of the inner most loop.2055 Operation *containingOp = candidateSliceOp->getParentOp();2056 if (containingOp != loops.back()) {2057 return rewriter.notifyMatchFailure(2058 candidateSliceOp,2059 "expected slice to be within body of inner-most loop");2060 }2061 2062 // 2. Check that the loop is perfectly nested.2063 if (!isPerfectlyNestedForLoops(loops)) {2064 return rewriter.notifyMatchFailure(2065 candidateSliceOp, "expected passed loops to be perfectly nested.");2066 }2067 2068 if (failed(checkAssumptionForFusingConsumer(candidateSliceOp)))2069 return failure();2070 Value sliceResult = candidateSliceOp.getResult();2071 2072 // 3. Fetch the corresponding output.2073 OpOperand &yieldOpOperand = (*sliceResult.getUses().begin());2074 unsigned resultNumber = yieldOpOperand.getOperandNumber();2075 2076 scf::ForOp topLevelForOp = cast<scf::ForOp>(loops.front().getOperation());2077 2078 return getConsumerFromLoopUses(rewriter, topLevelForOp, resultNumber);2079}2080 2081/// Fetch the first untiled consumer of a scf.forall's result which is yielded2082/// by a tensor.parallel_insert_slice.2083static FailureOr<OpOperand *>2084getUntiledConsumerFromSlice(RewriterBase &rewriter,2085 tensor::ParallelInsertSliceOp candidateSliceOp,2086 MutableArrayRef<LoopLikeOpInterface> loops) {2087 assert(!loops.empty() && "unexpected loops to be empty");2088 // 1. Check that the surrounding loop is a single scf.forall loop.2089 if (loops.size() != 1) {2090 return rewriter.notifyMatchFailure(2091 candidateSliceOp, "expected single surrounding scf.forall");2092 }2093 auto forallOp = dyn_cast<scf::ForallOp>(loops.front().getOperation());2094 if (!forallOp) {2095 return rewriter.notifyMatchFailure(2096 candidateSliceOp, "expected single surrounding scf.forall");2097 }2098 2099 // 2. Fetch the corresponding output2100 Value sliceDest = candidateSliceOp.getDest();2101 auto iterArg = dyn_cast<BlockArgument>(sliceDest);2102 if (!iterArg)2103 return failure();2104 if (iterArg.getOwner()->getParentOp() != forallOp)2105 return failure();2106 2107 unsigned resultNumber =2108 forallOp.getTiedOpResult(forallOp.getTiedOpOperand(iterArg))2109 .getResultNumber();2110 2111 return getConsumerFromLoopUses(rewriter, forallOp, resultNumber);2112}2113 2114/// A utility to fetch an untiled consumer of2115/// tensor.insert_slice/tensor.parallel_insert_slice.2116static FailureOr<SmallVector<OpOperand *>> getUntiledConsumerOperandsFromSlices(2117 RewriterBase &rewriter, ArrayRef<Operation *> sliceOps,2118 MutableArrayRef<LoopLikeOpInterface> loops) {2119 assert(!loops.empty() && "unexpected empty loops");2120 assert(!sliceOps.empty() && "unexpected empty list of candidate slices");2121 SmallVector<OpOperand *> fusedOperands;2122 for (auto sliceOp : sliceOps) {2123 FailureOr<OpOperand *> fusedOperand =2124 TypeSwitch<Operation *, FailureOr<OpOperand *>>(sliceOp)2125 .Case<tensor::InsertSliceOp, tensor::ParallelInsertSliceOp>(2126 [&](auto op) {2127 return getUntiledConsumerFromSlice(rewriter, op, loops);2128 })2129 .Default([&](Operation *op) {2130 return rewriter.notifyMatchFailure(op, "unhandled slice type");2131 });2132 if (failed(fusedOperand)) {2133 return failure();2134 }2135 if (!fusedOperands.empty() &&2136 fusedOperand.value()->getOwner() != fusedOperands.front()->getOwner()) {2137 return rewriter.notifyMatchFailure(2138 fusedOperand.value()->getOwner(),2139 "all candidate slices must be to the same consumer");2140 }2141 fusedOperands.push_back(fusedOperand.value());2142 }2143 return fusedOperands;2144}2145 2146template <typename InsertSliceOpTy>2147static tensor::InsertSliceOp cloneAsInsertSlice(RewriterBase &rewriter,2148 InsertSliceOpTy sliceOp);2149 2150template <>2151tensor::InsertSliceOp2152cloneAsInsertSlice<tensor::InsertSliceOp>(RewriterBase &rewriter,2153 tensor::InsertSliceOp insertSliceOp) {2154 return cast<tensor::InsertSliceOp>(2155 rewriter.clone(*insertSliceOp.getOperation()));2156}2157 2158template <>2159tensor::InsertSliceOp cloneAsInsertSlice<tensor::ParallelInsertSliceOp>(2160 RewriterBase &rewriter, tensor::ParallelInsertSliceOp insertSliceOp) {2161 return tensor::InsertSliceOp::create(2162 rewriter, insertSliceOp->getLoc(), insertSliceOp.getSource(),2163 insertSliceOp.getDest(), insertSliceOp.getMixedOffsets(),2164 insertSliceOp.getMixedSizes(), insertSliceOp.getMixedStrides());2165}2166 2167static SmallVector<tensor::InsertSliceOp>2168cloneAsInsertSlices(RewriterBase &rewriter,2169 ArrayRef<Operation *> candidateSlices) {2170 assert(!candidateSlices.empty() &&2171 "unexpected empty list of slices to clone");2172 SmallVector<tensor::InsertSliceOp> clonedSlices;2173 for (auto sliceOp : candidateSlices) {2174 TypeSwitch<Operation *>(sliceOp)2175 .Case<tensor::InsertSliceOp, tensor::ParallelInsertSliceOp>(2176 [&](auto op) {2177 auto clonedOp = cloneAsInsertSlice(rewriter, op);2178 clonedSlices.push_back(clonedOp);2179 })2180 // Assert here assuming this has already been checked.2181 .DefaultUnreachable(2182 "unexpected slice type while cloning as insert slice");2183 }2184 return clonedSlices;2185}2186 2187static FailureOr<scf::SCFFuseConsumerOfSliceResult>2188tileAndFuseConsumerOfSlicesImpl(RewriterBase &rewriter, Operation *consumerOp,2189 ArrayRef<OpOperand *> consumerOpOperands,2190 ArrayRef<Operation *> candidateSlices,2191 MutableArrayRef<LoopLikeOpInterface> loops) {2192 assert(!loops.empty() && "expected loops to be not empty");2193 2194 // 1. Check assumption for loop with `reorderOperations` disabled.2195 if (failed(checkAssumptionForLoop(loops.front(), consumerOp, false))) {2196 return rewriter.notifyMatchFailure(2197 loops.front(), "the first user of loop should not dominate any define "2198 "of consumer operand(s)");2199 }2200 2201 LoopLikeOpInterface outerMostLoop = loops.front();2202 LoopLikeOpInterface innerMostLoop = loops.back();2203 2204 OpBuilder::InsertionGuard g(rewriter);2205 // 2. Check consumer is not using scf loop's output as init.2206 auto dstOp = dyn_cast<DestinationStyleOpInterface>(consumerOp);2207 if (!dstOp)2208 return rewriter.notifyMatchFailure(consumerOp,2209 "consumer op is not DPS operation");2210 if (llvm::any_of(consumerOpOperands, [&](OpOperand *opOperand) {2211 return dstOp.isDpsInit(opOperand);2212 })) {2213 return rewriter.notifyMatchFailure(2214 consumerOp,2215 "consumer op taking the result of scf.for as init is not supported");2216 }2217 SmallVector<Value> newInits = llvm::to_vector(dstOp.getDpsInits());2218 2219 // 3. Move the whole loop structure right before firstUserOfLoop, the2220 // dominance should be already ensured by `checkAssumptionForLoop`.2221 FailureOr<Operation *> firstUserOfLoop = getFirstUserOfLoop(outerMostLoop);2222 if (failed(firstUserOfLoop)) {2223 return rewriter.notifyMatchFailure(2224 outerMostLoop, "could not find the first user of outer most loop");2225 }2226 rewriter.moveOpBefore(outerMostLoop, *firstUserOfLoop);2227 2228 // 4. Set insertion point before terminator op of the loop and create a new2229 // tensor.insert_slice. In the scf.for case this is a clone of the2230 // candidateSliceOp whereas in the scf.forall case this is created from the2231 // operands of tensor.parallel_insert_slice.2232 if (auto sliceOp =2233 dyn_cast<tensor::ParallelInsertSliceOp>(candidateSlices.front())) {2234 auto newForallOp = cast<scf::ForallOp>(innerMostLoop.getOperation());2235 rewriter.setInsertionPoint(newForallOp.getTerminator());2236 } else {2237 rewriter.setInsertionPoint(candidateSlices.front());2238 }2239 // 5.a. Clone all the candidate slices as equivalent insert slice ops.2240 SmallVector<tensor::InsertSliceOp> clonedInsertSlices =2241 cloneAsInsertSlices(rewriter, candidateSlices);2242 2243 // 5.b. Clone consumer op.2244 auto clonedConsumerOp = cast<TilingInterface>(rewriter.clone(*consumerOp));2245 SmallVector<unsigned> operandNumbers =2246 llvm::map_to_vector(consumerOpOperands, [](OpOperand *opOperand) {2247 return opOperand->getOperandNumber();2248 });2249 SmallVector<OpOperand *> clonedOpFusedOperandsList =2250 llvm::map_to_vector(operandNumbers, [&](unsigned operandNum) {2251 return &clonedConsumerOp->getOpOperand(operandNum);2252 });2253 2254 // 5.c. Replace all uses of the loop result with the result of the cloned2255 // tensor.insert_slice.2256 rewriter.modifyOpInPlace(clonedConsumerOp, [&]() {2257 for (auto [operandToReplace, clonedSliceOp] :2258 llvm::zip_equal(clonedOpFusedOperandsList, clonedInsertSlices)) {2259 operandToReplace->set(clonedSliceOp.getResult());2260 }2261 });2262 2263 // 6. Perform tiling of the cloned consumer and replace the operand at2264 // `operandNumber` with the source of the cloned tensor.insert_slice op.2265 FailureOr<TilingResult> tileAndFuseResult =2266 tensor::replaceInsertSlicesWithTiledConsumer(rewriter, clonedInsertSlices,2267 clonedOpFusedOperandsList);2268 if (failed(tileAndFuseResult)) {2269 return failure();2270 }2271 2272 auto tiledConsumerOp = cast<TilingInterface>(tileAndFuseResult->tiledOps[0]);2273 for (auto [operandNum, clonedSliceOp] :2274 llvm::zip_equal(operandNumbers, clonedInsertSlices)) {2275 rewriter.replaceAllUsesWith(tiledConsumerOp->getOperand(operandNum),2276 clonedSliceOp.getSource());2277 }2278 2279 // 7. Reconstruct [nested] loop with new inits.2280 YieldTiledValuesFn newYieldValuesFn =2281 [&](RewriterBase &innerRewriter, Location loc, ValueRange /*ivs*/,2282 ValueRange newRegionIterArgs, SmallVector<Value> &tiledResult,2283 SmallVector<SmallVector<OpFoldResult>> &tiledOffset,2284 SmallVector<SmallVector<OpFoldResult>> &tiledSizes) -> LogicalResult {2285 OpBuilder::InsertionGuard g(innerRewriter);2286 // 8. Set inner insertPoint right before tiled consumer op.2287 innerRewriter.setInsertionPoint(tiledConsumerOp);2288 2289 SmallVector<SmallVector<OpFoldResult>> allOffsets, allSizes;2290 for (auto candidateSliceOp : clonedInsertSlices) {2291 SmallVector<OpFoldResult> offsets = candidateSliceOp.getMixedOffsets();2292 SmallVector<OpFoldResult> sizes = candidateSliceOp.getMixedSizes();2293 SmallVector<OpFoldResult> strides = candidateSliceOp.getMixedStrides();2294 2295 // 9. Check all insert stride is 1.2296 if (!llvm::all_of(strides, isOneInteger)) {2297 return rewriter.notifyMatchFailure(2298 candidateSliceOp, "containingOp's result yield with stride");2299 }2300 2301 allOffsets.emplace_back(std::move(offsets));2302 allSizes.emplace_back(std::move(sizes));2303 }2304 2305 // 10. Try to get iter domain position from input position. Use2306 // clonedConsumerOp instead of tiledConsumerOp, because the iteration2307 // domain may require index computation based on the result size. The2308 // sizes and offsets should be the same either way, but using2309 // tiledConsumerOp could lead to some chained unnecessary extra index2310 // computation.2311 SmallVector<OpFoldResult> iterDomainOffsets, iterDomainSizes;2312 if (failed(clonedConsumerOp.getIterationDomainTileFromOperandTiles(2313 rewriter, operandNumbers, allOffsets, allSizes, iterDomainOffsets,2314 iterDomainSizes))) {2315 return rewriter.notifyMatchFailure(2316 clonedConsumerOp,2317 "can't get iter domain position from input position");2318 }2319 2320 // 11. Try to fetch the offset and size for all results of the cloned2321 // consumer. This would then be used to form the corresponding2322 // tensor.insert_slice/parallel_insert_slice later.2323 unsigned totalNumResultsOfConsumer = tiledConsumerOp->getNumResults();2324 SmallVector<SmallVector<OpFoldResult>> resultOffsets(2325 totalNumResultsOfConsumer);2326 SmallVector<SmallVector<OpFoldResult>> resultSizes(2327 totalNumResultsOfConsumer);2328 for (auto [idx, v] : llvm::enumerate(tiledConsumerOp->getResults())) {2329 if (failed(tiledConsumerOp.getResultTilePosition(2330 rewriter, idx, iterDomainOffsets, iterDomainSizes,2331 resultOffsets[idx], resultSizes[idx]))) {2332 return rewriter.notifyMatchFailure(2333 tiledConsumerOp,2334 "can't get result domain position from iter domain position");2335 }2336 }2337 2338 // 12. Create `extract_slice` for `iter_args` for DPS operation if2339 // necessary.2340 if (auto tiledDestStyleOp = dyn_cast<DestinationStyleOpInterface>(2341 tiledConsumerOp.getOperation())) {2342 rewriter.setInsertionPoint(tiledDestStyleOp);2343 for (const auto &&[index, newRegionArg] :2344 llvm::enumerate(newRegionIterArgs)) {2345 auto destSlice = tensor::ExtractSliceOp::create(2346 rewriter, loc, newRegionArg, resultOffsets[index],2347 resultSizes[index],2348 SmallVector<OpFoldResult>(resultOffsets[index].size(),2349 rewriter.getIndexAttr(1)));2350 // Make a copy of index to avoid a capturing structured binding, which2351 // is a C++20 extension.2352 auto dstNumber = index;2353 rewriter.modifyOpInPlace(tiledDestStyleOp, [&]() {2354 tiledDestStyleOp.getDpsInitsMutable()[dstNumber].set(destSlice);2355 });2356 }2357 }2358 2359 // 13. Prepare tiled offset and sizes for later `insert_slice` creation by2360 // caller.2361 Block *block = rewriter.getInsertionPoint()->getBlock();2362 rewriter.setInsertionPoint(block->getTerminator());2363 for (const auto &&[index, result] :2364 llvm::enumerate(tiledConsumerOp->getResults())) {2365 tiledResult.push_back(result);2366 tiledOffset.emplace_back(resultOffsets[index]);2367 tiledSizes.emplace_back(resultSizes[index]);2368 }2369 return success();2370 };2371 // 14. Add new inits to [nested] loops.2372 if (failed(addInitOperandsToLoopNest(rewriter, loops, newInits,2373 newYieldValuesFn))) {2374 return rewriter.notifyMatchFailure(tiledConsumerOp,2375 "unable to add new inits to nest loop");2376 }2377 2378 // 15. Replace the result of scf loop and consumer op with new loop's2379 // results.2380 2381 for (auto &&[oldResult, newResult] :2382 llvm::zip(consumerOp->getResults(),2383 loops.front()->getResults().take_back(newInits.size()))) {2384 rewriter.replaceAllUsesWith(oldResult, newResult);2385 }2386 2387 // 16. Need to erase the old scf loop and the cloned consumer op.2388 rewriter.eraseOp(clonedConsumerOp);2389 2390 SmallVector<OpOperand *> tiledAndFusedOpOperands =2391 llvm::map_to_vector(operandNumbers, [&](unsigned operandNum) {2392 return &tileAndFuseResult->tiledOps[0]->getOpOperand(operandNum);2393 });2394 auto consumerOpOperandsVec = llvm::to_vector(consumerOpOperands);2395 return scf::SCFFuseConsumerOfSliceResult{2396 std::move(consumerOpOperandsVec), std::move(tiledAndFusedOpOperands),2397 std::move(tileAndFuseResult->tiledOps)};2398}2399 2400/// Implementation of fusing consumer of a single slice by computing the2401/// slice of the consumer in-place for scf loop.2402FailureOr<scf::SCFFuseConsumerOfSliceResult>2403mlir::scf::tileAndFuseConsumerOfSlices(2404 RewriterBase &rewriter, ArrayRef<Operation *> candidateSlices,2405 MutableArrayRef<LoopLikeOpInterface> loops) {2406 if (candidateSlices.empty()) {2407 return rewriter.notifyMatchFailure(2408 rewriter.getUnknownLoc(),2409 "no candidate slices provided for consumer fusion");2410 }2411 // Return if `loops` is empty, return an error for now. Caller is expected2412 // to handle this case.2413 if (loops.empty()) {2414 return rewriter.notifyMatchFailure(2415 candidateSlices.front(),2416 "cannot call tile and fuse consumer with an empty loop nest");2417 }2418 2419 if (!(llvm::all_of(candidateSlices, llvm::IsaPred<tensor::InsertSliceOp>) ||2420 llvm::all_of(candidateSlices,2421 llvm::IsaPred<tensor::ParallelInsertSliceOp>))) {2422 return rewriter.notifyMatchFailure(2423 candidateSlices.front(),2424 "candidates slices need to be all `tensor.extract_slice`s or "2425 "`tensor.parallel_insert_slice`s");2426 }2427 2428 // Get the consumer of scf.for for the result yielded by2429 // tensor.insert_slice/parallel_insert_slice.2430 FailureOr<SmallVector<OpOperand *>> maybeConsumerOpOperands =2431 getUntiledConsumerOperandsFromSlices(rewriter, candidateSlices, loops);2432 if (failed(maybeConsumerOpOperands)) {2433 return rewriter.notifyMatchFailure(candidateSlices.front(),2434 "could not fetch consumer to fuse");2435 }2436 Operation *consumerOp = maybeConsumerOpOperands->front()->getOwner();2437 2438 return tileAndFuseConsumerOfSlicesImpl(rewriter, consumerOp,2439 maybeConsumerOpOperands.value(),2440 candidateSlices, loops);2441}2442 2443/// For a given `result` of a `forallOp` return the2444/// `tensor.parallel_insert_slice` op (or combining op) that is used to2445/// construct this result.2446static std::optional<Operation *>2447getProducingParallelInsertSlice(scf::ForallOp forallOp, OpResult result) {2448 if (result.getOwner() != forallOp)2449 return std::nullopt;2450 BlockArgument bbArg = forallOp.getTiedBlockArgument(result);2451 SmallVector<Operation *> combiningOps = forallOp.getCombiningOps(bbArg);2452 // If the number of combining ops is not 1, then this is unexpected. Return2453 // nullopt.2454 if (combiningOps.size() != 1)2455 return std::nullopt;2456 return combiningOps[0];2457}2458 2459/// For a given result of the loop nest that is a tiled loop nest, return the2460/// insert slice-like op that is used for consumer fusion2461static std::optional<Operation *>2462getProducingInsertSliceLikeOp(OpResult result,2463 ArrayRef<LoopLikeOpInterface> loops) {2464 assert(!loops.empty() && "Expected loops to be not empty");2465 LoopLikeOpInterface outerMostLoop = loops.front();2466 if (auto forallOp = dyn_cast<scf::ForallOp>(outerMostLoop.getOperation())) {2467 assert(loops.size() == 1 &&2468 "expected only a single loop when tiling using scf.forall");2469 return getProducingParallelInsertSlice(forallOp, result);2470 }2471 // Assume that the loop nest is a nested `scf.for` that is created through2472 // tiling and retrieve the `tensor.insert_slice` operation used to construct2473 // the result.2474 while (loops.size() != 1) {2475 LoopLikeOpInterface loop = loops.front();2476 if (result.getOwner() != loop)2477 return std::nullopt;2478 auto forOp = dyn_cast<scf::ForOp>(loop.getOperation());2479 if (!forOp)2480 return std::nullopt;2481 auto yieldOp = cast<scf::YieldOp>(forOp.getBody()->getTerminator());2482 auto innerForResult =2483 dyn_cast<OpResult>(yieldOp.getOperand(result.getResultNumber()));2484 if (!innerForResult)2485 return std::nullopt;2486 result = innerForResult;2487 loops = loops.drop_front();2488 }2489 LoopLikeOpInterface loop = loops.front();2490 if (result.getOwner() != loop)2491 return std::nullopt;2492 auto forOp = dyn_cast<scf::ForOp>(loop.getOperation());2493 if (!forOp)2494 return std::nullopt;2495 auto yieldOp = cast<scf::YieldOp>(forOp.getBody()->getTerminator());2496 auto insertSliceOp = yieldOp.getOperand(result.getResultNumber())2497 .getDefiningOp<tensor::InsertSliceOp>();2498 if (!insertSliceOp)2499 return std::nullopt;2500 return insertSliceOp;2501}2502 2503FailureOr<scf::SCFFuseConsumerOfSliceResult>2504mlir::scf::tileAndFuseConsumer(RewriterBase &rewriter, Operation *consumer,2505 MutableArrayRef<LoopLikeOpInterface> loops) {2506 if (!isa<TilingInterface>(consumer)) {2507 return rewriter.notifyMatchFailure(2508 consumer, "unhandled consumer that does not implement TilingInterface");2509 }2510 2511 // Return if `loops` is empty, return an error for now. Caller is expected2512 // to handle this case.2513 if (loops.empty()) {2514 return rewriter.notifyMatchFailure(2515 consumer, "cannot call tile and fuse consumer with an empty loop nest");2516 }2517 2518 LoopLikeOpInterface outermostLoop = loops.front();2519 2520 // Collect the operands of the consumer that come from the outermost loop of2521 // the loop nest.2522 SmallVector<OpOperand *> consumerFusableOperands;2523 for (OpOperand &opOperand : consumer->getOpOperands()) {2524 if (opOperand.get().getDefiningOp() == outermostLoop) {2525 consumerFusableOperands.push_back(&opOperand);2526 }2527 }2528 2529 // Nothing to fuse. Just return an empty set.2530 if (consumerFusableOperands.empty()) {2531 return mlir::scf::SCFFuseConsumerOfSliceResult{consumerFusableOperands,2532 SmallVector<OpOperand *>{},2533 SmallVector<Operation *>{}};2534 }2535 2536 // Collect the relevant tensor.insert_slice/tensor.parallel_insert_slices2537 // for fusion.2538 SmallVector<Operation *> candidateSlices;2539 candidateSlices.reserve(consumerFusableOperands.size());2540 for (OpOperand *opOperand : consumerFusableOperands) {2541 std::optional<Operation *> slice =2542 getProducingInsertSliceLikeOp(cast<OpResult>(opOperand->get()), loops);2543 if (!slice) {2544 return rewriter.notifyMatchFailure(2545 consumer,2546 "couldnt find producing insert-slice like operation for operand");2547 }2548 candidateSlices.push_back(slice.value());2549 }2550 return tileAndFuseConsumerOfSlicesImpl(2551 rewriter, consumer, consumerFusableOperands, candidateSlices, loops);2552}2553 2554//===----------------------------------------------------------------------===//2555// lowerToLoopsUsingSCFForOp implementation.2556//===----------------------------------------------------------------------===//2557 2558FailureOr<SmallVector<scf::ForOp>>2559mlir::scf::lowerToLoopsUsingSCFForOp(RewriterBase &rewriter,2560 TilingInterface op) {2561 // TODO: Handle cases where the op has results if needed.2562 if (op->getNumResults() > 0) {2563 return rewriter.notifyMatchFailure(2564 op, "unable to lower to loops operations with return values");2565 }2566 2567 SmallVector<Range> domain = op.getIterationDomain(rewriter);2568 SmallVector<Value> ivs;2569 SmallVector<scf::ForOp> loops;2570 Location loc = op.getLoc();2571 for (auto loopRange : domain) {2572 Value offsetVal =2573 getValueOrCreateConstantIndexOp(rewriter, loc, loopRange.offset);2574 Value sizeVal =2575 getValueOrCreateConstantIndexOp(rewriter, loc, loopRange.size);2576 Value strideVal =2577 getValueOrCreateConstantIndexOp(rewriter, loc, loopRange.stride);2578 auto loop = scf::ForOp::create(rewriter, op.getLoc(), offsetVal, sizeVal,2579 strideVal, ValueRange{});2580 loops.push_back(loop);2581 ivs.push_back(loop.getInductionVar());2582 rewriter.setInsertionPoint(loop.getBody()->getTerminator());2583 }2584 if (failed(op.generateScalarImplementation(rewriter, op.getLoc(), ivs))) {2585 return failure();2586 }2587 return loops;2588}2589