brintos

brintos / llvm-project-archived public Read only

0
0
Text · 107.7 KiB · 009c2c3 Raw
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