1 //===- Utils.cpp - Utilities to support the Linalg dialect ----------------===//
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-exception
6 //
7 //===----------------------------------------------------------------------===//
8 //
9 // This file implements utilities for the Linalg dialect.
10 //
11 //===----------------------------------------------------------------------===//
12 
13 #include "mlir/Dialect/Linalg/Utils/Utils.h"
14 
15 #include "mlir/Dialect/Affine/IR/AffineOps.h"
16 #include "mlir/Dialect/Linalg/IR/LinalgOps.h"
17 #include "mlir/Dialect/Linalg/IR/LinalgTypes.h"
18 #include "mlir/Dialect/SCF/SCF.h"
19 #include "mlir/Dialect/StandardOps/IR/Ops.h"
20 #include "mlir/Dialect/StandardOps/Utils/Utils.h"
21 #include "mlir/IR/AffineExpr.h"
22 #include "mlir/IR/AffineExprVisitor.h"
23 #include "mlir/IR/AffineMap.h"
24 #include "mlir/IR/Matchers.h"
25 #include "mlir/IR/OpImplementation.h"
26 #include "mlir/Pass/Pass.h"
27 #include "mlir/Transforms/LoopUtils.h"
28 #include "llvm/Support/Debug.h"
29 
30 #define DEBUG_TYPE "linalg-utils"
31 
32 using namespace mlir;
33 using namespace mlir::linalg;
34 using namespace mlir::scf;
35 
36 static bool isZero(Value v) {
37   if (auto cst = v.getDefiningOp<ConstantIndexOp>())
38     return cst.getValue() == 0;
39   return false;
40 }
41 
42 namespace {
43 
44 // Helper visitor to determine whether an AffineExpr is tiled.
45 // This is achieved by traversing every AffineDimExpr with position `pos` and
46 // checking whether the corresponding `tileSizes[pos]` is non-zero.
47 // This also enforces only positive coefficients occur in multiplications.
48 //
49 // Example:
50 //   `d0 + 2 * d1 + d3` is tiled by [0, 0, 0, 2] but not by [0, 0, 2, 0]
51 //
52 struct TileCheck : public AffineExprVisitor<TileCheck> {
53   TileCheck(ValueRange tileSizes) : isTiled(false), tileSizes(tileSizes) {}
54 
55   void visitDimExpr(AffineDimExpr expr) {
56     isTiled |= !isZero(tileSizes[expr.getPosition()]);
57   }
58   void visitAffineBinaryOpExpr(AffineBinaryOpExpr expr) {
59     visit(expr.getLHS());
60     visit(expr.getRHS());
61     if (expr.getKind() == mlir::AffineExprKind::Mul)
62       assert(expr.getRHS().cast<AffineConstantExpr>().getValue() > 0 &&
63              "nonpositive multiplying coefficient");
64   }
65   bool isTiled;
66   ValueRange tileSizes;
67 };
68 
69 } // namespace
70 
71 static bool isTiled(AffineExpr expr, ValueRange tileSizes) {
72   if (!expr)
73     return false;
74   TileCheck t(tileSizes);
75   t.visit(expr);
76   return t.isTiled;
77 }
78 
79 // Checks whether the `map  varies with respect to a non-zero `tileSize`.
80 static bool isTiled(AffineMap map, ValueRange tileSizes) {
81   if (!map)
82     return false;
83   for (unsigned r = 0; r < map.getNumResults(); ++r)
84     if (isTiled(map.getResult(r), tileSizes))
85       return true;
86   return false;
87 }
88 
89 Optional<RegionMatcher::BinaryOpKind>
90 RegionMatcher::matchAsScalarBinaryOp(GenericOp op) {
91   auto &region = op.region();
92   if (!llvm::hasSingleElement(region))
93     return llvm::None;
94 
95   Block &block = region.front();
96   if (block.getNumArguments() != 2 ||
97       !block.getArgument(0).getType().isSignlessIntOrFloat() ||
98       !block.getArgument(1).getType().isSignlessIntOrFloat())
99     return llvm::None;
100 
101   auto &ops = block.getOperations();
102   if (!llvm::hasSingleElement(block.without_terminator()))
103     return llvm::None;
104 
105   using mlir::matchers::m_Val;
106   auto a = m_Val(block.getArgument(0));
107   auto b = m_Val(block.getArgument(1));
108 
109   auto addPattern = m_Op<linalg::YieldOp>(m_Op<AddIOp>(a, b));
110   if (addPattern.match(&ops.back()))
111     return BinaryOpKind::IAdd;
112 
113   return llvm::None;
114 }
115 
116 bool mlir::linalg::isParallelIteratorType(Attribute attr) {
117   if (auto strAttr = attr.dyn_cast<StringAttr>()) {
118     return strAttr.getValue() == getParallelIteratorTypeName();
119   }
120   return false;
121 }
122 
123 bool mlir::linalg::isReductionIteratorType(Attribute attr) {
124   if (auto strAttr = attr.dyn_cast<StringAttr>()) {
125     return strAttr.getValue() == getReductionIteratorTypeName();
126   }
127   return false;
128 }
129 
130 bool mlir::linalg::isWindowIteratorType(Attribute attr) {
131   if (auto strAttr = attr.dyn_cast<StringAttr>()) {
132     return strAttr.getValue() == getWindowIteratorTypeName();
133   }
134   return false;
135 }
136 
137 /// Explicit instantiation of loop nest generator for different loop types.
138 template struct mlir::linalg::GenerateLoopNest<scf::ForOp>;
139 template struct mlir::linalg::GenerateLoopNest<scf::ParallelOp>;
140 template struct mlir::linalg::GenerateLoopNest<AffineForOp>;
141 template struct mlir::linalg::GenerateLoopNest<TiledLoopOp>;
142 
143 /// Given a list of subview ranges, extract individual values for lower, upper
144 /// bounds and steps and put them into the corresponding vectors.
145 static void unpackRanges(ArrayRef<Range> ranges, SmallVectorImpl<Value> &lbs,
146                          SmallVectorImpl<Value> &ubs,
147                          SmallVectorImpl<Value> &steps) {
148   for (Range range : ranges) {
149     lbs.emplace_back(range.offset);
150     ubs.emplace_back(range.size);
151     steps.emplace_back(range.stride);
152   }
153 }
154 
155 namespace mlir {
156 namespace linalg {
157 
158 /// If `size` comes from an AffineMinOp and one of the values of AffineMinOp
159 /// is a constant then return a new value set to the smallest such constant.
160 /// Otherwise returngetSmallestBoundingIndex nullptr.
161 IntegerAttr getSmallestBoundingIndex(Value size) {
162   Optional<int64_t> boundingConst = {};
163   if (auto affineMinOp = size.getDefiningOp<AffineMinOp>()) {
164     for (auto e : affineMinOp.getAffineMap().getResults())
165       if (auto cst = e.dyn_cast<AffineConstantExpr>())
166         boundingConst = boundingConst
167                             ? std::min(boundingConst.getValue(), cst.getValue())
168                             : cst.getValue();
169   } else if (auto constIndexOp = size.getDefiningOp<ConstantOp>()) {
170     if (constIndexOp.getType().isa<IndexType>())
171       boundingConst = constIndexOp.value().cast<IntegerAttr>().getInt();
172   } else if (auto affineApplyOp = size.getDefiningOp<AffineApplyOp>()) {
173     if (auto cExpr = affineApplyOp.getAffineMap()
174                          .getResult(0)
175                          .dyn_cast<AffineConstantExpr>())
176       boundingConst = cExpr.getValue();
177   } else if (auto dimOp = size.getDefiningOp<memref::DimOp>()) {
178     auto shape = dimOp.memrefOrTensor().getType().dyn_cast<ShapedType>();
179     if (auto constOp = dimOp.index().getDefiningOp<ConstantOp>()) {
180       if (auto indexAttr = constOp.value().dyn_cast<IntegerAttr>()) {
181         auto dimIndex = indexAttr.getInt();
182         if (!shape.isDynamicDim(dimIndex)) {
183           boundingConst = shape.getShape()[dimIndex];
184         }
185       }
186     }
187   }
188   if (boundingConst && *boundingConst >= 0)
189     return Builder(size.getContext()).getIndexAttr(*boundingConst);
190   return nullptr;
191 }
192 
193 /// Specialization to build an scf "for" nest.
194 template <>
195 void GenerateLoopNest<scf::ForOp>::doit(
196     OpBuilder &b, Location loc, ArrayRef<Range> loopRanges, LinalgOp linalgOp,
197     ArrayRef<Attribute> iteratorTypes,
198     function_ref<scf::ValueVector(OpBuilder &, Location, ValueRange,
199                                   ValueRange)>
200         bodyBuilderFn,
201     Optional<LinalgLoopDistributionOptions> distributionOptions,
202     ArrayRef<StringRef> distributionTypes) {
203   SmallVector<Value> iterArgInitValues = linalgOp.getOutputTensorOperands();
204   // Create procInfo so it dominates loops, if appropriate.
205   SmallVector<ProcInfo, 4> procInfo;
206   SmallVector<DistributionMethod, 0> distributionMethod;
207   if (distributionOptions.hasValue()) {
208     // Collect loop ranges for parallel dimensions.
209     SmallVector<Range, 2> parallelLoopRanges;
210     for (auto iteratorType : enumerate(iteratorTypes))
211       if (isParallelIteratorType(iteratorType.value()))
212         parallelLoopRanges.push_back(loopRanges[iteratorType.index()]);
213 
214     // Get their distribution schemes.
215     distributionMethod = distributionOptions->distributionMethod;
216     if (distributionMethod.size() < parallelLoopRanges.size())
217       parallelLoopRanges.resize(distributionMethod.size());
218     procInfo = distributionOptions->procInfo(b, loc, parallelLoopRanges);
219   }
220 
221   SmallVector<Value, 4> lbs, ubs, steps;
222   unpackRanges(loopRanges, lbs, ubs, steps);
223   LoopNest loopNest = mlir::scf::buildLoopNest(
224       b, loc, lbs, ubs, steps, iterArgInitValues, bodyBuilderFn);
225 
226   if (!distributionOptions || loopNest.loops.empty())
227     return;
228 
229   // Filter out scf.for loops that were created out of parallel dimensions.
230   SmallVector<scf::ForOp, 4> loops;
231   for (auto iteratorType : enumerate(iteratorTypes))
232     if (isParallelIteratorType(iteratorType.value()))
233       loops.push_back(loopNest.loops[iteratorType.index()]);
234 
235   // Distribute - only supports cyclic distribution for now.
236   for (auto it : llvm::zip(loops, procInfo, distributionMethod))
237     if (std::get<2>(it) == DistributionMethod::Cyclic)
238       mapLoopToProcessorIds(std::get<0>(it), std::get<1>(it).procId,
239                             std::get<1>(it).nprocs);
240 }
241 
242 /// Specialization to build affine "for" nest.
243 template <>
244 void GenerateLoopNest<AffineForOp>::doit(
245     OpBuilder &b, Location loc, ArrayRef<Range> loopRanges, LinalgOp linalgOp,
246     ArrayRef<Attribute> iteratorTypes,
247     function_ref<scf::ValueVector(OpBuilder &, Location, ValueRange,
248                                   ValueRange)>
249         bodyBuilderFn,
250     Optional<LinalgLoopDistributionOptions>, ArrayRef<StringRef>) {
251   SmallVector<Value> iterArgInitValues = linalgOp.getOutputTensorOperands();
252   assert(iterArgInitValues.empty() && "unexpected AffineForOp init values");
253   SmallVector<Value, 4> lbs, ubs, steps;
254   unpackRanges(loopRanges, lbs, ubs, steps);
255 
256   // Affine loops require constant steps.
257   SmallVector<int64_t, 4> constantSteps;
258   constantSteps.reserve(steps.size());
259   for (Value v : steps) {
260     auto op = v.getDefiningOp<ConstantIndexOp>();
261     assert(op && "Affine loops require constant steps");
262     constantSteps.push_back(op.getValue());
263   }
264 
265   mlir::buildAffineLoopNest(b, loc, lbs, ubs, constantSteps,
266                             [&](OpBuilder &b, Location loc, ValueRange ivs) {
267                               bodyBuilderFn(b, loc, ivs, {});
268                             });
269 }
270 
271 /// Specialization to build an linalg.tiled_loop
272 template <>
273 void GenerateLoopNest<TiledLoopOp>::doit(
274     OpBuilder &b, Location loc, ArrayRef<Range> loopRanges, LinalgOp linalgOp,
275     ArrayRef<Attribute> iteratorTypes,
276     function_ref<scf::ValueVector(OpBuilder &, Location, ValueRange,
277                                   ValueRange)>
278         bodyBuilderFn,
279     Optional<LinalgLoopDistributionOptions> distributionOptions,
280     ArrayRef<StringRef> distributionTypes) {
281   SmallVector<ProcInfo, 2> procInfo;
282   SmallVector<Value, 4> lbs, ubs, steps;
283   unpackRanges(loopRanges, lbs, ubs, steps);
284 
285   auto wrappedBuilderFn = [&](OpBuilder &nestedBuilder, Location nestedLoc,
286                               ValueRange ivs, ValueRange inputs,
287                               ValueRange outputs) {
288     SmallVector<Value> outputTensors = linalgOp.getOutputTensorOperands();
289     scf::ValueVector results =
290         bodyBuilderFn(nestedBuilder, nestedLoc, ivs, outputTensors);
291     nestedBuilder.create<linalg::YieldOp>(nestedLoc, results);
292   };
293 
294   SmallVector<Value> inputOperands = linalgOp.getInputOperands();
295   SmallVector<Value> outputOperands = linalgOp.getOutputOperands();
296   auto tiledLoop =
297       b.create<TiledLoopOp>(loc, lbs, ubs, steps, inputOperands, outputOperands,
298                             b.getArrayAttr(iteratorTypes), wrappedBuilderFn);
299   if (!distributionTypes.empty())
300     tiledLoop.setDistributionTypes(b, distributionTypes);
301 
302   // Replace inputs/outputs with the corresponding region args.
303   auto isInsideTiledLoop = [&](OpOperand &operand) {
304     return operand.getOwner()->getBlock() == tiledLoop.getBody();
305   };
306   for (auto it : llvm::zip(inputOperands, tiledLoop.getRegionInputArgs()))
307     std::get<0>(it).replaceUsesWithIf(std::get<1>(it), isInsideTiledLoop);
308   for (auto it : llvm::zip(outputOperands, tiledLoop.getRegionOutputArgs()))
309     std::get<0>(it).replaceUsesWithIf(std::get<1>(it), isInsideTiledLoop);
310 }
311 
312 /// Update the `lb`, `ub` and `step` to get per processor `lb`, `ub` and `step`.
313 void updateBoundsForCyclicDistribution(OpBuilder &b, Location loc, Value procId,
314                                        Value nprocs, Value &lb, Value &ub,
315                                        Value &step) {
316   AffineExpr d0, d1;
317   bindDims(b.getContext(), d0, d1);
318   AffineExpr s0 = getAffineSymbolExpr(0, b.getContext());
319   lb = makeComposedAffineApply(b, loc, d0 + d1 * s0, {lb, procId, step});
320   step = makeComposedAffineApply(b, loc, d0 * s0, {nprocs, step});
321 }
322 
323 /// Generates a loop nest consisting of scf.parallel and scf.for, depending
324 /// on the `iteratorTypes.` Consecutive parallel loops create a single
325 /// scf.parallel operation; each sequential loop creates a new scf.for
326 /// operation. The body of the innermost loop is populated by
327 /// `bodyBuilderFn` that accepts a range of induction variables for all
328 /// loops. `ivStorage` is used to store the partial list of induction
329 /// variables.
330 // TODO: this function can be made iterative instead. However, it
331 // will have at most as many recursive calls as nested loops, which rarely
332 // exceeds 10.
333 static void generateParallelLoopNest(
334     OpBuilder &b, Location loc, ValueRange lbs, ValueRange ubs,
335     ValueRange steps, ArrayRef<Attribute> iteratorTypes,
336     function_ref<void(OpBuilder &, Location, ValueRange)> bodyBuilderFn,
337     SmallVectorImpl<Value> &ivStorage,
338     ArrayRef<DistributionMethod> distributionMethod = {}) {
339   assert(lbs.size() == ubs.size());
340   assert(lbs.size() == steps.size());
341   assert(lbs.size() == iteratorTypes.size());
342 
343   // If there are no (more) loops to be generated, generate the body and be
344   // done with it.
345   if (iteratorTypes.empty()) {
346     bodyBuilderFn(b, loc, ivStorage);
347     return;
348   }
349 
350   // Find the outermost parallel loops and drop their types from the list.
351   unsigned nLoops = iteratorTypes.size();
352   unsigned nOuterPar =
353       nLoops - iteratorTypes.drop_while(isParallelIteratorType).size();
354 
355   // If there are no outer parallel loops, generate one sequential loop and
356   // recurse. Note that we wouldn't have dropped anything from `iteratorTypes`
357   // in this case.
358   if (nOuterPar == 0) {
359     LoopNest singleLoop = buildLoopNest(
360         b, loc, lbs.take_front(), ubs.take_front(), steps.take_front(),
361         [&](OpBuilder &b, Location loc, ValueRange ivs) {
362           ivStorage.append(ivs.begin(), ivs.end());
363           generateParallelLoopNest(b, loc, lbs.drop_front(), ubs.drop_front(),
364                                    steps.drop_front(),
365                                    iteratorTypes.drop_front(), bodyBuilderFn,
366                                    ivStorage, distributionMethod);
367         });
368     return;
369   }
370   if (distributionMethod.empty()) {
371     // Generate a single parallel loop-nest operation for all outermost
372     // parallel loops and recurse.
373     b.create<scf::ParallelOp>(
374         loc, lbs.take_front(nOuterPar), ubs.take_front(nOuterPar),
375         steps.take_front(nOuterPar),
376         [&](OpBuilder &nestedBuilder, Location nestedLoc, ValueRange localIvs) {
377           ivStorage.append(localIvs.begin(), localIvs.end());
378           generateParallelLoopNest(
379               nestedBuilder, nestedLoc, lbs.drop_front(nOuterPar),
380               ubs.drop_front(nOuterPar), steps.drop_front(nOuterPar),
381               iteratorTypes.drop_front(nOuterPar), bodyBuilderFn, ivStorage,
382               (distributionMethod.size() < nOuterPar)
383                   ? ArrayRef<DistributionMethod>()
384                   : distributionMethod.drop_front(nOuterPar));
385         });
386     return;
387   }
388 
389   // Process all consecutive similarly distributed loops simultaneously.
390   DistributionMethod methodToUse = distributionMethod[0];
391   unsigned numProcessed = 1;
392   for (unsigned i = 1; i < nOuterPar && i < distributionMethod.size(); ++i) {
393     if (distributionMethod[i] != methodToUse)
394       break;
395     numProcessed++;
396   }
397 
398   switch (methodToUse) {
399   case DistributionMethod::Cyclic: {
400     // Generate a single parallel loop-nest operation for all outermost
401     // parallel loops and recurse.
402     b.create<scf::ParallelOp>(
403         loc, lbs.take_front(numProcessed), ubs.take_front(numProcessed),
404         steps.take_front(numProcessed),
405         [&](OpBuilder &nestedBuilder, Location nestedLoc, ValueRange localIvs) {
406           ivStorage.append(localIvs.begin(), localIvs.end());
407           generateParallelLoopNest(
408               nestedBuilder, nestedLoc, lbs.drop_front(numProcessed),
409               ubs.drop_front(numProcessed), steps.drop_front(numProcessed),
410               iteratorTypes.drop_front(numProcessed), bodyBuilderFn, ivStorage,
411               (distributionMethod.size() < numProcessed)
412                   ? ArrayRef<DistributionMethod>()
413                   : distributionMethod.drop_front(numProcessed));
414         });
415     return;
416   }
417   case DistributionMethod::CyclicNumProcsGeNumIters: {
418     // Check (for the processed loops) that the iteration is in-bounds.
419     ArithBuilder ab(b, loc);
420     Value cond = ab.slt(lbs[0], ubs[0]);
421     for (unsigned i = 1; i < numProcessed; ++i)
422       cond = ab._and(cond, ab.slt(lbs[i], ubs[i]));
423     ivStorage.append(lbs.begin(), std::next(lbs.begin(), numProcessed));
424     b.create<scf::IfOp>(loc, cond, [&](OpBuilder &b, Location loc) {
425       generateParallelLoopNest(
426           b, loc, lbs.drop_front(numProcessed), ubs.drop_front(numProcessed),
427           steps.drop_front(numProcessed),
428           iteratorTypes.drop_front(numProcessed), bodyBuilderFn, ivStorage,
429           distributionMethod.drop_front(numProcessed));
430       b.create<scf::YieldOp>(loc, ValueRange{});
431     });
432     return;
433   }
434   case DistributionMethod::CyclicNumProcsEqNumIters:
435     // No check/loops needed here. Set the `%iv` to be the `%lb` and proceed
436     // with inner loop generation.
437     ivStorage.append(lbs.begin(), std::next(lbs.begin(), numProcessed));
438     generateParallelLoopNest(
439         b, loc, lbs.drop_front(numProcessed), ubs.drop_front(numProcessed),
440         steps.drop_front(numProcessed), iteratorTypes.drop_front(numProcessed),
441         bodyBuilderFn, ivStorage, distributionMethod.drop_front(numProcessed));
442     return;
443   }
444 }
445 
446 /// Specialization for generating a mix of parallel and sequential scf loops.
447 template <>
448 void GenerateLoopNest<scf::ParallelOp>::doit(
449     OpBuilder &b, Location loc, ArrayRef<Range> loopRanges, LinalgOp linalgOp,
450     ArrayRef<Attribute> iteratorTypes,
451     function_ref<scf::ValueVector(OpBuilder &, Location, ValueRange,
452                                   ValueRange)>
453         bodyBuilderFn,
454     Optional<LinalgLoopDistributionOptions> distributionOptions,
455     ArrayRef<StringRef> distributionTypes) {
456   SmallVector<Value> iterArgInitValues = linalgOp.getOutputTensorOperands();
457   assert(iterArgInitValues.empty() && "unexpected ParallelOp init values");
458   // This function may be passed more iterator types than ranges.
459   assert(iteratorTypes.size() >= loopRanges.size() &&
460          "expected iterator type for all ranges");
461   iteratorTypes = iteratorTypes.take_front(loopRanges.size());
462   SmallVector<Value, 8> lbsStorage, ubsStorage, stepsStorage, ivs;
463   unsigned numLoops = iteratorTypes.size();
464   ivs.reserve(numLoops);
465   lbsStorage.reserve(numLoops);
466   ubsStorage.reserve(numLoops);
467   stepsStorage.reserve(numLoops);
468 
469   // Get the loop lb, ub, and step.
470   unpackRanges(loopRanges, lbsStorage, ubsStorage, stepsStorage);
471 
472   // Modify the lb, ub, and step based on the distribution options.
473   SmallVector<DistributionMethod, 0> distributionMethod;
474   if (distributionOptions) {
475     auto &options = distributionOptions.getValue();
476     distributionMethod.assign(distributionOptions->distributionMethod.begin(),
477                               distributionOptions->distributionMethod.end());
478     SmallVector<Range, 2> parallelLoopRanges;
479     for (auto iteratorType : enumerate(iteratorTypes)) {
480       if (isParallelIteratorType(iteratorType.value()))
481         parallelLoopRanges.push_back(loopRanges[iteratorType.index()]);
482     }
483     if (distributionMethod.size() < parallelLoopRanges.size())
484       parallelLoopRanges.resize(distributionMethod.size());
485     SmallVector<ProcInfo, 2> procInfo =
486         options.procInfo(b, loc, parallelLoopRanges);
487     unsigned index = 0;
488     for (auto iteratorType : enumerate(iteratorTypes)) {
489       if (index >= procInfo.size())
490         break;
491       if (isParallelIteratorType(iteratorType.value())) {
492         unsigned i = iteratorType.index();
493         updateBoundsForCyclicDistribution(b, loc, procInfo[index].procId,
494                                           procInfo[index].nprocs, lbsStorage[i],
495                                           ubsStorage[i], stepsStorage[i]);
496         index++;
497       }
498     }
499   }
500   ValueRange lbs(lbsStorage), ubs(ubsStorage), steps(stepsStorage);
501   generateParallelLoopNest(
502       b, loc, lbs, ubs, steps, iteratorTypes,
503       [&](OpBuilder &b, Location loc, ValueRange ivs) {
504         bodyBuilderFn(b, loc, ivs, {});
505       },
506       ivs, distributionMethod);
507 
508   assert(ivs.size() == iteratorTypes.size() && "did not generate enough loops");
509 }
510 
511 SmallVector<Value, 4> makeTiledShapes(OpBuilder &b, Location loc,
512                                       LinalgOp linalgOp,
513                                       ArrayRef<Value> valuesToTile,
514                                       ValueRange ivs, ValueRange tileSizes,
515                                       ArrayRef<Value> sizeBounds) {
516   assert(ivs.size() == static_cast<size_t>(llvm::count_if(
517                            llvm::make_range(tileSizes.begin(), tileSizes.end()),
518                            [](Value v) { return !isZero(v); })) &&
519          "expected as many ivs as non-zero sizes");
520 
521   // Construct (potentially temporary) mins and maxes on which to apply maps
522   // that define tile subshapes.
523   SmallVector<Value, 8> lbs, subShapeSizes;
524   for (unsigned idx = 0, idxIvs = 0, e = tileSizes.size(); idx < e; ++idx) {
525     LLVM_DEBUG(llvm::dbgs() << "makeTiledShapes: for loop#" << idx << "\n");
526     bool isTiled = !isZero(tileSizes[idx]);
527     lbs.push_back(isTiled ? ivs[idxIvs++]
528                           : (Value)b.create<ConstantIndexOp>(loc, 0));
529     // Before composing, we need to make range a closed interval.
530     Value size = isTiled ? tileSizes[idx] : sizeBounds[idx];
531     AffineExpr d0 = getAffineDimExpr(0, b.getContext());
532     subShapeSizes.push_back(makeComposedAffineApply(b, loc, d0 - 1, size));
533     LLVM_DEBUG(llvm::dbgs() << "lb: " << lbs.back() << "\n");
534     LLVM_DEBUG(llvm::dbgs() << "size: " << subShapeSizes.back() << "\n");
535   }
536 
537   assert(static_cast<int64_t>(valuesToTile.size()) ==
538              linalgOp.getNumInputsAndOutputs() &&
539          "expected one value to tile for every operand");
540   MLIRContext *context = b.getContext();
541   SmallVector<Value, 4> tiledShapes;
542   tiledShapes.reserve(valuesToTile.size());
543   for (OpOperand *opOperand : linalgOp.getInputAndOutputOperands()) {
544     Value shapedOp = valuesToTile[opOperand->getOperandNumber()];
545     LLVM_DEBUG(llvm::dbgs() << "makeTiledShapes: for operand " << shapedOp);
546     int64_t rank = linalgOp.getRank(opOperand);
547     ArrayRef<int64_t> shape = linalgOp.getShape(opOperand);
548     AffineMap map = linalgOp.getTiedIndexingMap(opOperand);
549     // If the shape is not tiled, we can use it as is.
550     if (!isTiled(map, tileSizes)) {
551       tiledShapes.push_back(shapedOp);
552       LLVM_DEBUG(llvm::dbgs() << ": not tiled: use shape: "
553                               << opOperand->get().getType() << "\n");
554       continue;
555     }
556     LLVM_DEBUG(llvm::dbgs() << ": tiled: figure out subshape...\n");
557 
558     // Construct a new subview / subtensor for the tile.
559     SmallVector<OpFoldResult, 4> offsets, sizes, strides;
560     offsets.reserve(rank);
561     sizes.reserve(rank);
562     strides.reserve(rank);
563     for (unsigned r = 0; r < rank; ++r) {
564       LLVM_DEBUG(llvm::dbgs() << "makeTiledShapes: for dim#" << r);
565       if (!isTiled(map.getSubMap({r}), tileSizes)) {
566         offsets.push_back(b.getIndexAttr(0));
567         Value dim = b.createOrFold<memref::DimOp>(loc, shapedOp, r);
568         sizes.push_back(dim);
569         strides.push_back(b.getIndexAttr(1));
570         LLVM_DEBUG(llvm::dbgs() << ": not tiled: use size: " << dim << "\n");
571         continue;
572       }
573       LLVM_DEBUG(llvm::dbgs() << ": tiled: figure out subsize...\n");
574 
575       // Tiling creates a new slice at the proper index, the slice step is 1
576       // (i.e. the op does not subsample, stepping occurs in the loop).
577       auto m = map.getSubMap({r});
578       LLVM_DEBUG(llvm::dbgs() << "makeTiledShapes: submap: " << map << "\n");
579       auto offset = applyMapToValues(b, loc, m, lbs).front();
580       offsets.push_back(offset);
581       auto closedIntSize = applyMapToValues(b, loc, m, subShapeSizes).front();
582       // Resulting size needs to be made half open interval again.
583       AffineExpr s0 = getAffineSymbolExpr(0, b.getContext());
584       Value size = makeComposedAffineApply(b, loc, s0 + 1, closedIntSize);
585       LLVM_DEBUG(llvm::dbgs() << "makeTiledShapes: raw size: " << size << "\n");
586 
587       // The size of the subview / subtensor should be trimmed to avoid
588       // out-of-bounds accesses, unless we statically know the subshape size
589       // divides the shape size evenly.
590       int64_t shapeSize = shape[r];
591       auto sizeCst = size.getDefiningOp<ConstantIndexOp>();
592       if (ShapedType::isDynamic(shapeSize) || !sizeCst ||
593           (shapeSize % sizeCst.getValue()) != 0) {
594         LLVM_DEBUG(llvm::dbgs() << "makeTiledShapes: shapeSize=" << shapeSize
595                                 << ", size: " << size
596                                 << ": make sure in bound with affine.min\n");
597         AffineExpr dim0, dim1, dim2;
598         bindDims(context, dim0, dim1, dim2);
599         // Compute min(size, dim - offset) to avoid out-of-bounds accesses.
600         AffineMap minMap =
601             AffineMap::inferFromExprList(
602                 ArrayRef<ArrayRef<AffineExpr>>{{dim0, dim1 - dim2}})
603                 .front();
604         Value d = b.create<memref::DimOp>(loc, shapedOp, r);
605         SmallVector<Value, 4> operands{size, d, offset};
606         fullyComposeAffineMapAndOperands(&minMap, &operands);
607         size = b.create<AffineMinOp>(loc, b.getIndexType(), minMap, operands);
608       }
609 
610       sizes.push_back(size);
611       LLVM_DEBUG(llvm::dbgs()
612                  << "makeTiledShapes: new offset: " << offset << "\n");
613       LLVM_DEBUG(llvm::dbgs() << "makeTiledShapes: new size: " << size << "\n");
614       strides.push_back(b.getIndexAttr(1));
615     }
616 
617     if (opOperand->get().getType().isa<MemRefType>())
618       tiledShapes.push_back(
619           b.create<memref::SubViewOp>(loc, shapedOp, offsets, sizes, strides));
620     else
621       tiledShapes.push_back(
622           b.create<SubTensorOp>(loc, shapedOp, offsets, sizes, strides));
623   }
624 
625   return tiledShapes;
626 }
627 
628 } // namespace linalg
629 } // namespace mlir
630