MLIR 24.0.0git
Utils.cpp
Go to the documentation of this file.
1//===- Utils.cpp ---- Misc utilities for loop transformation ----------===//
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 miscellaneous loop transformation routines.
10//
11//===----------------------------------------------------------------------===//
12
20#include "mlir/IR/IRMapping.h"
25#include "llvm/ADT/APInt.h"
26#include "llvm/ADT/STLExtras.h"
27#include "llvm/ADT/SmallVector.h"
28#include "llvm/ADT/SmallVectorExtras.h"
29#include "llvm/Support/DebugLog.h"
30#include <cstdint>
31
32using namespace mlir;
33
34#define DEBUG_TYPE "scf-utils"
35
37 RewriterBase &rewriter, MutableArrayRef<scf::ForOp> loopNest,
38 ValueRange newIterOperands, const NewYieldValuesFn &newYieldValuesFn,
39 bool replaceIterOperandsUsesInLoop) {
40 if (loopNest.empty())
41 return {};
42 // This method is recursive (to make it more readable). Adding an
43 // assertion here to limit the recursion. (See
44 // https://discourse.llvm.org/t/rfc-update-to-mlir-developer-policy-on-recursion/62235)
45 assert(loopNest.size() <= 10 &&
46 "exceeded recursion limit when yielding value from loop nest");
47
48 // To yield a value from a perfectly nested loop nest, the following
49 // pattern needs to be created, i.e. starting with
50 //
51 // ```mlir
52 // scf.for .. {
53 // scf.for .. {
54 // scf.for .. {
55 // %value = ...
56 // }
57 // }
58 // }
59 // ```
60 //
61 // needs to be modified to
62 //
63 // ```mlir
64 // %0 = scf.for .. iter_args(%arg0 = %init) {
65 // %1 = scf.for .. iter_args(%arg1 = %arg0) {
66 // %2 = scf.for .. iter_args(%arg2 = %arg1) {
67 // %value = ...
68 // scf.yield %value
69 // }
70 // scf.yield %2
71 // }
72 // scf.yield %1
73 // }
74 // ```
75 //
76 // The inner most loop is handled using the `replaceWithAdditionalYields`
77 // that works on a single loop.
78 if (loopNest.size() == 1) {
79 auto innerMostLoop =
80 cast<scf::ForOp>(*loopNest.back().replaceWithAdditionalYields(
81 rewriter, newIterOperands, replaceIterOperandsUsesInLoop,
82 newYieldValuesFn));
83 return {innerMostLoop};
84 }
85 // The outer loops are modified by calling this method recursively
86 // - The return value of the inner loop is the value yielded by this loop.
87 // - The region iter args of this loop are the init_args for the inner loop.
88 SmallVector<scf::ForOp> newLoopNest;
90 [&](OpBuilder &innerBuilder, Location loc,
92 newLoopNest = replaceLoopNestWithNewYields(rewriter, loopNest.drop_front(),
93 innerNewBBArgs, newYieldValuesFn,
94 replaceIterOperandsUsesInLoop);
95 return llvm::map_to_vector(
96 newLoopNest.front().getResults().take_back(innerNewBBArgs.size()),
97 [](OpResult r) -> Value { return r; });
98 };
99 scf::ForOp outerMostLoop =
100 cast<scf::ForOp>(*loopNest.front().replaceWithAdditionalYields(
101 rewriter, newIterOperands, replaceIterOperandsUsesInLoop, fn));
102 newLoopNest.insert(newLoopNest.begin(), outerMostLoop);
103 return newLoopNest;
104}
105
106/// Outline a region with a single block into a new FuncOp.
107/// Assumes the FuncOp result types is the type of the yielded operands of the
108/// single block. This constraint makes it easy to determine the result.
109/// This method also clones the `arith::ConstantIndexOp` at the start of
110/// `outlinedFuncBody` to alloc simple canonicalizations. If `callOp` is
111/// provided, it will be set to point to the operation that calls the outlined
112/// function.
113// TODO: support more than single-block regions.
114// TODO: more flexible constant handling.
115FailureOr<func::FuncOp> mlir::outlineSingleBlockRegion(RewriterBase &rewriter,
116 Location loc,
117 Region &region,
118 StringRef funcName,
119 func::CallOp *callOp) {
120 assert(!funcName.empty() && "funcName cannot be empty");
121 if (!region.hasOneBlock())
122 return failure();
123
124 Block *originalBlock = &region.front();
125 Operation *originalTerminator = originalBlock->getTerminator();
126
127 // Outline before current function.
128 OpBuilder::InsertionGuard g(rewriter);
129 rewriter.setInsertionPoint(region.getParentOfType<FunctionOpInterface>());
130
131 SetVector<Value> captures;
132 getUsedValuesDefinedAbove(region, captures);
133
134 ValueRange outlinedValues(captures.getArrayRef());
135 SmallVector<Type> outlinedFuncArgTypes;
136 SmallVector<Location> outlinedFuncArgLocs;
137 // Region's arguments are exactly the first block's arguments as per
138 // Region::getArguments().
139 // Func's arguments are cat(regions's arguments, captures arguments).
140 for (BlockArgument arg : region.getArguments()) {
141 outlinedFuncArgTypes.push_back(arg.getType());
142 outlinedFuncArgLocs.push_back(arg.getLoc());
143 }
144 for (Value value : outlinedValues) {
145 outlinedFuncArgTypes.push_back(value.getType());
146 outlinedFuncArgLocs.push_back(value.getLoc());
147 }
148 FunctionType outlinedFuncType =
149 FunctionType::get(rewriter.getContext(), outlinedFuncArgTypes,
150 originalTerminator->getOperandTypes());
151 auto outlinedFunc =
152 func::FuncOp::create(rewriter, loc, funcName, outlinedFuncType);
153 Block *outlinedFuncBody = outlinedFunc.addEntryBlock();
154
155 // Merge blocks while replacing the original block operands.
156 // Warning: `mergeBlocks` erases the original block, reconstruct it later.
157 int64_t numOriginalBlockArguments = originalBlock->getNumArguments();
158 auto outlinedFuncBlockArgs = outlinedFuncBody->getArguments();
159 {
160 OpBuilder::InsertionGuard g(rewriter);
161 rewriter.setInsertionPointToEnd(outlinedFuncBody);
162 rewriter.mergeBlocks(
163 originalBlock, outlinedFuncBody,
164 outlinedFuncBlockArgs.take_front(numOriginalBlockArguments));
165 // Explicitly set up a new ReturnOp terminator.
166 rewriter.setInsertionPointToEnd(outlinedFuncBody);
167 func::ReturnOp::create(rewriter, loc, originalTerminator->getResultTypes(),
168 originalTerminator->getOperands());
169 }
170
171 // Reconstruct the block that was deleted and add a
172 // terminator(call_results).
173 Block *newBlock = rewriter.createBlock(
174 &region, region.begin(),
175 TypeRange{outlinedFuncArgTypes}.take_front(numOriginalBlockArguments),
176 ArrayRef<Location>(outlinedFuncArgLocs)
177 .take_front(numOriginalBlockArguments));
178 {
179 OpBuilder::InsertionGuard g(rewriter);
180 rewriter.setInsertionPointToEnd(newBlock);
181 SmallVector<Value> callValues;
182 llvm::append_range(callValues, newBlock->getArguments());
183 llvm::append_range(callValues, outlinedValues);
184 auto call = func::CallOp::create(rewriter, loc, outlinedFunc, callValues);
185 if (callOp)
186 *callOp = call;
187
188 // `originalTerminator` was moved to `outlinedFuncBody` and is still valid.
189 // Clone `originalTerminator` to take the callOp results then erase it from
190 // `outlinedFuncBody`.
191 IRMapping bvm;
192 bvm.map(originalTerminator->getOperands(), call->getResults());
193 rewriter.clone(*originalTerminator, bvm);
194 rewriter.eraseOp(originalTerminator);
195 }
196
197 // Lastly, explicit RAUW outlinedValues, only for uses within `outlinedFunc`.
198 // Clone the `arith::ConstantIndexOp` at the start of `outlinedFuncBody`.
199 for (auto it : llvm::zip(outlinedValues, outlinedFuncBlockArgs.take_back(
200 outlinedValues.size()))) {
201 Value orig = std::get<0>(it);
202 Value repl = std::get<1>(it);
203 {
204 OpBuilder::InsertionGuard g(rewriter);
205 rewriter.setInsertionPointToStart(outlinedFuncBody);
207 repl = rewriter.clone(*cst)->getResult(0);
208 }
209 }
210 orig.replaceUsesWithIf(repl, [&](OpOperand &opOperand) {
211 return outlinedFunc->isProperAncestor(opOperand.getOwner());
212 });
213 }
214
215 return outlinedFunc;
216}
217
218LogicalResult mlir::outlineIfOp(RewriterBase &b, scf::IfOp ifOp,
219 func::FuncOp *thenFn, StringRef thenFnName,
220 func::FuncOp *elseFn, StringRef elseFnName) {
221 IRRewriter rewriter(b);
222 Location loc = ifOp.getLoc();
223 FailureOr<func::FuncOp> outlinedFuncOpOrFailure;
224 if (thenFn && !ifOp.getThenRegion().empty()) {
225 outlinedFuncOpOrFailure = outlineSingleBlockRegion(
226 rewriter, loc, ifOp.getThenRegion(), thenFnName);
227 if (failed(outlinedFuncOpOrFailure))
228 return failure();
229 *thenFn = *outlinedFuncOpOrFailure;
230 }
231 if (elseFn && !ifOp.getElseRegion().empty()) {
232 outlinedFuncOpOrFailure = outlineSingleBlockRegion(
233 rewriter, loc, ifOp.getElseRegion(), elseFnName);
234 if (failed(outlinedFuncOpOrFailure))
235 return failure();
236 *elseFn = *outlinedFuncOpOrFailure;
237 }
238 return success();
239}
240
243 assert(rootOp != nullptr && "Root operation must not be a nullptr.");
244 bool rootEnclosesPloops = false;
245 for (Region &region : rootOp->getRegions()) {
246 for (Block &block : region.getBlocks()) {
247 for (Operation &op : block) {
248 bool enclosesPloops = getInnermostParallelLoops(&op, result);
249 rootEnclosesPloops |= enclosesPloops;
250 if (auto ploop = dyn_cast<scf::ParallelOp>(op)) {
251 rootEnclosesPloops = true;
252
253 // Collect parallel loop if it is an innermost one.
254 if (!enclosesPloops)
255 result.push_back(ploop);
256 }
257 }
258 }
259 }
260 return rootEnclosesPloops;
261}
262
263// Build the IR that performs ceil division of a positive value by a constant:
264// ceildiv(a, B) = divis(a + (B-1), B)
265// where divis is rounding-to-zero division.
266static Value ceilDivPositive(OpBuilder &builder, Location loc, Value dividend,
267 int64_t divisor) {
268 assert(divisor > 0 && "expected positive divisor");
269 assert(dividend.getType().isIntOrIndex() &&
270 "expected integer or index-typed value");
271
272 Value divisorMinusOneCst = arith::ConstantOp::create(
273 builder, loc, builder.getIntegerAttr(dividend.getType(), divisor - 1));
274 Value divisorCst = arith::ConstantOp::create(
275 builder, loc, builder.getIntegerAttr(dividend.getType(), divisor));
276 Value sum = arith::AddIOp::create(builder, loc, dividend, divisorMinusOneCst);
277 return arith::DivUIOp::create(builder, loc, sum, divisorCst);
278}
279
280// Build the IR that performs ceil division of a positive value by another
281// positive value:
282// ceildiv(a, b) = divis(a + (b - 1), b)
283// where divis is rounding-to-zero division.
284static Value ceilDivPositive(OpBuilder &builder, Location loc, Value dividend,
285 Value divisor) {
286 assert(dividend.getType().isIntOrIndex() &&
287 "expected integer or index-typed value");
288 Value cstOne = arith::ConstantOp::create(
289 builder, loc, builder.getOneAttr(dividend.getType()));
290 Value divisorMinusOne = arith::SubIOp::create(builder, loc, divisor, cstOne);
291 Value sum = arith::AddIOp::create(builder, loc, dividend, divisorMinusOne);
292 return arith::DivUIOp::create(builder, loc, sum, divisor);
293}
294
296 Block *loopBodyBlock, Value iv, uint64_t unrollFactor,
297 function_ref<Value(unsigned, Value, OpBuilder)> ivRemapFn,
298 function_ref<void(unsigned, Operation *, OpBuilder)> annotateFn,
299 ValueRange iterArgs, ValueRange yieldedValues,
300 IRMapping *clonedToSrcOpsMap) {
301
302 // Check if the op was cloned from another source op, and return it if found
303 // (or the same op if not found)
304 auto findOriginalSrcOp =
305 [](Operation *op, const IRMapping &clonedToSrcOpsMap) -> Operation * {
306 Operation *srcOp = op;
307 // If the source op derives from another op: traverse the chain to find the
308 // original source op
309 while (srcOp && clonedToSrcOpsMap.contains(srcOp))
310 srcOp = clonedToSrcOpsMap.lookup(srcOp);
311 return srcOp;
312 };
313
314 // Builder to insert unrolled bodies just before the terminator of the body of
315 // the loop.
316 auto builder = OpBuilder::atBlockTerminator(loopBodyBlock);
317
318 static const auto noopAnnotateFn = [](unsigned, Operation *, OpBuilder) {};
319 if (!annotateFn)
320 annotateFn = noopAnnotateFn;
321
322 // Keep a pointer to the last non-terminator operation in the original block
323 // so that we know what to clone (since we are doing this in-place).
324 Block::iterator srcBlockEnd = std::prev(loopBodyBlock->end(), 2);
325
326 // Unroll the contents of the loop body (append unrollFactor - 1 additional
327 // copies).
328 SmallVector<Value, 4> lastYielded(yieldedValues);
329
330 for (unsigned i = 1; i < unrollFactor; i++) {
331 // Prepare operand map.
332 IRMapping operandMap;
333 operandMap.map(iterArgs, lastYielded);
334
335 // If the induction variable is used, create a remapping to the value for
336 // this unrolled instance.
337 if (!iv.use_empty()) {
338 Value ivUnroll = ivRemapFn(i, iv, builder);
339 operandMap.map(iv, ivUnroll);
340 }
341
342 // Clone the original body of 'forOp'.
343 for (auto it = loopBodyBlock->begin(); it != std::next(srcBlockEnd); it++) {
344 Operation *srcOp = &(*it);
345 Operation *clonedOp = builder.clone(*srcOp, operandMap);
346 annotateFn(i, clonedOp, builder);
347 if (clonedToSrcOpsMap)
348 clonedToSrcOpsMap->map(clonedOp,
349 findOriginalSrcOp(srcOp, *clonedToSrcOpsMap));
350 }
351
352 // Update yielded values.
353 for (unsigned i = 0, e = lastYielded.size(); i < e; i++)
354 lastYielded[i] = operandMap.lookupOrDefault(yieldedValues[i]);
355 }
356
357 // Make sure we annotate the Ops in the original body. We do this last so that
358 // any annotations are not copied into the cloned Ops above.
359 for (auto it = loopBodyBlock->begin(); it != std::next(srcBlockEnd); it++)
360 annotateFn(0, &*it, builder);
361
362 // Update operands of the yield statement.
363 loopBodyBlock->getTerminator()->setOperands(lastYielded);
364}
365
366/// Unrolls 'forOp' by 'unrollFactor', returns the unrolled main loop and the
367/// epilogue loop, if the loop is unrolled.
368FailureOr<UnrolledLoopInfo> mlir::loopUnrollByFactor(
369 scf::ForOp forOp, uint64_t unrollFactor,
370 function_ref<void(unsigned, Operation *, OpBuilder)> annotateFn,
371 bool shouldPromoteIfSingleIteration) {
372 assert(unrollFactor > 0 && "expected positive unroll factor");
373
374 // Return if the loop body is empty.
375 if (llvm::hasSingleElement(forOp.getBody()->getOperations()))
376 return UnrolledLoopInfo{forOp, std::nullopt};
377
378 // Compute tripCount = ceilDiv((upperBound - lowerBound), step) and populate
379 // 'upperBoundUnrolled' and 'stepUnrolled' for static and dynamic cases.
380 OpBuilder boundsBuilder(forOp);
381 IRRewriter rewriter(forOp.getContext());
382 auto loc = forOp.getLoc();
383 Value step = forOp.getStep();
384 Value upperBoundUnrolled;
385 Value stepUnrolled;
386 bool generateEpilogueLoop = true;
387
388 std::optional<APInt> constTripCount = forOp.getStaticTripCount();
389 if (constTripCount) {
390 // Constant loop bounds computation.
391 bool isUnsignedLoop = forOp.getUnsignedCmp();
392 // For unsigned loops, bounds must be zero-extended: narrow integer types
393 // (e.g. i1, i2, i3) may have bit patterns that are negative in a signed
394 // context (e.g., i1 value 1 has getSExtValue() == -1, getZExtValue() == 1).
395 // Zero-extension is only safe when the unsigned value fits in int64_t, i.e.
396 // the type's bitwidth is < 64. Bail out for 64-bit unsigned loops.
397 if (isUnsignedLoop) {
398 if (auto intTy = dyn_cast<IntegerType>(forOp.getUpperBound().getType()))
399 if (intTy.getWidth() >= 64)
400 return failure();
401 }
402 auto getLoopBound = [&](Value v) -> int64_t {
403 auto apInt = getConstantAPIntValue(v);
404 assert(apInt && "expected constant loop bound");
405 return isUnsignedLoop ? static_cast<int64_t>(apInt->first.getZExtValue())
406 : apInt->first.getSExtValue();
407 };
408 int64_t lbCst = getLoopBound(forOp.getLowerBound());
409 int64_t ubCst = getLoopBound(forOp.getUpperBound());
410 int64_t stepCst = getLoopBound(step);
411 if (unrollFactor == 1) {
412 if (shouldPromoteIfSingleIteration && constTripCount->isOne() &&
413 failed(forOp.promoteIfSingleIteration(rewriter)))
414 return failure();
415 return UnrolledLoopInfo{forOp, std::nullopt};
416 }
417
418 uint64_t tripCount = constTripCount->getZExtValue();
419 uint64_t tripCountEvenMultiple = tripCount - tripCount % unrollFactor;
420 int64_t upperBoundUnrolledCst = lbCst + tripCountEvenMultiple * stepCst;
421 int64_t stepUnrolledCst = stepCst * unrollFactor;
422
423 // Create constant for 'upperBoundUnrolled' and set epilogue loop flag.
424 generateEpilogueLoop = upperBoundUnrolledCst < ubCst;
425 if (generateEpilogueLoop)
426 upperBoundUnrolled = arith::ConstantOp::create(
427 boundsBuilder, loc,
428 boundsBuilder.getIntegerAttr(forOp.getUpperBound().getType(),
429 upperBoundUnrolledCst));
430 else
431 upperBoundUnrolled = forOp.getUpperBound();
432
433 // Create constant for 'stepUnrolled'. When the main loop has zero
434 // iterations (tripCountEvenMultiple == 0), keep the original step.
435 // stepCst * unrollFactor may produce a value that, when truncated to the
436 // bound type's bitwidth during IntegerAttr construction, wraps to zero; a
437 // zero step causes constantTripCount to return nullopt instead of 0, which
438 // prevents the zero-trip main loop from being elided.
439 bool mainLoopHasNoIter = (tripCountEvenMultiple == 0);
440 bool stepUnchanged = (stepCst == stepUnrolledCst);
441 stepUnrolled =
442 (mainLoopHasNoIter || stepUnchanged)
443 ? step
444 : arith::ConstantOp::create(boundsBuilder, loc,
445 boundsBuilder.getIntegerAttr(
446 step.getType(), stepUnrolledCst));
447 } else {
448 // Dynamic loop bounds computation.
449 // TODO: Add dynamic asserts for negative lb/ub/step, or
450 // consider using ceilDiv from AffineApplyExpander.
451 auto lowerBound = forOp.getLowerBound();
452 auto upperBound = forOp.getUpperBound();
453 Value diff =
454 arith::SubIOp::create(boundsBuilder, loc, upperBound, lowerBound);
455 Value tripCount = ceilDivPositive(boundsBuilder, loc, diff, step);
456 Value unrollFactorCst = arith::ConstantOp::create(
457 boundsBuilder, loc,
458 boundsBuilder.getIntegerAttr(tripCount.getType(), unrollFactor));
459 Value tripCountRem =
460 arith::RemSIOp::create(boundsBuilder, loc, tripCount, unrollFactorCst);
461 // Compute tripCountEvenMultiple = tripCount - (tripCount % unrollFactor)
462 Value tripCountEvenMultiple =
463 arith::SubIOp::create(boundsBuilder, loc, tripCount, tripCountRem);
464 // Compute upperBoundUnrolled = lowerBound + tripCountEvenMultiple * step
465 upperBoundUnrolled = arith::AddIOp::create(
466 boundsBuilder, loc, lowerBound,
467 arith::MulIOp::create(boundsBuilder, loc, tripCountEvenMultiple, step));
468 // Scale 'step' by 'unrollFactor'.
469 stepUnrolled =
470 arith::MulIOp::create(boundsBuilder, loc, step, unrollFactorCst);
471 }
472
473 UnrolledLoopInfo resultLoops;
474
475 // Create epilogue clean up loop starting at 'upperBoundUnrolled'.
476 if (generateEpilogueLoop) {
477 OpBuilder epilogueBuilder(forOp->getContext());
478 epilogueBuilder.setInsertionPointAfter(forOp);
479 auto epilogueForOp = cast<scf::ForOp>(epilogueBuilder.clone(*forOp));
480 epilogueForOp.setLowerBound(upperBoundUnrolled);
481
482 // Update uses of loop results.
483 auto results = forOp.getResults();
484 auto epilogueResults = epilogueForOp.getResults();
485
486 for (auto e : llvm::zip(results, epilogueResults)) {
487 std::get<0>(e).replaceAllUsesWith(std::get<1>(e));
488 }
489 epilogueForOp->setOperands(epilogueForOp.getNumControlOperands(),
490 epilogueForOp.getInitArgs().size(), results);
491 if (!shouldPromoteIfSingleIteration ||
492 epilogueForOp.promoteIfSingleIteration(rewriter).failed())
493 resultLoops.epilogueLoopOp = epilogueForOp;
494 }
495
496 // Create unrolled loop.
497 forOp.setUpperBound(upperBoundUnrolled);
498 forOp.setStep(stepUnrolled);
499
500 auto iterArgs = ValueRange(forOp.getRegionIterArgs());
501 auto yieldedValues = forOp.getBody()->getTerminator()->getOperands();
502
504 forOp.getBody(), forOp.getInductionVar(), unrollFactor,
505 [&](unsigned i, Value iv, OpBuilder b) {
506 // iv' = iv + step * i;
507 auto stride = arith::MulIOp::create(
508 b, loc, step,
509 arith::ConstantOp::create(b, loc,
510 b.getIntegerAttr(iv.getType(), i)));
511 return arith::AddIOp::create(b, loc, iv, stride);
512 },
513 annotateFn, iterArgs, yieldedValues);
514 // Promote the loop body up if this has turned into a single iteration loop
515 // and `shouldPromoteIfSingleIteration` is true.
516 if (!shouldPromoteIfSingleIteration ||
517 forOp.promoteIfSingleIteration(rewriter).failed())
518 resultLoops.mainLoopOp = forOp;
519 return resultLoops;
520}
521
522/// Unrolls this loop completely.
523LogicalResult mlir::loopUnrollFull(scf::ForOp forOp) {
524 IRRewriter rewriter(forOp.getContext());
525 std::optional<APInt> mayBeConstantTripCount = forOp.getStaticTripCount();
526 if (!mayBeConstantTripCount.has_value())
527 return failure();
528 const APInt &tripCount = *mayBeConstantTripCount;
529 if (tripCount.isZero())
530 return success();
531 if (tripCount.isOne())
532 return forOp.promoteIfSingleIteration(rewriter);
533 return loopUnrollByFactor(forOp, tripCount.getZExtValue());
534}
535
536/// Check if bounds of all inner loops are defined outside of `forOp`
537/// and return false if not.
538static bool areInnerBoundsInvariant(scf::ForOp forOp) {
539 auto walkResult = forOp.walk([&](scf::ForOp innerForOp) {
540 if (!forOp.isDefinedOutsideOfLoop(innerForOp.getLowerBound()) ||
541 !forOp.isDefinedOutsideOfLoop(innerForOp.getUpperBound()) ||
542 !forOp.isDefinedOutsideOfLoop(innerForOp.getStep()))
543 return WalkResult::interrupt();
544
545 return WalkResult::advance();
546 });
547 return !walkResult.wasInterrupted();
548}
549
550/// Unrolls and jams this loop by the specified factor.
551LogicalResult mlir::loopUnrollJamByFactor(scf::ForOp forOp,
552 uint64_t unrollJamFactor) {
553 assert(unrollJamFactor > 0 && "unroll jam factor should be positive");
554
555 if (unrollJamFactor == 1)
556 return success();
557
558 // If any control operand of any inner loop of `forOp` is defined within
559 // `forOp`, no unroll jam.
560 if (!areInnerBoundsInvariant(forOp)) {
561 LDBG() << "failed to unroll and jam: inner bounds are not invariant";
562 return failure();
563 }
564
565 // Currently, for operations with results are not supported.
566 if (forOp->getNumResults() > 0) {
567 LDBG() << "failed to unroll and jam: unsupported loop with results";
568 return failure();
569 }
570
571 // Currently, only constant trip count that divided by the unroll factor is
572 // supported.
573 std::optional<APInt> tripCount = forOp.getStaticTripCount();
574 if (!tripCount.has_value()) {
575 // If the trip count is dynamic, do not unroll & jam.
576 LDBG() << "failed to unroll and jam: trip count could not be determined";
577 return failure();
578 }
579 uint64_t tripCountValue = tripCount->getZExtValue();
580 if (tripCountValue == 0)
581 return success();
582 if (unrollJamFactor > tripCountValue) {
583 LDBG() << "unroll and jam factor is greater than trip count, set factor to "
584 "trip "
585 "count";
586 unrollJamFactor = tripCountValue;
587 } else if (tripCountValue % unrollJamFactor != 0) {
588 LDBG() << "failed to unroll and jam: unsupported trip count that is not a "
589 "multiple of unroll jam factor";
590 return failure();
591 }
592
593 // Nothing in the loop body other than the terminator.
594 if (llvm::hasSingleElement(forOp.getBody()->getOperations()))
595 return success();
596
597 // Gather all sub-blocks to jam upon the loop being unrolled.
599 jbg.walk(forOp);
600 auto &subBlocks = jbg.subBlocks;
601
602 // Collect inner loops.
603 SmallVector<scf::ForOp> innerLoops;
604 forOp.walk([&](scf::ForOp innerForOp) { innerLoops.push_back(innerForOp); });
605
606 // `operandMaps[i - 1]` carries old->new operand mapping for the ith unrolled
607 // iteration. There are (`unrollJamFactor` - 1) iterations.
608 SmallVector<IRMapping> operandMaps(unrollJamFactor - 1);
609
610 // For any loop with iter_args, replace it with a new loop that has
611 // `unrollJamFactor` copies of its iterOperands, iter_args and yield
612 // operands.
613 SmallVector<scf::ForOp> newInnerLoops;
614 IRRewriter rewriter(forOp.getContext());
615 for (scf::ForOp oldForOp : innerLoops) {
616 SmallVector<Value> dupIterOperands, dupYieldOperands;
617 ValueRange oldIterOperands = oldForOp.getInits();
618 ValueRange oldIterArgs = oldForOp.getRegionIterArgs();
619 ValueRange oldYieldOperands =
620 cast<scf::YieldOp>(oldForOp.getBody()->getTerminator()).getOperands();
621 // Get additional iterOperands, iterArgs, and yield operands. We will
622 // fix iterOperands and yield operands after cloning of sub-blocks.
623 for (unsigned i = unrollJamFactor - 1; i >= 1; --i) {
624 dupIterOperands.append(oldIterOperands.begin(), oldIterOperands.end());
625 dupYieldOperands.append(oldYieldOperands.begin(), oldYieldOperands.end());
626 }
627 // Create a new loop with additional iterOperands, iter_args and yield
628 // operands. This new loop will take the loop body of the original loop.
629 bool forOpReplaced = oldForOp == forOp;
630 scf::ForOp newForOp =
631 cast<scf::ForOp>(*oldForOp.replaceWithAdditionalYields(
632 rewriter, dupIterOperands, /*replaceInitOperandUsesInLoop=*/false,
633 [&](OpBuilder &b, Location loc, ArrayRef<BlockArgument> newBbArgs) {
634 return dupYieldOperands;
635 }));
636 newInnerLoops.push_back(newForOp);
637 // `forOp` has been replaced with a new loop.
638 if (forOpReplaced)
639 forOp = newForOp;
640 // Update `operandMaps` for `newForOp` iterArgs and results.
641 ValueRange newIterArgs = newForOp.getRegionIterArgs();
642 unsigned oldNumIterArgs = oldIterArgs.size();
643 ValueRange newResults = newForOp.getResults();
644 unsigned oldNumResults = newResults.size() / unrollJamFactor;
645 assert(oldNumIterArgs == oldNumResults &&
646 "oldNumIterArgs must be the same as oldNumResults");
647 for (unsigned i = unrollJamFactor - 1; i >= 1; --i) {
648 for (unsigned j = 0; j < oldNumIterArgs; ++j) {
649 // `newForOp` has `unrollJamFactor` - 1 new sets of iterArgs and
650 // results. Update `operandMaps[i - 1]` to map old iterArgs and results
651 // to those in the `i`th new set.
652 operandMaps[i - 1].map(newIterArgs[j],
653 newIterArgs[i * oldNumIterArgs + j]);
654 operandMaps[i - 1].map(newResults[j],
655 newResults[i * oldNumResults + j]);
656 }
657 }
658 }
659
660 // Scale the step of loop being unroll-jammed by the unroll-jam factor.
661 rewriter.setInsertionPoint(forOp);
662 int64_t step = forOp.getConstantStep()->getSExtValue();
663 auto newStep = rewriter.createOrFold<arith::MulIOp>(
664 forOp.getLoc(), forOp.getStep(),
665 rewriter.createOrFold<arith::ConstantOp>(
666 forOp.getLoc(), rewriter.getIndexAttr(unrollJamFactor)));
667 forOp.setStep(newStep);
668 auto forOpIV = forOp.getInductionVar();
669
670 // Unroll and jam (appends unrollJamFactor - 1 additional copies).
671 for (unsigned i = unrollJamFactor - 1; i >= 1; --i) {
672 for (auto &subBlock : subBlocks) {
673 // Builder to insert unroll-jammed bodies. Insert right at the end of
674 // sub-block.
675 OpBuilder builder(subBlock.first->getBlock(), std::next(subBlock.second));
676
677 // If the induction variable is used, create a remapping to the value for
678 // this unrolled instance.
679 if (!forOpIV.use_empty()) {
680 // iv' = iv + i * step, i = 1 to unrollJamFactor-1.
681 auto ivTag = builder.createOrFold<arith::ConstantOp>(
682 forOp.getLoc(), builder.getIndexAttr(step * i));
683 auto ivUnroll =
684 builder.createOrFold<arith::AddIOp>(forOp.getLoc(), forOpIV, ivTag);
685 operandMaps[i - 1].map(forOpIV, ivUnroll);
686 }
687 // Clone the sub-block being unroll-jammed.
688 for (auto it = subBlock.first; it != std::next(subBlock.second); ++it)
689 builder.clone(*it, operandMaps[i - 1]);
690 }
691 // Fix iterOperands and yield op operands of newly created loops.
692 for (auto newForOp : newInnerLoops) {
693 unsigned oldNumIterOperands =
694 newForOp.getNumRegionIterArgs() / unrollJamFactor;
695 unsigned numControlOperands = newForOp.getNumControlOperands();
696 auto yieldOp = cast<scf::YieldOp>(newForOp.getBody()->getTerminator());
697 unsigned oldNumYieldOperands = yieldOp.getNumOperands() / unrollJamFactor;
698 assert(oldNumIterOperands == oldNumYieldOperands &&
699 "oldNumIterOperands must be the same as oldNumYieldOperands");
700 for (unsigned j = 0; j < oldNumIterOperands; ++j) {
701 // The `i`th duplication of an old iterOperand or yield op operand
702 // needs to be replaced with a mapped value from `operandMaps[i - 1]`
703 // if such mapped value exists.
704 newForOp.setOperand(numControlOperands + i * oldNumIterOperands + j,
705 operandMaps[i - 1].lookupOrDefault(
706 newForOp.getOperand(numControlOperands + j)));
707 yieldOp.setOperand(
708 i * oldNumYieldOperands + j,
709 operandMaps[i - 1].lookupOrDefault(yieldOp.getOperand(j)));
710 }
711 }
712 }
713
714 // Promote the loop body up if this has turned into a single iteration loop.
715 (void)forOp.promoteIfSingleIteration(rewriter);
716 return success();
717}
718
720 Location loc, OpFoldResult lb,
722 OpFoldResult step) {
723 Range normalizedLoopBounds;
724 normalizedLoopBounds.offset = rewriter.getIndexAttr(0);
725 normalizedLoopBounds.stride = rewriter.getIndexAttr(1);
726 AffineExpr s0, s1, s2;
727 bindSymbols(rewriter.getContext(), s0, s1, s2);
728 AffineExpr e = (s1 - s0).ceilDiv(s2);
729 normalizedLoopBounds.size =
730 affine::makeComposedFoldedAffineApply(rewriter, loc, e, {lb, ub, step});
731 return normalizedLoopBounds;
732}
733
736 OpFoldResult step) {
737 if (getType(lb).isIndex()) {
738 return emitNormalizedLoopBoundsForIndexType(rewriter, loc, lb, ub, step);
739 }
740 // For non-index types, generate `arith` instructions
741 // Check if the loop is already known to have a constant zero lower bound or
742 // a constant one step.
743 bool isZeroBased = false;
744 if (auto lbCst = getConstantIntValue(lb))
745 isZeroBased = lbCst.value() == 0;
746
747 bool isStepOne = false;
748 if (auto stepCst = getConstantIntValue(step))
749 isStepOne = stepCst.value() == 1;
750
751 Type rangeType = getType(lb);
752 assert(rangeType == getType(ub) && rangeType == getType(step) &&
753 "expected matching types");
754
755 // Compute the number of iterations the loop executes: ceildiv(ub - lb, step)
756 // assuming the step is strictly positive. Update the bounds and the step
757 // of the loop to go from 0 to the number of iterations, if necessary.
758 if (isZeroBased && isStepOne)
759 return {lb, ub, step};
760
761 OpFoldResult diff = ub;
762 if (!isZeroBased) {
763 diff = rewriter.createOrFold<arith::SubIOp>(
764 loc, getValueOrCreateConstantIntOp(rewriter, loc, ub),
765 getValueOrCreateConstantIntOp(rewriter, loc, lb));
766 }
767 OpFoldResult newUpperBound = diff;
768 if (!isStepOne) {
769 newUpperBound = rewriter.createOrFold<arith::CeilDivSIOp>(
770 loc, getValueOrCreateConstantIntOp(rewriter, loc, diff),
771 getValueOrCreateConstantIntOp(rewriter, loc, step));
772 }
773
774 OpFoldResult newLowerBound = rewriter.getZeroAttr(rangeType);
775 OpFoldResult newStep = rewriter.getOneAttr(rangeType);
776
777 return {newLowerBound, newUpperBound, newStep};
778}
779
781 Location loc,
782 Value normalizedIv,
783 OpFoldResult origLb,
784 OpFoldResult origStep) {
785 AffineExpr d0, s0, s1;
786 bindSymbols(rewriter.getContext(), s0, s1);
787 bindDims(rewriter.getContext(), d0);
788 AffineExpr e = d0 * s1 + s0;
790 rewriter, loc, e, ArrayRef<OpFoldResult>{normalizedIv, origLb, origStep});
791 Value denormalizedIvVal =
792 getValueOrCreateConstantIndexOp(rewriter, loc, denormalizedIv);
793 SmallPtrSet<Operation *, 1> preservedUses;
794 // If an `affine.apply` operation is generated for denormalization, the use
795 // of `origLb` in those ops must not be replaced. These arent not generated
796 // when `origLb == 0` and `origStep == 1`.
797 if (!isZeroInteger(origLb) || !isOneInteger(origStep)) {
798 if (Operation *preservedUse = denormalizedIvVal.getDefiningOp()) {
799 preservedUses.insert(preservedUse);
800 }
801 }
802 rewriter.replaceAllUsesExcept(normalizedIv, denormalizedIvVal, preservedUses);
803}
804
806 Value normalizedIv, OpFoldResult origLb,
807 OpFoldResult origStep) {
808 if (getType(origLb).isIndex()) {
809 return denormalizeInductionVariableForIndexType(rewriter, loc, normalizedIv,
810 origLb, origStep);
811 }
812 Value denormalizedIv;
814 bool isStepOne = isOneInteger(origStep);
815 bool isZeroBased = isZeroInteger(origLb);
816
817 Value scaled = normalizedIv;
818 if (!isStepOne) {
819 Value origStepValue =
820 getValueOrCreateConstantIntOp(rewriter, loc, origStep);
821 scaled = arith::MulIOp::create(rewriter, loc, normalizedIv, origStepValue);
822 preserve.insert(scaled.getDefiningOp());
823 }
824 denormalizedIv = scaled;
825 if (!isZeroBased) {
826 Value origLbValue = getValueOrCreateConstantIntOp(rewriter, loc, origLb);
827 denormalizedIv = arith::AddIOp::create(rewriter, loc, scaled, origLbValue);
828 preserve.insert(denormalizedIv.getDefiningOp());
829 }
830
831 rewriter.replaceAllUsesExcept(normalizedIv, denormalizedIv, preserve);
832}
833
835 ArrayRef<OpFoldResult> values) {
836 assert(!values.empty() && "unexecpted empty array");
837 AffineExpr s0, s1;
838 bindSymbols(rewriter.getContext(), s0, s1);
839 AffineExpr mul = s0 * s1;
840 OpFoldResult products = rewriter.getIndexAttr(1);
841 for (auto v : values) {
843 rewriter, loc, mul, ArrayRef<OpFoldResult>{products, v});
844 }
845 return products;
846}
847
848/// Helper function to multiply a sequence of values.
850 ArrayRef<Value> values) {
851 assert(!values.empty() && "unexpected empty list");
852 if (getType(values.front()).isIndex()) {
854 OpFoldResult product = getProductOfIndexes(rewriter, loc, ofrs);
855 return getValueOrCreateConstantIndexOp(rewriter, loc, product);
856 }
857 std::optional<Value> productOf;
858 for (auto v : values) {
859 auto vOne = getConstantIntValue(v);
860 if (vOne && vOne.value() == 1)
861 continue;
862 if (productOf)
863 productOf = arith::MulIOp::create(rewriter, loc, productOf.value(), v)
864 .getResult();
865 else
866 productOf = v;
867 }
868 if (!productOf) {
869 productOf = arith::ConstantOp::create(
870 rewriter, loc, rewriter.getOneAttr(getType(values.front())))
871 .getResult();
872 }
873 return productOf.value();
874}
875
876/// For each original loop, the value of the
877/// induction variable can be obtained by dividing the induction variable of
878/// the linearized loop by the total number of iterations of the loops nested
879/// in it modulo the number of iterations in this loop (remove the values
880/// related to the outer loops):
881/// iv_i = floordiv(iv_linear, product-of-loop-ranges-until-i) mod range_i.
882/// Compute these iteratively from the innermost loop by creating a "running
883/// quotient" of division by the range.
884static std::pair<SmallVector<Value>, SmallPtrSet<Operation *, 2>>
886 Value linearizedIv, ArrayRef<Value> ubs) {
887
888 if (linearizedIv.getType().isIndex()) {
889 Operation *delinearizedOp = affine::AffineDelinearizeIndexOp::create(
890 rewriter, loc, linearizedIv, ubs);
891 auto resultVals = llvm::map_to_vector(
892 delinearizedOp->getResults(), [](OpResult r) -> Value { return r; });
893 return {resultVals, SmallPtrSet<Operation *, 2>{delinearizedOp}};
894 }
895
896 SmallVector<Value> delinearizedIvs(ubs.size());
897 SmallPtrSet<Operation *, 2> preservedUsers;
898
899 llvm::BitVector isUbOne(ubs.size());
900 for (auto [index, ub] : llvm::enumerate(ubs)) {
901 auto ubCst = getConstantIntValue(ub);
902 if (ubCst && ubCst.value() == 1)
903 isUbOne.set(index);
904 }
905
906 // Prune the lead ubs that are all ones.
907 unsigned numLeadingOneUbs = 0;
908 for (auto [index, ub] : llvm::enumerate(ubs)) {
909 if (!isUbOne.test(index)) {
910 break;
911 }
912 delinearizedIvs[index] = arith::ConstantOp::create(
913 rewriter, loc, rewriter.getZeroAttr(ub.getType()));
914 numLeadingOneUbs++;
915 }
916
917 Value previous = linearizedIv;
918 for (unsigned i = numLeadingOneUbs, e = ubs.size(); i < e; ++i) {
919 unsigned idx = ubs.size() - (i - numLeadingOneUbs) - 1;
920 if (i != numLeadingOneUbs && !isUbOne.test(idx + 1)) {
921 previous = arith::DivSIOp::create(rewriter, loc, previous, ubs[idx + 1]);
922 preservedUsers.insert(previous.getDefiningOp());
923 }
924 Value iv = previous;
925 if (i != e - 1) {
926 if (!isUbOne.test(idx)) {
927 iv = arith::RemSIOp::create(rewriter, loc, previous, ubs[idx]);
928 preservedUsers.insert(iv.getDefiningOp());
929 } else {
930 iv = arith::ConstantOp::create(
931 rewriter, loc, rewriter.getZeroAttr(ubs[idx].getType()));
932 }
933 }
934 delinearizedIvs[idx] = iv;
935 }
936 return {delinearizedIvs, preservedUsers};
937}
938
939LogicalResult mlir::coalesceLoops(RewriterBase &rewriter,
941 if (loops.size() < 2)
942 return failure();
943
944 scf::ForOp innermost = loops.back();
945 scf::ForOp outermost = loops.front();
946
947 // Bail out if any loop has a known zero step, as normalization
948 // would result in a division by zero.
949 for (auto loop : loops) {
950 if (auto step = getConstantIntValue(loop.getStep())) {
951 if (step.value() == 0) {
952 return failure();
953 }
954 }
955 }
956 // 1. Make sure all loops iterate from 0 to upperBound with step 1. This
957 // allows the following code to assume upperBound is the number of iterations.
958 for (auto loop : loops) {
959 OpBuilder::InsertionGuard g(rewriter);
960 rewriter.setInsertionPoint(outermost);
961 Value lb = loop.getLowerBound();
962 Value ub = loop.getUpperBound();
963 Value step = loop.getStep();
964 auto newLoopRange =
965 emitNormalizedLoopBounds(rewriter, loop.getLoc(), lb, ub, step);
966
967 rewriter.modifyOpInPlace(loop, [&]() {
968 loop.setLowerBound(getValueOrCreateConstantIntOp(rewriter, loop.getLoc(),
969 newLoopRange.offset));
970 loop.setUpperBound(getValueOrCreateConstantIntOp(rewriter, loop.getLoc(),
971 newLoopRange.size));
972 loop.setStep(getValueOrCreateConstantIntOp(rewriter, loop.getLoc(),
973 newLoopRange.stride));
974 });
975 rewriter.setInsertionPointToStart(innermost.getBody());
976 denormalizeInductionVariable(rewriter, loop.getLoc(),
977 loop.getInductionVar(), lb, step);
978 }
979
980 // 2. Emit code computing the upper bound of the coalesced loop as product
981 // of the number of iterations of all loops.
982 OpBuilder::InsertionGuard g(rewriter);
983 rewriter.setInsertionPoint(outermost);
984 Location loc = outermost.getLoc();
985 SmallVector<Value> upperBounds = llvm::map_to_vector(
986 loops, [](auto loop) { return loop.getUpperBound(); });
987 Value upperBound = getProductOfIntsOrIndexes(rewriter, loc, upperBounds);
988 outermost.setUpperBound(upperBound);
989
990 // Insert delinearization at the start of the outermost loop body.
991 rewriter.setInsertionPointToStart(outermost.getBody());
992 auto [delinearizeIvs, preservedUsers] = delinearizeInductionVariable(
993 rewriter, loc, outermost.getInductionVar(), upperBounds);
994 rewriter.replaceAllUsesExcept(outermost.getInductionVar(), delinearizeIvs[0],
995 preservedUsers);
996
997 for (int i = loops.size() - 1; i > 0; --i) {
998 auto outerLoop = loops[i - 1];
999 auto innerLoop = loops[i];
1000
1001 Operation *innerTerminator = innerLoop.getBody()->getTerminator();
1002 auto yieldedVals = llvm::to_vector(innerTerminator->getOperands());
1003 assert(llvm::equal(outerLoop.getRegionIterArgs(), innerLoop.getInitArgs()));
1004 for (Value &yieldedVal : yieldedVals) {
1005 // The yielded value may be the induction variable of the inner loop,
1006 // which is about to be inlined and whose block argument is about to
1007 // be destroyed. Use its replacement value instead.
1008 if (yieldedVal == innerLoop.getInductionVar()) {
1009 yieldedVal = delinearizeIvs[i];
1010 continue;
1011 }
1012 // The yielded value may be an iteration argument of the inner loop
1013 // which is about to be inlined.
1014 auto iter = llvm::find(innerLoop.getRegionIterArgs(), yieldedVal);
1015 if (iter != innerLoop.getRegionIterArgs().end()) {
1016 unsigned iterArgIndex = iter - innerLoop.getRegionIterArgs().begin();
1017 // `outerLoop` iter args identical to the `innerLoop` init args.
1018 assert(iterArgIndex < innerLoop.getInitArgs().size());
1019 yieldedVal = innerLoop.getInitArgs()[iterArgIndex];
1020 }
1021 }
1022 rewriter.eraseOp(innerTerminator);
1023
1024 SmallVector<Value> innerBlockArgs;
1025 innerBlockArgs.push_back(delinearizeIvs[i]);
1026 llvm::append_range(innerBlockArgs, outerLoop.getRegionIterArgs());
1027 rewriter.inlineBlockBefore(innerLoop.getBody(), outerLoop.getBody(),
1028 Block::iterator(innerLoop), innerBlockArgs);
1029 rewriter.replaceOp(innerLoop, yieldedVals);
1030 }
1031 return success();
1032}
1033
1035 if (loops.empty()) {
1036 return failure();
1037 }
1038 IRRewriter rewriter(loops.front().getContext());
1039 return coalesceLoops(rewriter, loops);
1040}
1041
1042LogicalResult mlir::coalescePerfectlyNestedSCFForLoops(scf::ForOp op) {
1043 LogicalResult result(failure());
1045 getPerfectlyNestedLoops(loops, op);
1046
1047 // Look for a band of loops that can be coalesced, i.e. perfectly nested
1048 // loops with bounds defined above some loop.
1049
1050 // 1. For each loop, find above which parent loop its bounds operands are
1051 // defined.
1052 SmallVector<unsigned> operandsDefinedAbove(loops.size());
1053 for (unsigned i = 0, e = loops.size(); i < e; ++i) {
1054 operandsDefinedAbove[i] = i;
1055 for (unsigned j = 0; j < i; ++j) {
1056 SmallVector<Value> boundsOperands = {loops[i].getLowerBound(),
1057 loops[i].getUpperBound(),
1058 loops[i].getStep()};
1059 if (areValuesDefinedAbove(boundsOperands, loops[j].getRegion())) {
1060 operandsDefinedAbove[i] = j;
1061 break;
1062 }
1063 }
1064 }
1065
1066 // 2. For each inner loop check that the iter_args for the immediately outer
1067 // loop are the init for the immediately inner loop and that the yields of the
1068 // return of the inner loop is the yield for the immediately outer loop. Keep
1069 // track of where the chain starts from for each loop.
1070 SmallVector<unsigned> iterArgChainStart(loops.size());
1071 iterArgChainStart[0] = 0;
1072 for (unsigned i = 1, e = loops.size(); i < e; ++i) {
1073 // By default set the start of the chain to itself.
1074 iterArgChainStart[i] = i;
1075 auto outerloop = loops[i - 1];
1076 auto innerLoop = loops[i];
1077 if (outerloop.getNumRegionIterArgs() != innerLoop.getNumRegionIterArgs()) {
1078 continue;
1079 }
1080 if (!llvm::equal(outerloop.getRegionIterArgs(), innerLoop.getInitArgs())) {
1081 continue;
1082 }
1083 auto outerloopTerminator = outerloop.getBody()->getTerminator();
1084 if (!llvm::equal(outerloopTerminator->getOperands(),
1085 innerLoop.getResults())) {
1086 continue;
1087 }
1088 iterArgChainStart[i] = iterArgChainStart[i - 1];
1089 }
1090
1091 // 3. Identify bands of loops such that the operands of all of them are
1092 // defined above the first loop in the band. Traverse the nest bottom-up
1093 // so that modifications don't invalidate the inner loops.
1094 for (unsigned end = loops.size(); end > 0; --end) {
1095 unsigned start = 0;
1096 for (; start < end - 1; ++start) {
1097 auto maxPos =
1098 *std::max_element(std::next(operandsDefinedAbove.begin(), start),
1099 std::next(operandsDefinedAbove.begin(), end));
1100 if (maxPos > start)
1101 continue;
1102 if (iterArgChainStart[end - 1] > start)
1103 continue;
1104 auto band = llvm::MutableArrayRef(loops.data() + start, end - start);
1105 if (succeeded(coalesceLoops(band)))
1106 result = success();
1107 break;
1108 }
1109 // If a band was found and transformed, keep looking at the loops above
1110 // the outermost transformed loop.
1111 if (start != end - 1)
1112 end = start + 1;
1113 }
1114 return result;
1115}
1116
1118 RewriterBase &rewriter, scf::ParallelOp loops,
1119 ArrayRef<std::vector<unsigned>> combinedDimensions) {
1120 OpBuilder::InsertionGuard g(rewriter);
1121 rewriter.setInsertionPoint(loops);
1122 Location loc = loops.getLoc();
1123
1124 // Presort combined dimensions.
1125 auto sortedDimensions = llvm::to_vector<3>(combinedDimensions);
1126 for (auto &dims : sortedDimensions)
1127 llvm::sort(dims);
1128
1129 // Normalize ParallelOp's iteration pattern.
1130 SmallVector<Value, 3> normalizedUpperBounds;
1131 for (unsigned i = 0, e = loops.getNumLoops(); i < e; ++i) {
1132 OpBuilder::InsertionGuard g2(rewriter);
1133 rewriter.setInsertionPoint(loops);
1134 Value lb = loops.getLowerBound()[i];
1135 Value ub = loops.getUpperBound()[i];
1136 Value step = loops.getStep()[i];
1137 auto newLoopRange = emitNormalizedLoopBounds(rewriter, loc, lb, ub, step);
1138 normalizedUpperBounds.push_back(getValueOrCreateConstantIntOp(
1139 rewriter, loops.getLoc(), newLoopRange.size));
1140
1141 rewriter.setInsertionPointToStart(loops.getBody());
1142 denormalizeInductionVariable(rewriter, loc, loops.getInductionVars()[i], lb,
1143 step);
1144 }
1145
1146 // Combine iteration spaces.
1147 SmallVector<Value, 3> lowerBounds, upperBounds, steps;
1148 auto cst0 = arith::ConstantIndexOp::create(rewriter, loc, 0);
1149 auto cst1 = arith::ConstantIndexOp::create(rewriter, loc, 1);
1150 for (auto &sortedDimension : sortedDimensions) {
1151 Value newUpperBound = arith::ConstantIndexOp::create(rewriter, loc, 1);
1152 for (auto idx : sortedDimension) {
1153 newUpperBound = arith::MulIOp::create(rewriter, loc, newUpperBound,
1154 normalizedUpperBounds[idx]);
1155 }
1156 lowerBounds.push_back(cst0);
1157 steps.push_back(cst1);
1158 upperBounds.push_back(newUpperBound);
1159 }
1160
1161 // Create new ParallelLoop with conversions to the original induction values.
1162 // The loop below uses divisions to get the relevant range of values in the
1163 // new induction value that represent each range of the original induction
1164 // value. The remainders then determine based on that range, which iteration
1165 // of the original induction value this represents. This is a normalized value
1166 // that is un-normalized already by the previous logic.
1167 auto newPloop = scf::ParallelOp::create(
1168 rewriter, loc, lowerBounds, upperBounds, steps,
1169 [&](OpBuilder &insideBuilder, Location, ValueRange ploopIVs) {
1170 for (unsigned i = 0, e = combinedDimensions.size(); i < e; ++i) {
1171 Value previous = ploopIVs[i];
1172 unsigned numberCombinedDimensions = combinedDimensions[i].size();
1173 // Iterate over all except the last induction value.
1174 for (unsigned j = numberCombinedDimensions - 1; j > 0; --j) {
1175 unsigned idx = combinedDimensions[i][j];
1176
1177 // Determine the current induction value's current loop iteration
1178 Value iv = arith::RemSIOp::create(insideBuilder, loc, previous,
1179 normalizedUpperBounds[idx]);
1180 replaceAllUsesInRegionWith(loops.getBody()->getArgument(idx), iv,
1181 loops.getRegion());
1182
1183 // Remove the effect of the current induction value to prepare for
1184 // the next value.
1185 previous = arith::DivSIOp::create(insideBuilder, loc, previous,
1186 normalizedUpperBounds[idx]);
1187 }
1188
1189 // The final induction value is just the remaining value.
1190 unsigned idx = combinedDimensions[i][0];
1191 replaceAllUsesInRegionWith(loops.getBody()->getArgument(idx),
1192 previous, loops.getRegion());
1193 }
1194 });
1195
1196 // Replace the old loop with the new loop.
1197 loops.getBody()->back().erase();
1198 newPloop.getBody()->getOperations().splice(
1199 Block::iterator(newPloop.getBody()->back()),
1200 loops.getBody()->getOperations());
1201 loops.erase();
1202}
1203
1204// Hoist the ops within `outer` that appear before `inner`.
1205// Such ops include the ops that have been introduced by parametric tiling.
1206// Ops that come from triangular loops (i.e. that belong to the program slice
1207// rooted at `outer`) and ops that have side effects cannot be hoisted.
1208// Return failure when any op fails to hoist.
1209static LogicalResult hoistOpsBetween(scf::ForOp outer, scf::ForOp inner) {
1210 SetVector<Operation *> forwardSlice;
1212 options.filter = [&inner](Operation *op) {
1213 return op != inner.getOperation();
1214 };
1215 getForwardSlice(outer.getInductionVar(), &forwardSlice, options);
1216 LogicalResult status = success();
1218 for (auto &op : outer.getBody()->without_terminator()) {
1219 // Stop when encountering the inner loop.
1220 if (&op == inner.getOperation())
1221 break;
1222 // Skip over non-hoistable ops.
1223 if (forwardSlice.count(&op) > 0) {
1224 status = failure();
1225 continue;
1226 }
1227 // Skip intermediate scf::ForOp, these are not considered a failure.
1228 if (isa<scf::ForOp>(op))
1229 continue;
1230 // Skip other ops with regions.
1231 if (op.getNumRegions() > 0) {
1232 status = failure();
1233 continue;
1234 }
1235 // Skip if op has side effects.
1236 // TODO: loads to immutable memory regions are ok.
1237 if (!isMemoryEffectFree(&op)) {
1238 status = failure();
1239 continue;
1240 }
1241 toHoist.push_back(&op);
1242 }
1243 auto *outerForOp = outer.getOperation();
1244 for (auto *op : toHoist)
1245 op->moveBefore(outerForOp);
1246 return status;
1247}
1248
1249// Traverse the interTile and intraTile loops and try to hoist ops such that
1250// bands of perfectly nested loops are isolated.
1251// Return failure if either perfect interTile or perfect intraTile bands cannot
1252// be formed.
1253static LogicalResult tryIsolateBands(const TileLoops &tileLoops) {
1254 LogicalResult status = success();
1255 const Loops &interTile = tileLoops.first;
1256 const Loops &intraTile = tileLoops.second;
1257 auto size = interTile.size();
1258 assert(size == intraTile.size());
1259 if (size <= 1)
1260 return success();
1261 for (unsigned s = 1; s < size; ++s)
1262 status = succeeded(status) ? hoistOpsBetween(intraTile[0], intraTile[s])
1263 : failure();
1264 for (unsigned s = 1; s < size; ++s)
1265 status = succeeded(status) ? hoistOpsBetween(interTile[0], interTile[s])
1266 : failure();
1267 return status;
1268}
1269
1270/// Collect perfectly nested loops starting from `rootForOps`. Loops are
1271/// perfectly nested if each loop is the first and only non-terminator operation
1272/// in the parent loop. Collect at most `maxLoops` loops and append them to
1273/// `forOps`.
1274template <typename T>
1276 SmallVectorImpl<T> &forOps, T rootForOp,
1277 unsigned maxLoops = std::numeric_limits<unsigned>::max()) {
1278 for (unsigned i = 0; i < maxLoops; ++i) {
1279 forOps.push_back(rootForOp);
1280 Block &body = rootForOp.getRegion().front();
1281 if (body.begin() != std::prev(body.end(), 2))
1282 return;
1283
1284 rootForOp = dyn_cast<T>(&body.front());
1285 if (!rootForOp)
1286 return;
1287 }
1288}
1289
1290static Loops stripmineSink(scf::ForOp forOp, Value factor,
1291 ArrayRef<scf::ForOp> targets) {
1292 assert(!forOp.getUnsignedCmp() && "unsigned loops are not supported");
1293 auto originalStep = forOp.getStep();
1294 auto iv = forOp.getInductionVar();
1295
1296 OpBuilder b(forOp);
1297 forOp.setStep(arith::MulIOp::create(b, forOp.getLoc(), originalStep, factor));
1298
1299 Loops innerLoops;
1300 for (auto t : targets) {
1301 assert(!t.getUnsignedCmp() && "unsigned loops are not supported");
1302
1303 // Save information for splicing ops out of t when done
1304 auto begin = t.getBody()->begin();
1305 auto nOps = t.getBody()->getOperations().size();
1306
1307 // Insert newForOp before the terminator of `t`.
1308 auto b = OpBuilder::atBlockTerminator((t.getBody()));
1309 Value stepped = arith::AddIOp::create(b, t.getLoc(), iv, forOp.getStep());
1310 Value ub =
1311 arith::MinSIOp::create(b, t.getLoc(), forOp.getUpperBound(), stepped);
1312
1313 // Splice [begin, begin + nOps - 1) into `newForOp` and replace uses.
1314 auto newForOp = scf::ForOp::create(b, t.getLoc(), iv, ub, originalStep);
1315 newForOp.getBody()->getOperations().splice(
1316 newForOp.getBody()->getOperations().begin(),
1317 t.getBody()->getOperations(), begin, std::next(begin, nOps - 1));
1318 replaceAllUsesInRegionWith(iv, newForOp.getInductionVar(),
1319 newForOp.getRegion());
1320
1321 innerLoops.push_back(newForOp);
1322 }
1323
1324 return innerLoops;
1325}
1326
1327// Stripmines a `forOp` by `factor` and sinks it under a single `target`.
1328// Returns the new for operation, nested immediately under `target`.
1329template <typename SizeType>
1330static scf::ForOp stripmineSink(scf::ForOp forOp, SizeType factor,
1331 scf::ForOp target) {
1332 // TODO: Use cheap structural assertions that targets are nested under
1333 // forOp and that targets are not nested under each other when DominanceInfo
1334 // exposes the capability. It seems overkill to construct a whole function
1335 // dominance tree at this point.
1336 auto res = stripmineSink(forOp, factor, ArrayRef<scf::ForOp>(target));
1337 assert(res.size() == 1 && "Expected 1 inner forOp");
1338 return res[0];
1339}
1340
1342 ArrayRef<Value> sizes,
1343 ArrayRef<scf::ForOp> targets) {
1345 SmallVector<scf::ForOp, 8> currentTargets(targets);
1346 for (auto it : llvm::zip(forOps, sizes)) {
1347 auto step = stripmineSink(std::get<0>(it), std::get<1>(it), currentTargets);
1348 res.push_back(step);
1349 currentTargets = step;
1350 }
1351 return res;
1352}
1353
1355 scf::ForOp target) {
1357 for (auto loops : tile(forOps, sizes, ArrayRef<scf::ForOp>(target)))
1358 res.push_back(llvm::getSingleElement(loops));
1359 return res;
1360}
1361
1363 // Collect perfectly nested loops. If more size values provided than nested
1364 // loops available, truncate `sizes`.
1366 forOps.reserve(sizes.size());
1367 getPerfectlyNestedLoopsImpl(forOps, rootForOp, sizes.size());
1368 if (forOps.size() < sizes.size())
1369 sizes = sizes.take_front(forOps.size());
1370
1371 return ::tile(forOps, sizes, forOps.back());
1372}
1373
1375 scf::ForOp root) {
1376 getPerfectlyNestedLoopsImpl(nestedLoops, root);
1377}
1378
1380 ArrayRef<int64_t> sizes) {
1381 // Collect perfectly nested loops. If more size values provided than nested
1382 // loops available, truncate `sizes`.
1384 forOps.reserve(sizes.size());
1385 getPerfectlyNestedLoopsImpl(forOps, rootForOp, sizes.size());
1386 if (forOps.size() < sizes.size())
1387 sizes = sizes.take_front(forOps.size());
1388
1389 // The strip-mining transformation splices loop bodies into a new inner loop
1390 // without threading iter_args. If any of the collected loops carries
1391 // iter_args, the splice would produce invalid IR (yielded values from the
1392 // inner scope used in the outer terminator). Skip the transformation in
1393 // that case.
1394 if (llvm::any_of(forOps,
1395 [](scf::ForOp op) { return !op.getInitArgs().empty(); }))
1396 return {};
1397
1398 // Compute the tile sizes such that i-th outer loop executes size[i]
1399 // iterations. Given that the loop current executes
1400 // numIterations = ceildiv((upperBound - lowerBound), step)
1401 // iterations, we need to tile with size ceildiv(numIterations, size[i]).
1402 SmallVector<Value, 4> tileSizes;
1403 tileSizes.reserve(sizes.size());
1404 for (unsigned i = 0, e = sizes.size(); i < e; ++i) {
1405 assert(sizes[i] > 0 && "expected strictly positive size for strip-mining");
1406
1407 auto forOp = forOps[i];
1408 OpBuilder builder(forOp);
1409 auto loc = forOp.getLoc();
1410 Value diff = arith::SubIOp::create(builder, loc, forOp.getUpperBound(),
1411 forOp.getLowerBound());
1412 Value numIterations = ceilDivPositive(builder, loc, diff, forOp.getStep());
1413 Value iterationsPerBlock =
1414 ceilDivPositive(builder, loc, numIterations, sizes[i]);
1415 tileSizes.push_back(iterationsPerBlock);
1416 }
1417
1418 // Call parametric tiling with the given sizes.
1419 auto intraTile = tile(forOps, tileSizes, forOps.back());
1420 TileLoops tileLoops = std::make_pair(forOps, intraTile);
1421
1422 // TODO: for now we just ignore the result of band isolation.
1423 // In the future, mapping decisions may be impacted by the ability to
1424 // isolate perfectly nested bands.
1425 (void)tryIsolateBands(tileLoops);
1426
1427 return tileLoops;
1428}
1429
1431 scf::ForallOp source,
1432 RewriterBase &rewriter) {
1433 unsigned numTargetOuts = target.getNumResults();
1434 unsigned numSourceOuts = source.getNumResults();
1435
1436 // Create fused shared_outs.
1437 SmallVector<Value> fusedOuts;
1438 llvm::append_range(fusedOuts, target.getOutputs());
1439 llvm::append_range(fusedOuts, source.getOutputs());
1440
1441 // Create a new scf.forall op after the source loop.
1442 rewriter.setInsertionPointAfter(source);
1443 scf::ForallOp fusedLoop = scf::ForallOp::create(
1444 rewriter, source.getLoc(), source.getMixedLowerBound(),
1445 source.getMixedUpperBound(), source.getMixedStep(), fusedOuts,
1446 source.getMapping());
1447
1448 // Map control operands.
1449 IRMapping mapping;
1450 mapping.map(target.getInductionVars(), fusedLoop.getInductionVars());
1451 mapping.map(source.getInductionVars(), fusedLoop.getInductionVars());
1452
1453 // Map shared outs.
1454 mapping.map(target.getRegionIterArgs(),
1455 fusedLoop.getRegionIterArgs().take_front(numTargetOuts));
1456 mapping.map(source.getRegionIterArgs(),
1457 fusedLoop.getRegionIterArgs().take_back(numSourceOuts));
1458
1459 // Append everything except the terminator into the fused operation.
1460 rewriter.setInsertionPointToStart(fusedLoop.getBody());
1461 for (Operation &op : target.getBody()->without_terminator())
1462 rewriter.clone(op, mapping);
1463 for (Operation &op : source.getBody()->without_terminator())
1464 rewriter.clone(op, mapping);
1465
1466 // Fuse the old terminator in_parallel ops into the new one.
1467 scf::InParallelOp targetTerm = target.getTerminator();
1468 scf::InParallelOp sourceTerm = source.getTerminator();
1469 scf::InParallelOp fusedTerm = fusedLoop.getTerminator();
1470 rewriter.setInsertionPointToStart(fusedTerm.getBody());
1471 for (Operation &op : targetTerm.getYieldingOps())
1472 rewriter.clone(op, mapping);
1473 for (Operation &op : sourceTerm.getYieldingOps())
1474 rewriter.clone(op, mapping);
1475
1476 // Replace old loops by substituting their uses by results of the fused loop.
1477 rewriter.replaceOp(target, fusedLoop.getResults().take_front(numTargetOuts));
1478 rewriter.replaceOp(source, fusedLoop.getResults().take_back(numSourceOuts));
1479
1480 return fusedLoop;
1481}
1482
1484 scf::ForOp source,
1485 RewriterBase &rewriter) {
1486 assert(source.getUnsignedCmp() == target.getUnsignedCmp() &&
1487 "incompatible signedness");
1488 unsigned numTargetOuts = target.getNumResults();
1489 unsigned numSourceOuts = source.getNumResults();
1490
1491 // Create fused init_args, with target's init_args before source's init_args.
1492 SmallVector<Value> fusedInitArgs;
1493 llvm::append_range(fusedInitArgs, target.getInitArgs());
1494 llvm::append_range(fusedInitArgs, source.getInitArgs());
1495
1496 // Create a new scf.for op after the source loop (with scf.yield terminator
1497 // (without arguments) only in case its init_args is empty).
1498 rewriter.setInsertionPointAfter(source);
1499 scf::ForOp fusedLoop = scf::ForOp::create(
1500 rewriter, source.getLoc(), source.getLowerBound(), source.getUpperBound(),
1501 source.getStep(), fusedInitArgs, /*bodyBuilder=*/nullptr,
1502 source.getUnsignedCmp());
1503
1504 // Map original induction variables and operands to those of the fused loop.
1505 IRMapping mapping;
1506 mapping.map(target.getInductionVar(), fusedLoop.getInductionVar());
1507 mapping.map(target.getRegionIterArgs(),
1508 fusedLoop.getRegionIterArgs().take_front(numTargetOuts));
1509 mapping.map(source.getInductionVar(), fusedLoop.getInductionVar());
1510 mapping.map(source.getRegionIterArgs(),
1511 fusedLoop.getRegionIterArgs().take_back(numSourceOuts));
1512
1513 // Merge target's body into the new (fused) for loop and then source's body.
1514 rewriter.setInsertionPointToStart(fusedLoop.getBody());
1515 for (Operation &op : target.getBody()->without_terminator())
1516 rewriter.clone(op, mapping);
1517 for (Operation &op : source.getBody()->without_terminator())
1518 rewriter.clone(op, mapping);
1519
1520 // Build fused yield results by appropriately mapping original yield operands.
1521 SmallVector<Value> yieldResults;
1522 for (Value operand : target.getBody()->getTerminator()->getOperands())
1523 yieldResults.push_back(mapping.lookupOrDefault(operand));
1524 for (Value operand : source.getBody()->getTerminator()->getOperands())
1525 yieldResults.push_back(mapping.lookupOrDefault(operand));
1526 if (!yieldResults.empty())
1527 scf::YieldOp::create(rewriter, source.getLoc(), yieldResults);
1528
1529 // Replace old loops by substituting their uses by results of the fused loop.
1530 rewriter.replaceOp(target, fusedLoop.getResults().take_front(numTargetOuts));
1531 rewriter.replaceOp(source, fusedLoop.getResults().take_back(numSourceOuts));
1532
1533 return fusedLoop;
1534}
1535
1536FailureOr<scf::ForallOp> mlir::normalizeForallOp(RewriterBase &rewriter,
1537 scf::ForallOp forallOp) {
1538 SmallVector<OpFoldResult> lbs = forallOp.getMixedLowerBound();
1539 SmallVector<OpFoldResult> ubs = forallOp.getMixedUpperBound();
1540 SmallVector<OpFoldResult> steps = forallOp.getMixedStep();
1541
1542 if (forallOp.isNormalized())
1543 return forallOp;
1544
1545 OpBuilder::InsertionGuard g(rewriter);
1546 auto loc = forallOp.getLoc();
1547 rewriter.setInsertionPoint(forallOp);
1549 for (auto [lb, ub, step] : llvm::zip_equal(lbs, ubs, steps)) {
1550 Range normalizedLoopParams =
1551 emitNormalizedLoopBounds(rewriter, loc, lb, ub, step);
1552 newUbs.push_back(normalizedLoopParams.size);
1553 }
1554 (void)foldDynamicIndexList(newUbs);
1555
1556 // Use the normalized builder since the lower bounds are always 0 and the
1557 // steps are always 1.
1558 auto normalizedForallOp = scf::ForallOp::create(
1559 rewriter, loc, newUbs, forallOp.getOutputs(), forallOp.getMapping(),
1560 [](OpBuilder &, Location, ValueRange) {});
1561
1562 rewriter.inlineRegionBefore(forallOp.getBodyRegion(),
1563 normalizedForallOp.getBodyRegion(),
1564 normalizedForallOp.getBodyRegion().begin());
1565 // Remove the original empty block in the new loop.
1566 rewriter.eraseBlock(&normalizedForallOp.getBodyRegion().back());
1567
1568 rewriter.setInsertionPointToStart(normalizedForallOp.getBody());
1569 // Update the users of the original loop variables.
1570 for (auto [idx, iv] :
1571 llvm::enumerate(normalizedForallOp.getInductionVars())) {
1572 auto origLb = getValueOrCreateConstantIndexOp(rewriter, loc, lbs[idx]);
1573 auto origStep = getValueOrCreateConstantIndexOp(rewriter, loc, steps[idx]);
1574 denormalizeInductionVariable(rewriter, loc, iv, origLb, origStep);
1575 }
1576
1577 rewriter.replaceOp(forallOp, normalizedForallOp);
1578 return normalizedForallOp;
1579}
1580
1583 assert(!loops.empty() && "unexpected empty loop nest");
1584 if (loops.size() == 1)
1585 return isa_and_nonnull<scf::ForOp>(loops.front().getOperation());
1586 for (auto [outerLoop, innerLoop] :
1587 llvm::zip_equal(loops.drop_back(), loops.drop_front())) {
1588 auto outerFor = dyn_cast_or_null<scf::ForOp>(outerLoop.getOperation());
1589 auto innerFor = dyn_cast_or_null<scf::ForOp>(innerLoop.getOperation());
1590 if (!outerFor || !innerFor)
1591 return false;
1592 auto outerBBArgs = outerFor.getRegionIterArgs();
1593 auto innerIterArgs = innerFor.getInitArgs();
1594 if (outerBBArgs.size() != innerIterArgs.size())
1595 return false;
1596
1597 for (auto [outerBBArg, innerIterArg] :
1598 llvm::zip_equal(outerBBArgs, innerIterArgs)) {
1599 if (!llvm::hasSingleElement(outerBBArg.getUses()) ||
1600 innerIterArg != outerBBArg)
1601 return false;
1602 }
1603
1604 ValueRange outerYields =
1605 cast<scf::YieldOp>(outerFor.getBody()->getTerminator())->getOperands();
1606 ValueRange innerResults = innerFor.getResults();
1607 if (outerYields.size() != innerResults.size())
1608 return false;
1609 for (auto [outerYield, innerResult] :
1610 llvm::zip_equal(outerYields, innerResults)) {
1611 if (!llvm::hasSingleElement(innerResult.getUses()) ||
1612 outerYield != innerResult)
1613 return false;
1614 }
1615 }
1616 return true;
1617}
1618
1620mlir::getConstLoopBounds(mlir::LoopLikeOpInterface loopOp) {
1621 std::optional<SmallVector<OpFoldResult>> loBnds = loopOp.getLoopLowerBounds();
1622 std::optional<SmallVector<OpFoldResult>> upBnds = loopOp.getLoopUpperBounds();
1623 std::optional<SmallVector<OpFoldResult>> steps = loopOp.getLoopSteps();
1624 if (!loBnds || !upBnds || !steps)
1625 return {};
1627 for (auto [lb, ub, step] : llvm::zip(*loBnds, *upBnds, *steps)) {
1628 auto lbCst = getConstantIntValue(lb);
1629 auto ubCst = getConstantIntValue(ub);
1630 auto stepCst = getConstantIntValue(step);
1631 if (!lbCst || !ubCst || !stepCst)
1632 return {};
1633 loopRanges.emplace_back(*lbCst, *ubCst, *stepCst);
1634 }
1635 return loopRanges;
1636}
1637
1639mlir::getConstLoopTripCounts(mlir::LoopLikeOpInterface loopOp) {
1640 std::optional<SmallVector<OpFoldResult>> loBnds = loopOp.getLoopLowerBounds();
1641 std::optional<SmallVector<OpFoldResult>> upBnds = loopOp.getLoopUpperBounds();
1642 std::optional<SmallVector<OpFoldResult>> steps = loopOp.getLoopSteps();
1643 if (!loBnds || !upBnds || !steps)
1644 return {};
1646 for (auto [lb, ub, step] : llvm::zip(*loBnds, *upBnds, *steps)) {
1647 // TODO(#178506): Signedness is not handled correctly here.
1648 std::optional<llvm::APInt> numIter = constantTripCount(
1649 lb, ub, step, /*isSigned=*/true, scf::computeUbMinusLb);
1650 if (!numIter)
1651 return {};
1652 tripCounts.push_back(*numIter);
1653 }
1654 return tripCounts;
1655}
1656
1657FailureOr<scf::ParallelOp> mlir::parallelLoopUnrollByFactors(
1658 scf::ParallelOp op, ArrayRef<uint64_t> unrollFactors,
1659 RewriterBase &rewriter,
1660 function_ref<void(unsigned, Operation *, OpBuilder)> annotateFn,
1661 IRMapping *clonedToSrcOpsMap) {
1662 const unsigned numLoops = op.getNumLoops();
1663 assert(llvm::none_of(unrollFactors, [](uint64_t f) { return f == 0; }) &&
1664 "Expected positive unroll factors");
1665 assert((!unrollFactors.empty() && (unrollFactors.size() <= numLoops)) &&
1666 "Expected non-empty unroll factors of size <= to the number of loops");
1667
1668 // Bail out if no valid unroll factors were provided
1669 if (llvm::all_of(unrollFactors, [](uint64_t f) { return f == 1; }))
1670 return rewriter.notifyMatchFailure(
1671 op, "Unrolling not applied if all factors are 1");
1672
1673 // Return if the loop body is empty.
1674 if (llvm::hasSingleElement(op.getBody()->getOperations()))
1675 return rewriter.notifyMatchFailure(op, "Cannot unroll an empty loop body");
1676
1677 // If the provided unroll factors do not cover all the loop dims, they are
1678 // applied to the inner loop dimensions.
1679 const unsigned firstLoopDimIdx = numLoops - unrollFactors.size();
1680
1681 // Make sure that the unroll factors divide the iteration space evenly
1682 // TODO: Support unrolling loops with dynamic iteration spaces.
1684 if (tripCounts.empty())
1685 return rewriter.notifyMatchFailure(
1686 op, "Failed to compute constant trip counts for the loop. Note that "
1687 "dynamic loop sizes are not supported.");
1688
1689 for (unsigned dimIdx = firstLoopDimIdx; dimIdx < numLoops; dimIdx++) {
1690 const uint64_t unrollFactor = unrollFactors[dimIdx - firstLoopDimIdx];
1691 if (tripCounts[dimIdx].urem(unrollFactor) != 0)
1692 return rewriter.notifyMatchFailure(
1693 op, "Unroll factors don't divide the iteration space evenly");
1694 }
1695
1696 std::optional<SmallVector<OpFoldResult>> maybeFoldSteps = op.getLoopSteps();
1697 if (!maybeFoldSteps)
1698 return rewriter.notifyMatchFailure(op, "Failed to retrieve loop steps");
1700 for (auto step : *maybeFoldSteps)
1701 steps.push_back(static_cast<size_t>(*getConstantIntValue(step)));
1702
1703 for (unsigned dimIdx = firstLoopDimIdx; dimIdx < numLoops; dimIdx++) {
1704 const uint64_t unrollFactor = unrollFactors[dimIdx - firstLoopDimIdx];
1705 if (unrollFactor == 1)
1706 continue;
1707 const size_t origStep = steps[dimIdx];
1708 const int64_t newStep = origStep * unrollFactor;
1709 IRMapping clonedToSrcOpsMap;
1710
1711 ValueRange iterArgs = ValueRange(op.getRegionIterArgs());
1712 auto yieldedValues = op.getBody()->getTerminator()->getOperands();
1713
1715 op.getBody(), op.getInductionVars()[dimIdx], unrollFactor,
1716 [&](unsigned i, Value iv, OpBuilder b) {
1717 // iv' = iv + step * i;
1718 const AffineExpr expr = b.getAffineDimExpr(0) + (origStep * i);
1719 const auto map =
1720 b.getDimIdentityMap().dropResult(0).insertResult(expr, 0);
1721 return affine::AffineApplyOp::create(b, iv.getLoc(), map,
1722 ValueRange{iv});
1723 },
1724 /*annotateFn*/ annotateFn, iterArgs, yieldedValues, &clonedToSrcOpsMap);
1725
1726 // Update loop step
1727 auto prevInsertPoint = rewriter.saveInsertionPoint();
1728 rewriter.setInsertionPoint(op);
1729 op.getStepMutable()[dimIdx].assign(
1730 arith::ConstantIndexOp::create(rewriter, op.getLoc(), newStep));
1731 rewriter.restoreInsertionPoint(prevInsertPoint);
1732 }
1733 return op;
1734}
return success()
static OpFoldResult getProductOfIndexes(RewriterBase &rewriter, Location loc, ArrayRef< OpFoldResult > values)
Definition Utils.cpp:834
static LogicalResult tryIsolateBands(const TileLoops &tileLoops)
Definition Utils.cpp:1253
static void getPerfectlyNestedLoopsImpl(SmallVectorImpl< T > &forOps, T rootForOp, unsigned maxLoops=std::numeric_limits< unsigned >::max())
Collect perfectly nested loops starting from rootForOps.
Definition Utils.cpp:1275
static LogicalResult hoistOpsBetween(scf::ForOp outer, scf::ForOp inner)
Definition Utils.cpp:1209
static Range emitNormalizedLoopBoundsForIndexType(RewriterBase &rewriter, Location loc, OpFoldResult lb, OpFoldResult ub, OpFoldResult step)
Definition Utils.cpp:719
static Loops stripmineSink(scf::ForOp forOp, Value factor, ArrayRef< scf::ForOp > targets)
Definition Utils.cpp:1290
static Value ceilDivPositive(OpBuilder &builder, Location loc, Value dividend, int64_t divisor)
Definition Utils.cpp:266
static Value getProductOfIntsOrIndexes(RewriterBase &rewriter, Location loc, ArrayRef< Value > values)
Helper function to multiply a sequence of values.
Definition Utils.cpp:849
static std::pair< SmallVector< Value >, SmallPtrSet< Operation *, 2 > > delinearizeInductionVariable(RewriterBase &rewriter, Location loc, Value linearizedIv, ArrayRef< Value > ubs)
For each original loop, the value of the induction variable can be obtained by dividing the induction...
Definition Utils.cpp:885
static void denormalizeInductionVariableForIndexType(RewriterBase &rewriter, Location loc, Value normalizedIv, OpFoldResult origLb, OpFoldResult origStep)
Definition Utils.cpp:780
static bool areInnerBoundsInvariant(scf::ForOp forOp)
Check if bounds of all inner loops are defined outside of forOp and return false if not.
Definition Utils.cpp:538
static int64_t product(ArrayRef< int64_t > vals)
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
static llvm::ManagedStatic< PassManagerOptions > options
#define mul(a, b)
Base type for affine expression.
Definition AffineExpr.h:68
This class represents an argument of a Block.
Definition Value.h:306
Block represents an ordered list of Operations.
Definition Block.h:33
OpListType::iterator iterator
Definition Block.h:164
unsigned getNumArguments()
Definition Block.h:152
Operation & front()
Definition Block.h:177
Operation * getTerminator()
Get the terminator operation of this block.
Definition Block.cpp:249
BlockArgListType getArguments()
Definition Block.h:111
iterator end()
Definition Block.h:168
iterator begin()
Definition Block.h:167
IntegerAttr getIndexAttr(int64_t value)
Definition Builders.cpp:116
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
MLIRContext * getContext() const
Definition Builders.h:56
TypedAttr getOneAttr(Type type)
Definition Builders.cpp:351
This is a utility class for mapping one set of IR entities to another.
Definition IRMapping.h:26
auto lookupOrDefault(T from) const
Lookup a mapped value within the map.
Definition IRMapping.h:65
auto lookup(T from) const
Lookup a mapped value within the map.
Definition IRMapping.h:72
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
Definition IRMapping.h:30
bool contains(T from) const
Checks to see if a mapping for 'from' exists.
Definition IRMapping.h:51
This class coordinates rewriting a piece of IR outside of a pattern rewrite, providing a way to keep ...
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
RAII guard to reset the insertion point of the builder when destroyed.
Definition Builders.h:351
This class helps build Operations.
Definition Builders.h:210
InsertPoint saveInsertionPoint() const
Return a saved insertion point.
Definition Builders.h:388
Block * createBlock(Region *parent, Region::iterator insertPt={}, TypeRange argTypes={}, ArrayRef< Location > locs={})
Add new block with 'argTypes' arguments and set the insertion point to the end of it.
Definition Builders.cpp:439
Operation * clone(Operation &op, IRMapping &mapper)
Creates a deep copy of the specified operation, remapping any operands that use values outside of the...
Definition Builders.cpp:571
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
Definition Builders.h:434
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Definition Builders.h:401
static OpBuilder atBlockTerminator(Block *block, Listener *listener=nullptr)
Create a builder and set the insertion point to before the block terminator.
Definition Builders.h:255
void setInsertionPointToEnd(Block *block)
Sets the insertion point to the end of the specified block.
Definition Builders.h:439
void restoreInsertionPoint(InsertPoint ip)
Restore the insert point to a previously saved point.
Definition Builders.h:393
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
Definition Builders.h:528
void setInsertionPointAfter(Operation *op)
Sets the insertion point to the node after the specified operation, which will cause subsequent inser...
Definition Builders.h:415
This class represents a single result from folding an operation.
This class represents an operand of an operation.
Definition Value.h:254
This is a value defined by a result of an operation.
Definition Value.h:454
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
operand_type_range getOperandTypes()
Definition Operation.h:422
MutableArrayRef< Region > getRegions()
Returns the regions held by this operation.
Definition Operation.h:722
result_type_range getResultTypes()
Definition Operation.h:453
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
void setOperands(ValueRange operands)
Replace the current operands of this operation with the ones provided in 'operands'.
result_range getResults()
Definition Operation.h:440
Operation * clone(IRMapping &mapper, const CloneOptions &options=CloneOptions::all())
Create a deep copy of this operation, remapping any operands that use values outside of the operation...
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Definition Region.h:26
Block & front()
Definition Region.h:65
BlockArgListType getArguments()
Definition Region.h:94
iterator begin()
Definition Region.h:55
ParentT getParentOfType()
Find the first parent operation of the given type, or nullptr if there is no ancestor operation.
Definition Region.h:221
bool hasOneBlock()
Return true if this region has exactly one block.
Definition Region.h:68
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
virtual void eraseBlock(Block *block)
This method erases all operations in a block.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
void replaceAllUsesExcept(Value from, Value to, Operation *exceptedUser)
Find uses of from and replace them with to except if the user is exceptedUser.
virtual void inlineBlockBefore(Block *source, Block *dest, Block::iterator before, ValueRange argValues={})
Inline the operations of block 'source' into block 'dest' before the given position.
void mergeBlocks(Block *source, Block *dest, ValueRange argValues={})
Inline the operations of block 'source' into the end of block 'dest'.
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
void inlineRegionBefore(Region &region, Region &parent, Region::iterator before)
Move the blocks that belong to "region" before the given position in another region "parent".
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isIndex() const
Definition Types.cpp:56
bool isIntOrIndex() const
Return true if this is an integer (of any signedness) or an index type.
Definition Types.cpp:114
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
bool use_empty() const
Returns true if this value has no uses.
Definition Value.h:208
void replaceUsesWithIf(Value newValue, function_ref< bool(OpOperand &)> shouldReplace)
Replace all uses of 'this' value with 'newValue' if the given callback returns true.
Definition Value.cpp:91
Type getType() const
Return the type of this value.
Definition Value.h:105
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
static WalkResult advance()
Definition WalkResult.h:47
static WalkResult interrupt()
Definition WalkResult.h:46
Specialization of arith.constant op that returns an integer of index type.
Definition Arith.h:114
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
Operation * getOwner() const
Return the owner of this operand.
Definition UseDefLists.h:38
OpFoldResult makeComposedFoldedAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Constructs an AffineApplyOp that applies map to operands after composing the map with the maps of any...
std::optional< llvm::APSInt > computeUbMinusLb(Value lb, Value ub, bool isSigned)
Helper function to compute the difference between two values.
Definition SCF.cpp:115
Include the generated interface declarations.
void getPerfectlyNestedLoops(SmallVectorImpl< scf::ForOp > &nestedLoops, scf::ForOp root)
Get perfectly nested sequence of loops starting at root of loop nest (the first op being another Affi...
Definition Utils.cpp:1374
bool isPerfectlyNestedForLoops(MutableArrayRef< LoopLikeOpInterface > loops)
Check if the provided loops are perfectly nested for-loops.
Definition Utils.cpp:1581
FailureOr< UnrolledLoopInfo > loopUnrollByFactor(scf::ForOp forOp, uint64_t unrollFactor, function_ref< void(unsigned, Operation *, OpBuilder)> annotateFn=nullptr, bool shouldPromoteIfSingleIteration=true)
Unrolls this for operation by the specified unroll factor.
Definition Utils.cpp:368
LogicalResult outlineIfOp(RewriterBase &b, scf::IfOp ifOp, func::FuncOp *thenFn, StringRef thenFnName, func::FuncOp *elseFn, StringRef elseFnName)
Outline the then and/or else regions of ifOp as follows:
Definition Utils.cpp:218
void replaceAllUsesInRegionWith(Value orig, Value replacement, Region &region)
Replace all uses of orig within the given region with replacement.
SmallVector< scf::ForOp > replaceLoopNestWithNewYields(RewriterBase &rewriter, MutableArrayRef< scf::ForOp > loopNest, ValueRange newIterOperands, const NewYieldValuesFn &newYieldValuesFn, bool replaceIterOperandsUsesInLoop=true)
Update a perfectly nested loop nest to yield new values from the innermost loop and propagating it up...
Definition Utils.cpp:36
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
std::function< SmallVector< Value >( OpBuilder &b, Location loc, ArrayRef< BlockArgument > newBbArgs)> NewYieldValuesFn
A function that returns the additional yielded values during replaceWithAdditionalYields.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:307
LogicalResult coalescePerfectlyNestedSCFForLoops(scf::ForOp op)
Walk an affine.for to find a band to coalesce.
Definition Utils.cpp:1042
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
Definition AffineExpr.h:311
void generateUnrolledLoop(Block *loopBodyBlock, Value iv, uint64_t unrollFactor, function_ref< Value(unsigned, Value, OpBuilder)> ivRemapFn, function_ref< void(unsigned, Operation *, OpBuilder)> annotateFn, ValueRange iterArgs, ValueRange yieldedValues, IRMapping *clonedToSrcOpsMap=nullptr)
Generate unrolled copies of an scf loop's 'loopBodyBlock', with 'iterArgs' and 'yieldedValues' as the...
Definition Utils.cpp:295
Value getValueOrCreateConstantIntOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
Definition Utils.cpp:105
LogicalResult loopUnrollFull(scf::ForOp forOp)
Unrolls this loop completely.
Definition Utils.cpp:523
llvm::SmallVector< llvm::APInt > getConstLoopTripCounts(mlir::LoopLikeOpInterface loopOp)
Get constant trip counts for each of the induction variables of the given loop operation.
Definition Utils.cpp:1639
std::pair< Loops, Loops > TileLoops
Definition Utils.h:150
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
llvm::SmallVector< std::tuple< int64_t, int64_t, int64_t > > getConstLoopBounds(mlir::LoopLikeOpInterface loopOp)
Get constant loop bounds and steps for each of the induction variables of the given loop operation,...
Definition Utils.cpp:1620
void collapseParallelLoops(RewriterBase &rewriter, scf::ParallelOp loops, ArrayRef< std::vector< unsigned > > combinedDimensions)
Take the ParallelLoop and for each set of dimension indices, combine them into a single dimension.
Definition Utils.cpp:1117
llvm::SetVector< T, Vector, Set, N > SetVector
Definition LLVM.h:125
std::optional< std::pair< APInt, bool > > getConstantAPIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
SliceOptions ForwardSliceOptions
Loops tilePerfectlyNested(scf::ForOp rootForOp, ArrayRef< Value > sizes)
Tile a nest of scf::ForOp loops rooted at rootForOp with the given (parametric) sizes.
Definition Utils.cpp:1362
LogicalResult loopUnrollJamByFactor(scf::ForOp forOp, uint64_t unrollFactor)
Unrolls and jams this scf.for operation by the specified unroll factor.
Definition Utils.cpp:551
bool getInnermostParallelLoops(Operation *rootOp, SmallVectorImpl< scf::ParallelOp > &result)
Get a list of innermost parallel loops contained in rootOp.
Definition Utils.cpp:241
bool isZeroInteger(OpFoldResult v)
Return "true" if v is an integer value/attribute with constant value 0.
void bindSymbols(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to SymbolExpr at positions: [0 .
Definition AffineExpr.h:325
FailureOr< scf::ParallelOp > parallelLoopUnrollByFactors(scf::ParallelOp op, ArrayRef< uint64_t > unrollFactors, RewriterBase &rewriter, function_ref< void(unsigned, Operation *, OpBuilder)> annotateFn=nullptr, IRMapping *clonedToSrcOpsMap=nullptr)
Unroll this scf::Parallel loop by the specified unroll factors.
Definition Utils.cpp:1657
void getUsedValuesDefinedAbove(Region &region, Region &limit, SetVector< Value > &values)
Fill values with a list of values defined at the ancestors of the limit region and used within region...
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
Definition Utils.cpp:114
SmallVector< Loops, 8 > tile(ArrayRef< scf::ForOp > forOps, ArrayRef< Value > sizes, ArrayRef< scf::ForOp > targets)
Performs tiling fo imperfectly nested loops (with interchange) by strip-mining the forOps by sizes an...
Definition Utils.cpp:1341
FailureOr< func::FuncOp > outlineSingleBlockRegion(RewriterBase &rewriter, Location loc, Region &region, StringRef funcName, func::CallOp *callOp=nullptr)
Outline a region with a single block into a new FuncOp.
Definition Utils.cpp:115
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
bool areValuesDefinedAbove(Range values, Region &limit)
Check if all values in the provided range are defined above the limit region.
Definition RegionUtils.h:26
void denormalizeInductionVariable(RewriterBase &rewriter, Location loc, Value normalizedIv, OpFoldResult origLb, OpFoldResult origStep)
Get back the original induction variable values after loop normalization.
Definition Utils.cpp:805
scf::ForallOp fuseIndependentSiblingForallLoops(scf::ForallOp target, scf::ForallOp source, RewriterBase &rewriter)
Given two scf.forall loops, target and source, fuses target into source.
Definition Utils.cpp:1430
LogicalResult coalesceLoops(MutableArrayRef< scf::ForOp > loops)
Replace a perfect nest of "for" loops with a single linearized loop.
Definition Utils.cpp:1034
scf::ForOp fuseIndependentSiblingForLoops(scf::ForOp target, scf::ForOp source, RewriterBase &rewriter)
Given two scf.for loops, target and source, fuses target into source.
Definition Utils.cpp:1483
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
TileLoops extractFixedOuterLoops(scf::ForOp rootFOrOp, ArrayRef< int64_t > sizes)
Definition Utils.cpp:1379
Range emitNormalizedLoopBounds(RewriterBase &rewriter, Location loc, OpFoldResult lb, OpFoldResult ub, OpFoldResult step)
Materialize bounds and step of a zero-based and unit-step loop derived by normalizing the specified b...
Definition Utils.cpp:734
SmallVector< scf::ForOp, 8 > Loops
Tile a nest of standard for loops rooted at rootForOp by finding such parametric tile sizes that the ...
Definition Utils.h:149
bool isOneInteger(OpFoldResult v)
Return true if v is an IntegerAttr with value 1.
std::optional< APInt > constantTripCount(OpFoldResult lb, OpFoldResult ub, OpFoldResult step, bool isSigned, llvm::function_ref< std::optional< llvm::APSInt >(Value, Value, bool)> computeUbMinusLb)
Return the number of iterations for a loop with a lower bound lb, upper bound ub and step step,...
LogicalResult foldDynamicIndexList(SmallVectorImpl< OpFoldResult > &ofrs, bool onlyNonNegative=false, bool onlyNonZero=false)
Returns "success" when any of the elements in ofrs is a constant value.
FailureOr< scf::ForallOp > normalizeForallOp(RewriterBase &rewriter, scf::ForallOp forallOp)
Normalize an scf.forall operation.
Definition Utils.cpp:1536
void getForwardSlice(Operation *op, SetVector< Operation * > *forwardSlice, const ForwardSliceOptions &options={})
Fills forwardSlice with the computed forward slice (i.e.
void walk(Operation *op)
SmallVector< std::pair< Block::iterator, Block::iterator > > subBlocks
Represents a range (offset, size, and stride) where each element of the triple may be dynamic or stat...
OpFoldResult stride
OpFoldResult size
OpFoldResult offset
std::optional< scf::ForOp > epilogueLoopOp
Definition Utils.h:108
std::optional< scf::ForOp > mainLoopOp
Definition Utils.h:107
Eliminates variable at the specified position using Fourier-Motzkin variable elimination.