1//===-- UnrollLoop.cpp - Loop unrolling utilities -------------------------===//
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 some loop unrolling utilities. It does not define any
10// actual pass or policy, but provides a single function to perform loop
11// unrolling.
12//
13// The process of unrolling can produce extraneous basic blocks linked with
14// unconditional branches. This will be corrected in the future.
15//
16//===----------------------------------------------------------------------===//
17
18#include "llvm/ADT/ArrayRef.h"
19#include "llvm/ADT/DenseMap.h"
20#include "llvm/ADT/MapVector.h"
21#include "llvm/ADT/STLExtras.h"
22#include "llvm/ADT/ScopedHashTable.h"
23#include "llvm/ADT/SetVector.h"
24#include "llvm/ADT/SmallVector.h"
25#include "llvm/ADT/Statistic.h"
26#include "llvm/ADT/StringRef.h"
27#include "llvm/ADT/Twine.h"
28#include "llvm/Analysis/AliasAnalysis.h"
29#include "llvm/Analysis/AssumptionCache.h"
30#include "llvm/Analysis/DomTreeUpdater.h"
31#include "llvm/Analysis/InstructionSimplify.h"
32#include "llvm/Analysis/LoopInfo.h"
33#include "llvm/Analysis/LoopIterator.h"
34#include "llvm/Analysis/MemorySSA.h"
35#include "llvm/Analysis/OptimizationRemarkEmitter.h"
36#include "llvm/Analysis/ScalarEvolution.h"
37#include "llvm/IR/BasicBlock.h"
38#include "llvm/IR/CFG.h"
39#include "llvm/IR/Constants.h"
40#include "llvm/IR/DebugInfoMetadata.h"
41#include "llvm/IR/DebugLoc.h"
42#include "llvm/IR/DiagnosticInfo.h"
43#include "llvm/IR/Dominators.h"
44#include "llvm/IR/Function.h"
45#include "llvm/IR/IRBuilder.h"
46#include "llvm/IR/Instruction.h"
47#include "llvm/IR/Instructions.h"
48#include "llvm/IR/IntrinsicInst.h"
49#include "llvm/IR/Metadata.h"
50#include "llvm/IR/PatternMatch.h"
51#include "llvm/IR/Use.h"
52#include "llvm/IR/User.h"
53#include "llvm/IR/ValueHandle.h"
54#include "llvm/IR/ValueMap.h"
55#include "llvm/Support/Casting.h"
56#include "llvm/Support/CommandLine.h"
57#include "llvm/Support/Debug.h"
58#include "llvm/Support/GenericDomTree.h"
59#include "llvm/Support/raw_ostream.h"
60#include "llvm/Transforms/Utils/BasicBlockUtils.h"
61#include "llvm/Transforms/Utils/Cloning.h"
62#include "llvm/Transforms/Utils/Local.h"
63#include "llvm/Transforms/Utils/LoopSimplify.h"
64#include "llvm/Transforms/Utils/LoopUtils.h"
65#include "llvm/Transforms/Utils/SimplifyIndVar.h"
66#include "llvm/Transforms/Utils/UnrollLoop.h"
67#include "llvm/Transforms/Utils/ValueMapper.h"
68#include <assert.h>
69#include <cmath>
70#include <numeric>
71#include <vector>
72
73namespace llvm {
74class DataLayout;
75class Value;
76} // namespace llvm
77
78using namespace llvm;
79
80#define DEBUG_TYPE "loop-unroll"
81
82// TODO: Should these be here or in LoopUnroll?
83STATISTIC(NumCompletelyUnrolled, "Number of loops completely unrolled");
84STATISTIC(NumUnrolled, "Number of loops unrolled (completely or otherwise)");
85STATISTIC(NumUnrolledNotLatch, "Number of loops unrolled without a conditional "
86 "latch (completely or otherwise)");
87
88static cl::opt<bool>
89UnrollRuntimeEpilog("unroll-runtime-epilog", cl::init(Val: false), cl::Hidden,
90 cl::desc("Allow runtime unrolled loops to be unrolled "
91 "with epilog instead of prolog."));
92
93static cl::opt<bool> UnrollUniformWeights(
94 "unroll-uniform-weights", cl::init(Val: false), cl::Hidden,
95 cl::desc("If new branch weights must be found, work harder to keep them "
96 "uniform."));
97
98static cl::opt<bool>
99UnrollVerifyDomtree("unroll-verify-domtree", cl::Hidden,
100 cl::desc("Verify domtree after unrolling"),
101#ifdef EXPENSIVE_CHECKS
102 cl::init(true)
103#else
104 cl::init(Val: false)
105#endif
106 );
107
108static cl::opt<bool>
109UnrollVerifyLoopInfo("unroll-verify-loopinfo", cl::Hidden,
110 cl::desc("Verify loopinfo after unrolling"),
111#ifdef EXPENSIVE_CHECKS
112 cl::init(true)
113#else
114 cl::init(Val: false)
115#endif
116 );
117
118static cl::opt<bool> UnrollAddParallelReductions(
119 "unroll-add-parallel-reductions", cl::init(Val: false), cl::Hidden,
120 cl::desc("Allow unrolling to add parallel reduction phis."));
121
122/// Check if unrolling created a situation where we need to insert phi nodes to
123/// preserve LCSSA form.
124/// \param Blocks is a vector of basic blocks representing unrolled loop.
125/// \param L is the outer loop.
126/// It's possible that some of the blocks are in L, and some are not. In this
127/// case, if there is a use is outside L, and definition is inside L, we need to
128/// insert a phi-node, otherwise LCSSA will be broken.
129/// The function is just a helper function for llvm::UnrollLoop that returns
130/// true if this situation occurs, indicating that LCSSA needs to be fixed.
131static bool needToInsertPhisForLCSSA(Loop *L,
132 const std::vector<BasicBlock *> &Blocks,
133 LoopInfo *LI) {
134 for (BasicBlock *BB : Blocks) {
135 if (LI->getLoopFor(BB) == L)
136 continue;
137 for (Instruction &I : *BB) {
138 for (Use &U : I.operands()) {
139 if (const auto *Def = dyn_cast<Instruction>(Val&: U)) {
140 Loop *DefLoop = LI->getLoopFor(BB: Def->getParent());
141 if (!DefLoop)
142 continue;
143 if (DefLoop->contains(L))
144 return true;
145 }
146 }
147 }
148 }
149 return false;
150}
151
152/// Adds ClonedBB to LoopInfo, creates a new loop for ClonedBB if necessary
153/// and adds a mapping from the original loop to the new loop to NewLoops.
154/// Returns nullptr if no new loop was created and a pointer to the
155/// original loop OriginalBB was part of otherwise.
156const Loop* llvm::addClonedBlockToLoopInfo(BasicBlock *OriginalBB,
157 BasicBlock *ClonedBB, LoopInfo *LI,
158 NewLoopsMap &NewLoops) {
159 // Figure out which loop New is in.
160 const Loop *OldLoop = LI->getLoopFor(BB: OriginalBB);
161 assert(OldLoop && "Should (at least) be in the loop being unrolled!");
162
163 Loop *&NewLoop = NewLoops[OldLoop];
164 if (!NewLoop) {
165 // Found a new sub-loop.
166 assert(OriginalBB == OldLoop->getHeader() &&
167 "Header should be first in RPO");
168
169 NewLoop = LI->AllocateLoop();
170 Loop *NewLoopParent = NewLoops.lookup(Val: OldLoop->getParentLoop());
171
172 if (NewLoopParent)
173 NewLoopParent->addChildLoop(NewChild: NewLoop);
174 else
175 LI->addTopLevelLoop(New: NewLoop);
176
177 NewLoop->addBasicBlockToLoop(NewBB: ClonedBB, LI&: *LI);
178 return OldLoop;
179 } else {
180 NewLoop->addBasicBlockToLoop(NewBB: ClonedBB, LI&: *LI);
181 return nullptr;
182 }
183}
184
185/// The function chooses which type of unroll (epilog or prolog) is more
186/// profitabale.
187/// Epilog unroll is more profitable when there is PHI that starts from
188/// constant. In this case epilog will leave PHI start from constant,
189/// but prolog will convert it to non-constant.
190///
191/// loop:
192/// PN = PHI [I, Latch], [CI, PreHeader]
193/// I = foo(PN)
194/// ...
195///
196/// Epilog unroll case.
197/// loop:
198/// PN = PHI [I2, Latch], [CI, PreHeader]
199/// I1 = foo(PN)
200/// I2 = foo(I1)
201/// ...
202/// Prolog unroll case.
203/// NewPN = PHI [PrologI, Prolog], [CI, PreHeader]
204/// loop:
205/// PN = PHI [I2, Latch], [NewPN, PreHeader]
206/// I1 = foo(PN)
207/// I2 = foo(I1)
208/// ...
209///
210static bool isEpilogProfitable(Loop *L) {
211 BasicBlock *PreHeader = L->getLoopPreheader();
212 BasicBlock *Header = L->getHeader();
213 assert(PreHeader && Header);
214 for (const PHINode &PN : Header->phis()) {
215 if (isa<ConstantInt>(Val: PN.getIncomingValueForBlock(BB: PreHeader)))
216 return true;
217 }
218 return false;
219}
220
221struct LoadValue {
222 Instruction *DefI = nullptr;
223 unsigned Generation = 0;
224 LoadValue() = default;
225 LoadValue(Instruction *Inst, unsigned Generation)
226 : DefI(Inst), Generation(Generation) {}
227};
228
229class StackNode {
230 ScopedHashTable<const SCEV *, LoadValue>::ScopeTy LoadScope;
231 unsigned CurrentGeneration;
232 unsigned ChildGeneration;
233 DomTreeNode *Node;
234 DomTreeNode::const_iterator ChildIter;
235 DomTreeNode::const_iterator EndIter;
236 bool Processed = false;
237
238public:
239 StackNode(ScopedHashTable<const SCEV *, LoadValue> &AvailableLoads,
240 unsigned cg, DomTreeNode *N, DomTreeNode::const_iterator Child,
241 DomTreeNode::const_iterator End)
242 : LoadScope(AvailableLoads), CurrentGeneration(cg), ChildGeneration(cg),
243 Node(N), ChildIter(Child), EndIter(End) {}
244 // Accessors.
245 unsigned currentGeneration() const { return CurrentGeneration; }
246 unsigned childGeneration() const { return ChildGeneration; }
247 void childGeneration(unsigned generation) { ChildGeneration = generation; }
248 DomTreeNode *node() { return Node; }
249 DomTreeNode::const_iterator childIter() const { return ChildIter; }
250
251 DomTreeNode *nextChild() {
252 DomTreeNode *Child = *ChildIter;
253 ++ChildIter;
254 return Child;
255 }
256
257 DomTreeNode::const_iterator end() const { return EndIter; }
258 bool isProcessed() const { return Processed; }
259 void process() { Processed = true; }
260};
261
262Value *getMatchingValue(LoadValue LV, LoadInst *LI, unsigned CurrentGeneration,
263 BatchAAResults &BAA,
264 function_ref<MemorySSA *()> GetMSSA) {
265 if (!LV.DefI)
266 return nullptr;
267 if (LV.DefI->getType() != LI->getType())
268 return nullptr;
269 if (LV.Generation != CurrentGeneration) {
270 MemorySSA *MSSA = GetMSSA();
271 if (!MSSA)
272 return nullptr;
273 auto *EarlierMA = MSSA->getMemoryAccess(I: LV.DefI);
274 MemoryAccess *LaterDef =
275 MSSA->getWalker()->getClobberingMemoryAccess(I: LI, AA&: BAA);
276 if (!MSSA->dominates(A: LaterDef, B: EarlierMA))
277 return nullptr;
278 }
279 return LV.DefI;
280}
281
282void loadCSE(Loop *L, DominatorTree &DT, ScalarEvolution &SE, LoopInfo &LI,
283 BatchAAResults &BAA, function_ref<MemorySSA *()> GetMSSA) {
284 ScopedHashTable<const SCEV *, LoadValue> AvailableLoads;
285 SmallVector<std::unique_ptr<StackNode>> NodesToProcess;
286 DomTreeNode *HeaderD = DT.getNode(BB: L->getHeader());
287 NodesToProcess.emplace_back(Args: new StackNode(AvailableLoads, 0, HeaderD,
288 HeaderD->begin(), HeaderD->end()));
289
290 unsigned CurrentGeneration = 0;
291 while (!NodesToProcess.empty()) {
292 StackNode *NodeToProcess = &*NodesToProcess.back();
293
294 CurrentGeneration = NodeToProcess->currentGeneration();
295
296 if (!NodeToProcess->isProcessed()) {
297 // Process the node.
298
299 // If this block has a single predecessor, then the predecessor is the
300 // parent
301 // of the domtree node and all of the live out memory values are still
302 // current in this block. If this block has multiple predecessors, then
303 // they could have invalidated the live-out memory values of our parent
304 // value. For now, just be conservative and invalidate memory if this
305 // block has multiple predecessors.
306 if (!NodeToProcess->node()->getBlock()->getSinglePredecessor())
307 ++CurrentGeneration;
308 for (auto &I : make_early_inc_range(Range&: *NodeToProcess->node()->getBlock())) {
309
310 auto *Load = dyn_cast<LoadInst>(Val: &I);
311 if (!Load || !Load->isSimple()) {
312 if (I.mayWriteToMemory())
313 CurrentGeneration++;
314 continue;
315 }
316
317 const SCEV *PtrSCEV = SE.getSCEV(V: Load->getPointerOperand());
318 LoadValue LV = AvailableLoads.lookup(Key: PtrSCEV);
319 if (Value *M =
320 getMatchingValue(LV, LI: Load, CurrentGeneration, BAA, GetMSSA)) {
321 if (LI.replacementPreservesLCSSAForm(From: Load, To: M)) {
322 Load->replaceAllUsesWith(V: M);
323 Load->eraseFromParent();
324 }
325 } else {
326 AvailableLoads.insert(Key: PtrSCEV, Val: LoadValue(Load, CurrentGeneration));
327 }
328 }
329 NodeToProcess->childGeneration(generation: CurrentGeneration);
330 NodeToProcess->process();
331 } else if (NodeToProcess->childIter() != NodeToProcess->end()) {
332 // Push the next child onto the stack.
333 DomTreeNode *Child = NodeToProcess->nextChild();
334 if (!L->contains(BB: Child->getBlock()))
335 continue;
336 NodesToProcess.emplace_back(
337 Args: new StackNode(AvailableLoads, NodeToProcess->childGeneration(), Child,
338 Child->begin(), Child->end()));
339 } else {
340 // It has been processed, and there are no more children to process,
341 // so delete it and pop it off the stack.
342 NodesToProcess.pop_back();
343 }
344 }
345}
346
347/// Perform some cleanup and simplifications on loops after unrolling. It is
348/// useful to simplify the IV's in the new loop, as well as do a quick
349/// simplify/dce pass of the instructions.
350void llvm::simplifyLoopAfterUnroll(Loop *L, bool SimplifyIVs, LoopInfo *LI,
351 ScalarEvolution *SE, DominatorTree *DT,
352 AssumptionCache *AC,
353 const TargetTransformInfo *TTI,
354 ArrayRef<BasicBlock *> Blocks,
355 AAResults *AA) {
356 using namespace llvm::PatternMatch;
357
358 // Simplify any new induction variables in the partially unrolled loop.
359 if (SE && SimplifyIVs) {
360 SmallVector<WeakTrackingVH, 16> DeadInsts;
361 simplifyLoopIVs(L, SE, DT, LI, TTI, Dead&: DeadInsts);
362
363 // Aggressively clean up dead instructions that simplifyLoopIVs already
364 // identified. Any remaining should be cleaned up below.
365 while (!DeadInsts.empty()) {
366 Value *V = DeadInsts.pop_back_val();
367 if (Instruction *Inst = dyn_cast_or_null<Instruction>(Val: V))
368 RecursivelyDeleteTriviallyDeadInstructions(V: Inst);
369 }
370
371 if (AA) {
372 std::unique_ptr<MemorySSA> MSSA = nullptr;
373 BatchAAResults BAA(*AA);
374 loadCSE(L, DT&: *DT, SE&: *SE, LI&: *LI, BAA, GetMSSA: [L, AA, DT, &MSSA]() -> MemorySSA * {
375 if (!MSSA)
376 MSSA.reset(p: new MemorySSA(*L, AA, DT));
377 return &*MSSA;
378 });
379 }
380 }
381
382 // At this point, the code is well formed. Perform constprop, instsimplify,
383 // and dce.
384 SmallVector<WeakTrackingVH, 16> DeadInsts;
385 for (BasicBlock *BB : Blocks) {
386 // Remove repeated debug instructions after loop unrolling.
387 if (BB->getParent()->getSubprogram())
388 RemoveRedundantDbgInstrs(BB);
389
390 for (Instruction &Inst : llvm::make_early_inc_range(Range&: *BB)) {
391 if (Value *V = simplifyInstruction(
392 I: &Inst, Q: {BB->getDataLayout(), nullptr, DT, AC}))
393 if (LI->replacementPreservesLCSSAForm(From: &Inst, To: V))
394 Inst.replaceAllUsesWith(V);
395 if (isInstructionTriviallyDead(I: &Inst))
396 DeadInsts.emplace_back(Args: &Inst);
397
398 // Fold ((add X, C1), C2) to (add X, C1+C2). This is very common in
399 // unrolled loops, and handling this early allows following code to
400 // identify the IV as a "simple recurrence" without first folding away
401 // a long chain of adds.
402 {
403 Value *X;
404 const APInt *C1, *C2;
405 if (match(V: &Inst, P: m_Add(L: m_Add(L: m_Value(V&: X), R: m_APInt(Res&: C1)), R: m_APInt(Res&: C2)))) {
406 auto *InnerI = dyn_cast<Instruction>(Val: Inst.getOperand(i: 0));
407 auto *InnerOBO = cast<OverflowingBinaryOperator>(Val: Inst.getOperand(i: 0));
408 bool SignedOverflow;
409 APInt NewC = C1->sadd_ov(RHS: *C2, Overflow&: SignedOverflow);
410 Inst.setOperand(i: 0, Val: X);
411 Inst.setOperand(i: 1, Val: ConstantInt::get(Ty: Inst.getType(), V: NewC));
412 Inst.setHasNoUnsignedWrap(Inst.hasNoUnsignedWrap() &&
413 InnerOBO->hasNoUnsignedWrap());
414 Inst.setHasNoSignedWrap(Inst.hasNoSignedWrap() &&
415 InnerOBO->hasNoSignedWrap() &&
416 !SignedOverflow);
417 if (InnerI && isInstructionTriviallyDead(I: InnerI))
418 DeadInsts.emplace_back(Args&: InnerI);
419 }
420 }
421 }
422 // We can't do recursive deletion until we're done iterating, as we might
423 // have a phi which (potentially indirectly) uses instructions later in
424 // the block we're iterating through.
425 RecursivelyDeleteTriviallyDeadInstructions(DeadInsts);
426 }
427}
428
429// If LoopUnroll has proven OriginalLoopProb is incorrect for some iterations
430// of the original loop, adjust latch probabilities in the unrolled loop to
431// maintain the original total frequency of the original loop body.
432//
433// OriginalLoopProb is practical but imprecise
434// -------------------------------------------
435//
436// The latch branch weights that LLVM originally adds to a loop encode one latch
437// probability, OriginalLoopProb, applied uniformly across the loop's infinite
438// set of theoretically possible iterations. While this uniform latch
439// probability serves as a practical statistic summarizing the trip counts
440// observed during profiling, it is imprecise. Specifically, unless it is zero,
441// it is impossible for it to be the actual probability observed at every
442// individual iteration. To see why, consider that the only way to actually
443// observe at run time that the latch probability remains non-zero is to profile
444// at least one loop execution that has an infinite number of iterations. I do
445// not know how to profile an infinite number of loop iterations, and most loops
446// I work with are always finite.
447//
448// LoopUnroll proves OriginalLoopProb is incorrect
449// ------------------------------------------------
450//
451// LoopUnroll reorganizes the original loop so that loop iterations are no
452// longer all implemented by the same code, and then it analyzes some of those
453// loop iteration implementations independently of others. In particular, it
454// converts some of their conditional latches to unconditional. That is, by
455// examining code structure without any profile data, LoopUnroll proves that the
456// actual latch probability at the end of such an iteration is either 1 or 0.
457// When an individual iteration's actual latch probability is 1 or 0, that means
458// it always behaves the same, so it is impossible to observe it as having any
459// other probability. The original uniform latch probability is rarely 1 or 0
460// because, when applied to all possible iterations, that would yield an
461// estimated trip count of infinity or 1, respectively.
462//
463// Thus, the new probabilities of 1 or 0 are proven corrections to
464// OriginalLoopProb for individual iterations in the original loop. However,
465// LoopUnroll often is able to perform these corrections for only some
466// iterations, leaving other iterations with OriginalLoopProb, and thus
467// corrupting the aggregate effect on the total frequency of the original loop
468// body.
469//
470// Adjusting latch probabilities
471// -----------------------------
472//
473// This function ensures that the total frequency of the original loop body,
474// summed across all its occurrences in the unrolled loop after the
475// aforementioned latch conversions, is the same as in the original loop. To do
476// so, it adjusts probabilities on the remaining conditional latches. However,
477// it cannot derive the new probabilities directly from the original uniform
478// latch probability because the latter has been proven incorrect for some
479// original loop iterations.
480//
481// There are often many sets of latch probabilities that can produce the
482// original total loop body frequency. If there are many remaining conditional
483// latches and !UnrollUniformWeights, this function just quickly hacks a few of
484// their probabilities to restore the original total loop body frequency.
485// Otherwise, it tries harder to determine less arbitrary probabilities.
486static void fixProbContradiction(Loop *L, UnrollLoopOptions ULO,
487 OptimizationRemarkEmitter *ORE,
488 BranchProbability OriginalLoopProb,
489 bool CompletelyUnroll,
490 std::vector<unsigned> &IterCounts,
491 const std::vector<BasicBlock *> &CondLatches,
492 std::vector<BasicBlock *> &CondLatchNexts) {
493 // Runtime unrolling is handled later in LoopUnroll not here.
494 //
495 // There are two scenarios in which LoopUnroll sets ProbUpdateRequired to true
496 // because it needs to update probabilities that were originally
497 // OriginalLoopProb, but only in one scenario has LoopUnroll proven
498 // OriginalLoopProb incorrect for iterations within the original loop:
499 // - If ULO.Runtime, LoopUnroll adds new guards that enforce new reaching
500 // conditions for new loop iteration implementations (e.g., one unrolled
501 // loop iteration executes only if at least ULO.Count original loop
502 // iterations remain). Those reaching conditions dictate how conditional
503 // latches can be converted to unconditional (e.g., within an unrolled loop
504 // iteration, there is no need to recheck the number of remaining original
505 // loop iterations). None of this reorganization alters the set of possible
506 // original loop iteration counts or proves OriginalLoopProb incorrect for
507 // any of the original loop iterations. Thus, LoopUnroll derives
508 // probabilities for the new guards and latches directly from
509 // OriginalLoopProb based on the probabilities that their reaching
510 // conditions would occur in the original loop. Doing so maintains the
511 // total frequency of the original loop body.
512 // - If !ULO.Runtime, LoopUnroll initially adds new loop iteration
513 // implementations, which have the same latch probabilities as in the
514 // original loop because there are no new guards that change their reaching
515 // conditions. Sometimes, LoopUnroll is then done, and so does not set
516 // ProbUpdateRequired to true. Other times, LoopUnroll then proves that
517 // some latches are unconditional, directly contradicting OriginalLoopProb
518 // for the corresponding original loop iterations. That reduces the set of
519 // possible original loop iteration counts, possibly producing a finite set
520 // if it manages to eliminate the backedge. LoopUnroll has to choose a new
521 // set of latch probabilities that produce the same total loop body
522 // frequency.
523 //
524 // This function addresses the second scenario only.
525 if (ULO.Runtime)
526 return;
527
528 // If CondLatches.empty(), there are no latch branches with probabilities we
529 // can adjust. That should mean that the actual trip count is always exactly
530 // the number of remaining unrolled iterations, and so OriginalLoopProb should
531 // have yielded that trip count as the original loop body frequency. Of
532 // course, OriginalLoopProb could be based on inaccurate profile data, but
533 // there is nothing we can do about that here.
534 if (CondLatches.empty())
535 return;
536
537 // If the original latch probability is 1, the original frequency is infinity.
538 // Leaving all remaining probabilities set to 1 might or might not get us
539 // there (e.g., a completely unrolled loop cannot be infinite), but it is the
540 // closest we can come.
541 assert(!OriginalLoopProb.isUnknown() &&
542 "Expected to have loop probability to fix");
543 if (OriginalLoopProb.isOne())
544 return;
545
546 // FreqDesired is the frequency implied by the original loop probability.
547 double FreqDesired = 1 / (1 - OriginalLoopProb.toDouble());
548
549 // Get the probability at CondLatches[I].
550 auto GetProb = [&](unsigned I) {
551 CondBrInst *B = cast<CondBrInst>(Val: CondLatches[I]->getTerminator());
552 bool FirstTargetIsNext = B->getSuccessor(i: 0) == CondLatchNexts[I];
553 return getBranchProbability(B, ForFirstTarget: FirstTargetIsNext).toDouble();
554 };
555
556 // Set the probability at CondLatches[I] to Prob.
557 auto SetProb = [&](unsigned I, double Prob) {
558 CondBrInst *B = cast<CondBrInst>(Val: CondLatches[I]->getTerminator());
559 bool FirstTargetIsNext = B->getSuccessor(i: 0) == CondLatchNexts[I];
560 setBranchProbability(B, P: BranchProbability::getBranchProbability(Prob),
561 ForFirstTarget: FirstTargetIsNext);
562 };
563
564 // Set all probabilities in CondLatches to Prob.
565 auto SetAllProbs = [&](double Prob) {
566 for (unsigned I = 0, E = CondLatches.size(); I < E; ++I)
567 SetProb(I, Prob);
568 };
569
570 // If UnrollUniformWeights or n <= 2, we choose the simplest probability model
571 // we can think of: every remaining conditional branch instruction has the
572 // same probability, Prob, of continuing to the next iteration. This model
573 // has several helpful properties:
574 // - There is only one search parameter, Prob.
575 // - We have no reason to think one latch branch's probability should be
576 // higher or lower than another, and so this model makes them all the same.
577 // In the worst cases, we thus avoid setting just some probabilities to 0 or
578 // 1, which can unrealistically make some code appear unreachable. There
579 // are cases where they *all* must become 0 or 1 to achieve the total
580 // frequency of original loop body, and our model does permit that.
581 // - The frequency, FreqOne, of the original loop body in a single iteration
582 // of the unrolled loop is computed by a simple polynomial, where p=Prob,
583 // n=CondLatches.size(), and c_i=IterCounts[i]:
584 //
585 // FreqOne = Sum(i=0..n)(c_i * p^i)
586 //
587 // - If the backedge has been eliminated:
588 // - FreqOne is the total frequency of the original loop body in the
589 // unrolled loop.
590 // - If Prob == 1, the total frequency of the original loop body is exactly
591 // the number of remaining loop iterations, as expected because every
592 // remaining loop iteration always then executes.
593 // - If the backedge remains:
594 // - Sum(i=0..inf)(FreqOne * p^(n*i)) = FreqOne / (1 - p^n) is the total
595 // frequency of the original loop body in the unrolled loop, regardless of
596 // whether the backedge is conditional or unconditional.
597 // - As Prob approaches 1, the total frequency of the original loop body
598 // approaches infinity, as expected because the loop approaches never
599 // exiting.
600 // - For n <= 2, we can use simple formulas to solve the above polynomial
601 // equations exactly for p without performing a search.
602 // - For n > 2, evaluating each point in the search space, using ComputeFreq
603 // below, requires about as few instructions as we could hope for. That is,
604 // the probability is constant across the conditional branches, so the only
605 // computation is across conditional branches and any backedge, as required
606 // for any model for Prob.
607 // - Prob == 1 produces the maximum possible total frequency for the original
608 // loop body, as described above. Prob == 0 produces the minimum, 0.
609 // Increasing or decreasing Prob monotonically increases or decreases the
610 // frequency, respectively. Thus, for every possible frequency, there
611 // exists some Prob that can produce it, and we can easily use bisection to
612 // search the problem space.
613
614 // When iterating for a solution, we stop early if we find probabilities
615 // that produce a Freq whose relative difference from FreqDesired is small
616 // (FreqPrec). Otherwise, we expect to compute a solution at least that
617 // accurate (but surely far more accurate).
618 const double FreqPrec = 1e-6;
619
620 // Compute the new frequency produced by using Prob throughout CondLatches.
621 auto ComputeFreq = [&](double Prob) {
622 double ProbReaching = 1; // p^0
623 double FreqOne = IterCounts[0]; // c_0*p^0
624 for (unsigned I = 0, E = CondLatches.size(); I < E; ++I) {
625 ProbReaching *= Prob; // p^(I+1)
626 FreqOne += IterCounts[I + 1] * ProbReaching; // c_(I+1)*p^(I+1)
627 }
628 double ProbReachingBackedge = CompletelyUnroll ? 0 : ProbReaching;
629 assert(FreqOne > 0 && "Expected at least one iteration before first latch");
630 if (ProbReachingBackedge == 1)
631 return std::numeric_limits<double>::infinity();
632 return FreqOne / (1 - ProbReachingBackedge);
633 };
634
635 // Compute the probability that, used at CondLaches[0] where
636 // CondLatches.size() == 1, gets as close as possible to FreqDesired.
637 auto ComputeProbForLinear = [&]() {
638 // The polynomial is linear (0 = A*p + B), so just solve it.
639 double A = IterCounts[1] + (CompletelyUnroll ? 0 : FreqDesired);
640 double B = IterCounts[0] - FreqDesired;
641 assert(A > 0 && "Expected iterations after last conditional latch");
642 double Prob = -B / A;
643 // If it computes an invalid Prob, FreqDesired is impossibly low or high.
644 // Otherwise, Prob should produce nearly FreqDesired.
645 assert((Prob < 0 || Prob > 1 ||
646 fabs(ComputeFreq(Prob) - FreqDesired) / FreqDesired < FreqPrec) &&
647 "Expected accurate frequency when linear case is possible");
648 Prob = std::max(a: Prob, b: 0.);
649 Prob = std::min(a: Prob, b: 1.);
650 return Prob;
651 };
652
653 // Compute the probability that, used throughout CondLatches where
654 // CondLatches.size() == 2, gets as close as possible to FreqDesired.
655 auto ComputeProbForQuadratic = [&]() {
656 // The polynomial is quadratic (0 = A*p^2 + B*p + C), so just solve it.
657 double A = IterCounts[2] + (CompletelyUnroll ? 0 : FreqDesired);
658 double B = IterCounts[1];
659 double C = IterCounts[0] - FreqDesired;
660 assert(A > 0 && "Expected iterations after last conditional latch");
661 double Prob = (-B + sqrt(x: B * B - 4 * A * C)) / (2 * A);
662 // If it computes an invalid Prob, FreqDesired is impossibly low or high.
663 // Otherwise, Prob should produce nearly FreqDesired.
664 assert((Prob < 0 || Prob > 1 ||
665 fabs(ComputeFreq(Prob) - FreqDesired) / FreqDesired < FreqPrec) &&
666 "Expected accurate frequency when quadratic case is possible");
667 Prob = std::max(a: Prob, b: 0.);
668 Prob = std::min(a: Prob, b: 1.);
669 return Prob;
670 };
671
672 // Adjust the probability at CondLatches[ComputeIdx] to get as close as
673 // possible to FreqDesired without replacing probabilities elsewhere in
674 // CondLatches. Return the new total frequency.
675 //
676 // Given a CondLatches index I, then for a single unrolled loop iteration:
677 // - ProbBefore or ProbAfter is the probability that control flow can pass
678 // through every CondLatches[J] for J < I or J > I, respectively.
679 // - FreqBefore or FreqAfter is the total frequency accumulated before or
680 // after CondLatches[I], respectively, while the probability at
681 // CondLatches[I] is treated as 1.
682 //
683 // If ComputeIdx == 0, then ComputeProb will set those values for I == 0 and
684 // ignore the current values. If ComputeIdx > 0, then it expects those values
685 // to already be set for I == ComputeIdx - 1, and it will set them for I ==
686 // ComputeIdx.
687 auto AdjustProb = [&](unsigned ComputeIdx, double &ProbBefore,
688 double &ProbAfter, double &FreqBefore,
689 double &FreqAfter) {
690 assert(ComputeIdx < CondLatches.size() &&
691 "Expected valid CondLatches index");
692
693 // Compute or update ProbBefore, ProbAfter, FreqBefore, and FreqAfter.
694 auto ComputeAfter = [&]() {
695 ProbAfter = 1;
696 FreqAfter = IterCounts[ComputeIdx + 1];
697 for (unsigned I = ComputeIdx + 1, E = CondLatches.size(); I < E; ++I) {
698 double Prob = GetProb(I);
699 ProbAfter *= Prob;
700 // After Prob == 0, ProbAfter and FreqAfter won't change, so save time.
701 if (Prob == 0)
702 break;
703 FreqAfter += IterCounts[I + 1] * ProbAfter;
704 }
705 };
706 if (ComputeIdx == 0) {
707 ProbBefore = 1;
708 FreqBefore = IterCounts[0];
709 ComputeAfter();
710 } else {
711 // Rather than iterating all of CondLatches again, we fix up the
712 // previously computed values.
713 double ProbOld = GetProb(ComputeIdx);
714 if (ProbOld > 0) {
715 FreqAfter -= IterCounts[ComputeIdx] * ProbBefore;
716 ProbAfter /= ProbOld;
717 FreqAfter /= ProbOld;
718 } else {
719 // We cannot divide out the old zero probability. We short-circuited
720 // the iteration at that zero in the previous ComputeAfter call, so now
721 // we pick up where we left off.
722 ComputeAfter();
723 }
724 ProbBefore *= GetProb(ComputeIdx - 1);
725 FreqBefore += IterCounts[ComputeIdx] * ProbBefore;
726 }
727
728 // Compute the required probability, and limit it to a valid probability (0
729 // <= p <= 1). See the FreqCompute formula below for how to derive the
730 // ProbCompute formula.
731 double ProbReachingBackedge = CompletelyUnroll ? 0 : ProbBefore * ProbAfter;
732 double ProbComputeNumerator = FreqDesired - FreqBefore;
733 double ProbComputeDenominator =
734 FreqAfter + FreqDesired * ProbReachingBackedge;
735 double ProbCompute = -1; // Init expected to be unused.
736 if (ProbComputeNumerator <= 0) {
737 // FreqBefore has already reached or surpassed FreqDesired, so add no more
738 // frequency. It is possible that ProbComputeDenominator == 0 here
739 // because some latch probability (maybe the original) was set to zero, so
740 // this check avoids setting ProbCompute=1 (in the else if below) and
741 // division by zero where the numerator <= 0 (in the else below).
742 ProbCompute = 0;
743 } else if (ProbComputeDenominator == 0) {
744 // Analytically, this case seems impossible. It would occur if either:
745 // - Both FreqAfter and FreqDesired are zero. But the latter would cause
746 // ProbComputeNumerator < 0, which we catch above, and FreqDesired
747 // should always be >= 1 anyway.
748 // - There are no iterations after CondLatches[ComputeIdx], not even via
749 // a backedge, so that both FreqAfter and ProbReachingBackedge are zero.
750 // But iterations should exist after even the last conditional latch.
751 // - Some latch probability (maybe the original) was set to zero so that
752 // both FreqAfter and ProbReachingBackedge are zero. But that should
753 // not have happened because, according to the above
754 // ProbComputeNumerator check, we have not yet reached FreqDesired
755 // (which, if the original latch probability is zero, is just 1 and thus
756 // always reached or surpassed).
757 //
758 // Numerically, perhaps this case is possible. We interpret it to mean we
759 // need more frequency (ProbComputeNumerator > 0) but have no way to get
760 // any (ProbComputeDenominator is analytically too small to distinguish it
761 // from 0 in floating point), suggesting infinite probability is needed,
762 // but 1 is the maximum valid probability and thus the best we can do.
763 //
764 // TODO: Cover this case in the test suite if you can.
765 ProbCompute = 1;
766 } else {
767 ProbCompute = ProbComputeNumerator / ProbComputeDenominator;
768 ProbCompute = std::max(a: ProbCompute, b: 0.);
769 ProbCompute = std::min(a: ProbCompute, b: 1.);
770 }
771 SetProb(ComputeIdx, ProbCompute);
772
773 // Compute the resulting total frequency.
774 double FreqCompute = -1; // Init expected to be unused.
775 if (ProbReachingBackedge * ProbCompute == 1) {
776 // Analytically, this case seems impossible. It requires that there is a
777 // backedge and that FreqDesired == infinity so that every conditional
778 // latch's probability had to be set to 1. But FreqDesired == infinity
779 // means OriginalLoopProb.isOne(), which we guarded against earlier.
780 //
781 // Numerically, perhaps this case is possible. We interpret it to mean
782 // that analytically the probability has to be so near 1 that, in floating
783 // point, the frequency is computed as infinite.
784 //
785 // TODO: Cover this case in the test suite if you can.
786 FreqCompute = std::numeric_limits<double>::infinity();
787 if (ORE) {
788 ORE->emit(RemarkBuilder: [&]() {
789 return OptimizationRemark(DEBUG_TYPE, "InfiniteFrequency",
790 L->getStartLoc(), L->getHeader());
791 });
792 }
793 } else {
794 assert(FreqBefore > 0 &&
795 "Expected at least one iteration before first latch");
796 // In this equation, if we replace the left-hand side with FreqDesired and
797 // then solve for ProbCompute, we get the ProbCompute formula above.
798 FreqCompute = (FreqBefore + FreqAfter * ProbCompute) /
799 (1 - ProbReachingBackedge * ProbCompute);
800 }
801 assert(FreqCompute > 0 && "Expected valid frequency");
802 return FreqCompute;
803 };
804
805 // Determine and set branch weights.
806 //
807 // Prob < 0 and Prob > 1 cannot be represented as branch weights. We might
808 // compute such a Prob if FreqDesired is impossible (e.g., due to inaccurate
809 // profile data) for the maximum trip count we have determined when completely
810 // unrolling. In that case, so just go with whichever is closest.
811 if (CondLatches.size() == 1) {
812 SetAllProbs(ComputeProbForLinear());
813 } else if (CondLatches.size() == 2) {
814 SetAllProbs(ComputeProbForQuadratic());
815 } else if (!UnrollUniformWeights) {
816 // The polynomial is too complex for a simple formula, and the quick and
817 // dirty fix has been selected. Adjust probabilities starting from the
818 // first latch, which has the most influence on the total frequency, so
819 // starting there should minimize the number of latches that have to be
820 // visited. We do have to iterate because the first latch alone might not
821 // be enough. For example, we might need to set all probabilities to 1 if
822 // the frequency is the unroll factor.
823 double ProbBefore = -1, ProbAfter = -1; // Inits expected to be unused.
824 double FreqBefore = -1, FreqAfter = -1; // Inits expected to be unused.
825 for (unsigned I = 0; I != CondLatches.size(); ++I) {
826 double Freq = AdjustProb(I, ProbBefore, ProbAfter, FreqBefore, FreqAfter);
827 if (fabs(x: Freq - FreqDesired) / FreqDesired < FreqPrec)
828 break;
829 }
830 } else {
831 // The polynomial is too complex for a simple formula, and uniform branch
832 // weights have been selected, so bisect.
833 double ProbMin = -1, ProbMax = -1; // Inits expected to be unused.
834 double ProbPrev = -1; // Inits expected to be unused.
835 auto TryProb = [&](double Prob) {
836 ProbPrev = Prob;
837 double FreqDelta = ComputeFreq(Prob) - FreqDesired;
838 if (fabs(x: FreqDelta) / FreqDesired < FreqPrec)
839 return 0;
840 if (FreqDelta < 0) {
841 ProbMin = Prob;
842 return -1;
843 }
844 ProbMax = Prob;
845 return 1;
846 };
847 // If Prob == 0 is too small and Prob == 1 is too large, bisect between
848 // them. Accuracy (relative difference) is controlled by FreqPrec above.
849 // However, to place a hard upper limit on the search time, we stop
850 // bisecting when Prob stops changing (ProbDelta) by much (ProbPrec). In
851 // this case, we compute an absolute difference not a relative difference,
852 // which could produce more search time for smaller probabilities.
853 if (TryProb(0.) < 0 && TryProb(1.) > 0) {
854 assert(ProbMin == 0 && ProbMax == 1 &&
855 "expected probability bounds to be initialized");
856 const double ProbPrec = 1e-12;
857 double Prob, ProbDelta;
858 do {
859 Prob = (ProbMin + ProbMax) / 2;
860 ProbDelta = Prob - ProbPrev;
861 } while (TryProb(Prob) != 0 && fabs(x: ProbDelta) > ProbPrec);
862 }
863 SetAllProbs(ProbPrev);
864 }
865
866 // FIXME: We have not considered non-latch loop exits:
867 // - Their original probabilities are not considered in our calculation of
868 // FreqDesired.
869 // - Their probabilities are not considered in our probability model used to
870 // determine new probabilities for remaining conditional branches.
871 // - If they are conditional and LoopUnroll converts them to unconditional,
872 // LoopUnroll has proven their original probabilities are incorrect for some
873 // original loop iterations, but that does not cause ProbUpdateRequired to
874 // be set to true.
875 //
876 // To adjust FreqDesired and our probability model correctly for a non-latch
877 // loop exit, we would need to compute the original probability that the exit
878 // is reached from the loop header (in contrast, we currently assume that
879 // probability is 1 in the case of a latch exit) and the probability that the
880 // exit is taken if it is conditional (use the branch's old or new weights for
881 // FreqDesired or the probability model, respectively). Does computing the
882 // reaching probability require a CFG traversal, or is there some existing
883 // library that can do it? Prior discussions suggest some such libraries are
884 // difficult to use within LoopUnroll:
885 // <https://github.com/llvm/llvm-project/pull/164799#issuecomment-3438681519>.
886 // For now, we just let our corrected probabilities be less accurate in that
887 // scenario. Alternatively, we could refuse to correct probabilities at all
888 // in that scenario, but that seems worse.
889}
890
891/// Unroll the given loop by Count. The loop must be in LCSSA form. Unrolling
892/// can only fail when the loop's latch block is not terminated by a conditional
893/// branch instruction. However, if the trip count (and multiple) are not known,
894/// loop unrolling will mostly produce more code that is no faster.
895///
896/// If Runtime is true then UnrollLoop will try to insert a prologue or
897/// epilogue that ensures the latch has a trip multiple of Count. UnrollLoop
898/// will not runtime-unroll the loop if computing the run-time trip count will
899/// be expensive and AllowExpensiveTripCount is false.
900///
901/// The LoopInfo Analysis that is passed will be kept consistent.
902///
903/// This utility preserves LoopInfo. It will also preserve ScalarEvolution and
904/// DominatorTree if they are non-null.
905///
906/// If RemainderLoop is non-null, it will receive the remainder loop (if
907/// required and not fully unrolled).
908LoopUnrollResult
909llvm::UnrollLoop(Loop *L, UnrollLoopOptions ULO, LoopInfo *LI,
910 ScalarEvolution *SE, DominatorTree *DT, AssumptionCache *AC,
911 const TargetTransformInfo *TTI, OptimizationRemarkEmitter *ORE,
912 bool PreserveLCSSA, Loop **RemainderLoop, AAResults *AA) {
913 assert(DT && "DomTree is required");
914
915 if (!L->getLoopPreheader()) {
916 LLVM_DEBUG(dbgs() << " Can't unroll; loop preheader-insertion failed.\n");
917 return LoopUnrollResult::Unmodified;
918 }
919
920 if (!L->getLoopLatch()) {
921 LLVM_DEBUG(dbgs() << " Can't unroll; loop exit-block-insertion failed.\n");
922 return LoopUnrollResult::Unmodified;
923 }
924
925 // Loops with indirectbr cannot be cloned.
926 if (!L->isSafeToClone()) {
927 LLVM_DEBUG(dbgs() << " Can't unroll; Loop body cannot be cloned.\n");
928 return LoopUnrollResult::Unmodified;
929 }
930
931 if (L->getHeader()->hasAddressTaken()) {
932 // The loop-rotate pass can be helpful to avoid this in many cases.
933 LLVM_DEBUG(
934 dbgs() << " Won't unroll loop: address of header block is taken.\n");
935 return LoopUnrollResult::Unmodified;
936 }
937
938 assert(ULO.Count > 0);
939
940 // All these values should be taken only after peeling because they might have
941 // changed.
942 BasicBlock *Preheader = L->getLoopPreheader();
943 BasicBlock *Header = L->getHeader();
944 BasicBlock *LatchBlock = L->getLoopLatch();
945 SmallVector<BasicBlock *, 4> ExitBlocks;
946 L->getExitBlocks(ExitBlocks);
947
948 const unsigned MaxTripCount = SE->getSmallConstantMaxTripCount(L);
949 const bool MaxOrZero = SE->isBackedgeTakenCountMaxOrZero(L);
950 std::optional<unsigned> OriginalTripCount =
951 llvm::getLoopEstimatedTripCount(L);
952 BranchProbability OriginalLoopProb = llvm::getLoopProbability(L);
953
954 // Effectively "DCE" unrolled iterations that are beyond the max tripcount
955 // and will never be executed.
956 if (MaxTripCount && ULO.Count > MaxTripCount)
957 ULO.Count = MaxTripCount;
958
959 struct ExitInfo {
960 unsigned TripCount;
961 unsigned TripMultiple;
962 unsigned BreakoutTrip;
963 bool ExitOnTrue;
964 BasicBlock *FirstExitingBlock = nullptr;
965 SmallVector<BasicBlock *> ExitingBlocks;
966 };
967 MapVector<BasicBlock *, ExitInfo> ExitInfos;
968 SmallVector<BasicBlock *, 4> ExitingBlocks;
969 L->getExitingBlocks(ExitingBlocks);
970 for (auto *ExitingBlock : ExitingBlocks) {
971 // The folding code is not prepared to deal with non-branch instructions
972 // right now.
973 auto *BI = dyn_cast<CondBrInst>(Val: ExitingBlock->getTerminator());
974 if (!BI)
975 continue;
976
977 ExitInfo &Info = ExitInfos[ExitingBlock];
978 Info.TripCount = SE->getSmallConstantTripCount(L, ExitingBlock);
979 Info.TripMultiple = SE->getSmallConstantTripMultiple(L, ExitingBlock);
980 if (Info.TripCount != 0) {
981 Info.BreakoutTrip = Info.TripCount % ULO.Count;
982 Info.TripMultiple = 0;
983 } else {
984 Info.BreakoutTrip = Info.TripMultiple =
985 (unsigned)std::gcd(m: ULO.Count, n: Info.TripMultiple);
986 }
987 Info.ExitOnTrue = !L->contains(BB: BI->getSuccessor(i: 0));
988 Info.ExitingBlocks.push_back(Elt: ExitingBlock);
989 LLVM_DEBUG(dbgs() << " Exiting block %" << ExitingBlock->getName()
990 << ": TripCount=" << Info.TripCount
991 << ", TripMultiple=" << Info.TripMultiple
992 << ", BreakoutTrip=" << Info.BreakoutTrip << "\n");
993 }
994
995 // Are we eliminating the loop control altogether? Note that we can know
996 // we're eliminating the backedge without knowing exactly which iteration
997 // of the unrolled body exits.
998 const bool CompletelyUnroll = ULO.Count == MaxTripCount;
999
1000 const bool PreserveOnlyFirst = CompletelyUnroll && MaxOrZero;
1001
1002 // There's no point in performing runtime unrolling if this unroll count
1003 // results in a full unroll.
1004 if (CompletelyUnroll)
1005 ULO.Runtime = false;
1006
1007 // Go through all exits of L and see if there are any phi-nodes there. We just
1008 // conservatively assume that they're inserted to preserve LCSSA form, which
1009 // means that complete unrolling might break this form. We need to either fix
1010 // it in-place after the transformation, or entirely rebuild LCSSA. TODO: For
1011 // now we just recompute LCSSA for the outer loop, but it should be possible
1012 // to fix it in-place.
1013 bool NeedToFixLCSSA =
1014 PreserveLCSSA && CompletelyUnroll &&
1015 any_of(Range&: ExitBlocks,
1016 P: [](const BasicBlock *BB) { return isa<PHINode>(Val: BB->begin()); });
1017
1018 // The current loop unroll pass can unroll loops that have
1019 // (1) single latch; and
1020 // (2a) latch is unconditional; or
1021 // (2b) latch is conditional and is an exiting block
1022 // FIXME: The implementation can be extended to work with more complicated
1023 // cases, e.g. loops with multiple latches.
1024 Instruction *LatchTerm = LatchBlock->getTerminator();
1025
1026 // A conditional branch which exits the loop, which can be optimized to an
1027 // unconditional branch in the unrolled loop in some cases.
1028 bool LatchIsExiting = L->isLoopExiting(BB: LatchBlock);
1029 if (!isa<UncondBrInst>(Val: LatchTerm) &&
1030 !(isa<CondBrInst>(Val: LatchTerm) && LatchIsExiting)) {
1031 LLVM_DEBUG(
1032 dbgs() << "Can't unroll; a conditional latch must exit the loop");
1033 return LoopUnrollResult::Unmodified;
1034 }
1035
1036 bool EpilogProfitability =
1037 UnrollRuntimeEpilog.getNumOccurrences() ? UnrollRuntimeEpilog
1038 : isEpilogProfitable(L);
1039
1040 if (ULO.Runtime &&
1041 !UnrollRuntimeLoopRemainder(
1042 L, Count: ULO.Count, AllowExpensiveTripCount: ULO.AllowExpensiveTripCount, UseEpilogRemainder: EpilogProfitability,
1043 UnrollRemainder: ULO.UnrollRemainder, ForgetAllSCEV: ULO.ForgetAllSCEV, LI, SE, DT, AC, TTI,
1044 PreserveLCSSA, SCEVExpansionBudget: ULO.SCEVExpansionBudget, RuntimeUnrollMultiExit: ULO.RuntimeUnrollMultiExit,
1045 ResultLoop: RemainderLoop, OriginalTripCount, OriginalLoopProb)) {
1046 if (ULO.Force)
1047 ULO.Runtime = false;
1048 else {
1049 LLVM_DEBUG(dbgs() << "Won't unroll; remainder loop could not be "
1050 "generated when assuming runtime trip count\n");
1051 return LoopUnrollResult::Unmodified;
1052 }
1053 }
1054
1055 using namespace ore;
1056
1057 // Determine whether this loop originated from the vectorizer so we can
1058 // produce more informative remarks.
1059 StringRef LoopKind = getLoopVectorizeKindPrefix(L);
1060
1061 // Report the unrolling decision.
1062 if (CompletelyUnroll) {
1063 LLVM_DEBUG(dbgs() << "COMPLETELY UNROLLING loop %" << Header->getName()
1064 << " with trip count " << ULO.Count << "!\n");
1065 if (ORE)
1066 ORE->emit(RemarkBuilder: [&]() {
1067 return OptimizationRemark(DEBUG_TYPE, "FullyUnrolled", L->getStartLoc(),
1068 L->getHeader())
1069 << "completely unrolled " + LoopKind.str() + "loop with "
1070 << NV("UnrollCount", ULO.Count) << " iterations";
1071 });
1072 } else {
1073 LLVM_DEBUG({
1074 dbgs() << "UNROLLING loop %" << Header->getName() << " by " << ULO.Count;
1075 if (ULO.Runtime) {
1076 dbgs() << " with run-time trip count";
1077 if (ULO.UnrollRemainder)
1078 dbgs() << " (remainder unrolled)";
1079 }
1080 dbgs() << "!\n";
1081 });
1082
1083 if (ORE)
1084 ORE->emit(RemarkBuilder: [&]() {
1085 OptimizationRemark Diag(DEBUG_TYPE, "PartialUnrolled", L->getStartLoc(),
1086 L->getHeader());
1087 Diag << "unrolled " + LoopKind.str() + "loop by a factor of "
1088 << NV("UnrollCount", ULO.Count);
1089 if (ULO.Runtime)
1090 Diag << " with run-time trip count"
1091 << (ULO.UnrollRemainder ? " (remainder unrolled)" : "");
1092 return Diag;
1093 });
1094 }
1095
1096 // We are going to make changes to this loop. SCEV may be keeping cached info
1097 // about it, in particular about backedge taken count. The changes we make
1098 // are guaranteed to invalidate this information for our loop. It is tempting
1099 // to only invalidate the loop being unrolled, but it is incorrect as long as
1100 // all exiting branches from all inner loops have impact on the outer loops,
1101 // and if something changes inside them then any of outer loops may also
1102 // change. When we forget outermost loop, we also forget all contained loops
1103 // and this is what we need here.
1104 if (SE) {
1105 if (ULO.ForgetAllSCEV)
1106 SE->forgetAllLoops();
1107 else {
1108 SE->forgetTopmostLoop(L);
1109 SE->forgetBlockAndLoopDispositions();
1110 }
1111 }
1112
1113 if (!LatchIsExiting)
1114 ++NumUnrolledNotLatch;
1115
1116 // For the first iteration of the loop, we should use the precloned values for
1117 // PHI nodes. Insert associations now.
1118 ValueToValueMapTy LastValueMap;
1119 std::vector<PHINode*> OrigPHINode;
1120 for (BasicBlock::iterator I = Header->begin(); isa<PHINode>(Val: I); ++I) {
1121 OrigPHINode.push_back(x: cast<PHINode>(Val&: I));
1122 }
1123
1124 // Collect phi nodes for reductions for which we can introduce multiple
1125 // parallel reduction phis and compute the final reduction result after the
1126 // loop. This requires a single exit block after unrolling. This is ensured by
1127 // restricting to single-block loops where the unrolled iterations are known
1128 // to not exit.
1129 DenseMap<PHINode *, RecurrenceDescriptor> Reductions;
1130 bool CanAddAdditionalAccumulators =
1131 (UnrollAddParallelReductions.getNumOccurrences() > 0
1132 ? UnrollAddParallelReductions
1133 : ULO.AddAdditionalAccumulators) &&
1134 !CompletelyUnroll && L->getNumBlocks() == 1 &&
1135 (ULO.Runtime ||
1136 (ExitInfos.contains(Key: Header) && ((ExitInfos[Header].TripCount != 0 &&
1137 ExitInfos[Header].BreakoutTrip == 0))));
1138
1139 // Limit parallelizing reductions to unroll counts of 4 or less for now.
1140 // TODO: The number of parallel reductions should depend on the number of
1141 // execution units. We also don't have to add a parallel reduction phi per
1142 // unrolled iteration, but could for example add a parallel phi for every 2
1143 // unrolled iterations.
1144 if (CanAddAdditionalAccumulators && ULO.Count <= 4) {
1145 for (PHINode &Phi : Header->phis()) {
1146 auto RdxDesc = canParallelizeReductionWhenUnrolling(Phi, L, SE);
1147 if (!RdxDesc)
1148 continue;
1149
1150 // Only handle duplicate phis for a single reduction for now.
1151 // TODO: Handle any number of reductions
1152 if (!Reductions.empty())
1153 continue;
1154
1155 Reductions[&Phi] = *RdxDesc;
1156 }
1157 }
1158
1159 std::vector<BasicBlock *> Headers;
1160 std::vector<BasicBlock *> Latches;
1161 Headers.push_back(x: Header);
1162 Latches.push_back(x: LatchBlock);
1163
1164 // The current on-the-fly SSA update requires blocks to be processed in
1165 // reverse postorder so that LastValueMap contains the correct value at each
1166 // exit.
1167 LoopBlocksDFS DFS(L);
1168 DFS.perform(LI);
1169
1170 // Stash the DFS iterators before adding blocks to the loop.
1171 LoopBlocksDFS::RPOIterator BlockBegin = DFS.beginRPO();
1172 LoopBlocksDFS::RPOIterator BlockEnd = DFS.endRPO();
1173
1174 std::vector<BasicBlock*> UnrolledLoopBlocks = L->getBlocks();
1175
1176 // Snapshot the blocks to be cloned after remainder generation, which may
1177 // delete blocks from L, but before cloning adds new blocks to L.
1178 const std::vector<BasicBlock *> PostRemainderLoopBlocks = L->getBlocks();
1179
1180 // Loop Unrolling might create new loops. While we do preserve LoopInfo, we
1181 // might break loop-simplified form for these loops (as they, e.g., would
1182 // share the same exit blocks). We'll keep track of loops for which we can
1183 // break this so that later we can re-simplify them.
1184 SmallSetVector<Loop *, 4> LoopsToSimplify;
1185 LoopsToSimplify.insert_range(R&: *L);
1186
1187 // When a FSDiscriminator is enabled, we don't need to add the multiply
1188 // factors to the discriminators.
1189 if (Header->getParent()->shouldEmitDebugInfoForProfiling() &&
1190 !EnableFSDiscriminator)
1191 for (BasicBlock *BB : L->getBlocks())
1192 for (Instruction &I : *BB)
1193 if (!I.isDebugOrPseudoInst())
1194 if (const DILocation *DIL = I.getDebugLoc()) {
1195 auto NewDIL = DIL->cloneByMultiplyingDuplicationFactor(DF: ULO.Count);
1196 if (NewDIL)
1197 I.setDebugLoc(*NewDIL);
1198 else
1199 LLVM_DEBUG(dbgs()
1200 << "Failed to create new discriminator: "
1201 << DIL->getFilename() << " Line: " << DIL->getLine());
1202 }
1203
1204 // Identify what noalias metadata is inside the loop: if it is inside the
1205 // loop, the associated metadata must be cloned for each iteration.
1206 SmallVector<MDNode *, 6> LoopLocalNoAliasDeclScopes;
1207 identifyNoAliasScopesToClone(BBs: L->getBlocks(), NoAliasDeclScopes&: LoopLocalNoAliasDeclScopes);
1208
1209 // We place the unrolled iterations immediately after the original loop
1210 // latch. This is a reasonable default placement if we don't have block
1211 // frequencies, and if we do, well the layout will be adjusted later.
1212 auto BlockInsertPt = std::next(x: LatchBlock->getIterator());
1213 SmallVector<Instruction *> PartialReductions;
1214 for (unsigned It = 1; It != ULO.Count; ++It) {
1215 SmallVector<BasicBlock *, 8> NewBlocks;
1216 SmallDenseMap<const Loop *, Loop *, 4> NewLoops;
1217 NewLoops[L] = L;
1218
1219 for (LoopBlocksDFS::RPOIterator BB = BlockBegin; BB != BlockEnd; ++BB) {
1220 ValueToValueMapTy VMap;
1221 BasicBlock *New = CloneBasicBlock(BB: *BB, VMap, NameSuffix: "." + Twine(It));
1222 Header->getParent()->insert(Position: BlockInsertPt, BB: New);
1223
1224 assert((*BB != Header || LI->getLoopFor(*BB) == L) &&
1225 "Header should not be in a sub-loop");
1226 // Tell LI about New.
1227 const Loop *OldLoop = addClonedBlockToLoopInfo(OriginalBB: *BB, ClonedBB: New, LI, NewLoops);
1228 if (OldLoop)
1229 LoopsToSimplify.insert(X: NewLoops[OldLoop]);
1230
1231 if (*BB == Header) {
1232 // Loop over all of the PHI nodes in the block, changing them to use
1233 // the incoming values from the previous block.
1234 for (PHINode *OrigPHI : OrigPHINode) {
1235 PHINode *NewPHI = cast<PHINode>(Val&: VMap[OrigPHI]);
1236 Value *InVal = NewPHI->getIncomingValueForBlock(BB: LatchBlock);
1237
1238 // Use cloned phis as parallel phis for partial reductions, which will
1239 // get combined to the final reduction result after the loop.
1240 if (Reductions.contains(Val: OrigPHI)) {
1241 // Collect partial reduction results.
1242 if (PartialReductions.empty())
1243 PartialReductions.push_back(Elt: cast<Instruction>(Val: InVal));
1244 PartialReductions.push_back(Elt: cast<Instruction>(Val&: VMap[InVal]));
1245
1246 // Update the start value for the cloned phis to use the identity
1247 // value for the reduction.
1248 const RecurrenceDescriptor &RdxDesc = Reductions[OrigPHI];
1249 NewPHI->setIncomingValueForBlock(
1250 BB: L->getLoopPreheader(),
1251 V: getRecurrenceIdentity(K: RdxDesc.getRecurrenceKind(),
1252 Tp: OrigPHI->getType(),
1253 FMF: RdxDesc.getFastMathFlags()));
1254
1255 // Update NewPHI to use the cloned value for the iteration and move
1256 // to header.
1257 NewPHI->replaceUsesOfWith(From: InVal, To: VMap[InVal]);
1258 NewPHI->moveBefore(InsertPos: OrigPHI->getIterator());
1259 continue;
1260 }
1261
1262 if (Instruction *InValI = dyn_cast<Instruction>(Val: InVal))
1263 if (It > 1 && L->contains(Inst: InValI))
1264 InVal = LastValueMap[InValI];
1265 VMap[OrigPHI] = InVal;
1266 NewPHI->eraseFromParent();
1267 }
1268
1269 // Eliminate copies of the loop heart intrinsic, if any.
1270 if (ULO.Heart) {
1271 auto it = VMap.find(Val: ULO.Heart);
1272 assert(it != VMap.end());
1273 Instruction *heartCopy = cast<Instruction>(Val&: it->second);
1274 heartCopy->eraseFromParent();
1275 VMap.erase(I: it);
1276 }
1277 }
1278
1279 // Remap source location atom instance. Do this now, rather than
1280 // when we remap instructions, because remap is called once we've
1281 // cloned all blocks (all the clones would get the same atom
1282 // number).
1283 if (!VMap.AtomMap.empty())
1284 for (Instruction &I : *New)
1285 RemapSourceAtom(I: &I, VM&: VMap);
1286
1287 // Update our running map of newest clones
1288 LastValueMap[*BB] = New;
1289 for (ValueToValueMapTy::iterator VI = VMap.begin(), VE = VMap.end();
1290 VI != VE; ++VI)
1291 LastValueMap[VI->first] = VI->second;
1292
1293 // Add phi entries for newly created values to all exit blocks.
1294 for (BasicBlock *Succ : successors(BB: *BB)) {
1295 if (L->contains(BB: Succ))
1296 continue;
1297 for (PHINode &PHI : Succ->phis()) {
1298 Value *Incoming = PHI.getIncomingValueForBlock(BB: *BB);
1299 ValueToValueMapTy::iterator It = LastValueMap.find(Val: Incoming);
1300 if (It != LastValueMap.end())
1301 Incoming = It->second;
1302 PHI.addIncoming(V: Incoming, BB: New);
1303 SE->forgetLcssaPhiWithNewPredecessor(L, V: &PHI);
1304 }
1305 }
1306 // Keep track of new headers and latches as we create them, so that
1307 // we can insert the proper branches later.
1308 if (*BB == Header)
1309 Headers.push_back(x: New);
1310 if (*BB == LatchBlock)
1311 Latches.push_back(x: New);
1312
1313 // Keep track of the exiting block and its successor block contained in
1314 // the loop for the current iteration.
1315 auto ExitInfoIt = ExitInfos.find(Key: *BB);
1316 if (ExitInfoIt != ExitInfos.end())
1317 ExitInfoIt->second.ExitingBlocks.push_back(Elt: New);
1318
1319 NewBlocks.push_back(Elt: New);
1320 UnrolledLoopBlocks.push_back(x: New);
1321
1322 // Update DomTree: since we just copy the loop body, and each copy has a
1323 // dedicated entry block (copy of the header block), this header's copy
1324 // dominates all copied blocks. That means, dominance relations in the
1325 // copied body are the same as in the original body.
1326 if (*BB == Header)
1327 DT->addNewBlock(BB: New, DomBB: Latches[It - 1]);
1328 else {
1329 auto BBDomNode = DT->getNode(BB: *BB);
1330 auto BBIDom = BBDomNode->getIDom();
1331 BasicBlock *OriginalBBIDom = BBIDom->getBlock();
1332 DT->addNewBlock(
1333 BB: New, DomBB: cast<BasicBlock>(Val&: LastValueMap[cast<Value>(Val: OriginalBBIDom)]));
1334 }
1335 }
1336
1337 // Remap all instructions in the most recent iteration.
1338 // Key Instructions: Nothing to do - we've already remapped the atoms.
1339 remapInstructionsInBlocks(Blocks: NewBlocks, VMap&: LastValueMap);
1340 for (BasicBlock *NewBlock : NewBlocks)
1341 for (Instruction &I : *NewBlock)
1342 if (auto *II = dyn_cast<AssumeInst>(Val: &I))
1343 AC->registerAssumption(CI: II);
1344
1345 {
1346 // Identify what other metadata depends on the cloned version. After
1347 // cloning, replace the metadata with the corrected version for both
1348 // memory instructions and noalias intrinsics.
1349 std::string ext = (Twine("It") + Twine(It)).str();
1350 cloneAndAdaptNoAliasScopes(NoAliasDeclScopes: LoopLocalNoAliasDeclScopes, NewBlocks,
1351 Context&: Header->getContext(), Ext: ext);
1352 }
1353 }
1354
1355 // Loop over the PHI nodes in the original block, setting incoming values.
1356 for (PHINode *PN : OrigPHINode) {
1357 if (CompletelyUnroll) {
1358 // The RAUW below disconnects the original PHI from its users.
1359 // Invalidate cached SCEVs while the def-use chain is still intact.
1360 if (SE)
1361 SE->forgetValue(V: PN);
1362 PN->replaceAllUsesWith(V: PN->getIncomingValueForBlock(BB: Preheader));
1363 PN->eraseFromParent();
1364 } else if (ULO.Count > 1) {
1365 if (Reductions.contains(Val: PN))
1366 continue;
1367
1368 Value *InVal = PN->removeIncomingValue(BB: LatchBlock, DeletePHIIfEmpty: false);
1369 // If this value was defined in the loop, take the value defined by the
1370 // last iteration of the loop.
1371 if (Instruction *InValI = dyn_cast<Instruction>(Val: InVal)) {
1372 if (L->contains(Inst: InValI))
1373 InVal = LastValueMap[InVal];
1374 }
1375 assert(Latches.back() == LastValueMap[LatchBlock] && "bad last latch");
1376 PN->addIncoming(V: InVal, BB: Latches.back());
1377 }
1378 }
1379
1380 // Connect latches of the unrolled iterations to the headers of the next
1381 // iteration. Currently they point to the header of the same iteration.
1382 for (unsigned i = 0, e = Latches.size(); i != e; ++i) {
1383 unsigned j = (i + 1) % e;
1384 Latches[i]->getTerminator()->replaceSuccessorWith(OldBB: Headers[i], NewBB: Headers[j]);
1385 }
1386
1387 // Remove loop metadata copied from the original loop latch to branches that
1388 // are no longer latches.
1389 for (unsigned I = 0, E = Latches.size() - (CompletelyUnroll ? 0 : 1); I < E;
1390 ++I)
1391 Latches[I]->getTerminator()->setMetadata(KindID: LLVMContext::MD_loop, Node: nullptr);
1392
1393 // Update dominators of blocks we might reach through exits.
1394 // Immediate dominator of such block might change, because we add more
1395 // routes which can lead to the exit: we can now reach it from the copied
1396 // iterations too.
1397 if (ULO.Count > 1) {
1398 for (auto *BB : PostRemainderLoopBlocks) {
1399 auto *BBDomNode = DT->getNode(BB);
1400 SmallVector<BasicBlock *, 16> ChildrenToUpdate;
1401 for (auto *ChildDomNode : BBDomNode->children()) {
1402 auto *ChildBB = ChildDomNode->getBlock();
1403 if (!L->contains(BB: ChildBB))
1404 ChildrenToUpdate.push_back(Elt: ChildBB);
1405 }
1406 // The new idom of the block will be the nearest common dominator
1407 // of all copies of the previous idom. This is equivalent to the
1408 // nearest common dominator of the previous idom and the first latch,
1409 // which dominates all copies of the previous idom.
1410 BasicBlock *NewIDom = DT->findNearestCommonDominator(A: BB, B: LatchBlock);
1411 for (auto *ChildBB : ChildrenToUpdate)
1412 DT->changeImmediateDominator(BB: ChildBB, NewBB: NewIDom);
1413 }
1414 }
1415
1416 assert(!UnrollVerifyDomtree ||
1417 DT->verify(DominatorTree::VerificationLevel::Fast));
1418
1419 SmallVector<DominatorTree::UpdateType> DTUpdates;
1420 auto SetDest = [&](BasicBlock *Src, bool WillExit, bool ExitOnTrue) {
1421 auto *Term = cast<CondBrInst>(Val: Src->getTerminator());
1422 const unsigned Idx = ExitOnTrue ^ WillExit;
1423 BasicBlock *Dest = Term->getSuccessor(i: Idx);
1424 BasicBlock *DeadSucc = Term->getSuccessor(i: 1-Idx);
1425
1426 // Remove predecessors from all non-Dest successors.
1427 DeadSucc->removePredecessor(Pred: Src, /* KeepOneInputPHIs */ true);
1428
1429 // Replace the conditional branch with an unconditional one.
1430 auto *BI = UncondBrInst::Create(Target: Dest, InsertBefore: Term->getIterator());
1431 BI->setDebugLoc(Term->getDebugLoc());
1432 Term->eraseFromParent();
1433
1434 DTUpdates.emplace_back(Args: DominatorTree::Delete, Args&: Src, Args&: DeadSucc);
1435 };
1436
1437 auto WillExit = [&](const ExitInfo &Info, unsigned i, unsigned j,
1438 bool IsLatch) -> std::optional<bool> {
1439 if (CompletelyUnroll) {
1440 if (PreserveOnlyFirst) {
1441 if (i == 0)
1442 return std::nullopt;
1443 return j == 0;
1444 }
1445 // Complete (but possibly inexact) unrolling
1446 if (j == 0)
1447 return true;
1448 if (Info.TripCount && j != Info.TripCount)
1449 return false;
1450 return std::nullopt;
1451 }
1452
1453 if (ULO.Runtime) {
1454 // If runtime unrolling inserts a prologue, information about non-latch
1455 // exits may be stale.
1456 if (IsLatch && j != 0)
1457 return false;
1458 return std::nullopt;
1459 }
1460
1461 if (j != Info.BreakoutTrip &&
1462 (Info.TripMultiple == 0 || j % Info.TripMultiple != 0)) {
1463 // If we know the trip count or a multiple of it, we can safely use an
1464 // unconditional branch for some iterations.
1465 return false;
1466 }
1467 return std::nullopt;
1468 };
1469
1470 // Fold branches for iterations where we know that they will exit or not
1471 // exit. In the case of an iteration's latch, if we thus find
1472 // *OriginalLoopProb is incorrect, set ProbUpdateRequired to true.
1473 bool ProbUpdateRequired = false;
1474 for (auto &Pair : ExitInfos) {
1475 ExitInfo &Info = Pair.second;
1476 for (unsigned i = 0, e = Info.ExitingBlocks.size(); i != e; ++i) {
1477 // The branch destination.
1478 unsigned j = (i + 1) % e;
1479 bool IsLatch = Pair.first == LatchBlock;
1480 std::optional<bool> KnownWillExit = WillExit(Info, i, j, IsLatch);
1481 if (!KnownWillExit) {
1482 if (!Info.FirstExitingBlock)
1483 Info.FirstExitingBlock = Info.ExitingBlocks[i];
1484 continue;
1485 }
1486
1487 // We don't fold known-exiting branches for non-latch exits here,
1488 // because this ensures that both all loop blocks and all exit blocks
1489 // remain reachable in the CFG.
1490 // TODO: We could fold these branches, but it would require much more
1491 // sophisticated updates to LoopInfo.
1492 if (*KnownWillExit && !IsLatch) {
1493 if (!Info.FirstExitingBlock)
1494 Info.FirstExitingBlock = Info.ExitingBlocks[i];
1495 continue;
1496 }
1497
1498 // For a latch, record any OriginalLoopProb contradiction.
1499 if (!OriginalLoopProb.isUnknown() && IsLatch) {
1500 BranchProbability ActualProb = *KnownWillExit
1501 ? BranchProbability::getZero()
1502 : BranchProbability::getOne();
1503 ProbUpdateRequired |= OriginalLoopProb != ActualProb;
1504 }
1505
1506 SetDest(Info.ExitingBlocks[i], *KnownWillExit, Info.ExitOnTrue);
1507 }
1508 }
1509
1510 DomTreeUpdater DTU(DT, DomTreeUpdater::UpdateStrategy::Lazy);
1511 DomTreeUpdater *DTUToUse = &DTU;
1512 if (ExitingBlocks.size() == 1 && ExitInfos.size() == 1) {
1513 // Manually update the DT if there's a single exiting node. In that case
1514 // there's a single exit node and it is sufficient to update the nodes
1515 // immediately dominated by the original exiting block. They will become
1516 // dominated by the first exiting block that leaves the loop after
1517 // unrolling. Note that the CFG inside the loop does not change, so there's
1518 // no need to update the DT inside the unrolled loop.
1519 DTUToUse = nullptr;
1520 auto &[OriginalExit, Info] = *ExitInfos.begin();
1521 if (!Info.FirstExitingBlock)
1522 Info.FirstExitingBlock = Info.ExitingBlocks.back();
1523 for (auto *C : to_vector(Range: DT->getNode(BB: OriginalExit)->children())) {
1524 if (L->contains(BB: C->getBlock()))
1525 continue;
1526 C->setIDom(DT->getNode(BB: Info.FirstExitingBlock));
1527 }
1528 } else {
1529 DTU.applyUpdates(Updates: DTUpdates);
1530 }
1531
1532 // When completely unrolling, the last latch becomes unreachable.
1533 if (!LatchIsExiting && CompletelyUnroll) {
1534 // There is no need to update the DT here, because there must be a unique
1535 // latch. Hence if the latch is not exiting it must directly branch back to
1536 // the original loop header and does not dominate any nodes.
1537 assert(LatchBlock->getSingleSuccessor() && "Loop with multiple latches?");
1538 changeToUnreachable(I: Latches.back()->getTerminator(), PreserveLCSSA);
1539 }
1540
1541 // After merging adjacent blocks in Latches below:
1542 // - CondLatches will list the blocks from Latches that are still terminated
1543 // with conditional branches.
1544 // - For 1 <= I < CondLatches.size(), IterCounts[I] will store the number of
1545 // the original loop iterations through which control flows from
1546 // CondLatches[I-1] to CondLatches[I].
1547 // - For I == 0 or I == CondLatches.size(), IterCounts[I] will store the
1548 // number of the original loop iterations through which control can flow
1549 // before CondLatches.front() or after CondLatches.back(), respectively,
1550 // without taking the unrolled loop's backedge, if any.
1551 // - CondLatchNexts[I] will store the CondLatches[I] branch target for the
1552 // next of the original loop's iterations (as opposed to the exit target).
1553 assert(ULO.Count == Latches.size() &&
1554 "Expected one latch block per unrolled iteration");
1555 std::vector<unsigned> IterCounts(1, 0);
1556 std::vector<BasicBlock *> CondLatches;
1557 std::vector<BasicBlock *> CondLatchNexts;
1558 IterCounts.reserve(n: Latches.size() + 1);
1559 CondLatches.reserve(n: Latches.size());
1560 CondLatchNexts.reserve(n: Latches.size());
1561
1562 // Merge adjacent basic blocks, if possible.
1563 for (auto [I, Latch] : enumerate(First&: Latches)) {
1564 ++IterCounts.back();
1565 assert((isa<UncondBrInst, CondBrInst>(Latch->getTerminator()) ||
1566 (CompletelyUnroll && !LatchIsExiting && Latch == Latches.back())) &&
1567 "Need a branch as terminator, except when fully unrolling with "
1568 "unconditional latch");
1569 if (auto *Term = dyn_cast<UncondBrInst>(Val: Latch->getTerminator())) {
1570 BasicBlock *Dest = Term->getSuccessor();
1571 BasicBlock *Fold = Dest->getUniquePredecessor();
1572 if (MergeBlockIntoPredecessor(BB: Dest, /*DTU=*/DTUToUse, LI,
1573 /*MSSAU=*/nullptr, /*MemDep=*/nullptr,
1574 /*PredecessorWithTwoSuccessors=*/false,
1575 DT: DTUToUse ? nullptr : DT)) {
1576 // Dest has been folded into Fold. Update our worklists accordingly.
1577 llvm::replace(Range&: Latches, OldValue: Dest, NewValue: Fold);
1578 llvm::erase(C&: UnrolledLoopBlocks, V: Dest);
1579 }
1580 } else if (isa<CondBrInst>(Val: Latch->getTerminator())) {
1581 IterCounts.push_back(x: 0);
1582 CondLatches.push_back(x: Latch);
1583 CondLatchNexts.push_back(x: Headers[(I + 1) % Latches.size()]);
1584 }
1585 }
1586
1587 // Fix probabilities we contradicted above.
1588 if (ProbUpdateRequired) {
1589 fixProbContradiction(L, ULO, ORE, OriginalLoopProb, CompletelyUnroll,
1590 IterCounts, CondLatches, CondLatchNexts);
1591 }
1592
1593 // If there are partial reductions, create code in the exit block to compute
1594 // the final result and update users of the final result.
1595 if (!PartialReductions.empty()) {
1596 BasicBlock *ExitBlock = L->getExitBlock();
1597 assert(ExitBlock &&
1598 "Can only introduce parallel reduction phis with single exit block");
1599 assert(Reductions.size() == 1 &&
1600 "currently only a single reduction is supported");
1601 Value *FinalRdxValue = PartialReductions.back();
1602 Value *RdxResult = nullptr;
1603 for (PHINode &Phi : ExitBlock->phis()) {
1604 if (Phi.getIncomingValueForBlock(BB: L->getLoopLatch()) != FinalRdxValue)
1605 continue;
1606 if (!RdxResult) {
1607 RdxResult = PartialReductions.front();
1608 IRBuilder Builder(ExitBlock->getFirstNonPHIIt());
1609 Builder.setFastMathFlags(Reductions.begin()->second.getFastMathFlags());
1610 RecurKind RK = Reductions.begin()->second.getRecurrenceKind();
1611 for (Instruction *RdxPart : drop_begin(RangeOrContainer&: PartialReductions)) {
1612 if (RecurrenceDescriptor::isMinMaxRecurrenceKind(Kind: RK))
1613 RdxResult = createMinMaxOp(Builder, RK, Left: RdxResult, Right: RdxPart);
1614 else
1615 RdxResult = Builder.CreateBinOp(
1616 Opc: (Instruction::BinaryOps)RecurrenceDescriptor::getOpcode(Kind: RK),
1617 LHS: RdxPart, RHS: RdxResult, Name: "bin.rdx");
1618 }
1619 NeedToFixLCSSA = true;
1620 for (Instruction *RdxPart : PartialReductions)
1621 RdxPart->dropPoisonGeneratingFlags();
1622 }
1623
1624 Phi.replaceAllUsesWith(V: RdxResult);
1625 }
1626 }
1627
1628 if (DTUToUse) {
1629 // Apply updates to the DomTree.
1630 DT = &DTU.getDomTree();
1631 }
1632 assert(!UnrollVerifyDomtree ||
1633 DT->verify(DominatorTree::VerificationLevel::Fast));
1634
1635 Loop *OuterL = L->getParentLoop();
1636 std::vector<BasicBlock *> Blocks;
1637 // Update LoopInfo if the loop is completely removed.
1638 if (CompletelyUnroll) {
1639 Blocks = L->getBlocks();
1640 LI->erase(L);
1641 // We shouldn't try to use `L` anymore.
1642 L = nullptr;
1643 }
1644
1645 // At this point, the code is well formed. We now simplify the unrolled loop,
1646 // doing constant propagation and dead code elimination as we go.
1647 simplifyLoopAfterUnroll(
1648 L, SimplifyIVs: !CompletelyUnroll && ULO.Count > 1, LI, SE, DT, AC, TTI,
1649 Blocks: CompletelyUnroll ? ArrayRef<BasicBlock *>(Blocks) : L->getBlocks(), AA);
1650
1651 NumCompletelyUnrolled += CompletelyUnroll;
1652 ++NumUnrolled;
1653
1654 if (!CompletelyUnroll) {
1655 // Update metadata for the loop's branch weights and estimated trip count:
1656 // - If ULO.Runtime, UnrollRuntimeLoopRemainder sets the guard branch
1657 // weights, latch branch weights, and estimated trip count of the
1658 // remainder loop it creates. It also sets the branch weights for the
1659 // unrolled loop guard it creates. The branch weights for the unrolled
1660 // loop latch are adjusted below. FIXME: Handle prologue loops.
1661 // - Otherwise, if unrolled loop iteration latches become unconditional,
1662 // branch weights are adjusted by the fixProbContradiction call above.
1663 // - Otherwise, the original loop's branch weights are correct for the
1664 // unrolled loop, so do not adjust them.
1665 // - In all cases, the unrolled loop's estimated trip count is set below.
1666 //
1667 // As an example of the last case, consider what happens if the unroll count
1668 // is 4 for a loop with an estimated trip count of 10 when we do not create
1669 // a remainder loop and all iterations' latches remain conditional. Each
1670 // unrolled iteration's latch still has the same probability of exiting the
1671 // loop as it did when in the original loop, and thus it should still have
1672 // the same branch weights. Each unrolled iteration's non-zero probability
1673 // of exiting already appropriately reduces the probability of reaching the
1674 // remaining iterations just as it did in the original loop. Trying to also
1675 // adjust the branch weights of the final unrolled iteration's latch (i.e.,
1676 // the backedge for the unrolled loop as a whole) to reflect its new trip
1677 // count of 3 will erroneously further reduce its block frequencies.
1678 // However, in case an analysis later needs to estimate the trip count of
1679 // the unrolled loop as a whole without considering the branch weights for
1680 // each unrolled iteration's latch within it, we store the new trip count as
1681 // separate metadata.
1682 if (!OriginalLoopProb.isUnknown() && ULO.Runtime && EpilogProfitability) {
1683 assert((CondLatches.size() == 1 &&
1684 (ProbUpdateRequired || OriginalLoopProb.isOne())) &&
1685 "Expected ULO.Runtime to give unrolled loop 1 conditional latch, "
1686 "the backedge, requiring a probability update unless infinite");
1687 // Where p is always the probability of executing at least 1 more
1688 // iteration, the probability for at least n more iterations is p^n.
1689 setLoopProbability(L, P: OriginalLoopProb.pow(N: ULO.Count));
1690 }
1691 if (OriginalTripCount) {
1692 unsigned NewTripCount = *OriginalTripCount / ULO.Count;
1693 if (!ULO.Runtime && *OriginalTripCount % ULO.Count)
1694 ++NewTripCount;
1695 setLoopEstimatedTripCount(L, EstimatedTripCount: NewTripCount);
1696 }
1697 }
1698
1699 // LoopInfo should not be valid, confirm that.
1700 if (UnrollVerifyLoopInfo)
1701 LI->verify();
1702
1703 // After complete unrolling most of the blocks should be contained in OuterL.
1704 // However, some of them might happen to be out of OuterL (e.g. if they
1705 // precede a loop exit). In this case we might need to insert PHI nodes in
1706 // order to preserve LCSSA form.
1707 // We don't need to check this if we already know that we need to fix LCSSA
1708 // form.
1709 // TODO: For now we just recompute LCSSA for the outer loop in this case, but
1710 // it should be possible to fix it in-place.
1711 if (PreserveLCSSA && OuterL && CompletelyUnroll && !NeedToFixLCSSA)
1712 NeedToFixLCSSA |= ::needToInsertPhisForLCSSA(L: OuterL, Blocks: UnrolledLoopBlocks, LI);
1713
1714 // Make sure that loop-simplify form is preserved. We want to simplify
1715 // at least one layer outside of the loop that was unrolled so that any
1716 // changes to the parent loop exposed by the unrolling are considered.
1717 if (OuterL) {
1718 // OuterL includes all loops for which we can break loop-simplify, so
1719 // it's sufficient to simplify only it (it'll recursively simplify inner
1720 // loops too).
1721 if (NeedToFixLCSSA) {
1722 // LCSSA must be performed on the outermost affected loop. The unrolled
1723 // loop's last loop latch is guaranteed to be in the outermost loop
1724 // after LoopInfo's been updated by LoopInfo::erase.
1725 Loop *LatchLoop = LI->getLoopFor(BB: Latches.back());
1726 Loop *FixLCSSALoop = OuterL;
1727 if (!FixLCSSALoop->contains(L: LatchLoop))
1728 while (FixLCSSALoop->getParentLoop() != LatchLoop)
1729 FixLCSSALoop = FixLCSSALoop->getParentLoop();
1730
1731 formLCSSARecursively(L&: *FixLCSSALoop, DT: *DT, LI, SE);
1732 } else if (PreserveLCSSA) {
1733 assert(OuterL->isLCSSAForm(*DT) &&
1734 "Loops should be in LCSSA form after loop-unroll.");
1735 }
1736
1737 // TODO: That potentially might be compile-time expensive. We should try
1738 // to fix the loop-simplified form incrementally.
1739 simplifyLoop(L: OuterL, DT, LI, SE, AC, MSSAU: nullptr, PreserveLCSSA);
1740 } else {
1741 // Simplify loops for which we might've broken loop-simplify form.
1742 for (Loop *SubLoop : LoopsToSimplify)
1743 simplifyLoop(L: SubLoop, DT, LI, SE, AC, MSSAU: nullptr, PreserveLCSSA);
1744 }
1745
1746 return CompletelyUnroll ? LoopUnrollResult::FullyUnrolled
1747 : LoopUnrollResult::PartiallyUnrolled;
1748}
1749
1750/// Given an llvm.loop loop id metadata node, returns the loop hint metadata
1751/// node with the given name (for example, "llvm.loop.unroll.count"). If no
1752/// such metadata node exists, then nullptr is returned.
1753MDNode *llvm::GetUnrollMetadata(MDNode *LoopID, StringRef Name) {
1754 // First operand should refer to the loop id itself.
1755 assert(LoopID->getNumOperands() > 0 && "requires at least one operand");
1756 assert(LoopID->getOperand(0) == LoopID && "invalid loop id");
1757
1758 for (MDNode *MD :
1759 make_isa_range<MDNode>(Range: llvm::drop_begin(RangeOrContainer: LoopID->operands()))) {
1760 MDString *S = dyn_cast<MDString>(Val: MD->getOperand(I: 0));
1761 if (!S)
1762 continue;
1763
1764 if (Name == S->getString())
1765 return MD;
1766 }
1767 return nullptr;
1768}
1769
1770// Returns the loop hint metadata node with the given name (for example,
1771// "llvm.loop.unroll.count"). If no such metadata node exists, then nullptr is
1772// returned.
1773MDNode *llvm::getUnrollMetadataForLoop(const Loop *L, StringRef Name) {
1774 if (MDNode *LoopID = L->getLoopID())
1775 return GetUnrollMetadata(LoopID, Name);
1776 return nullptr;
1777}
1778
1779std::optional<RecurrenceDescriptor>
1780llvm::canParallelizeReductionWhenUnrolling(PHINode &Phi, Loop *L,
1781 ScalarEvolution *SE) {
1782 RecurrenceDescriptor RdxDesc;
1783 if (!RecurrenceDescriptor::isReductionPHI(Phi: &Phi, TheLoop: L, RedDes&: RdxDesc,
1784 /*DemandedBits=*/DB: nullptr,
1785 /*AC=*/nullptr, /*DT=*/nullptr, SE))
1786 return std::nullopt;
1787 if (RdxDesc.hasUsesOutsideReductionChain())
1788 return std::nullopt;
1789 RecurKind RK = RdxDesc.getRecurrenceKind();
1790 static const auto ValidRKs = {
1791 RecurKind::Add, RecurKind::Mul, RecurKind::Or,
1792 RecurKind::And, RecurKind::Xor, RecurKind::SMin,
1793 RecurKind::SMax, RecurKind::UMin, RecurKind::UMax,
1794 RecurKind::FAdd, RecurKind::FMul, RecurKind::FMin,
1795 RecurKind::FMax, RecurKind::FMinNum, RecurKind::FMaxNum,
1796 RecurKind::FMinimum, RecurKind::FMaximum, RecurKind::FMinimumNum,
1797 RecurKind::FMaximumNum, RecurKind::FMulAdd};
1798 // Skip unsupported reductions, including sub, any-of and find-last.
1799 // TODO: Handle sub, any-of and find-last reductions.
1800 if (!any_of(Range: ValidRKs, P: equal_to(Arg&: RK)))
1801 return std::nullopt;
1802
1803 if (RdxDesc.hasExactFPMath())
1804 return std::nullopt;
1805
1806 if (RdxDesc.IntermediateStore)
1807 return std::nullopt;
1808
1809 BasicBlock *Latch = L->getLoopLatch();
1810 if (!Latch)
1811 return std::nullopt;
1812 Instruction *LatchInst =
1813 cast<Instruction>(Val: Phi.getIncomingValueForBlock(BB: Latch));
1814 // Don't unroll reductions with constant ops; those can be folded to a
1815 // single induction update. For calls (e.g. fmuladd or min/max
1816 // intrinsics), the called function is itself a Constant operand and is
1817 // not a reduction operand, so restrict the check to the argument list.
1818 auto Ops = isa<CallBase>(Val: LatchInst) ? cast<CallBase>(Val: LatchInst)->args()
1819 : LatchInst->operands();
1820 if (any_of(Range&: Ops, P: IsaPred<Constant>))
1821 return std::nullopt;
1822
1823 if (!is_contained(Range: LatchInst->operands(), Element: &Phi))
1824 return std::nullopt;
1825
1826 return RdxDesc;
1827}
1828