LLVM 24.0.0git
ScalarEvolution.cpp
Go to the documentation of this file.
1//===- ScalarEvolution.cpp - Scalar Evolution Analysis --------------------===//
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 contains the implementation of the scalar evolution analysis
10// engine, which is used primarily to analyze expressions involving induction
11// variables in loops.
12//
13// There are several aspects to this library. First is the representation of
14// scalar expressions, which are represented as subclasses of the SCEV class.
15// These classes are used to represent certain types of subexpressions that we
16// can handle. We only create one SCEV of a particular shape, so
17// pointer-comparisons for equality are legal.
18//
19// One important aspect of the SCEV objects is that they are never cyclic, even
20// if there is a cycle in the dataflow for an expression (ie, a PHI node). If
21// the PHI node is one of the idioms that we can represent (e.g., a polynomial
22// recurrence) then we represent it directly as a recurrence node, otherwise we
23// represent it as a SCEVUnknown node.
24//
25// In addition to being able to represent expressions of various types, we also
26// have folders that are used to build the *canonical* representation for a
27// particular expression. These folders are capable of using a variety of
28// rewrite rules to simplify the expressions.
29//
30// Once the folders are defined, we can implement the more interesting
31// higher-level code, such as the code that recognizes PHI nodes of various
32// types, computes the execution count of a loop, etc.
33//
34// TODO: We should use these routines and value representations to implement
35// dependence analysis!
36//
37//===----------------------------------------------------------------------===//
38//
39// There are several good references for the techniques used in this analysis.
40//
41// Chains of recurrences -- a method to expedite the evaluation
42// of closed-form functions
43// Olaf Bachmann, Paul S. Wang, Eugene V. Zima
44//
45// On computational properties of chains of recurrences
46// Eugene V. Zima
47//
48// Symbolic Evaluation of Chains of Recurrences for Loop Optimization
49// Robert A. van Engelen
50//
51// Efficient Symbolic Analysis for Optimizing Compilers
52// Robert A. van Engelen
53//
54// Using the chains of recurrences algebra for data dependence testing and
55// induction variable substitution
56// MS Thesis, Johnie Birch
57//
58//===----------------------------------------------------------------------===//
59
61#include "llvm/ADT/APInt.h"
62#include "llvm/ADT/ArrayRef.h"
63#include "llvm/ADT/DenseMap.h"
65#include "llvm/ADT/FoldingSet.h"
66#include "llvm/ADT/STLExtras.h"
67#include "llvm/ADT/ScopeExit.h"
68#include "llvm/ADT/Sequence.h"
71#include "llvm/ADT/Statistic.h"
73#include "llvm/ADT/StringRef.h"
83#include "llvm/Config/llvm-config.h"
84#include "llvm/IR/Argument.h"
85#include "llvm/IR/BasicBlock.h"
86#include "llvm/IR/CFG.h"
87#include "llvm/IR/Constant.h"
89#include "llvm/IR/Constants.h"
90#include "llvm/IR/DataLayout.h"
92#include "llvm/IR/Dominators.h"
93#include "llvm/IR/Function.h"
94#include "llvm/IR/GlobalAlias.h"
95#include "llvm/IR/GlobalValue.h"
97#include "llvm/IR/InstrTypes.h"
98#include "llvm/IR/Instruction.h"
101#include "llvm/IR/Intrinsics.h"
102#include "llvm/IR/LLVMContext.h"
103#include "llvm/IR/Operator.h"
104#include "llvm/IR/PatternMatch.h"
105#include "llvm/IR/Type.h"
106#include "llvm/IR/Use.h"
107#include "llvm/IR/User.h"
108#include "llvm/IR/Value.h"
109#include "llvm/IR/Verifier.h"
111#include "llvm/Pass.h"
112#include "llvm/Support/Casting.h"
115#include "llvm/Support/Debug.h"
121#include <algorithm>
122#include <cassert>
123#include <climits>
124#include <cstdint>
125#include <cstdlib>
126#include <map>
127#include <memory>
128#include <numeric>
129#include <optional>
130#include <tuple>
131#include <utility>
132#include <vector>
133
134using namespace llvm;
135using namespace PatternMatch;
136using namespace SCEVPatternMatch;
137
138#define DEBUG_TYPE "scalar-evolution"
139
140STATISTIC(NumExitCountsComputed,
141 "Number of loop exits with predictable exit counts");
142STATISTIC(NumExitCountsNotComputed,
143 "Number of loop exits without predictable exit counts");
144STATISTIC(NumBruteForceTripCountsComputed,
145 "Number of loops with trip counts computed by force");
146
147#ifdef EXPENSIVE_CHECKS
148bool llvm::VerifySCEV = true;
149#else
150bool llvm::VerifySCEV = false;
151#endif
152
154 MaxBruteForceIterations("scalar-evolution-max-iterations", cl::ReallyHidden,
155 cl::desc("Maximum number of iterations SCEV will "
156 "symbolically execute a constant "
157 "derived loop"),
158 cl::init(100));
159
161 "verify-scev", cl::Hidden, cl::location(VerifySCEV),
162 cl::desc("Verify ScalarEvolution's backedge taken counts (slow)"));
164 "verify-scev-strict", cl::Hidden,
165 cl::desc("Enable stricter verification with -verify-scev is passed"));
166
168 "scev-verify-ir", cl::Hidden,
169 cl::desc("Verify IR correctness when making sensitive SCEV queries (slow)"),
170 cl::init(false));
171
173 "scev-mulops-inline-threshold", cl::Hidden,
174 cl::desc("Threshold for inlining multiplication operands into a SCEV"),
175 cl::init(32));
176
178 "scev-addops-inline-threshold", cl::Hidden,
179 cl::desc("Threshold for inlining addition operands into a SCEV"),
180 cl::init(500));
181
183 "scalar-evolution-max-scev-compare-depth", cl::Hidden,
184 cl::desc("Maximum depth of recursive SCEV complexity comparisons"),
185 cl::init(32));
186
188 "scalar-evolution-max-scev-operations-implication-depth", cl::Hidden,
189 cl::desc("Maximum depth of recursive SCEV operations implication analysis"),
190 cl::init(2));
191
193 "scalar-evolution-max-value-compare-depth", cl::Hidden,
194 cl::desc("Maximum depth of recursive value complexity comparisons"),
195 cl::init(2));
196
198 MaxArithDepth("scalar-evolution-max-arith-depth", cl::Hidden,
199 cl::desc("Maximum depth of recursive arithmetics"),
200 cl::init(32));
201
203 "scalar-evolution-max-constant-evolving-depth", cl::Hidden,
204 cl::desc("Maximum depth of recursive constant evolving"), cl::init(32));
205
207 MaxCastDepth("scalar-evolution-max-cast-depth", cl::Hidden,
208 cl::desc("Maximum depth of recursive SExt/ZExt/Trunc"),
209 cl::init(8));
210
212 MaxAddRecSize("scalar-evolution-max-add-rec-size", cl::Hidden,
213 cl::desc("Max coefficients in AddRec during evolving"),
214 cl::init(8));
215
217 HugeExprThreshold("scalar-evolution-huge-expr-threshold", cl::Hidden,
218 cl::desc("Size of the expression which is considered huge"),
219 cl::init(4096));
220
222 "scev-range-iter-threshold", cl::Hidden,
223 cl::desc("Threshold for switching to iteratively computing SCEV ranges"),
224 cl::init(32));
225
227 "scalar-evolution-max-loop-guard-collection-depth", cl::Hidden,
228 cl::desc("Maximum depth for recursive loop guard collection"), cl::init(1));
229
230static cl::opt<bool>
231ClassifyExpressions("scalar-evolution-classify-expressions",
232 cl::Hidden, cl::init(true),
233 cl::desc("When printing analysis, include information on every instruction"));
234
236 "scalar-evolution-use-expensive-range-sharpening", cl::Hidden,
237 cl::init(false),
238 cl::desc("Use more powerful methods of sharpening expression ranges. May "
239 "be costly in terms of compile time"));
240
241static cl::opt<bool>
242 EnableFiniteLoopControl("scalar-evolution-finite-loop", cl::Hidden,
243 cl::desc("Handle <= and >= in finite loops"),
244 cl::init(true));
245
247 "scalar-evolution-use-context-for-no-wrap-flag-strenghening", cl::Hidden,
248 cl::desc("Infer nuw/nsw flags using context where suitable"),
249 cl::init(true));
250
251//===----------------------------------------------------------------------===//
252// SCEV class definitions
253//===----------------------------------------------------------------------===//
254
256 // Leaf nodes are always their own canonical.
257 switch (getSCEVType()) {
258 case scConstant:
259 case scVScale:
260 case scUnknown:
261 CanonicalSCEV = this;
262 return;
263 default:
264 break;
265 }
266
267 // For all other expressions, check whether any immediate operand has a
268 // different canonical. Since operands are always created before their parent,
269 // their canonical pointers are already set — no recursion needed.
270 bool Changed = false;
272 for (SCEVUse Op : operands()) {
273 CanonOps.push_back(Op->getCanonical());
274 Changed |= CanonOps.back() != Op;
275 }
276
277 if (!Changed) {
278 CanonicalSCEV = this;
279 return;
280 }
281
282 // Rebuild the expression from the canonical operands, stripping use flags.
283 CanonicalSCEV = SE.getWithOperands(this, CanonOps);
284}
285
286//===----------------------------------------------------------------------===//
287// Implementation of the SCEV class.
288//
289
290#if !defined(NDEBUG) || defined(LLVM_ENABLE_DUMP)
292 print(dbgs());
293 dbgs() << '\n';
294}
295#endif
296
297void SCEV::print(raw_ostream &OS) const {
298 switch (getSCEVType()) {
299 case scConstant:
300 cast<SCEVConstant>(this)->getValue()->printAsOperand(OS, false);
301 return;
302 case scVScale:
303 OS << "vscale";
304 return;
305 case scPtrToAddr: {
306 const SCEVCastExpr *PtrCast = cast<SCEVCastExpr>(this);
307 SCEVUse Op = PtrCast->getOperand();
308 OS << "(ptrtoaddr " << *Op->getType() << " " << Op << " to "
309 << *PtrCast->getType() << ")";
310 return;
311 }
312 case scTruncate: {
313 const SCEVTruncateExpr *Trunc = cast<SCEVTruncateExpr>(this);
314 SCEVUse Op = Trunc->getOperand();
315 OS << "(trunc " << *Op->getType() << " " << Op << " to "
316 << *Trunc->getType() << ")";
317 return;
318 }
319 case scZeroExtend: {
321 SCEVUse Op = ZExt->getOperand();
322 OS << "(zext " << *Op->getType() << " " << Op << " to " << *ZExt->getType()
323 << ")";
324 return;
325 }
326 case scSignExtend: {
328 SCEVUse Op = SExt->getOperand();
329 OS << "(sext " << *Op->getType() << " " << Op << " to " << *SExt->getType()
330 << ")";
331 return;
332 }
333 case scAddRecExpr: {
334 const SCEVAddRecExpr *AR = cast<SCEVAddRecExpr>(this);
335 OS << "{" << AR->getOperand(0);
336 for (unsigned i = 1, e = AR->getNumOperands(); i != e; ++i)
337 OS << ",+," << AR->getOperand(i);
338 OS << "}<";
339 if (AR->hasNoUnsignedWrap())
340 OS << "nuw><";
341 if (AR->hasNoSignedWrap())
342 OS << "nsw><";
343 if (AR->hasNoSelfWrap() && !AR->hasNoUnsignedWrap() &&
344 !AR->hasNoSignedWrap())
345 OS << "nw><";
346 AR->getLoop()->getHeader()->printAsOperand(OS, /*PrintType=*/false);
347 OS << ">";
348 return;
349 }
350 case scAddExpr:
351 case scMulExpr:
352 case scUMaxExpr:
353 case scSMaxExpr:
354 case scUMinExpr:
355 case scSMinExpr:
357 const SCEVNAryExpr *NAry = cast<SCEVNAryExpr>(this);
358 const char *OpStr = nullptr;
359 switch (NAry->getSCEVType()) {
360 case scAddExpr: OpStr = " + "; break;
361 case scMulExpr: OpStr = " * "; break;
362 case scUMaxExpr: OpStr = " umax "; break;
363 case scSMaxExpr: OpStr = " smax "; break;
364 case scUMinExpr:
365 OpStr = " umin ";
366 break;
367 case scSMinExpr:
368 OpStr = " smin ";
369 break;
371 OpStr = " umin_seq ";
372 break;
373 default:
374 llvm_unreachable("There are no other nary expression types.");
375 }
376 OS << "(" << llvm::interleaved(NAry->operands(), OpStr) << ")";
377 switch (NAry->getSCEVType()) {
378 case scAddExpr:
379 case scMulExpr:
380 if (NAry->hasNoUnsignedWrap())
381 OS << "<nuw>";
382 if (NAry->hasNoSignedWrap())
383 OS << "<nsw>";
384 break;
385 default:
386 // Nothing to print for other nary expressions.
387 break;
388 }
389 return;
390 }
391 case scUDivExpr: {
392 const SCEVUDivExpr *UDiv = cast<SCEVUDivExpr>(this);
393 OS << "(" << UDiv->getLHS() << " /u " << UDiv->getRHS() << ")";
394 return;
395 }
396 case scUnknown:
397 cast<SCEVUnknown>(this)->getValue()->printAsOperand(OS, false);
398 return;
400 OS << "***COULDNOTCOMPUTE***";
401 return;
402 }
403 llvm_unreachable("Unknown SCEV kind!");
404}
405
407 switch (getSCEVType()) {
408 case scConstant:
409 case scVScale:
410 case scUnknown:
411 return {};
412 case scPtrToAddr:
413 case scTruncate:
414 case scZeroExtend:
415 case scSignExtend:
416 return cast<SCEVCastExpr>(this)->operands();
417 case scAddRecExpr:
418 case scAddExpr:
419 case scMulExpr:
420 case scUMaxExpr:
421 case scSMaxExpr:
422 case scUMinExpr:
423 case scSMinExpr:
425 return cast<SCEVNAryExpr>(this)->operands();
426 case scUDivExpr:
427 return cast<SCEVUDivExpr>(this)->operands();
429 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
430 }
431 llvm_unreachable("Unknown SCEV kind!");
432}
433
434bool SCEV::isZero() const { return match(this, m_scev_Zero()); }
435
436bool SCEV::isOne() const { return match(this, m_scev_One()); }
437
438bool SCEV::isAllOnesValue() const { return match(this, m_scev_AllOnes()); }
439
442 if (!Mul) return false;
443
444 // If there is a constant factor, it will be first.
445 const SCEVConstant *SC = dyn_cast<SCEVConstant>(Mul->getOperand(0));
446 if (!SC) return false;
447
448 // Return true if the value is negative, this matches things like (-42 * V).
449 return SC->getAPInt().isNegative();
450}
451
454
456 return S->getSCEVType() == scCouldNotCompute;
457}
458
460 auto &Entry = ConstantSCEVs[V];
461 if (Entry)
462 return Entry;
463
466 ID.AddPointer(V);
468 if (SCEVConstant *S =
469 static_cast<SCEVConstant *>(UniqueSCEVs.lookup(ID, Token)))
470 return Entry = S;
471 SCEVConstant *S =
472 new (SCEVAllocator) SCEVConstant(ID.Intern(SCEVAllocator), V);
473 UniqueSCEVs.insert(S, Token);
474 S->computeAndSetCanonical(*this);
475 return Entry = S;
476}
477
479 return getConstant(ConstantInt::get(getContext(), Val));
480}
481
482const SCEV *
485 // TODO: Avoid implicit trunc?
486 // See https://github.com/llvm/llvm-project/issues/112510.
487 return getConstant(
488 ConstantInt::get(ITy, V, isSigned, /*ImplicitTrunc=*/true));
489}
490
494 ID.AddPointer(Ty);
496 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
497 return S;
498 SCEV *S = new (SCEVAllocator) SCEVVScale(ID.Intern(SCEVAllocator), Ty);
499 UniqueSCEVs.insert(S, Token);
500 S->computeAndSetCanonical(*this);
501 return S;
502}
503
505 SCEVFlags Flags) {
506 const SCEV *Res = getConstant(Ty, EC.getKnownMinValue());
507 if (EC.isScalable())
508 Res = getMulExpr(Res, getVScale(Ty), Flags);
509 return Res;
510}
511
513 SCEVUse op, Type *ty)
514 : SCEV(ID, SCEVTy, computeExpressionSize(op), ty), Op(op) {}
515
516SCEVPtrToAddrExpr::SCEVPtrToAddrExpr(const FoldingSetNodeIDRef ID,
517 const SCEV *Op, Type *ITy)
518 : SCEVCastExpr(ID, scPtrToAddr, Op, ITy) {
519 assert(getOperand()->getType()->isPointerTy() && getType()->isIntegerTy() &&
520 "Must be a non-bit-width-changing pointer-to-integer cast!");
521}
522
527
528SCEVTruncateExpr::SCEVTruncateExpr(const FoldingSetNodeIDRef ID, SCEVUse op,
529 Type *ty)
531 assert(getOperand()->getType()->isIntOrPtrTy() && getType()->isIntOrPtrTy() &&
532 "Cannot truncate non-integer value!");
533}
534
535SCEVZeroExtendExpr::SCEVZeroExtendExpr(const FoldingSetNodeIDRef ID, SCEVUse op,
536 Type *ty)
538 assert(getOperand()->getType()->isIntOrPtrTy() && getType()->isIntOrPtrTy() &&
539 "Cannot zero extend non-integer value!");
540}
541
542SCEVSignExtendExpr::SCEVSignExtendExpr(const FoldingSetNodeIDRef ID, SCEVUse op,
543 Type *ty)
545 assert(getOperand()->getType()->isIntOrPtrTy() && getType()->isIntOrPtrTy() &&
546 "Cannot sign extend non-integer value!");
547}
548
550 // Clear this SCEVUnknown from various maps.
551 SE->forgetMemoizedResults({this});
552
553 // Remove this SCEVUnknown from the uniquing map.
554 SE->UniqueSCEVs.erase(this);
555
556 // Release the value.
557 setValPtr(nullptr);
558}
559
560void SCEVUnknown::allUsesReplacedWith(Value *New) {
561 // Clear this SCEVUnknown from various maps.
562 SE->forgetMemoizedResults({this});
563
564 // Remove this SCEVUnknown from the uniquing map.
565 SE->UniqueSCEVs.erase(this);
566
567 // Replace the value pointer in case someone is still using this SCEVUnknown.
568 setValPtr(New);
569}
570
571//===----------------------------------------------------------------------===//
572// SCEV Utilities
573//===----------------------------------------------------------------------===//
574
575/// Compare the two values \p LV and \p RV in terms of their "complexity" where
576/// "complexity" is a partial (and somewhat ad-hoc) relation used to order
577/// operands in SCEV expressions.
578static int CompareValueComplexity(const LoopInfo *const LI, Value *LV,
579 Value *RV, unsigned Depth) {
581 return 0;
582
583 // Order pointer values after integer values. This helps SCEVExpander form
584 // GEPs.
585 bool LIsPointer = LV->getType()->isPointerTy(),
586 RIsPointer = RV->getType()->isPointerTy();
587 if (LIsPointer != RIsPointer)
588 return (int)LIsPointer - (int)RIsPointer;
589
590 // Compare getValueID values.
591 unsigned LID = LV->getValueID(), RID = RV->getValueID();
592 if (LID != RID)
593 return (int)LID - (int)RID;
594
595 // Sort arguments by their position.
596 if (const auto *LA = dyn_cast<Argument>(LV)) {
597 const auto *RA = cast<Argument>(RV);
598 unsigned LArgNo = LA->getArgNo(), RArgNo = RA->getArgNo();
599 return (int)LArgNo - (int)RArgNo;
600 }
601
602 if (const auto *LGV = dyn_cast<GlobalValue>(LV)) {
603 const auto *RGV = cast<GlobalValue>(RV);
604
605 if (auto L = LGV->getLinkage() - RGV->getLinkage())
606 return L;
607
608 const auto IsGVNameSemantic = [&](const GlobalValue *GV) {
609 auto LT = GV->getLinkage();
610 return !(GlobalValue::isPrivateLinkage(LT) ||
612 };
613
614 // Use the names to distinguish the two values, but only if the
615 // names are semantically important.
616 if (IsGVNameSemantic(LGV) && IsGVNameSemantic(RGV))
617 return LGV->getName().compare(RGV->getName());
618 }
619
620 // For instructions, compare their loop depth, and their operand count. This
621 // is pretty loose.
622 if (const auto *LInst = dyn_cast<Instruction>(LV)) {
623 const auto *RInst = cast<Instruction>(RV);
624
625 // Compare loop depths.
626 const BasicBlock *LParent = LInst->getParent(),
627 *RParent = RInst->getParent();
628 if (LParent != RParent) {
629 unsigned LDepth = LI->getLoopDepth(LParent),
630 RDepth = LI->getLoopDepth(RParent);
631 if (LDepth != RDepth)
632 return (int)LDepth - (int)RDepth;
633 }
634
635 // Compare the number of operands.
636 unsigned LNumOps = LInst->getNumOperands(),
637 RNumOps = RInst->getNumOperands();
638 if (LNumOps != RNumOps)
639 return (int)LNumOps - (int)RNumOps;
640
641 for (unsigned Idx : seq(LNumOps)) {
642 int Result = CompareValueComplexity(LI, LInst->getOperand(Idx),
643 RInst->getOperand(Idx), Depth + 1);
644 if (Result != 0)
645 return Result;
646 }
647 }
648
649 return 0;
650}
651
652// Return negative, zero, or positive, if LHS is less than, equal to, or greater
653// than RHS, respectively. A three-way result allows recursive comparisons to be
654// more efficient.
655// If the max analysis depth was reached, return std::nullopt, assuming we do
656// not know if they are equivalent for sure.
657static std::optional<int>
658CompareSCEVComplexity(const LoopInfo *const LI, const SCEV *LHS,
659 const SCEV *RHS, DominatorTree &DT, unsigned Depth = 0) {
660 // Fast-path: SCEVs are uniqued so we can do a quick equality check.
661 if (LHS == RHS)
662 return 0;
663
664 // Primarily, sort the SCEVs by their getSCEVType().
665 SCEVTypes LType = LHS->getSCEVType(), RType = RHS->getSCEVType();
666 if (LType != RType)
667 return (int)LType - (int)RType;
668
670 return std::nullopt;
671
672 // Aside from the getSCEVType() ordering, the particular ordering
673 // isn't very important except that it's beneficial to be consistent,
674 // so that (a + b) and (b + a) don't end up as different expressions.
675 switch (LType) {
676 case scUnknown: {
677 const SCEVUnknown *LU = cast<SCEVUnknown>(LHS);
678 const SCEVUnknown *RU = cast<SCEVUnknown>(RHS);
679
680 int X =
681 CompareValueComplexity(LI, LU->getValue(), RU->getValue(), Depth + 1);
682 return X;
683 }
684
685 case scConstant: {
688
689 // Compare constant values.
690 const APInt &LA = LC->getAPInt();
691 const APInt &RA = RC->getAPInt();
692 unsigned LBitWidth = LA.getBitWidth(), RBitWidth = RA.getBitWidth();
693 if (LBitWidth != RBitWidth)
694 return (int)LBitWidth - (int)RBitWidth;
695 return LA.ult(RA) ? -1 : 1;
696 }
697
698 case scVScale: {
699 const auto *LTy = cast<IntegerType>(cast<SCEVVScale>(LHS)->getType());
700 const auto *RTy = cast<IntegerType>(cast<SCEVVScale>(RHS)->getType());
701 return LTy->getBitWidth() - RTy->getBitWidth();
702 }
703
704 case scAddRecExpr: {
707
708 // There is always a dominance between two recs that are used by one SCEV,
709 // so we can safely sort recs by loop header dominance. We require such
710 // order in getAddExpr.
711 const Loop *LLoop = LA->getLoop(), *RLoop = RA->getLoop();
712 if (LLoop != RLoop) {
713 const BasicBlock *LHead = LLoop->getHeader(), *RHead = RLoop->getHeader();
714 assert(LHead != RHead && "Two loops share the same header?");
715 if (DT.dominates(LHead, RHead))
716 return 1;
717 assert(DT.dominates(RHead, LHead) &&
718 "No dominance between recurrences used by one SCEV?");
719 return -1;
720 }
721
722 [[fallthrough]];
723 }
724
725 case scTruncate:
726 case scZeroExtend:
727 case scSignExtend:
728 case scPtrToAddr:
729 case scAddExpr:
730 case scMulExpr:
731 case scUDivExpr:
732 case scSMaxExpr:
733 case scUMaxExpr:
734 case scSMinExpr:
735 case scUMinExpr:
737 ArrayRef<SCEVUse> LOps = LHS->operands();
738 ArrayRef<SCEVUse> ROps = RHS->operands();
739
740 // Lexicographically compare n-ary-like expressions.
741 unsigned LNumOps = LOps.size(), RNumOps = ROps.size();
742 if (LNumOps != RNumOps)
743 return (int)LNumOps - (int)RNumOps;
744
745 for (unsigned i = 0; i != LNumOps; ++i) {
746 auto X = CompareSCEVComplexity(LI, LOps[i].getPointer(),
747 ROps[i].getPointer(), DT, Depth + 1);
748 if (X != 0)
749 return X;
750 }
751 return 0;
752 }
753
755 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
756 }
757 llvm_unreachable("Unknown SCEV kind!");
758}
759
760/// Given a list of SCEV objects, order them by their complexity, and group
761/// objects of the same complexity together by value. When this routine is
762/// finished, we know that any duplicates in the vector are consecutive and that
763/// complexity is monotonically increasing.
764///
765/// Note that we go take special precautions to ensure that we get deterministic
766/// results from this routine. In other words, we don't want the results of
767/// this to depend on where the addresses of various SCEV objects happened to
768/// land in memory.
770 DominatorTree &DT) {
771 if (Ops.size() < 2) return; // Noop
772
773 // Whether LHS has provably less complexity than RHS.
774 auto IsLessComplex = [&](SCEVUse LHS, SCEVUse RHS) {
775 auto Complexity = CompareSCEVComplexity(LI, LHS, RHS, DT);
776 return Complexity && *Complexity < 0;
777 };
778 if (Ops.size() == 2) {
779 // This is the common case, which also happens to be trivially simple.
780 // Special case it.
781 SCEVUse &LHS = Ops[0], &RHS = Ops[1];
782 if (IsLessComplex(RHS, LHS))
783 std::swap(LHS, RHS);
784 return;
785 }
786
787 // Do the rough sort by complexity.
789 Ops, [&](SCEVUse LHS, SCEVUse RHS) { return IsLessComplex(LHS, RHS); });
790
791 // Now that we are sorted by complexity, group elements of the same
792 // complexity. Note that this is, at worst, N^2, but the vector is likely to
793 // be extremely short in practice. Note that we take this approach because we
794 // do not want to depend on the addresses of the objects we are grouping.
795 for (unsigned i = 0, e = Ops.size(); i != e-2; ++i) {
796 const SCEV *S = Ops[i];
797 unsigned Complexity = S->getSCEVType();
798
799 // If there are any objects of the same complexity and same value as this
800 // one, group them.
801 for (unsigned j = i+1; j != e && Ops[j]->getSCEVType() == Complexity; ++j) {
802 if (Ops[j] == S) { // Found a duplicate.
803 // Move it to immediately after i'th element.
804 std::swap(Ops[i+1], Ops[j]);
805 ++i; // no need to rescan it.
806 if (i == e-2) return; // Done!
807 }
808 }
809 }
810}
811
812/// Returns true if \p Ops contains a huge SCEV (the subtree of S contains at
813/// least HugeExprThreshold nodes).
815 return any_of(Ops, [](const SCEV *S) {
817 });
818}
819
820/// Performs a number of common optimizations on the passed \p Ops. If the
821/// whole expression reduces down to a single operand, it will be returned.
822///
823/// The following optimizations are performed:
824/// * Fold constants using the \p Fold function.
825/// * Remove identity constants satisfying \p IsIdentity.
826/// * If a constant satisfies \p IsAbsorber, return it.
827/// * Sort operands by complexity.
828template <typename FoldT, typename IsIdentityT, typename IsAbsorberT>
829static const SCEV *
831 SmallVectorImpl<SCEVUse> &Ops, FoldT Fold,
832 IsIdentityT IsIdentity, IsAbsorberT IsAbsorber) {
833 const SCEVConstant *Folded = nullptr;
834 for (unsigned Idx = 0; Idx < Ops.size();) {
835 const SCEV *Op = Ops[Idx];
836 if (const auto *C = dyn_cast<SCEVConstant>(Op)) {
837 if (!Folded)
838 Folded = C;
839 else
840 Folded = cast<SCEVConstant>(
841 SE.getConstant(Fold(Folded->getAPInt(), C->getAPInt())));
842 Ops.erase(Ops.begin() + Idx);
843 continue;
844 }
845 ++Idx;
846 }
847
848 if (Ops.empty()) {
849 assert(Folded && "Must have folded value");
850 return Folded;
851 }
852
853 if (Folded && IsAbsorber(Folded->getAPInt()))
854 return Folded;
855
856 GroupByComplexity(Ops, &LI, DT);
857 if (Folded && !IsIdentity(Folded->getAPInt()))
858 Ops.insert(Ops.begin(), Folded);
859
860 return Ops.size() == 1 ? Ops[0] : nullptr;
861}
862
863//===----------------------------------------------------------------------===//
864// Simple SCEV method implementations
865//===----------------------------------------------------------------------===//
866
867/// Compute BC(It, K). The result has width W. Assume, K > 0.
868static const SCEV *BinomialCoefficient(const SCEV *It, unsigned K,
869 ScalarEvolution &SE,
870 Type *ResultTy) {
871 // Handle the simplest case efficiently.
872 if (K == 1)
873 return SE.getTruncateOrZeroExtend(It, ResultTy);
874
875 // We are using the following formula for BC(It, K):
876 //
877 // BC(It, K) = (It * (It - 1) * ... * (It - K + 1)) / K!
878 //
879 // Suppose, W is the bitwidth of the return value. We must be prepared for
880 // overflow. Hence, we must assure that the result of our computation is
881 // equal to the accurate one modulo 2^W. Unfortunately, division isn't
882 // safe in modular arithmetic.
883 //
884 // However, this code doesn't use exactly that formula; the formula it uses
885 // is something like the following, where T is the number of factors of 2 in
886 // K! (i.e. trailing zeros in the binary representation of K!), and ^ is
887 // exponentiation:
888 //
889 // BC(It, K) = (It * (It - 1) * ... * (It - K + 1)) / 2^T / (K! / 2^T)
890 //
891 // This formula is trivially equivalent to the previous formula. However,
892 // this formula can be implemented much more efficiently. The trick is that
893 // K! / 2^T is odd, and exact division by an odd number *is* safe in modular
894 // arithmetic. To do exact division in modular arithmetic, all we have
895 // to do is multiply by the inverse. Therefore, this step can be done at
896 // width W.
897 //
898 // The next issue is how to safely do the division by 2^T. The way this
899 // is done is by doing the multiplication step at a width of at least W + T
900 // bits. This way, the bottom W+T bits of the product are accurate. Then,
901 // when we perform the division by 2^T (which is equivalent to a right shift
902 // by T), the bottom W bits are accurate. Extra bits are okay; they'll get
903 // truncated out after the division by 2^T.
904 //
905 // In comparison to just directly using the first formula, this technique
906 // is much more efficient; using the first formula requires W * K bits,
907 // but this formula less than W + K bits. Also, the first formula requires
908 // a division step, whereas this formula only requires multiplies and shifts.
909 //
910 // It doesn't matter whether the subtraction step is done in the calculation
911 // width or the input iteration count's width; if the subtraction overflows,
912 // the result must be zero anyway. We prefer here to do it in the width of
913 // the induction variable because it helps a lot for certain cases; CodeGen
914 // isn't smart enough to ignore the overflow, which leads to much less
915 // efficient code if the width of the subtraction is wider than the native
916 // register width.
917 //
918 // (It's possible to not widen at all by pulling out factors of 2 before
919 // the multiplication; for example, K=2 can be calculated as
920 // It/2*(It+(It*INT_MIN/INT_MIN)+-1). However, it requires
921 // extra arithmetic, so it's not an obvious win, and it gets
922 // much more complicated for K > 3.)
923
924 // Protection from insane SCEVs; this bound is conservative,
925 // but it probably doesn't matter.
926 if (K > 1000)
927 return SE.getCouldNotCompute();
928
929 unsigned W = SE.getTypeSizeInBits(ResultTy);
930
931 // Calculate K! / 2^T and T; we divide out the factors of two before
932 // multiplying for calculating K! / 2^T to avoid overflow.
933 // Other overflow doesn't matter because we only care about the bottom
934 // W bits of the result.
935 APInt OddFactorial(W, 1);
936 unsigned T = 1;
937 for (unsigned i = 3; i <= K; ++i) {
938 unsigned TwoFactors = countr_zero(i);
939 T += TwoFactors;
940 OddFactorial *= (i >> TwoFactors);
941 }
942
943 // We need at least W + T bits for the multiplication step
944 unsigned CalculationBits = W + T;
945
946 // Calculate 2^T, at width T+W.
947 APInt DivFactor = APInt::getOneBitSet(CalculationBits, T);
948
949 // Calculate the multiplicative inverse of K! / 2^T;
950 // this multiplication factor will perform the exact division by
951 // K! / 2^T.
952 APInt MultiplyFactor = OddFactorial.multiplicativeInverse();
953
954 // Calculate the product, at width T+W
955 IntegerType *CalculationTy = IntegerType::get(SE.getContext(),
956 CalculationBits);
957 const SCEV *Dividend = SE.getTruncateOrZeroExtend(It, CalculationTy);
958 for (unsigned i = 1; i != K; ++i) {
959 const SCEV *S = SE.getMinusSCEV(It, SE.getConstant(It->getType(), i));
960 Dividend = SE.getMulExpr(Dividend,
961 SE.getTruncateOrZeroExtend(S, CalculationTy));
962 }
963
964 // Divide by 2^T
965 const SCEV *DivResult = SE.getUDivExpr(Dividend, SE.getConstant(DivFactor));
966
967 // Truncate the result, and divide by K! / 2^T.
968
969 return SE.getMulExpr(SE.getConstant(MultiplyFactor),
970 SE.getTruncateOrZeroExtend(DivResult, ResultTy));
971}
972
973/// Return the value of this chain of recurrences at the specified iteration
974/// number. We can evaluate this recurrence by multiplying each element in the
975/// chain by the binomial coefficient corresponding to it. In other words, we
976/// can evaluate {A,+,B,+,C,+,D} as:
977///
978/// A*BC(It, 0) + B*BC(It, 1) + C*BC(It, 2) + D*BC(It, 3)
979///
980/// where BC(It, k) stands for binomial coefficient.
982 ScalarEvolution &SE) const {
983 return evaluateAtIteration(operands(), It, SE);
984}
985
987 const SCEV *It, ScalarEvolution &SE,
988 SCEVFlags UseFlags) {
989 assert(Operands.size() > 0);
990 assert((Operands.size() == 2 || UseFlags == SCEV::FlagNone) &&
991 "use-specific flags only supported for affine AddRecs");
992 SCEVUse Result = Operands[0].getPointer();
993 for (unsigned i = 1, e = Operands.size(); i != e; ++i) {
994 // The computation is correct in the face of overflow provided that the
995 // multiplication is performed _after_ the evaluation of the binomial
996 // coefficient.
997 const SCEV *Coeff = BinomialCoefficient(It, i, SE, Result->getType());
998 if (isa<SCEVCouldNotCompute>(Coeff))
999 return Coeff;
1000
1001 SCEVUse Mul = SE.getMulExpr(Operands[i].getPointer(), Coeff,
1002 {SCEV::FlagNone, UseFlags});
1003 Result = SE.getAddExpr(Result, Mul, {SCEV::FlagNone, UseFlags});
1004 }
1005 return Result;
1006}
1007
1009 const SCEV *BTC = SE.getBackedgeTakenCount(getLoop());
1010 if (isa<SCEVCouldNotCompute>(BTC))
1011 return BTC;
1012 // The loop reaches iteration BTC, so the value this recurrence computes there
1013 // is the value it had, and that did not wrap.
1014 return evaluateAtIteration(operands(), BTC, SE,
1016 : SCEV::FlagNone);
1017}
1018
1019//===----------------------------------------------------------------------===//
1020// SCEV Expression folder implementations
1021//===----------------------------------------------------------------------===//
1022
1023/// The SCEVCastSinkingRewriter takes a scalar evolution expression,
1024/// which computes a pointer-typed value, and rewrites the whole expression
1025/// tree so that *all* the computations are done on integers, and the only
1026/// pointer-typed operands in the expression are SCEVUnknown.
1027/// The CreatePtrCast callback is invoked to create the actual conversion
1028/// (ptrtoint or ptrtoaddr) at the SCEVUnknown leaves.
1030 : public SCEVRewriteVisitor<SCEVCastSinkingRewriter> {
1032 using ConversionFn = function_ref<const SCEV *(const SCEVUnknown *)>;
1033 Type *TargetTy;
1034 ConversionFn CreatePtrCast;
1035
1036public:
1038 ConversionFn CreatePtrCast)
1039 : Base(SE), TargetTy(TargetTy), CreatePtrCast(std::move(CreatePtrCast)) {}
1040
1041 static const SCEV *rewrite(const SCEV *Scev, ScalarEvolution &SE,
1042 Type *TargetTy, ConversionFn CreatePtrCast) {
1043 SCEVCastSinkingRewriter Rewriter(SE, TargetTy, std::move(CreatePtrCast));
1044 return Rewriter.visit(Scev);
1045 }
1046
1047 const SCEV *visit(const SCEV *S) {
1048 Type *STy = S->getType();
1049 // If the expression is not pointer-typed, just keep it as-is.
1050 if (!STy->isPointerTy())
1051 return S;
1052 // Else, recursively sink the cast down into it.
1053 return Base::visit(S);
1054 }
1055
1056 const SCEV *visitAddExpr(const SCEVAddExpr *Expr) {
1057 // Preserve wrap flags on rewritten SCEVAddExpr, which the default
1058 // implementation drops.
1060 bool Changed = false;
1061 for (SCEVUse Op : Expr->operands()) {
1062 Operands.push_back(visit(Op.getPointer()));
1063 Changed |= Op.getPointer() != Operands.back();
1064 }
1065 return !Changed ? Expr : SE.getAddExpr(Operands, Expr->getNoWrapFlags());
1066 }
1067
1068 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
1069 assert(Expr->getType()->isPointerTy() &&
1070 "Should only reach pointer-typed SCEVUnknown's.");
1071 // Perform some basic constant folding. If the operand of the cast is a
1072 // null pointer, don't create a cast SCEV expression (that will be left
1073 // as-is), but produce a zero constant.
1075 return SE.getZero(TargetTy);
1076 return CreatePtrCast(Expr);
1077 }
1078};
1079
1081 assert(Op->getType()->isPointerTy() && "Op must be a pointer");
1082
1083 // Treat pointers with unstable representation conservatively, since the
1084 // address bits may change.
1085 if (DL.hasUnstableRepresentation(Op->getType()))
1086 return getCouldNotCompute();
1087
1088 Type *Ty = DL.getAddressType(Op->getType());
1089
1090 // Use the rewriter to sink the cast down to SCEVUnknown leaves.
1091 // The rewriter handles null pointer constant folding.
1093 Op, *this, Ty, [this, Ty](const SCEVUnknown *U) {
1096 ID.AddPointer(U);
1097 ID.AddPointer(Ty);
1099 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
1100 return S;
1101 SCEV *S = new (SCEVAllocator)
1102 SCEVPtrToAddrExpr(ID.Intern(SCEVAllocator), U, Ty);
1103 UniqueSCEVs.insert(S, Token);
1104 S->computeAndSetCanonical(*this);
1105 registerUser(S, {U});
1106 return static_cast<const SCEV *>(S);
1107 });
1108 assert(IntOp->getType()->isIntegerTy() &&
1109 "We must have succeeded in sinking the cast, "
1110 "and ending up with an integer-typed expression!");
1111 return IntOp;
1112}
1113
1115 unsigned Depth) {
1116 assert(getTypeSizeInBits(Op->getType()) > getTypeSizeInBits(Ty) &&
1117 "This is not a truncating conversion!");
1118 assert(isSCEVable(Ty) &&
1119 "This is not a conversion to a SCEVable type!");
1120 assert(!Op->getType()->isPointerTy() && "Can't truncate pointer!");
1121 Ty = getEffectiveSCEVType(Ty);
1122
1125 ID.AddPointer(Op.getOpaqueValue());
1126 ID.AddPointer(Ty);
1128 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
1129 return S;
1130
1131 // Fold if the operand is constant.
1132 if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(Op))
1133 return getConstant(
1134 cast<ConstantInt>(ConstantExpr::getTrunc(SC->getValue(), Ty)));
1135
1136 // trunc(trunc(x)) --> trunc(x)
1138 return getTruncateExpr(ST->getOperand(), Ty, Depth + 1);
1139
1140 // trunc(sext(x)) --> sext(x) if widening or trunc(x) if narrowing
1142 return getTruncateOrSignExtend(SS->getOperand(), Ty, Depth + 1);
1143
1144 // trunc(zext(x)) --> zext(x) if widening or trunc(x) if narrowing
1146 return getTruncateOrZeroExtend(SZ->getOperand(), Ty, Depth + 1);
1147
1148 if (Depth > MaxCastDepth) {
1149 SCEV *S =
1150 new (SCEVAllocator) SCEVTruncateExpr(ID.Intern(SCEVAllocator), Op, Ty);
1151 UniqueSCEVs.insert(S, Token);
1152 S->computeAndSetCanonical(*this);
1153 registerUser(S, Op);
1154 return S;
1155 }
1156
1157 // trunc(x1 + ... + xN) --> trunc(x1) + ... + trunc(xN) and
1158 // trunc(x1 * ... * xN) --> trunc(x1) * ... * trunc(xN),
1159 // if after transforming we have at most one truncate, not counting truncates
1160 // that replace other casts.
1162 auto *CommOp = cast<SCEVCommutativeExpr>(Op);
1164 unsigned numTruncs = 0;
1165 for (unsigned i = 0, e = CommOp->getNumOperands(); i != e && numTruncs < 2;
1166 ++i) {
1167 const SCEV *S = getTruncateExpr(CommOp->getOperand(i), Ty, Depth + 1);
1168 if (!isa<SCEVIntegralCastExpr>(CommOp->getOperand(i)) &&
1170 numTruncs++;
1171 Operands.push_back(S);
1172 }
1173 if (numTruncs < 2) {
1174 if (isa<SCEVAddExpr>(Op))
1175 return getAddExpr(Operands);
1176 if (isa<SCEVMulExpr>(Op))
1177 return getMulExpr(Operands);
1178 llvm_unreachable("Unexpected SCEV type for Op.");
1179 }
1180 // Although we checked in the beginning that ID is not in the cache, it is
1181 // possible that during recursion and different modification ID was inserted
1182 // into the cache. So if we find it, just return it.
1183 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
1184 return S;
1185 }
1186
1187 // If the input value is a chrec scev, truncate the chrec's operands.
1188 if (const SCEVAddRecExpr *AddRec = dyn_cast<SCEVAddRecExpr>(Op)) {
1190 for (const SCEV *Op : AddRec->operands())
1191 Operands.push_back(getTruncateExpr(Op, Ty, Depth + 1));
1192 return getAddRecExpr(Operands, AddRec->getLoop(), SCEV::FlagNone);
1193 }
1194
1195 // Return zero if truncating to known zeros.
1196 uint32_t MinTrailingZeros = getMinTrailingZeros(Op);
1197 if (MinTrailingZeros >= getTypeSizeInBits(Ty))
1198 return getZero(Ty);
1199
1200 // The cast wasn't folded; create an explicit cast node. We can reuse
1201 // the existing insert position since if we get here, we won't have
1202 // made any changes which would invalidate it.
1203 SCEV *S = new (SCEVAllocator) SCEVTruncateExpr(ID.Intern(SCEVAllocator),
1204 Op, Ty);
1205 UniqueSCEVs.insert(S, Token);
1206 S->computeAndSetCanonical(*this);
1207 registerUser(S, Op);
1208 return S;
1209}
1210
1211// Get the limit of a recurrence such that incrementing by Step cannot cause
1212// signed overflow as long as the value of the recurrence within the
1213// loop does not exceed this limit before incrementing.
1214static const SCEV *getSignedOverflowLimitForStep(const SCEV *Step,
1215 ICmpInst::Predicate *Pred,
1216 ScalarEvolution *SE) {
1217 unsigned BitWidth = SE->getTypeSizeInBits(Step->getType());
1218 if (SE->isKnownPositive(Step)) {
1219 *Pred = ICmpInst::ICMP_SLT;
1221 SE->getSignedRangeMax(Step));
1222 }
1223 if (SE->isKnownNegative(Step)) {
1224 *Pred = ICmpInst::ICMP_SGT;
1226 SE->getSignedRangeMin(Step));
1227 }
1228 return nullptr;
1229}
1230
1231// Get the limit of a recurrence such that incrementing by Step cannot cause
1232// unsigned overflow as long as the value of the recurrence within the loop does
1233// not exceed this limit before incrementing.
1235 ICmpInst::Predicate *Pred,
1236 ScalarEvolution *SE) {
1237 unsigned BitWidth = SE->getTypeSizeInBits(Step->getType());
1238 *Pred = ICmpInst::ICMP_ULT;
1239
1241 SE->getUnsignedRangeMax(Step));
1242}
1243
1244namespace {
1245
1246struct ExtendOpTraitsBase {
1247 typedef const SCEV *(ScalarEvolution::*GetExtendExprTy)(SCEVUse, Type *,
1248 unsigned);
1249};
1250
1251// Used to make code generic over signed and unsigned overflow.
1252template <typename ExtendOp> struct ExtendOpTraits {
1253 // Members present:
1254 //
1255 // static const SCEVFlags WrapType;
1256 //
1257 // static const ExtendOpTraitsBase::GetExtendExprTy GetExtendExpr;
1258 //
1259 // static const SCEV *getOverflowLimitForStep(const SCEV *Step,
1260 // ICmpInst::Predicate *Pred,
1261 // ScalarEvolution *SE);
1262};
1263
1264template <>
1265struct ExtendOpTraits<SCEVSignExtendExpr> : public ExtendOpTraitsBase {
1266 static const SCEVFlags WrapType = SCEV::FlagNSW;
1267
1268 static const GetExtendExprTy GetExtendExpr;
1269
1270 static const SCEV *getOverflowLimitForStep(const SCEV *Step,
1271 ICmpInst::Predicate *Pred,
1272 ScalarEvolution *SE) {
1273 return getSignedOverflowLimitForStep(Step, Pred, SE);
1274 }
1275};
1276
1277const ExtendOpTraitsBase::GetExtendExprTy ExtendOpTraits<
1279
1280template <>
1281struct ExtendOpTraits<SCEVZeroExtendExpr> : public ExtendOpTraitsBase {
1282 static const SCEVFlags WrapType = SCEV::FlagNUW;
1283
1284 static const GetExtendExprTy GetExtendExpr;
1285
1286 static const SCEV *getOverflowLimitForStep(const SCEV *Step,
1287 ICmpInst::Predicate *Pred,
1288 ScalarEvolution *SE) {
1289 return getUnsignedOverflowLimitForStep(Step, Pred, SE);
1290 }
1291};
1292
1293const ExtendOpTraitsBase::GetExtendExprTy ExtendOpTraits<
1295
1296} // end anonymous namespace
1297
1298// The recurrence AR has been shown to have no signed/unsigned wrap or something
1299// close to it. Typically, if we can prove NSW/NUW for AR, then we can just as
1300// easily prove NSW/NUW for its preincrement or postincrement sibling. This
1301// allows normalizing a sign/zero extended AddRec as such: {sext/zext(Step +
1302// Start),+,Step} => {(Step + sext/zext(Start),+,Step} As a result, the
1303// expression "Step + sext/zext(PreIncAR)" is congruent with
1304// "sext/zext(PostIncAR)"
1305template <typename ExtendOpTy>
1307 ScalarEvolution *SE, unsigned Depth) {
1308 auto WrapType = ExtendOpTraits<ExtendOpTy>::WrapType;
1309 auto GetExtendExpr = ExtendOpTraits<ExtendOpTy>::GetExtendExpr;
1310
1311 const Loop *L = AR->getLoop();
1312 const SCEV *Start = AR->getStart();
1313 const SCEV *Step = AR->getStepRecurrence(*SE);
1314
1315 // Check for a simple looking step prior to loop entry.
1316 const SCEVAddExpr *SA = dyn_cast<SCEVAddExpr>(Start);
1317 if (!SA)
1318 return nullptr;
1319
1320 // Create an AddExpr for "PreStart" after subtracting Step. Full SCEV
1321 // subtraction is expensive. For this purpose, perform a quick and dirty
1322 // difference, by checking for Step in the operand list. Note, that
1323 // SA might have repeated ops, like %a + %a + ..., so only remove one.
1324 SmallVector<SCEVUse, 4> DiffOps(SA->operands());
1325 for (auto It = DiffOps.begin(); It != DiffOps.end(); ++It)
1326 if (*It == Step) {
1327 DiffOps.erase(It);
1328 break;
1329 }
1330
1331 if (DiffOps.size() == SA->getNumOperands())
1332 return nullptr;
1333
1334 // Try to prove `WrapType` (SCEV::FlagNSW or SCEV::FlagNUW) on `PreStart` +
1335 // `Step`:
1336
1337 // 1. NSW/NUW flags on the step increment.
1338 auto PreStartFlags =
1340 const SCEV *PreStart = SE->getAddExpr(DiffOps, PreStartFlags);
1342 SE->getAddRecExpr(PreStart, Step, L, SCEV::FlagNone));
1343
1344 // "{S,+,X} is <nsw>/<nuw>" and "the backedge is taken at least once" implies
1345 // "S+X does not sign/unsign-overflow".
1346 //
1347
1348 const SCEV *BECount = SE->getBackedgeTakenCount(L);
1349 if (PreAR && any(PreAR->getNoWrapFlags(WrapType)) &&
1350 !isa<SCEVCouldNotCompute>(BECount) && SE->isKnownPositive(BECount))
1351 return PreStart;
1352
1353 // 2. Direct overflow check on the step operation's expression.
1354 unsigned BitWidth = SE->getTypeSizeInBits(AR->getType());
1355 Type *WideTy = IntegerType::get(SE->getContext(), BitWidth * 2);
1356 const SCEV *OperandExtendedStart =
1357 SE->getAddExpr((SE->*GetExtendExpr)(PreStart, WideTy, Depth),
1358 (SE->*GetExtendExpr)(Step, WideTy, Depth));
1359 if ((SE->*GetExtendExpr)(Start, WideTy, Depth) == OperandExtendedStart) {
1360 if (PreAR && any(AR->getNoWrapFlags(WrapType))) {
1361 // If we know `AR` == {`PreStart`+`Step`,+,`Step`} is `WrapType` (FlagNSW
1362 // or FlagNUW) and that `PreStart` + `Step` is `WrapType` too, then
1363 // `PreAR` == {`PreStart`,+,`Step`} is also `WrapType`. Cache this fact.
1364 SE->setNoWrapFlags(const_cast<SCEVAddRecExpr *>(PreAR), WrapType);
1365 }
1366 return PreStart;
1367 }
1368
1369 // 3. Loop precondition.
1371 const SCEV *OverflowLimit =
1372 ExtendOpTraits<ExtendOpTy>::getOverflowLimitForStep(Step, &Pred, SE);
1373
1374 if (OverflowLimit &&
1375 SE->isLoopEntryGuardedByCond(L, Pred, PreStart, OverflowLimit))
1376 return PreStart;
1377
1378 return nullptr;
1379}
1380
1381// Get the normalized zero or sign extended expression for this AddRec's Start.
1382template <typename ExtendOpTy>
1383static const SCEV *getExtendAddRecStart(const SCEVAddRecExpr *AR, Type *Ty,
1384 ScalarEvolution *SE,
1385 unsigned Depth) {
1386 auto GetExtendExpr = ExtendOpTraits<ExtendOpTy>::GetExtendExpr;
1387
1388 const SCEV *PreStart = getPreStartForExtend<ExtendOpTy>(AR, SE, Depth);
1389 if (!PreStart)
1390 return (SE->*GetExtendExpr)(AR->getStart(), Ty, Depth);
1391
1392 return SE->getAddExpr((SE->*GetExtendExpr)(AR->getStepRecurrence(*SE), Ty,
1393 Depth),
1394 (SE->*GetExtendExpr)(PreStart, Ty, Depth));
1395}
1396
1397// Try to prove away overflow by looking at "nearby" add recurrences. A
1398// motivating example for this rule: if we know `{0,+,4}` is `ult` `-1` and it
1399// does not itself wrap then we can conclude that `{1,+,4}` is `nuw`.
1400//
1401// Formally:
1402//
1403// {S,+,X} == {S-T,+,X} + T
1404// => Ext({S,+,X}) == Ext({S-T,+,X} + T)
1405//
1406// If ({S-T,+,X} + T) does not overflow ... (1)
1407//
1408// RHS == Ext({S-T,+,X} + T) == Ext({S-T,+,X}) + Ext(T)
1409//
1410// If {S-T,+,X} does not overflow ... (2)
1411//
1412// RHS == Ext({S-T,+,X}) + Ext(T) == {Ext(S-T),+,Ext(X)} + Ext(T)
1413// == {Ext(S-T)+Ext(T),+,Ext(X)}
1414//
1415// If (S-T)+T does not overflow ... (3)
1416//
1417// RHS == {Ext(S-T)+Ext(T),+,Ext(X)} == {Ext(S-T+T),+,Ext(X)}
1418// == {Ext(S),+,Ext(X)} == LHS
1419//
1420// Thus, if (1), (2) and (3) are true for some T, then
1421// Ext({S,+,X}) == {Ext(S),+,Ext(X)}
1422//
1423// (3) is implied by (1) -- "(S-T)+T does not overflow" is simply "({S-T,+,X}+T)
1424// does not overflow" restricted to the 0th iteration. Therefore we only need
1425// to check for (1) and (2).
1426//
1427// In the current context, S is `Start`, X is `Step`, Ext is `ExtendOpTy` and T
1428// is `Delta` (defined below).
1429template <typename ExtendOpTy>
1430bool ScalarEvolution::proveNoWrapByVaryingStart(const SCEV *Start,
1431 const SCEV *Step,
1432 const Loop *L) {
1433 auto WrapType = ExtendOpTraits<ExtendOpTy>::WrapType;
1434
1435 // We restrict `Start` to a constant to prevent SCEV from spending too much
1436 // time here. It is correct (but more expensive) to continue with a
1437 // non-constant `Start` and do a general SCEV subtraction to compute
1438 // `PreStart` below.
1439 const SCEVConstant *StartC = dyn_cast<SCEVConstant>(Start);
1440 if (!StartC)
1441 return false;
1442
1443 APInt StartAI = StartC->getAPInt();
1444
1445 for (unsigned Delta : {-2, -1, 1, 2}) {
1446 const SCEV *PreStart = getConstant(StartAI - Delta);
1447 const auto *PreAR = static_cast<SCEVAddRecExpr *>(
1448 findExistingSCEVInCache(scAddRecExpr, {PreStart, Step}, L));
1449
1450 // Give up if we don't already have the add recurrence we need because
1451 // actually constructing an add recurrence is relatively expensive.
1452 if (PreAR && any(PreAR->getNoWrapFlags(WrapType))) { // proves (2)
1453 const SCEV *DeltaS = getConstant(StartC->getType(), Delta);
1455 const SCEV *Limit = ExtendOpTraits<ExtendOpTy>::getOverflowLimitForStep(
1456 DeltaS, &Pred, this);
1457 if (Limit && isKnownPredicate(Pred, PreAR, Limit)) // proves (1)
1458 return true;
1459 }
1460 }
1461
1462 return false;
1463}
1464
1465// Finds an integer D for an expression (C + x + y + ...) such that the top
1466// level addition in (D + (C - D + x + y + ...)) would not wrap (signed or
1467// unsigned) and the number of trailing zeros of (C - D + x + y + ...) is
1468// maximized, where C is the \p ConstantTerm, x, y, ... are arbitrary SCEVs, and
1469// the (C + x + y + ...) expression is \p WholeAddExpr.
1471 const SCEVConstant *ConstantTerm,
1472 const SCEVAddExpr *WholeAddExpr) {
1473 const APInt &C = ConstantTerm->getAPInt();
1474 const unsigned BitWidth = C.getBitWidth();
1475 // Find number of trailing zeros of (x + y + ...) w/o the C first:
1476 uint32_t TZ = BitWidth;
1477 for (unsigned I = 1, E = WholeAddExpr->getNumOperands(); I < E && TZ; ++I)
1478 TZ = std::min(TZ, SE.getMinTrailingZeros(WholeAddExpr->getOperand(I)));
1479 if (TZ) {
1480 // Set D to be as many least significant bits of C as possible while still
1481 // guaranteeing that adding D to (C - D + x + y + ...) won't cause a wrap:
1482 return TZ < BitWidth ? C.trunc(TZ).zext(BitWidth) : C;
1483 }
1484 return APInt(BitWidth, 0);
1485}
1486
1487// Finds an integer D for an affine AddRec expression {C,+,x} such that the top
1488// level addition in (D + {C-D,+,x}) would not wrap (signed or unsigned) and the
1489// number of trailing zeros of (C - D + x * n) is maximized, where C is the \p
1490// ConstantStart, x is an arbitrary \p Step, and n is the loop trip count.
1492 const APInt &ConstantStart,
1493 const SCEV *Step) {
1494 const unsigned BitWidth = ConstantStart.getBitWidth();
1495 const uint32_t TZ = SE.getMinTrailingZeros(Step);
1496 if (TZ)
1497 return TZ < BitWidth ? ConstantStart.trunc(TZ).zext(BitWidth)
1498 : ConstantStart;
1499 return APInt(BitWidth, 0);
1500}
1501
1503 const ScalarEvolution::FoldID &ID, const SCEV *S,
1506 &FoldCacheUser) {
1507 auto I = FoldCache.insert({ID, S});
1508 if (!I.second) {
1509 // Remove FoldCacheUser entry for ID when replacing an existing FoldCache
1510 // entry.
1511 auto &UserIDs = FoldCacheUser[I.first->second];
1512 assert(count(UserIDs, ID) == 1 && "unexpected duplicates in UserIDs");
1513 for (unsigned I = 0; I != UserIDs.size(); ++I)
1514 if (UserIDs[I] == ID) {
1515 std::swap(UserIDs[I], UserIDs.back());
1516 break;
1517 }
1518 UserIDs.pop_back();
1519 I.first->second = S;
1520 }
1521 FoldCacheUser[S].push_back(ID);
1522}
1523
1525 unsigned Depth) {
1526 assert(getTypeSizeInBits(Op->getType()) < getTypeSizeInBits(Ty) &&
1527 "This is not an extending conversion!");
1528 assert(isSCEVable(Ty) &&
1529 "This is not a conversion to a SCEVable type!");
1530 assert(!Op->getType()->isPointerTy() && "Can't extend pointer!");
1531 Ty = getEffectiveSCEVType(Ty);
1532
1533 FoldID ID(scZeroExtend, Op, Ty);
1534 if (const SCEV *S = FoldCache.lookup(ID))
1535 return S;
1536
1537 const SCEV *S = getZeroExtendExprImpl(Op, Ty, Depth);
1539 insertFoldCacheEntry(ID, S, FoldCache, FoldCacheUser);
1540 return S;
1541}
1542
1544 unsigned Depth) {
1545 assert(getTypeSizeInBits(Op->getType()) < getTypeSizeInBits(Ty) &&
1546 "This is not an extending conversion!");
1547 assert(isSCEVable(Ty) && "This is not a conversion to a SCEVable type!");
1548 assert(!Op->getType()->isPointerTy() && "Can't extend pointer!");
1549
1550 // Fold if the operand is constant.
1551 if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(Op))
1552 return getConstant(SC->getAPInt().zext(getTypeSizeInBits(Ty)));
1553
1554 // zext(zext(x)) --> zext(x)
1556 return getZeroExtendExpr(SZ->getOperand(), Ty, Depth + 1);
1557
1558 // If the operand is an affine AddRec with the no-unsigned-wrap flag, the
1559 // zero-extension distributes over the recurrence.
1560 const SCEV *Start, *Step;
1561 const Loop *L;
1562 if (Depth <= MaxCastDepth &&
1563 match(Op, m_scev_AffineAddRec(m_SCEV(Start), m_SCEV(Step), m_Loop(L)))) {
1564 const auto *AR = cast<SCEVAddRecExpr>(Op);
1565 if (AR->hasNoUnsignedWrap()) {
1566 Start = getExtendAddRecStart<SCEVZeroExtendExpr>(AR, Ty, this, Depth + 1);
1567 Step = getZeroExtendExpr(Step, Ty, Depth + 1);
1568 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1569 }
1570 }
1571
1572 // Before doing any expensive analysis, check to see if we've already
1573 // computed a SCEV for this Op and Ty.
1576 ID.AddPointer(Op.getOpaqueValue());
1577 ID.AddPointer(Ty);
1579 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
1580 return S;
1581 if (Depth > MaxCastDepth) {
1582 SCEV *S = new (SCEVAllocator) SCEVZeroExtendExpr(ID.Intern(SCEVAllocator),
1583 Op, Ty);
1584 UniqueSCEVs.insert(S, Token);
1585 S->computeAndSetCanonical(*this);
1586 registerUser(S, Op);
1587 return S;
1588 }
1589
1590 // zext(trunc(x)) --> zext(x) or x or trunc(x)
1592 // It's possible the bits taken off by the truncate were all zero bits. If
1593 // so, we should be able to simplify this further.
1594 const SCEV *X = ST->getOperand();
1596 unsigned TruncBits = getTypeSizeInBits(ST->getType());
1597 unsigned NewBits = getTypeSizeInBits(Ty);
1598 if (CR.truncate(TruncBits).zeroExtend(NewBits).contains(
1599 CR.zextOrTrunc(NewBits)))
1600 return getTruncateOrZeroExtend(X, Ty, Depth);
1601 }
1602
1603 // If the input value is a chrec scev, and we can prove that the value
1604 // did not overflow the old, smaller, value, we can zero extend all of the
1605 // operands (often constants). This allows analysis of something like
1606 // this: for (unsigned char X = 0; X < 100; ++X) { int Y = X; }
1607 if (match(Op, m_scev_AffineAddRec(m_SCEV(Start), m_SCEV(Step), m_Loop(L)))) {
1608 const auto *AR = cast<SCEVAddRecExpr>(Op);
1609 unsigned BitWidth = getTypeSizeInBits(AR->getType());
1610
1611 // The no-unsigned-wrap case is handled before the uniquing lookup above.
1612
1613 // Check whether the backedge-taken count is SCEVCouldNotCompute.
1614 // Note that this serves two purposes: It filters out loops that are
1615 // simply not analyzable, and it covers the case where this code is
1616 // being called from within backedge-taken count analysis, such that
1617 // attempting to ask for the backedge-taken count would likely result
1618 // in infinite recursion. In the later case, the analysis code will
1619 // cope with a conservative value, and it will take care to purge
1620 // that value once it has finished.
1621 const SCEV *MaxBECount = getConstantMaxBackedgeTakenCount(L);
1622 if (!isa<SCEVCouldNotCompute>(MaxBECount)) {
1623 // Manually compute the final value for AR, checking for overflow.
1624
1625 // Check whether the backedge-taken count can be losslessly casted to
1626 // the addrec's type. The count is always unsigned.
1627 const SCEV *CastedMaxBECount =
1628 getTruncateOrZeroExtend(MaxBECount, Start->getType(), Depth);
1629 const SCEV *RecastedMaxBECount = getTruncateOrZeroExtend(
1630 CastedMaxBECount, MaxBECount->getType(), Depth);
1631 if (MaxBECount == RecastedMaxBECount) {
1632 Type *WideTy = IntegerType::get(getContext(), BitWidth * 2);
1633 // Check whether Start+Step*MaxBECount has no unsigned overflow.
1634 const SCEV *ZMul =
1635 getMulExpr(CastedMaxBECount, Step, SCEV::FlagNone, Depth + 1);
1636 const SCEV *ZAdd = getZeroExtendExpr(
1637 getAddExpr(Start, ZMul, SCEV::FlagNone, Depth + 1), WideTy,
1638 Depth + 1);
1639 const SCEV *WideStart = getZeroExtendExpr(Start, WideTy, Depth + 1);
1640 const SCEV *WideMaxBECount =
1641 getZeroExtendExpr(CastedMaxBECount, WideTy, Depth + 1);
1642 const SCEV *OperandExtendedAdd =
1643 getAddExpr(WideStart,
1644 getMulExpr(WideMaxBECount,
1645 getZeroExtendExpr(Step, WideTy, Depth + 1),
1646 SCEV::FlagNone, Depth + 1),
1647 SCEV::FlagNone, Depth + 1);
1648 if (ZAdd == OperandExtendedAdd) {
1649 // Cache knowledge of AR NUW, which is propagated to this AddRec.
1650 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNUW);
1651 // Return the expression with the addrec on the outside.
1652 Start =
1654 Step = getZeroExtendExpr(Step, Ty, Depth + 1);
1655 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1656 }
1657 // Similar to above, only this time treat the step value as signed.
1658 // This covers loops that count down.
1659 OperandExtendedAdd =
1660 getAddExpr(WideStart,
1661 getMulExpr(WideMaxBECount,
1662 getSignExtendExpr(Step, WideTy, Depth + 1),
1663 SCEV::FlagNone, Depth + 1),
1664 SCEV::FlagNone, Depth + 1);
1665 if (ZAdd == OperandExtendedAdd) {
1666 // Cache knowledge of AR NW, which is propagated to this AddRec.
1667 // Negative step causes unsigned wrap, but it still can't self-wrap.
1668 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNW);
1669 // Return the expression with the addrec on the outside.
1670 Start =
1672 Step = getSignExtendExpr(Step, Ty, Depth + 1);
1673 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1674 }
1675 }
1676 }
1677
1678 // Normally, in the cases we can prove no-overflow via a
1679 // backedge guarding condition, we can also compute a backedge
1680 // taken count for the loop. The exceptions are assumptions and
1681 // guards present in the loop -- SCEV is not great at exploiting
1682 // these to compute max backedge taken counts, but can still use
1683 // these to prove lack of overflow. Use this fact to avoid
1684 // doing extra work that may not pay off.
1685 if (!isa<SCEVCouldNotCompute>(MaxBECount) || HasGuards ||
1686 !AC.assumptions().empty()) {
1687
1688 auto NewFlags = proveNoUnsignedWrapViaInduction(AR);
1689 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), NewFlags);
1690 if (AR->hasNoUnsignedWrap()) {
1691 // Same as nuw case above - duplicated here to avoid a compile time
1692 // issue. It's not clear that the order of checks does matter, but
1693 // it's one of two issue possible causes for a change which was
1694 // reverted. Be conservative for the moment.
1695 Start =
1697 Step = getZeroExtendExpr(Step, Ty, Depth + 1);
1698 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1699 }
1700
1701 // For a negative step, we can extend the operands iff doing so only
1702 // traverses values in the range zext([0,UINT_MAX]).
1703 if (isKnownNegative(Step)) {
1704 const SCEV *N =
1708 // Cache knowledge of AR NW, which is propagated to this
1709 // AddRec. Negative step causes unsigned wrap, but it
1710 // still can't self-wrap.
1711 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNW);
1712 // Return the expression with the addrec on the outside.
1713 Start =
1715 Step = getSignExtendExpr(Step, Ty, Depth + 1);
1716 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1717 }
1718 }
1719 }
1720
1721 // zext({C,+,Step}) --> (zext(D) + zext({C-D,+,Step}))<nuw><nsw>
1722 // if D + (C - D + Step * n) could be proven to not unsigned wrap
1723 // where D maximizes the number of trailing zeros of (C - D + Step * n)
1724 if (const auto *SC = dyn_cast<SCEVConstant>(Start)) {
1725 const APInt &C = SC->getAPInt();
1726 const APInt &D = extractConstantWithoutWrapping(*this, C, Step);
1727 if (D != 0) {
1728 const SCEV *SZExtD = getZeroExtendExpr(getConstant(D), Ty, Depth);
1729 const SCEV *SResidual =
1730 getAddRecExpr(getConstant(C - D), Step, L, AR->getNoWrapFlags());
1731 const SCEV *SZExtR = getZeroExtendExpr(SResidual, Ty, Depth + 1);
1732 return getAddExpr(SZExtD, SZExtR, SCEV::FlagNSW | SCEV::FlagNUW,
1733 Depth + 1);
1734 }
1735 }
1736
1737 if (proveNoWrapByVaryingStart<SCEVZeroExtendExpr>(Start, Step, L)) {
1738 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNUW);
1739 Start = getExtendAddRecStart<SCEVZeroExtendExpr>(AR, Ty, this, Depth + 1);
1740 Step = getZeroExtendExpr(Step, Ty, Depth + 1);
1741 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1742 }
1743 }
1744
1745 // zext(A % B) --> zext(A) % zext(B)
1746 {
1747 const SCEV *LHS;
1748 const SCEV *RHS;
1749 if (match(Op, m_scev_URem(m_SCEV(LHS), m_SCEV(RHS), *this)))
1750 return getURemExpr(getZeroExtendExpr(LHS, Ty, Depth + 1),
1751 getZeroExtendExpr(RHS, Ty, Depth + 1));
1752 }
1753
1754 // zext(A / B) --> zext(A) / zext(B).
1755 if (auto *Div = dyn_cast<SCEVUDivExpr>(Op))
1756 return getUDivExpr(getZeroExtendExpr(Div->getLHS(), Ty, Depth + 1),
1757 getZeroExtendExpr(Div->getRHS(), Ty, Depth + 1));
1758
1759 if (auto *SA = dyn_cast<SCEVAddExpr>(Op)) {
1760 // zext((A + B + ...)<nuw>) --> (zext(A) + zext(B) + ...)<nuw>
1761 if (SA->hasNoUnsignedWrap()) {
1762 // If the addition does not unsign overflow then we can, by definition,
1763 // commute the zero extension with the addition operation.
1765 for (SCEVUse Op : SA->operands())
1766 Ops.push_back(getZeroExtendExpr(Op, Ty, Depth + 1));
1767 return getAddExpr(Ops, SCEV::FlagNUW, Depth + 1);
1768 }
1769
1770 const APInt *C, *C2;
1771 // zext (C + A)<nsw> -> (sext(C) + sext(A))<nsw> if zext (C + A)<nsw> >=s 0.
1772 // Currently the non-negative check is done manually, as isKnownNonNegative
1773 // is too expensive.
1774 if (SA->hasNoSignedWrap() &&
1776 m_scev_SMax(m_scev_APInt(C2), m_SCEV()))) &&
1777 C->isNegative() && !C->isMinSignedValue() && C2->sge(C->abs())) {
1778 assert(isKnownNonNegative(SA) && "incorrectly determined non-negative");
1779 return getAddExpr(getSignExtendExpr(SA->getOperand(0), Ty, Depth + 1),
1780 getSignExtendExpr(SA->getOperand(1), Ty, Depth + 1),
1781 SCEV::FlagNSW, Depth + 1);
1782 }
1783
1784 // zext(C + x + y + ...) --> (zext(D) + zext((C - D) + x + y + ...))
1785 // if D + (C - D + x + y + ...) could be proven to not unsigned wrap
1786 // where D maximizes the number of trailing zeros of (C - D + x + y + ...)
1787 //
1788 // Often address arithmetics contain expressions like
1789 // (zext (add (shl X, C1), C2)), for instance, (zext (5 + (4 * X))).
1790 // This transformation is useful while proving that such expressions are
1791 // equal or differ by a small constant amount, see LoadStoreVectorizer pass.
1792 if (const auto *SC = dyn_cast<SCEVConstant>(SA->getOperand(0))) {
1793 const APInt &D = extractConstantWithoutWrapping(*this, SC, SA);
1794 if (D != 0) {
1795 const SCEV *SZExtD = getZeroExtendExpr(getConstant(D), Ty, Depth);
1796 const SCEV *SResidual =
1798 const SCEV *SZExtR = getZeroExtendExpr(SResidual, Ty, Depth + 1);
1799 return getAddExpr(SZExtD, SZExtR, (SCEV::FlagNSW | SCEV::FlagNUW),
1800 Depth + 1);
1801 }
1802 }
1803 }
1804
1805 if (auto *SM = dyn_cast<SCEVMulExpr>(Op)) {
1806 // zext((A * B * ...)<nuw>) --> (zext(A) * zext(B) * ...)<nuw>
1807 if (SM->hasNoUnsignedWrap()) {
1808 // If the multiply does not unsign overflow then we can, by definition,
1809 // commute the zero extension with the multiply operation.
1811 for (SCEVUse Op : SM->operands())
1812 Ops.push_back(getZeroExtendExpr(Op, Ty, Depth + 1));
1813 return getMulExpr(Ops, SCEV::FlagNUW, Depth + 1);
1814 }
1815
1816 // zext(2^K * (trunc X to iN)) to iM ->
1817 // 2^K * (zext(trunc X to i{N-K}) to iM)<nuw>
1818 //
1819 // Proof:
1820 //
1821 // zext(2^K * (trunc X to iN)) to iM
1822 // = zext((trunc X to iN) << K) to iM
1823 // = zext((trunc X to i{N-K}) << K)<nuw> to iM
1824 // (because shl removes the top K bits)
1825 // = zext((2^K * (trunc X to i{N-K}))<nuw>) to iM
1826 // = (2^K * (zext(trunc X to i{N-K}) to iM))<nuw>.
1827 //
1828 const APInt *C;
1829 const SCEV *TruncRHS;
1830 if (match(SM,
1831 m_scev_Mul(m_scev_APInt(C), m_scev_Trunc(m_SCEV(TruncRHS)))) &&
1832 C->isPowerOf2()) {
1833 int NewTruncBits =
1834 getTypeSizeInBits(SM->getOperand(1)->getType()) - C->logBase2();
1835 Type *NewTruncTy = IntegerType::get(getContext(), NewTruncBits);
1836 return getMulExpr(
1837 getZeroExtendExpr(SM->getOperand(0), Ty),
1838 getZeroExtendExpr(getTruncateExpr(TruncRHS, NewTruncTy), Ty),
1839 SCEV::FlagNUW, Depth + 1);
1840 }
1841 }
1842
1843 // zext(umin(x, y)) -> umin(zext(x), zext(y))
1844 // zext(umax(x, y)) -> umax(zext(x), zext(y))
1848 for (SCEVUse Operand : MinMax->operands())
1849 Operands.push_back(getZeroExtendExpr(Operand, Ty));
1851 return getUMinExpr(Operands);
1852 return getUMaxExpr(Operands);
1853 }
1854
1855 // zext(umin_seq(x, y)) -> umin_seq(zext(x), zext(y))
1857 assert(isa<SCEVSequentialUMinExpr>(MinMax) && "Not supported!");
1859 for (SCEVUse Operand : MinMax->operands())
1860 Operands.push_back(getZeroExtendExpr(Operand, Ty));
1861 return getUMinExpr(Operands, /*Sequential*/ true);
1862 }
1863
1864 // The cast wasn't folded; create an explicit cast node.
1865 // Recompute the insert position, as it may have been invalidated.
1866 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
1867 return S;
1868 SCEV *S = new (SCEVAllocator) SCEVZeroExtendExpr(ID.Intern(SCEVAllocator),
1869 Op, Ty);
1870 UniqueSCEVs.insert(S, Token);
1871 S->computeAndSetCanonical(*this);
1872 registerUser(S, Op);
1873 return S;
1874}
1875
1877 unsigned Depth) {
1878 assert(getTypeSizeInBits(Op->getType()) < getTypeSizeInBits(Ty) &&
1879 "This is not an extending conversion!");
1880 assert(isSCEVable(Ty) &&
1881 "This is not a conversion to a SCEVable type!");
1882 assert(!Op->getType()->isPointerTy() && "Can't extend pointer!");
1883 Ty = getEffectiveSCEVType(Ty);
1884
1885 FoldID ID(scSignExtend, Op, Ty);
1886 if (const SCEV *S = FoldCache.lookup(ID))
1887 return S;
1888
1889 const SCEV *S = getSignExtendExprImpl(Op, Ty, Depth);
1891 insertFoldCacheEntry(ID, S, FoldCache, FoldCacheUser);
1892 return S;
1893}
1894
1896 unsigned Depth) {
1897 assert(getTypeSizeInBits(Op->getType()) < getTypeSizeInBits(Ty) &&
1898 "This is not an extending conversion!");
1899 assert(isSCEVable(Ty) && "This is not a conversion to a SCEVable type!");
1900 assert(!Op->getType()->isPointerTy() && "Can't extend pointer!");
1901 Ty = getEffectiveSCEVType(Ty);
1902
1903 // Fold if the operand is constant.
1904 if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(Op))
1905 return getConstant(SC->getAPInt().sext(getTypeSizeInBits(Ty)));
1906
1907 // sext(sext(x)) --> sext(x)
1909 return getSignExtendExpr(SS->getOperand(), Ty, Depth + 1);
1910
1911 // sext(zext(x)) --> zext(x)
1913 return getZeroExtendExpr(SZ->getOperand(), Ty, Depth + 1);
1914
1915 // If the operand is an affine AddRec with the no-signed-wrap flag, the
1916 // sign-extension distributes over the recurrence.
1917 const SCEV *Start, *Step;
1918 const Loop *L;
1919 if (Depth <= MaxCastDepth &&
1920 match(Op, m_scev_AffineAddRec(m_SCEV(Start), m_SCEV(Step), m_Loop(L)))) {
1921 const auto *AR = cast<SCEVAddRecExpr>(Op);
1922 if (AR->hasNoSignedWrap()) {
1923 Start = getExtendAddRecStart<SCEVSignExtendExpr>(AR, Ty, this, Depth + 1);
1924 Step = getSignExtendExpr(Step, Ty, Depth + 1);
1925 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1926 }
1927 }
1928
1929 // Before doing any expensive analysis, check to see if we've already
1930 // computed a SCEV for this Op and Ty.
1933 ID.AddPointer(Op.getOpaqueValue());
1934 ID.AddPointer(Ty);
1936 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
1937 return S;
1938 // Limit recursion depth.
1939 if (Depth > MaxCastDepth) {
1940 SCEV *S = new (SCEVAllocator) SCEVSignExtendExpr(ID.Intern(SCEVAllocator),
1941 Op, Ty);
1942 UniqueSCEVs.insert(S, Token);
1943 S->computeAndSetCanonical(*this);
1944 registerUser(S, Op);
1945 return S;
1946 }
1947
1948 // sext(trunc(x)) --> sext(x) or x or trunc(x)
1950 // It's possible the bits taken off by the truncate were all sign bits. If
1951 // so, we should be able to simplify this further.
1952 const SCEV *X = ST->getOperand();
1954 unsigned TruncBits = getTypeSizeInBits(ST->getType());
1955 unsigned NewBits = getTypeSizeInBits(Ty);
1956 if (CR.truncate(TruncBits).signExtend(NewBits).contains(
1957 CR.sextOrTrunc(NewBits)))
1958 return getTruncateOrSignExtend(X, Ty, Depth);
1959 }
1960
1961 if (auto *SA = dyn_cast<SCEVAddExpr>(Op)) {
1962 // sext((A + B + ...)<nsw>) --> (sext(A) + sext(B) + ...)<nsw>
1963 if (SA->hasNoSignedWrap()) {
1964 // If the addition does not sign overflow then we can, by definition,
1965 // commute the sign extension with the addition operation.
1967 for (SCEVUse Op : SA->operands())
1968 Ops.push_back(getSignExtendExpr(Op, Ty, Depth + 1));
1969 return getAddExpr(Ops, SCEV::FlagNSW, Depth + 1);
1970 }
1971
1972 // sext(C + x + y + ...) --> (sext(D) + sext((C - D) + x + y + ...))
1973 // if D + (C - D + x + y + ...) could be proven to not signed wrap
1974 // where D maximizes the number of trailing zeros of (C - D + x + y + ...)
1975 //
1976 // For instance, this will bring two seemingly different expressions:
1977 // 1 + sext(5 + 20 * %x + 24 * %y) and
1978 // sext(6 + 20 * %x + 24 * %y)
1979 // to the same form:
1980 // 2 + sext(4 + 20 * %x + 24 * %y)
1981 if (const auto *SC = dyn_cast<SCEVConstant>(SA->getOperand(0))) {
1982 const APInt &D = extractConstantWithoutWrapping(*this, SC, SA);
1983 if (D != 0) {
1984 const SCEV *SSExtD = getSignExtendExpr(getConstant(D), Ty, Depth);
1985 const SCEV *SResidual =
1987 const SCEV *SSExtR = getSignExtendExpr(SResidual, Ty, Depth + 1);
1988 return getAddExpr(SSExtD, SSExtR, (SCEV::FlagNSW | SCEV::FlagNUW),
1989 Depth + 1);
1990 }
1991 }
1992 }
1993 // If the input value is a chrec scev, and we can prove that the value
1994 // did not overflow the old, smaller, value, we can sign extend all of the
1995 // operands (often constants). This allows analysis of something like
1996 // this: for (signed char X = 0; X < 100; ++X) { int Y = X; }
1997 if (match(Op, m_scev_AffineAddRec(m_SCEV(Start), m_SCEV(Step), m_Loop(L)))) {
1998 const auto *AR = cast<SCEVAddRecExpr>(Op);
1999 unsigned BitWidth = getTypeSizeInBits(AR->getType());
2000
2001 // The no-signed-wrap case is handled before the uniquing lookup above.
2002
2003 // Check whether the backedge-taken count is SCEVCouldNotCompute.
2004 // Note that this serves two purposes: It filters out loops that are
2005 // simply not analyzable, and it covers the case where this code is
2006 // being called from within backedge-taken count analysis, such that
2007 // attempting to ask for the backedge-taken count would likely result
2008 // in infinite recursion. In the later case, the analysis code will
2009 // cope with a conservative value, and it will take care to purge
2010 // that value once it has finished.
2011 const SCEV *MaxBECount = getConstantMaxBackedgeTakenCount(L);
2012 if (!isa<SCEVCouldNotCompute>(MaxBECount)) {
2013 // Manually compute the final value for AR, checking for
2014 // overflow.
2015
2016 // Check whether the backedge-taken count can be losslessly casted to
2017 // the addrec's type. The count is always unsigned.
2018 const SCEV *CastedMaxBECount =
2019 getTruncateOrZeroExtend(MaxBECount, Start->getType(), Depth);
2020 const SCEV *RecastedMaxBECount = getTruncateOrZeroExtend(
2021 CastedMaxBECount, MaxBECount->getType(), Depth);
2022 if (MaxBECount == RecastedMaxBECount) {
2023 Type *WideTy = IntegerType::get(getContext(), BitWidth * 2);
2024 // Check whether Start+Step*MaxBECount has no signed overflow.
2025 const SCEV *SMul =
2026 getMulExpr(CastedMaxBECount, Step, SCEV::FlagNone, Depth + 1);
2027 const SCEV *SAdd = getSignExtendExpr(
2028 getAddExpr(Start, SMul, SCEV::FlagNone, Depth + 1), WideTy,
2029 Depth + 1);
2030 const SCEV *WideStart = getSignExtendExpr(Start, WideTy, Depth + 1);
2031 const SCEV *WideMaxBECount =
2032 getZeroExtendExpr(CastedMaxBECount, WideTy, Depth + 1);
2033 const SCEV *OperandExtendedAdd =
2034 getAddExpr(WideStart,
2035 getMulExpr(WideMaxBECount,
2036 getSignExtendExpr(Step, WideTy, Depth + 1),
2037 SCEV::FlagNone, Depth + 1),
2038 SCEV::FlagNone, Depth + 1);
2039 if (SAdd == OperandExtendedAdd) {
2040 // Cache knowledge of AR NSW, which is propagated to this AddRec.
2041 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNSW);
2042 // Return the expression with the addrec on the outside.
2043 Start =
2045 Step = getSignExtendExpr(Step, Ty, Depth + 1);
2046 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
2047 }
2048 // Similar to above, only this time treat the step value as unsigned.
2049 // This covers loops that count up with an unsigned step.
2050 OperandExtendedAdd =
2051 getAddExpr(WideStart,
2052 getMulExpr(WideMaxBECount,
2053 getZeroExtendExpr(Step, WideTy, Depth + 1),
2054 SCEV::FlagNone, Depth + 1),
2055 SCEV::FlagNone, Depth + 1);
2056 if (SAdd == OperandExtendedAdd) {
2057 // If AR wraps around then
2058 //
2059 // abs(Step) * MaxBECount > unsigned-max(AR->getType())
2060 // => SAdd != OperandExtendedAdd
2061 //
2062 // Thus (AR is not NW => SAdd != OperandExtendedAdd) <=>
2063 // (SAdd == OperandExtendedAdd => AR is NW)
2064
2065 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNW);
2066
2067 // Return the expression with the addrec on the outside.
2068 Start =
2070 Step = getZeroExtendExpr(Step, Ty, Depth + 1);
2071 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
2072 }
2073 }
2074 }
2075
2076 auto NewFlags = proveNoSignedWrapViaInduction(AR);
2077 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), NewFlags);
2078 if (AR->hasNoSignedWrap()) {
2079 // Same as nsw case above - duplicated here to avoid a compile time
2080 // issue. It's not clear that the order of checks does matter, but
2081 // it's one of two issue possible causes for a change which was
2082 // reverted. Be conservative for the moment.
2083 Start = getExtendAddRecStart<SCEVSignExtendExpr>(AR, Ty, this, Depth + 1);
2084 Step = getSignExtendExpr(Step, Ty, Depth + 1);
2085 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
2086 }
2087
2088 // sext({C,+,Step}) --> (sext(D) + sext({C-D,+,Step}))<nuw><nsw>
2089 // if D + (C - D + Step * n) could be proven to not signed wrap
2090 // where D maximizes the number of trailing zeros of (C - D + Step * n)
2091 if (const auto *SC = dyn_cast<SCEVConstant>(Start)) {
2092 const APInt &C = SC->getAPInt();
2093 const APInt &D = extractConstantWithoutWrapping(*this, C, Step);
2094 if (D != 0) {
2095 const SCEV *SSExtD = getSignExtendExpr(getConstant(D), Ty, Depth);
2096 const SCEV *SResidual =
2097 getAddRecExpr(getConstant(C - D), Step, L, AR->getNoWrapFlags());
2098 const SCEV *SSExtR = getSignExtendExpr(SResidual, Ty, Depth + 1);
2099 return getAddExpr(SSExtD, SSExtR, (SCEV::FlagNSW | SCEV::FlagNUW),
2100 Depth + 1);
2101 }
2102 }
2103
2104 if (proveNoWrapByVaryingStart<SCEVSignExtendExpr>(Start, Step, L)) {
2105 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNSW);
2106 Start = getExtendAddRecStart<SCEVSignExtendExpr>(AR, Ty, this, Depth + 1);
2107 Step = getSignExtendExpr(Step, Ty, Depth + 1);
2108 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
2109 }
2110 }
2111
2112 // If the input value is provably positive and we could not simplify
2113 // away the sext build a zext instead.
2115 return getZeroExtendExpr(Op, Ty, Depth + 1);
2116
2117 // sext(smin(x, y)) -> smin(sext(x), sext(y))
2118 // sext(smax(x, y)) -> smax(sext(x), sext(y))
2122 for (SCEVUse Operand : MinMax->operands())
2123 Operands.push_back(getSignExtendExpr(Operand, Ty));
2125 return getSMinExpr(Operands);
2126 return getSMaxExpr(Operands);
2127 }
2128
2129 // The cast wasn't folded; create an explicit cast node.
2130 // Recompute the insert position, as it may have been invalidated.
2131 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
2132 return S;
2133 SCEV *S = new (SCEVAllocator) SCEVSignExtendExpr(ID.Intern(SCEVAllocator),
2134 Op, Ty);
2135 UniqueSCEVs.insert(S, Token);
2136 S->computeAndSetCanonical(*this);
2137 registerUser(S, Op);
2138 return S;
2139}
2140
2142 switch (Kind) {
2143 case scTruncate:
2144 return getTruncateExpr(Op, Ty);
2145 case scZeroExtend:
2146 return getZeroExtendExpr(Op, Ty);
2147 case scSignExtend:
2148 return getSignExtendExpr(Op, Ty);
2149 case scPtrToAddr: {
2150 const SCEV *Expr = getPtrToAddrExpr(Op);
2151 assert(Expr->getType() == Ty && "requested type must match");
2152 return Expr;
2153 }
2154 default:
2155 llvm_unreachable("Not a SCEV cast expression!");
2156 }
2157}
2158
2159/// getAnyExtendExpr - Return a SCEV for the given operand extended with
2160/// unspecified bits out to the given type.
2162 assert(getTypeSizeInBits(Op->getType()) < getTypeSizeInBits(Ty) &&
2163 "This is not an extending conversion!");
2164 assert(isSCEVable(Ty) &&
2165 "This is not a conversion to a SCEVable type!");
2166 Ty = getEffectiveSCEVType(Ty);
2167
2168 // Sign-extend negative constants.
2169 if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(Op))
2170 if (SC->getAPInt().isNegative())
2171 return getSignExtendExpr(Op, Ty);
2172
2173 // Peel off a truncate cast.
2175 const SCEV *NewOp = T->getOperand();
2176 if (getTypeSizeInBits(NewOp->getType()) < getTypeSizeInBits(Ty))
2177 return getAnyExtendExpr(NewOp, Ty);
2178 return getTruncateOrNoop(NewOp, Ty);
2179 }
2180
2181 // Next try a zext cast. If the cast is folded, use it.
2182 const SCEV *ZExt = getZeroExtendExpr(Op, Ty);
2183 if (!isa<SCEVZeroExtendExpr>(ZExt))
2184 return ZExt;
2185
2186 // Next try a sext cast. If the cast is folded, use it.
2187 const SCEV *SExt = getSignExtendExpr(Op, Ty);
2188 if (!isa<SCEVSignExtendExpr>(SExt))
2189 return SExt;
2190
2191 // Force the cast to be folded into the operands of an addrec.
2192 if (const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(Op)) {
2194 for (const SCEV *Op : AR->operands())
2195 Ops.push_back(getAnyExtendExpr(Op, Ty));
2196 return getAddRecExpr(Ops, AR->getLoop(), SCEV::FlagNW);
2197 }
2198
2199 // If the expression is obviously signed, use the sext cast value.
2200 if (isa<SCEVSMaxExpr>(Op))
2201 return SExt;
2202
2203 // Absent any other information, use the zext cast value.
2204 return ZExt;
2205}
2206
2207/// Process the given Ops list, which is a list of operands to be added under
2208/// the given scale, update the given map. This is a helper function for
2209/// getAddRecExpr. As an example of what it does, given a sequence of operands
2210/// that would form an add expression like this:
2211///
2212/// m + n + 13 + (A * (o + p + (B * (q + m + 29)))) + r + (-1 * r)
2213///
2214/// where A and B are constants, update the map with these values:
2215///
2216/// (m, 1+A*B), (n, 1), (o, A), (p, A), (q, A*B), (r, 0)
2217///
2218/// and add 13 + A*B*29 to AccumulatedConstant.
2219/// This will allow getAddRecExpr to produce this:
2220///
2221/// 13+A*B*29 + n + (m * (1+A*B)) + ((o + p) * A) + (q * A*B)
2222///
2223/// This form often exposes folding opportunities that are hidden in
2224/// the original operand list.
2225///
2226/// Return true iff it appears that any interesting folding opportunities
2227/// may be exposed. This helps getAddRecExpr short-circuit extra work in
2228/// the common case where no interesting opportunities are present, and
2229/// is also used as a check to avoid infinite recursion.
2232 APInt &AccumulatedConstant,
2234 const APInt &Scale,
2235 ScalarEvolution &SE) {
2236 bool Interesting = false;
2237
2238 // Iterate over the add operands. They are sorted, with constants first.
2239 unsigned i = 0;
2240 while (const SCEVConstant *C = dyn_cast<SCEVConstant>(Ops[i])) {
2241 ++i;
2242 // Pull a buried constant out to the outside.
2243 if (Scale != 1 || AccumulatedConstant != 0 || C->getValue()->isZero())
2244 Interesting = true;
2245 AccumulatedConstant += Scale * C->getAPInt();
2246 }
2247
2248 // Next comes everything else. We're especially interested in multiplies
2249 // here, but they're in the middle, so just visit the rest with one loop.
2250 for (; i != Ops.size(); ++i) {
2252 if (Mul && isa<SCEVConstant>(Mul->getOperand(0))) {
2253 APInt NewScale =
2254 Scale * cast<SCEVConstant>(Mul->getOperand(0))->getAPInt();
2255 if (Mul->getNumOperands() == 2 && isa<SCEVAddExpr>(Mul->getOperand(1))) {
2256 // A multiplication of a constant with another add; recurse.
2257 const SCEVAddExpr *Add = cast<SCEVAddExpr>(Mul->getOperand(1));
2258 Interesting |= CollectAddOperandsWithScales(
2259 M, NewOps, AccumulatedConstant, Add->operands(), NewScale, SE);
2260 } else {
2261 // A multiplication of a constant with some other value. Update
2262 // the map.
2263 SmallVector<SCEVUse, 4> MulOps(drop_begin(Mul->operands()));
2264 const SCEV *Key = SE.getMulExpr(MulOps);
2265 auto Pair = M.insert({Key, NewScale});
2266 if (Pair.second) {
2267 NewOps.push_back(Pair.first->first);
2268 } else {
2269 Pair.first->second += NewScale;
2270 // The map already had an entry for this value, which may indicate
2271 // a folding opportunity.
2272 Interesting = true;
2273 }
2274 }
2275 } else {
2276 // An ordinary operand. Update the map.
2277 auto Pair = M.insert({Ops[i], Scale});
2278 if (Pair.second) {
2279 NewOps.push_back(Pair.first->first);
2280 } else {
2281 Pair.first->second += Scale;
2282 // The map already had an entry for this value, which may indicate
2283 // a folding opportunity.
2284 Interesting = true;
2285 }
2286 }
2287 }
2288
2289 return Interesting;
2290}
2291
2293 const SCEV *LHS, const SCEV *RHS,
2294 const Instruction *CtxI) {
2295 auto Operation = [this, BinOp](SCEVUse L, SCEVUse R) -> const SCEV * {
2296 switch (BinOp) {
2297 default:
2298 llvm_unreachable("Unsupported binary op");
2299 case Instruction::Add:
2300 return getAddExpr(L, R);
2301 case Instruction::Sub:
2302 return getMinusSCEV(L, R);
2303 case Instruction::Mul:
2304 return getMulExpr(L, R);
2305 }
2306 };
2307
2308 const SCEV *(ScalarEvolution::*Extension)(SCEVUse, Type *, unsigned) =
2311
2312 // Check ext(LHS op RHS) == ext(LHS) op ext(RHS)
2313 auto *NarrowTy = cast<IntegerType>(LHS->getType());
2314 auto *WideTy =
2315 IntegerType::get(NarrowTy->getContext(), NarrowTy->getBitWidth() * 2);
2316
2317 const SCEV *A = (this->*Extension)(Operation(LHS, RHS), WideTy, 0);
2318 const SCEV *LHSB = (this->*Extension)(LHS, WideTy, 0);
2319 const SCEV *RHSB = (this->*Extension)(RHS, WideTy, 0);
2320 const SCEV *B = Operation(LHSB, RHSB);
2321 if (A == B)
2322 return true;
2323 // Can we use context to prove the fact we need?
2324 if (!CtxI)
2325 return false;
2326 // TODO: Support mul.
2327 if (BinOp == Instruction::Mul)
2328 return false;
2329 auto *RHSC = dyn_cast<SCEVConstant>(RHS);
2330 // TODO: Lift this limitation.
2331 if (!RHSC)
2332 return false;
2333 APInt C = RHSC->getAPInt();
2334 unsigned NumBits = C.getBitWidth();
2335 bool IsSub = (BinOp == Instruction::Sub);
2336 bool IsNegativeConst = (Signed && C.isNegative());
2337 // Compute the direction and magnitude by which we need to check overflow.
2338 bool OverflowDown = IsSub ^ IsNegativeConst;
2339 APInt Magnitude = C;
2340 if (IsNegativeConst) {
2341 if (C == APInt::getSignedMinValue(NumBits))
2342 // TODO: SINT_MIN on inversion gives the same negative value, we don't
2343 // want to deal with that.
2344 return false;
2345 Magnitude = -C;
2346 }
2347
2349 if (OverflowDown) {
2350 // To avoid overflow down, we need to make sure that MIN + Magnitude <= LHS.
2351 APInt Min = Signed ? APInt::getSignedMinValue(NumBits)
2352 : APInt::getMinValue(NumBits);
2353 APInt Limit = Min + Magnitude;
2354 return isKnownPredicateAt(Pred, getConstant(Limit), LHS, CtxI);
2355 } else {
2356 // To avoid overflow up, we need to make sure that LHS <= MAX - Magnitude.
2357 APInt Max = Signed ? APInt::getSignedMaxValue(NumBits)
2358 : APInt::getMaxValue(NumBits);
2359 APInt Limit = Max - Magnitude;
2360 return isKnownPredicateAt(Pred, LHS, getConstant(Limit), CtxI);
2361 }
2362}
2363
2365 const OverflowingBinaryOperator *OBO) {
2366 // It cannot be done any better.
2367 if (OBO->hasNoUnsignedWrap() && OBO->hasNoSignedWrap())
2368 return std::nullopt;
2369
2371
2372 if (OBO->hasNoUnsignedWrap())
2374 if (OBO->hasNoSignedWrap())
2376
2377 bool Deduced = false;
2378
2380 const SCEV *LHS = getSCEV(OBO->getOperand(0));
2381 const SCEV *RHS = getSCEV(OBO->getOperand(1));
2382
2383 bool CanUseNSW = true;
2384 const APInt *ShiftAmt;
2385 // Treat `shl %a, C` as `mul %a, 1 << C`.
2386 if (match(OBO, m_Shl(m_Value(), m_APInt(ShiftAmt)))) {
2387 unsigned BitWidth = ShiftAmt->getBitWidth();
2388 if (ShiftAmt->uge(BitWidth))
2389 return std::nullopt;
2390 // NSW only transfers if the shift amount is < BitWidth - 1, as INT_MIN * -1
2391 // overflows.
2392 CanUseNSW = ShiftAmt->ult(BitWidth - 1);
2393 Opcode = Instruction::Mul;
2395 } else if (Opcode != Instruction::Add && Opcode != Instruction::Sub &&
2396 Opcode != Instruction::Mul) {
2397 return std::nullopt;
2398 }
2399
2400 const Instruction *CtxI =
2402 if (!OBO->hasNoUnsignedWrap() &&
2403 willNotOverflow(Opcode, /* Signed */ false, LHS, RHS, CtxI)) {
2405 Deduced = true;
2406 }
2407
2408 if (CanUseNSW && !OBO->hasNoSignedWrap() &&
2409 willNotOverflow(Opcode, /* Signed */ true, LHS, RHS, CtxI)) {
2411 Deduced = true;
2412 }
2413
2414 if (Deduced)
2415 return Flags;
2416 return std::nullopt;
2417}
2418
2419// We're trying to construct a SCEV of type `Type' with `Ops' as operands and
2420// `OldFlags' as can't-wrap behavior. Infer a more aggressive set of
2421// can't-overflow flags for the operation if possible.
2424 using namespace std::placeholders;
2425
2426 using OBO = OverflowingBinaryOperator;
2427
2428 bool CanAnalyze =
2430 (void)CanAnalyze;
2431 assert(CanAnalyze && "don't call from other places!");
2432
2433 SCEVFlags SignOrUnsignMask = SCEV::FlagNUW | SCEV::FlagNSW;
2434 SCEVFlags SignOrUnsignWrap =
2435 ScalarEvolution::maskFlags(Flags, SignOrUnsignMask);
2436
2437 // If FlagNSW is true and all the operands are non-negative, infer FlagNUW.
2438 auto IsKnownNonNegative = [&](SCEVUse U) {
2439 return SE->isKnownNonNegative(U);
2440 };
2441
2442 if (SignOrUnsignWrap == SCEV::FlagNSW && all_of(Ops, IsKnownNonNegative))
2443 Flags = ScalarEvolution::setFlags(Flags, SignOrUnsignMask);
2444
2445 SignOrUnsignWrap = ScalarEvolution::maskFlags(Flags, SignOrUnsignMask);
2446
2447 if (SignOrUnsignWrap != SignOrUnsignMask &&
2448 (Type == scAddExpr || Type == scMulExpr) && Ops.size() == 2 &&
2449 isa<SCEVConstant>(Ops[0])) {
2450
2451 auto Opcode = [&] {
2452 switch (Type) {
2453 case scAddExpr:
2454 return Instruction::Add;
2455 case scMulExpr:
2456 return Instruction::Mul;
2457 default:
2458 llvm_unreachable("Unexpected SCEV op.");
2459 }
2460 }();
2461
2462 const APInt &C = cast<SCEVConstant>(Ops[0])->getAPInt();
2463
2464 // (A <opcode> C) --> (A <opcode> C)<nsw> if the op doesn't sign overflow.
2465 if (!(SignOrUnsignWrap & SCEV::FlagNSW)) {
2466 auto NSWRegion =
2467 ConstantRange::makeExactNoWrapRegion(Opcode, C, OBO::NoSignedWrap);
2468 if (NSWRegion.contains(SE->getSignedRange(Ops[1])))
2470 }
2471
2472 // (A <opcode> C) --> (A <opcode> C)<nuw> if the op doesn't unsign overflow.
2473 if (!(SignOrUnsignWrap & SCEV::FlagNUW)) {
2474 auto NUWRegion =
2475 ConstantRange::makeExactNoWrapRegion(Opcode, C, OBO::NoUnsignedWrap);
2476 if (NUWRegion.contains(SE->getUnsignedRange(Ops[1])))
2478 }
2479 }
2480
2481 // <0,+,nonnegative><nw> is also nuw
2482 // TODO: Add corresponding nsw case
2484 !ScalarEvolution::hasFlags(Flags, SCEV::FlagNUW) && Ops.size() == 2 &&
2485 Ops[0]->isZero() && IsKnownNonNegative(Ops[1]))
2487
2488 // both (udiv X, Y) * Y and Y * (udiv X, Y) are always NUW
2490 Ops.size() == 2) {
2491 if (auto *UDiv = dyn_cast<SCEVUDivExpr>(Ops[0]))
2492 if (UDiv->getOperand(1) == Ops[1])
2494 if (auto *UDiv = dyn_cast<SCEVUDivExpr>(Ops[1]))
2495 if (UDiv->getOperand(1) == Ops[0])
2497 }
2498
2499 return Flags;
2500}
2501
2503 return isLoopInvariant(S, L) && properlyDominates(S, L->getHeader());
2504}
2505
2506/// Get a canonical add expression, or something simpler if possible.
2508 SCEVFlagsPair Flags, unsigned Depth) {
2509 SCEVFlags ExprFlags = Flags.ExprFlags;
2510 SCEVFlags UseFlags = Flags.UseFlags;
2511 assert(!(ExprFlags & ~(SCEV::FlagNUW | SCEV::FlagNSW)) &&
2512 "only nuw or nsw allowed");
2513 assert(!(UseFlags & ~(SCEV::FlagNUW | SCEV::FlagNSW)) &&
2514 "only nuw or nsw allowed");
2515 assert(!Ops.empty() && "Cannot get empty add!");
2516 if (Ops.size() == 1) return Ops[0];
2517#ifndef NDEBUG
2518 Type *ETy = getEffectiveSCEVType(Ops[0]->getType());
2519 for (unsigned i = 1, e = Ops.size(); i != e; ++i)
2520 assert(getEffectiveSCEVType(Ops[i]->getType()) == ETy &&
2521 "SCEVAddExpr operand types don't match!");
2522 unsigned NumPtrs = count_if(
2523 Ops, [](const SCEV *Op) { return Op->getType()->isPointerTy(); });
2524 assert(NumPtrs <= 1 && "add has at most one pointer operand");
2525#endif
2526
2527 const SCEV *Folded = constantFoldAndGroupOps(
2528 *this, LI, DT, Ops,
2529 [](const APInt &C1, const APInt &C2) { return C1 + C2; },
2530 [](const APInt &C) { return C.isZero(); }, // identity
2531 [](const APInt &C) { return false; }); // absorber
2532 if (Folded)
2533 return Folded;
2534
2535#ifndef NDEBUG
2536 // Keep track of operands after constant folding, for verification when adding
2537 // use-specific flags.
2538 const SmallVector<SCEVUse, 8> OrigOps(Ops.begin(), Ops.end());
2539#endif
2540
2541 unsigned Idx = isa<SCEVConstant>(Ops[0]) ? 1 : 0;
2542
2543 // Delay expensive flag strengthening until necessary.
2544 auto ComputeFlags = [this, ExprFlags](ArrayRef<SCEVUse> Ops) {
2545 return StrengthenNoWrapFlags(this, scAddExpr, Ops, ExprFlags);
2546 };
2547
2548 // Limit recursion calls depth.
2550 return {getOrCreateAddExpr(Ops, ComputeFlags(Ops)), UseFlags};
2551
2552 if (SCEV *S = findExistingSCEVInCache(scAddExpr, Ops)) {
2553 // Don't strengthen flags if we have no new information.
2554 SCEVAddExpr *Add = static_cast<SCEVAddExpr *>(S);
2555 if (Add->getNoWrapFlags(ExprFlags) != ExprFlags)
2556 Add->setNoWrapFlags(ComputeFlags(Ops));
2557 return {S, UseFlags};
2558 }
2559
2560 // Okay, check to see if the same value occurs in the operand list more than
2561 // once. If so, merge them together into an multiply expression. Since we
2562 // sorted the list, these values are required to be adjacent.
2563 Type *Ty = Ops[0]->getType();
2564 bool FoundMatch = false;
2565 for (unsigned i = 0, e = Ops.size(); i != e-1; ++i)
2566 if (Ops[i] == Ops[i+1]) { // X + Y + Y --> X + Y*2
2567 // Scan ahead to count how many equal operands there are.
2568 unsigned Count = 2;
2569 while (i+Count != e && Ops[i+Count] == Ops[i])
2570 ++Count;
2571 // Merge the values into a multiply.
2572 SCEVUse Scale = getConstant(Ty, Count);
2573 const SCEV *Mul = getMulExpr(Scale, Ops[i], SCEV::FlagNone, Depth + 1);
2574 if (Ops.size() == Count)
2575 return Mul;
2576 Ops[i] = Mul;
2577 Ops.erase(Ops.begin()+i+1, Ops.begin()+i+Count);
2578 --i; e -= Count - 1;
2579 FoundMatch = true;
2580 }
2581 if (FoundMatch)
2582 return getAddExpr(Ops, ExprFlags, Depth + 1);
2583
2584 // Check for truncates. If all the operands are truncated from the same
2585 // type, see if factoring out the truncate would permit the result to be
2586 // folded. eg., n*trunc(x) + m*trunc(y) --> trunc(trunc(m)*x + trunc(n)*y)
2587 // if the contents of the resulting outer trunc fold to something simple.
2588 auto FindTruncSrcType = [&]() -> Type * {
2589 // We're ultimately looking to fold an addrec of truncs and muls of only
2590 // constants and truncs, so if we find any other types of SCEV
2591 // as operands of the addrec then we bail and return nullptr here.
2592 // Otherwise, we return the type of the operand of a trunc that we find.
2593 if (auto *T = dyn_cast<SCEVTruncateExpr>(Ops[Idx]))
2594 return T->getOperand()->getType();
2595 if (const auto *Mul = dyn_cast<SCEVMulExpr>(Ops[Idx])) {
2596 SCEVUse LastOp = Mul->getOperand(Mul->getNumOperands() - 1);
2597 if (const auto *T = dyn_cast<SCEVTruncateExpr>(LastOp))
2598 return T->getOperand()->getType();
2599 }
2600 return nullptr;
2601 };
2602 if (auto *SrcType = FindTruncSrcType()) {
2603 SmallVector<SCEVUse, 8> LargeOps;
2604 bool Ok = true;
2605 // Check all the operands to see if they can be represented in the
2606 // source type of the truncate.
2607 for (const SCEV *Op : Ops) {
2609 if (T->getOperand()->getType() != SrcType) {
2610 Ok = false;
2611 break;
2612 }
2613 LargeOps.push_back(T->getOperand());
2614 } else if (const SCEVConstant *C = dyn_cast<SCEVConstant>(Op)) {
2615 LargeOps.push_back(getAnyExtendExpr(C, SrcType));
2616 } else if (const SCEVMulExpr *M = dyn_cast<SCEVMulExpr>(Op)) {
2617 SmallVector<SCEVUse, 8> LargeMulOps;
2618 for (unsigned j = 0, f = M->getNumOperands(); j != f && Ok; ++j) {
2619 if (const SCEVTruncateExpr *T =
2620 dyn_cast<SCEVTruncateExpr>(M->getOperand(j))) {
2621 if (T->getOperand()->getType() != SrcType) {
2622 Ok = false;
2623 break;
2624 }
2625 LargeMulOps.push_back(T->getOperand());
2626 } else if (const auto *C = dyn_cast<SCEVConstant>(M->getOperand(j))) {
2627 LargeMulOps.push_back(getAnyExtendExpr(C, SrcType));
2628 } else {
2629 Ok = false;
2630 break;
2631 }
2632 }
2633 if (Ok)
2634 LargeOps.push_back(
2635 getMulExpr(LargeMulOps, SCEV::FlagNone, Depth + 1));
2636 } else {
2637 Ok = false;
2638 break;
2639 }
2640 }
2641 if (Ok) {
2642 // Evaluate the expression in the larger type.
2643 const SCEV *Fold = getAddExpr(LargeOps, SCEV::FlagNone, Depth + 1);
2644 // If it folds to something simple, use it. Otherwise, don't.
2645 if (isa<SCEVConstant>(Fold) || isa<SCEVUnknown>(Fold))
2646 return getTruncateExpr(Fold, Ty);
2647 }
2648 }
2649
2650 if (Ops.size() == 2) {
2651 // Check if we have an expression of the form ((X + C1) - C2), where C1 and
2652 // C2 can be folded in a way that allows retaining wrapping flags of (X +
2653 // C1).
2654 const SCEV *A = Ops[0];
2655 const SCEV *B = Ops[1];
2656 auto *AddExpr = dyn_cast<SCEVAddExpr>(B);
2657 auto *C = dyn_cast<SCEVConstant>(A);
2658 if (AddExpr && C && isa<SCEVConstant>(AddExpr->getOperand(0))) {
2659 auto C1 = cast<SCEVConstant>(AddExpr->getOperand(0))->getAPInt();
2660 auto C2 = C->getAPInt();
2661 SCEVFlags PreservedFlags = SCEV::FlagNone;
2662
2663 APInt ConstAdd = C1 + C2;
2664 auto AddFlags = AddExpr->getNoWrapFlags();
2665 // Adding a smaller constant is NUW if the original AddExpr was NUW.
2667 ConstAdd.ule(C1)) {
2668 PreservedFlags =
2670 }
2671
2672 // Adding a constant with the same sign and small magnitude is NSW, if the
2673 // original AddExpr was NSW.
2675 C1.isSignBitSet() == ConstAdd.isSignBitSet() &&
2676 ConstAdd.abs().ule(C1.abs())) {
2677 PreservedFlags =
2679 }
2680
2681 if (PreservedFlags != SCEV::FlagNone) {
2682 SmallVector<SCEVUse, 4> NewOps(AddExpr->operands());
2683 NewOps[0] = getConstant(ConstAdd);
2684 return getAddExpr(NewOps, PreservedFlags);
2685 }
2686 }
2687
2688 // Try to push the constant operand into a ZExt: A + zext (-A + B) -> zext
2689 // (B), if trunc (A) + -A + B does not unsigned-wrap.
2690 const SCEVAddExpr *InnerAdd;
2691 if (match(B, m_scev_ZExt(m_scev_Add(InnerAdd)))) {
2692 const SCEV *NarrowA = getTruncateExpr(A, InnerAdd->getType());
2693 if (NarrowA == getNegativeSCEV(InnerAdd->getOperand(0)) &&
2694 getZeroExtendExpr(NarrowA, B->getType()) == A &&
2695 hasFlags(StrengthenNoWrapFlags(this, scAddExpr, {NarrowA, InnerAdd},
2697 SCEV::FlagNUW)) {
2698 return getZeroExtendExpr(getAddExpr(NarrowA, InnerAdd), B->getType());
2699 }
2700 }
2701 }
2702
2703 // Canonicalize (-1 * urem X, Y) + X --> (Y * X/Y)
2704 const SCEV *Y;
2705 if (Ops.size() == 2 &&
2706 match(Ops[0],
2708 m_scev_URem(m_scev_Specific(Ops[1]), m_SCEV(Y), *this))))
2709 return getMulExpr(Y, getUDivExpr(Ops[1], Y));
2710
2711 // Skip past any other cast SCEVs.
2712 while (Idx < Ops.size() && Ops[Idx]->getSCEVType() < scAddExpr)
2713 ++Idx;
2714
2715 // If there are add operands they would be next.
2716 if (Idx < Ops.size()) {
2717 bool DeletedAdd = false;
2718 // If the original flags and all inlined SCEVAddExprs are NUW, use the
2719 // common NUW flag for expression after inlining. Other flags cannot be
2720 // preserved, because they may depend on the original order of operations.
2721 SCEVFlags CommonFlags = maskFlags(ExprFlags, SCEV::FlagNUW);
2722 while (const SCEVAddExpr *Add = dyn_cast<SCEVAddExpr>(Ops[Idx])) {
2723 if (Ops.size() > AddOpsInlineThreshold ||
2724 Add->getNumOperands() > AddOpsInlineThreshold)
2725 break;
2726 // If we have an add, expand the add operands onto the end of the operands
2727 // list.
2728 Ops.erase(Ops.begin()+Idx);
2729 append_range(Ops, Add->operands());
2730 DeletedAdd = true;
2731 CommonFlags = maskFlags(CommonFlags, Add->getNoWrapFlags());
2732 }
2733
2734 // If we deleted at least one add, we added operands to the end of the list,
2735 // and they are not necessarily sorted. Recurse to resort and resimplify
2736 // any operands we just acquired.
2737 if (DeletedAdd)
2738 return getAddExpr(Ops, CommonFlags, Depth + 1);
2739 }
2740
2741 // Skip over the add expression until we get to a multiply.
2742 while (Idx < Ops.size() && Ops[Idx]->getSCEVType() < scMulExpr)
2743 ++Idx;
2744
2745 // Check to see if there are any folding opportunities present with
2746 // operands multiplied by constant values.
2747 if (Idx < Ops.size() && isa<SCEVMulExpr>(Ops[Idx])) {
2748 uint64_t BitWidth = getTypeSizeInBits(Ty);
2751 APInt AccumulatedConstant(BitWidth, 0);
2752 if (CollectAddOperandsWithScales(M, NewOps, AccumulatedConstant,
2753 Ops, APInt(BitWidth, 1), *this)) {
2754 struct APIntCompare {
2755 bool operator()(const APInt &LHS, const APInt &RHS) const {
2756 return LHS.ult(RHS);
2757 }
2758 };
2759
2760 // Some interesting folding opportunity is present, so its worthwhile to
2761 // re-generate the operands list. Group the operands by constant scale,
2762 // to avoid multiplying by the same constant scale multiple times.
2763 std::map<APInt, SmallVector<SCEVUse, 4>, APIntCompare> MulOpLists;
2764 for (SCEVUse NewOp : NewOps)
2765 MulOpLists[M.find(NewOp)->second].push_back(NewOp);
2766 // Re-generate the operands list.
2767 Ops.clear();
2768 if (AccumulatedConstant != 0)
2769 Ops.push_back(getConstant(AccumulatedConstant));
2770 for (auto &MulOp : MulOpLists) {
2771 if (MulOp.first == 1) {
2772 Ops.push_back(getAddExpr(MulOp.second, SCEV::FlagNone, Depth + 1));
2773 } else if (MulOp.first != 0) {
2774 Ops.push_back(
2775 getMulExpr(getConstant(MulOp.first),
2776 getAddExpr(MulOp.second, SCEV::FlagNone, Depth + 1),
2777 SCEV::FlagNone, Depth + 1));
2778 }
2779 }
2780 if (Ops.empty())
2781 return getZero(Ty);
2782 if (Ops.size() == 1)
2783 return Ops[0];
2784 return getAddExpr(Ops, SCEV::FlagNone, Depth + 1);
2785 }
2786 }
2787
2788 // Given a SCEVMulExpr and an operand index, return the product of all
2789 // operands except the one at OpIdx.
2790 auto StripFactor = [&](const SCEVMulExpr *M, unsigned OpIdx) -> SCEVUse {
2791 if (M->getNumOperands() == 2)
2792 return M->getOperand(OpIdx == 0);
2793 SmallVector<SCEVUse, 4> Remaining(M->operands().take_front(OpIdx));
2794 append_range(Remaining, M->operands().drop_front(OpIdx + 1));
2795 return getMulExpr(Remaining, SCEV::FlagNone, Depth + 1);
2796 };
2797
2798 // If we are adding something to a multiply expression, make sure the
2799 // something is not already an operand of the multiply. If so, merge it into
2800 // the multiply.
2801 for (; Idx < Ops.size() && isa<SCEVMulExpr>(Ops[Idx]); ++Idx) {
2802 const SCEVMulExpr *Mul = cast<SCEVMulExpr>(Ops[Idx]);
2803 for (unsigned MulOp = 0, e = Mul->getNumOperands(); MulOp != e; ++MulOp) {
2804 // Scan all terms to find every occurrence of common factor MulOpSCEV
2805 // and fold them in one shot:
2806 // A1*X + A2*X + ... + An*X --> X * (A1 + A2 + ... + An)
2807 const SCEV *MulOpSCEV = Mul->getOperand(MulOp);
2808 if (isa<SCEVConstant>(MulOpSCEV))
2809 continue;
2810
2811 // Cofactors: 1 for bare addends matching MulOpSCEV, or the
2812 // remaining product for multiply terms containing MulOpSCEV.
2813 SmallVector<SCEVUse, 4> Cofactors;
2814 SmallVector<unsigned, 4> DeadIndices;
2815 for (unsigned AddOp = 0, e = Ops.size(); AddOp != e; ++AddOp) {
2816 if (MulOpSCEV == Ops[AddOp]) {
2817 // W + X + (X * Y * Z) --> W + (X * ((Y*Z)+1))
2818 Cofactors.push_back(getOne(Ty));
2819 DeadIndices.push_back(AddOp);
2820 continue;
2821 }
2822
2823 if (AddOp <= Idx || !isa<SCEVMulExpr>(Ops[AddOp]))
2824 continue;
2825
2826 const SCEVMulExpr *OtherMul = cast<SCEVMulExpr>(Ops[AddOp]);
2827 for (unsigned OMulOp = 0, OE = OtherMul->getNumOperands(); OMulOp != OE;
2828 ++OMulOp) {
2829 if (OtherMul->getOperand(OMulOp) == MulOpSCEV) {
2830 // (A*B*C) + (A*D*E) --> A * (B*C + D*E)
2831 Cofactors.push_back(StripFactor(OtherMul, OMulOp));
2832 DeadIndices.push_back(AddOp);
2833 break;
2834 }
2835 }
2836 }
2837
2838 // Fold all collected cofactors with the anchor multiply's cofactor:
2839 // MulOpSCEV * (Cofactor_1 + ... + Cofactor_n + AnchorCofactor)
2840 if (!Cofactors.empty()) {
2841 Cofactors.push_back(StripFactor(Mul, MulOp));
2842
2843 SCEVUse InnerSum = getAddExpr(Cofactors, SCEV::FlagNone, Depth + 1);
2844 SCEVUse OuterMul =
2845 getMulExpr(MulOpSCEV, InnerSum, SCEV::FlagNone, Depth + 1);
2846
2847 // DeadIndices does not include Idx (the anchor), hence +1.
2848 if (Ops.size() == DeadIndices.size() + 1)
2849 return OuterMul;
2850
2851 // Erase Ops[Idx] first, then erase DeadIndices in reverse order.
2852 // The -1 adjustment accounts for the shift from removing Idx;
2853 // reverse order means each erasure only shifts later positions,
2854 // which have already been processed.
2855 Ops.erase(Ops.begin() + Idx);
2856 for (unsigned Dead : reverse(DeadIndices))
2857 Ops.erase(Ops.begin() + (Dead > Idx ? Dead - 1 : Dead));
2858
2859 Ops.push_back(OuterMul);
2860 return getAddExpr(Ops, SCEV::FlagNone, Depth + 1);
2861 }
2862 }
2863 }
2864
2865 // If there are any add recurrences in the operands list, see if any other
2866 // added values are loop invariant. If so, we can fold them into the
2867 // recurrence.
2868 while (Idx < Ops.size() && Ops[Idx]->getSCEVType() < scAddRecExpr)
2869 ++Idx;
2870
2871 // Scan over all recurrences, trying to fold loop invariants into them.
2872 for (; Idx < Ops.size() && isa<SCEVAddRecExpr>(Ops[Idx]); ++Idx) {
2873 // Scan all of the other operands to this add and add them to the vector if
2874 // they are loop invariant w.r.t. the recurrence.
2876 const SCEVAddRecExpr *AddRec = cast<SCEVAddRecExpr>(Ops[Idx]);
2877 const Loop *AddRecLoop = AddRec->getLoop();
2878 for (unsigned i = 0, e = Ops.size(); i != e; ++i)
2879 if (isAvailableAtLoopEntry(Ops[i], AddRecLoop)) {
2880 LIOps.push_back(Ops[i]);
2881 Ops.erase(Ops.begin()+i);
2882 --i; --e;
2883 }
2884
2885 // If we found some loop invariants, fold them into the recurrence.
2886 if (!LIOps.empty()) {
2887 // Compute nowrap flags for the addition of the loop-invariant ops and
2888 // the addrec. Temporarily push it as an operand for that purpose. These
2889 // flags are valid in the scope of the addrec only.
2890 LIOps.push_back(AddRec);
2891 SCEVFlags Flags = ComputeFlags(LIOps);
2892 LIOps.pop_back();
2893
2894 // NLI + LI + {Start,+,Step} --> NLI + {LI+Start,+,Step}
2895 LIOps.push_back(AddRec->getStart());
2896
2897 SmallVector<SCEVUse, 4> AddRecOps(AddRec->operands());
2898
2899 // It is not in general safe to propagate flags valid on an add within
2900 // the addrec scope to one outside it. We must prove that the inner
2901 // scope is guaranteed to execute if the outer one does to be able to
2902 // safely propagate. We know the program is undefined if poison is
2903 // produced on the inner scoped addrec. We also know that *for this use*
2904 // the outer scoped add can't overflow (because of the flags we just
2905 // computed for the inner scoped add) without the program being undefined.
2906 // Proving that entry to the outer scope neccesitates entry to the inner
2907 // scope, thus proves the program undefined if the flags would be violated
2908 // in the outer scope.
2909 SCEVFlags AddFlags = Flags;
2910 if (AddFlags != SCEV::FlagNone) {
2911 auto *DefI = getDefiningScopeBound(LIOps);
2912 auto *ReachI = &*AddRecLoop->getHeader()->begin();
2913 if (!isGuaranteedToTransferExecutionTo(DefI, ReachI))
2914 AddFlags = SCEV::FlagNone;
2915 }
2916 AddRecOps[0] = getAddExpr(LIOps, AddFlags, Depth + 1);
2917
2918 // Build the new addrec. Propagate the NUW and NSW flags if both the
2919 // outer add and the inner addrec are guaranteed to have no overflow.
2920 // Always propagate NW.
2921 Flags = AddRec->getNoWrapFlags(setFlags(Flags, SCEV::FlagNW));
2922 const SCEV *NewRec = getAddRecExpr(AddRecOps, AddRecLoop, Flags);
2923
2924 // If all of the other operands were loop invariant, we are done.
2925 if (Ops.size() == 1) return NewRec;
2926
2927 // Otherwise, add the folded AddRec by the non-invariant parts.
2928 for (unsigned i = 0;; ++i)
2929 if (Ops[i] == AddRec) {
2930 Ops[i] = NewRec;
2931 break;
2932 }
2933 return getAddExpr(Ops, SCEV::FlagNone, Depth + 1);
2934 }
2935
2936 // Okay, if there weren't any loop invariants to be folded, check to see if
2937 // there are multiple AddRec's with the same loop induction variable being
2938 // added together. If so, we can fold them.
2939 for (unsigned OtherIdx = Idx+1;
2940 OtherIdx < Ops.size() && isa<SCEVAddRecExpr>(Ops[OtherIdx]);
2941 ++OtherIdx) {
2942 // We expect the AddRecExpr's to be sorted in reverse dominance order,
2943 // so that the 1st found AddRecExpr is dominated by all others.
2944 assert(DT.dominates(
2945 cast<SCEVAddRecExpr>(Ops[OtherIdx])->getLoop()->getHeader(),
2946 AddRec->getLoop()->getHeader()) &&
2947 "AddRecExprs are not sorted in reverse dominance order?");
2948 if (AddRecLoop == cast<SCEVAddRecExpr>(Ops[OtherIdx])->getLoop()) {
2949 // Other + {A,+,B}<L> + {C,+,D}<L> --> Other + {A+C,+,B+D}<L>
2950 SmallVector<SCEVUse, 4> AddRecOps(AddRec->operands());
2951 for (; OtherIdx != Ops.size() && isa<SCEVAddRecExpr>(Ops[OtherIdx]);
2952 ++OtherIdx) {
2953 const auto *OtherAddRec = cast<SCEVAddRecExpr>(Ops[OtherIdx]);
2954 if (OtherAddRec->getLoop() == AddRecLoop) {
2955 for (unsigned i = 0, e = OtherAddRec->getNumOperands();
2956 i != e; ++i) {
2957 if (i >= AddRecOps.size()) {
2958 append_range(AddRecOps, OtherAddRec->operands().drop_front(i));
2959 break;
2960 }
2961 AddRecOps[i] =
2962 getAddExpr(AddRecOps[i], OtherAddRec->getOperand(i),
2963 SCEV::FlagNone, Depth + 1);
2964 }
2965 Ops.erase(Ops.begin() + OtherIdx); --OtherIdx;
2966 }
2967 }
2968 // Step size has changed, so we cannot guarantee no self-wraparound.
2969 Ops[Idx] = getAddRecExpr(AddRecOps, AddRecLoop, SCEV::FlagNone);
2970 return getAddExpr(Ops, SCEV::FlagNone, Depth + 1);
2971 }
2972 }
2973
2974 // Otherwise couldn't fold anything into this recurrence. Move onto the
2975 // next one.
2976 }
2977
2978 // Okay, it looks like we really DO need an add expr. Check to see if we
2979 // already have one, otherwise create a new one.
2980 assert((UseFlags == SCEV::FlagNone || equal(OrigOps, Ops)) &&
2981 "Tried to add SCEVUse flags after operands changed");
2982 return {getOrCreateAddExpr(Ops, ComputeFlags(Ops)), UseFlags};
2983}
2984
2985const SCEV *ScalarEvolution::getOrCreateAddExpr(ArrayRef<SCEVUse> Ops,
2986 SCEVFlags Flags) {
2989 for (SCEVUse Op : Ops)
2990 ID.AddPointer(Op.getOpaqueValue());
2992 SCEVAddExpr *S = static_cast<SCEVAddExpr *>(UniqueSCEVs.lookup(ID, Token));
2993 if (!S) {
2994 SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
2996 S = new (SCEVAllocator)
2997 SCEVAddExpr(ID.Intern(SCEVAllocator), O, Ops.size());
2998 UniqueSCEVs.insert(S, Token);
2999 S->computeAndSetCanonical(*this);
3000 registerUser(S, Ops);
3001 }
3002 S->setNoWrapFlags(Flags);
3003 return S;
3004}
3005
3006const SCEV *ScalarEvolution::getOrCreateAddRecExpr(ArrayRef<SCEVUse> Ops,
3007 const Loop *L,
3008 SCEVFlags Flags) {
3009 FoldingSetNodeID ID;
3010 ID.AddInteger(scAddRecExpr);
3011 for (SCEVUse Op : Ops)
3012 ID.AddPointer(Op.getOpaqueValue());
3013 ID.AddPointer(L);
3014 FoldingSetInsertToken Token;
3015 SCEVAddRecExpr *S =
3016 static_cast<SCEVAddRecExpr *>(UniqueSCEVs.lookup(ID, Token));
3017 if (!S) {
3018 SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
3020 S = new (SCEVAllocator)
3021 SCEVAddRecExpr(ID.Intern(SCEVAllocator), O, Ops.size(), L);
3022 UniqueSCEVs.insert(S, Token);
3023 S->computeAndSetCanonical(*this);
3024 LoopUsers[L].push_back(S);
3025 registerUser(S, Ops);
3026 }
3027 setNoWrapFlags(S, Flags);
3028 return S;
3029}
3030
3031const SCEV *ScalarEvolution::getOrCreateMulExpr(ArrayRef<SCEVUse> Ops,
3032 SCEVFlags Flags) {
3033 FoldingSetNodeID ID;
3034 ID.AddInteger(scMulExpr);
3035 for (SCEVUse Op : Ops)
3036 ID.AddPointer(Op.getOpaqueValue());
3037 FoldingSetInsertToken Token;
3038 SCEVMulExpr *S = static_cast<SCEVMulExpr *>(UniqueSCEVs.lookup(ID, Token));
3039 if (!S) {
3040 SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
3042 S = new (SCEVAllocator) SCEVMulExpr(ID.Intern(SCEVAllocator),
3043 O, Ops.size());
3044 UniqueSCEVs.insert(S, Token);
3045 S->computeAndSetCanonical(*this);
3046 registerUser(S, Ops);
3047 }
3048 S->setNoWrapFlags(Flags);
3049 return S;
3050}
3051
3052const SCEV *ScalarEvolution::getOrCreateUDivExpr(SCEVUse LHS, SCEVUse RHS) {
3053 FoldingSetNodeID ID;
3054 ID.AddInteger(scUDivExpr);
3055 ID.AddPointer(LHS.getOpaqueValue());
3056 ID.AddPointer(RHS.getOpaqueValue());
3057 FoldingSetInsertToken Token;
3058 SCEV *S = UniqueSCEVs.lookup(ID, Token);
3059 if (!S) {
3060 S = new (SCEVAllocator) SCEVUDivExpr(ID.Intern(SCEVAllocator), LHS, RHS);
3061 UniqueSCEVs.insert(S, Token);
3062 S->computeAndSetCanonical(*this);
3063 registerUser(S, {LHS, RHS});
3064 }
3065 return S;
3066}
3067
3068static uint64_t umul_ov(uint64_t i, uint64_t j, bool &Overflow) {
3069 uint64_t k = i*j;
3070 if (j > 1 && k / j != i) Overflow = true;
3071 return k;
3072}
3073
3074/// Compute the result of "n choose k", the binomial coefficient. If an
3075/// intermediate computation overflows, Overflow will be set and the return will
3076/// be garbage. Overflow is not cleared on absence of overflow.
3077static uint64_t Choose(uint64_t n, uint64_t k, bool &Overflow) {
3078 // We use the multiplicative formula:
3079 // n(n-1)(n-2)...(n-(k-1)) / k(k-1)(k-2)...1 .
3080 // At each iteration, we take the n-th term of the numeral and divide by the
3081 // (k-n)th term of the denominator. This division will always produce an
3082 // integral result, and helps reduce the chance of overflow in the
3083 // intermediate computations. However, we can still overflow even when the
3084 // final result would fit.
3085
3086 if (n == 0 || n == k) return 1;
3087 if (k > n) return 0;
3088
3089 if (k > n/2)
3090 k = n-k;
3091
3092 uint64_t r = 1;
3093 for (uint64_t i = 1; i <= k; ++i) {
3094 r = umul_ov(r, n-(i-1), Overflow);
3095 r /= i;
3096 }
3097 return r;
3098}
3099
3100/// Determine if any of the operands in this SCEV are a constant or if
3101/// any of the add or multiply expressions in this SCEV contain a constant.
3102static bool containsConstantInAddMulChain(const SCEV *StartExpr) {
3103 struct FindConstantInAddMulChain {
3104 bool FoundConstant = false;
3105
3106 bool follow(const SCEV *S) {
3107 FoundConstant |= isa<SCEVConstant>(S);
3108 return isa<SCEVAddExpr>(S) || isa<SCEVMulExpr>(S);
3109 }
3110
3111 bool isDone() const {
3112 return FoundConstant;
3113 }
3114 };
3115
3116 FindConstantInAddMulChain F;
3118 ST.visitAll(StartExpr);
3119 return F.FoundConstant;
3120}
3121
3122/// Get a canonical multiply expression, or something simpler if possible.
3124 SCEVFlagsPair Flags, unsigned Depth) {
3125 SCEVFlags ExprFlags = Flags.ExprFlags;
3126 SCEVFlags UseFlags = Flags.UseFlags;
3127 assert(ExprFlags == maskFlags(ExprFlags, SCEV::FlagNUW | SCEV::FlagNSW) &&
3128 "only nuw or nsw allowed");
3129 assert(UseFlags == maskFlags(UseFlags, SCEV::FlagNUW | SCEV::FlagNSW) &&
3130 "only nuw or nsw allowed");
3131 assert(!Ops.empty() && "Cannot get empty mul!");
3132 if (Ops.size() == 1) return Ops[0];
3133#ifndef NDEBUG
3134 Type *ETy = Ops[0]->getType();
3135 assert(!ETy->isPointerTy());
3136 for (unsigned i = 1, e = Ops.size(); i != e; ++i)
3137 assert(Ops[i]->getType() == ETy &&
3138 "SCEVMulExpr operand types don't match!");
3139#endif
3140
3141 const SCEV *Folded = constantFoldAndGroupOps(
3142 *this, LI, DT, Ops,
3143 [](const APInt &C1, const APInt &C2) { return C1 * C2; },
3144 [](const APInt &C) { return C.isOne(); }, // identity
3145 [](const APInt &C) { return C.isZero(); }); // absorber
3146 if (Folded)
3147 return Folded;
3148
3149#ifndef NDEBUG
3150 // Keep track of operands after constant folding, for verification when adding
3151 // use-specific flags.
3152 const SmallVector<SCEVUse, 8> OrigOps(Ops.begin(), Ops.end());
3153#endif
3154
3155 // Delay expensive flag strengthening until necessary.
3156 auto ComputeFlags = [this, ExprFlags](const ArrayRef<SCEVUse> Ops) {
3157 return StrengthenNoWrapFlags(this, scMulExpr, Ops, ExprFlags);
3158 };
3159
3160 // Limit recursion calls depth.
3162 return {getOrCreateMulExpr(Ops, ComputeFlags(Ops)), UseFlags};
3163
3164 if (SCEV *S = findExistingSCEVInCache(scMulExpr, Ops)) {
3165 // Don't strengthen flags if we have no new information.
3166 SCEVMulExpr *Mul = static_cast<SCEVMulExpr *>(S);
3167 if (Mul->getNoWrapFlags(ExprFlags) != ExprFlags)
3168 Mul->setNoWrapFlags(ComputeFlags(Ops));
3169 return {S, UseFlags};
3170 }
3171
3172 if (const SCEVConstant *LHSC = dyn_cast<SCEVConstant>(Ops[0])) {
3173 if (Ops.size() == 2) {
3174 // C1*(C2+V) -> C1*C2 + C1*V
3175 // If any of Add's ops are Adds or Muls with a constant, apply this
3176 // transformation as well.
3177 //
3178 // TODO: There are some cases where this transformation is not
3179 // profitable; for example, Add = (C0 + X) * Y + Z. Maybe the scope of
3180 // this transformation should be narrowed down.
3181 const SCEV *Op0, *Op1;
3182 if (match(Ops[1], m_scev_Add(m_SCEV(Op0), m_SCEV(Op1))) &&
3184 const SCEV *LHS = getMulExpr(LHSC, Op0, SCEV::FlagNone, Depth + 1);
3185 const SCEV *RHS = getMulExpr(LHSC, Op1, SCEV::FlagNone, Depth + 1);
3186 return getAddExpr(LHS, RHS, SCEV::FlagNone, Depth + 1);
3187 }
3188
3189 if (Ops[0]->isAllOnesValue()) {
3190 // If we have a mul by -1 of an add, try distributing the -1 among the
3191 // add operands.
3192 if (const SCEVAddExpr *Add = dyn_cast<SCEVAddExpr>(Ops[1])) {
3194 bool AnyFolded = false;
3195 for (const SCEV *AddOp : Add->operands()) {
3196 const SCEV *Mul =
3197 getMulExpr(Ops[0], SCEVUse(AddOp), SCEV::FlagNone, Depth + 1);
3198 if (!isa<SCEVMulExpr>(Mul)) AnyFolded = true;
3199 NewOps.push_back(Mul);
3200 }
3201 if (AnyFolded)
3202 return getAddExpr(NewOps, SCEV::FlagNone, Depth + 1);
3203 } else if (const auto *AddRec = dyn_cast<SCEVAddRecExpr>(Ops[1])) {
3204 // Negation preserves a recurrence's no self-wrap property.
3206 for (const SCEV *AddRecOp : AddRec->operands())
3207 Operands.push_back(getMulExpr(Ops[0], SCEVUse(AddRecOp),
3208 SCEV::FlagNone, Depth + 1));
3209 // Let M be the minimum representable signed value. AddRec with nsw
3210 // multiplied by -1 can have signed overflow if and only if it takes a
3211 // value of M: M * (-1) would stay M and (M + 1) * (-1) would be the
3212 // maximum signed value. In all other cases signed overflow is
3213 // impossible.
3214 auto FlagsMask = SCEV::FlagNW;
3215 if (AddRec->hasNoSignedWrap()) {
3216 auto MinInt =
3217 APInt::getSignedMinValue(getTypeSizeInBits(AddRec->getType()));
3218 if (getSignedRangeMin(AddRec) != MinInt)
3220 }
3221 return getAddRecExpr(Operands, AddRec->getLoop(),
3222 AddRec->getNoWrapFlags(FlagsMask));
3223 }
3224 }
3225
3226 // Try to push the constant operand into a ZExt: C * zext (A + B) ->
3227 // zext (C*A + C*B) if trunc (C) * (A + B) does not unsigned-wrap.
3228 const SCEVAddExpr *InnerAdd;
3229 if (match(Ops[1], m_scev_ZExt(m_scev_Add(InnerAdd)))) {
3230 const SCEV *NarrowC = getTruncateExpr(LHSC, InnerAdd->getType());
3231 if (isa<SCEVConstant>(InnerAdd->getOperand(0)) &&
3232 getZeroExtendExpr(NarrowC, Ops[1]->getType()) == LHSC &&
3233 hasFlags(StrengthenNoWrapFlags(this, scMulExpr, {NarrowC, InnerAdd},
3235 SCEV::FlagNUW)) {
3236 const SCEV *Res =
3237 getMulExpr(NarrowC, InnerAdd, SCEV::FlagNUW, Depth + 1);
3238 return getZeroExtendExpr(Res, Ops[1]->getType(), Depth + 1);
3239 };
3240 }
3241
3242 // Try to fold (C1 * D /u C2) -> C1/C2 * D, if C1 and C2 are powers-of-2,
3243 // D is a multiple of C2, and C1 is a multiple of C2. If C2 is a multiple
3244 // of C1, fold to (D /u (C2 /u C1)).
3245 const SCEV *D;
3246 APInt C1V = LHSC->getAPInt();
3247 // (C1 * D /u C2) == -1 * -C1 * D /u C2 when C1 != INT_MIN. Don't treat -1
3248 // as -1 * 1, as it won't enable additional folds.
3249 if (C1V.isNegative() && !C1V.isMinSignedValue() && !C1V.isAllOnes())
3250 C1V = C1V.abs();
3251 const SCEVConstant *C2;
3252 if (C1V.isPowerOf2() &&
3254 C2->getAPInt().isPowerOf2() &&
3255 C1V.logBase2() <= getMinTrailingZeros(D)) {
3256 const SCEV *NewMul = nullptr;
3257 if (C1V.uge(C2->getAPInt())) {
3258 NewMul = getMulExpr(getUDivExpr(getConstant(C1V), C2), D);
3259 } else if (C2->getAPInt().logBase2() <= getMinTrailingZeros(D)) {
3260 assert(C1V.ugt(1) && "C1 <= 1 should have been folded earlier");
3261 NewMul = getUDivExpr(D, getUDivExpr(C2, getConstant(C1V)));
3262 }
3263 if (NewMul)
3264 return C1V == LHSC->getAPInt() ? NewMul : getNegativeSCEV(NewMul);
3265 }
3266 }
3267 }
3268
3269 // Skip over the add expression until we get to a multiply.
3270 unsigned Idx = 0;
3271 while (Idx < Ops.size() && Ops[Idx]->getSCEVType() < scMulExpr)
3272 ++Idx;
3273
3274 // If there are mul operands inline them all into this expression.
3275 if (Idx < Ops.size()) {
3276 bool DeletedMul = false;
3277 while (const SCEVMulExpr *Mul = dyn_cast<SCEVMulExpr>(Ops[Idx])) {
3278 if (Ops.size() > MulOpsInlineThreshold)
3279 break;
3280 // If we have an mul, expand the mul operands onto the end of the
3281 // operands list.
3282 Ops.erase(Ops.begin()+Idx);
3283 append_range(Ops, Mul->operands());
3284 DeletedMul = true;
3285 }
3286
3287 // If we deleted at least one mul, we added operands to the end of the
3288 // list, and they are not necessarily sorted. Recurse to resort and
3289 // resimplify any operands we just acquired.
3290 if (DeletedMul)
3291 return getMulExpr(Ops, SCEV::FlagNone, Depth + 1);
3292 }
3293
3294 // If there are any add recurrences in the operands list, see if any other
3295 // added values are loop invariant. If so, we can fold them into the
3296 // recurrence.
3297 while (Idx < Ops.size() && Ops[Idx]->getSCEVType() < scAddRecExpr)
3298 ++Idx;
3299
3300 // Scan over all recurrences, trying to fold loop invariants into them.
3301 for (; Idx < Ops.size() && isa<SCEVAddRecExpr>(Ops[Idx]); ++Idx) {
3302 // Scan all of the other operands to this mul and add them to the vector
3303 // if they are loop invariant w.r.t. the recurrence.
3305 const SCEVAddRecExpr *AddRec = cast<SCEVAddRecExpr>(Ops[Idx]);
3306 for (unsigned i = 0, e = Ops.size(); i != e; ++i)
3307 if (isAvailableAtLoopEntry(Ops[i], AddRec->getLoop())) {
3308 LIOps.push_back(Ops[i]);
3309 Ops.erase(Ops.begin()+i);
3310 --i; --e;
3311 }
3312
3313 // If we found some loop invariants, fold them into the recurrence.
3314 if (!LIOps.empty()) {
3315 // NLI * LI * {Start,+,Step} --> NLI * {LI*Start,+,LI*Step}
3317 NewOps.reserve(AddRec->getNumOperands());
3318 const SCEV *Scale = getMulExpr(LIOps, SCEV::FlagNone, Depth + 1);
3319
3320 // If both the mul and addrec are nuw, we can preserve nuw.
3321 // If both the mul and addrec are nsw, we can only preserve nsw if either
3322 // a) they are also nuw, or
3323 // b) all multiplications of addrec operands with scale are nsw.
3324 SCEVFlags Flags = AddRec->getNoWrapFlags(ComputeFlags({Scale, AddRec}));
3325
3326 for (unsigned i = 0, e = AddRec->getNumOperands(); i != e; ++i) {
3327 NewOps.push_back(getMulExpr(Scale, AddRec->getOperand(i),
3328 SCEV::FlagNone, Depth + 1));
3329
3330 if (hasFlags(Flags, SCEV::FlagNSW) && !hasFlags(Flags, SCEV::FlagNUW)) {
3332 Instruction::Mul, getSignedRange(Scale),
3334 if (!NSWRegion.contains(getSignedRange(AddRec->getOperand(i))))
3335 Flags = clearFlags(Flags, SCEV::FlagNSW);
3336 }
3337 }
3338
3339 const SCEV *NewRec = getAddRecExpr(NewOps, AddRec->getLoop(), Flags);
3340
3341 // If all of the other operands were loop invariant, we are done.
3342 if (Ops.size() == 1) return NewRec;
3343
3344 // Otherwise, multiply the folded AddRec by the non-invariant parts.
3345 for (unsigned i = 0;; ++i)
3346 if (Ops[i] == AddRec) {
3347 Ops[i] = NewRec;
3348 break;
3349 }
3350 return getMulExpr(Ops, SCEV::FlagNone, Depth + 1);
3351 }
3352
3353 // Okay, if there weren't any loop invariants to be folded, check to see
3354 // if there are multiple AddRec's with the same loop induction variable
3355 // being multiplied together. If so, we can fold them.
3356
3357 // {A1,+,A2,+,...,+,An}<L> * {B1,+,B2,+,...,+,Bn}<L>
3358 // = {x=1 in [ sum y=x..2x [ sum z=max(y-x, y-n)..min(x,n) [
3359 // choose(x, 2x)*choose(2x-y, x-z)*A_{y-z}*B_z
3360 // ]]],+,...up to x=2n}.
3361 // Note that the arguments to choose() are always integers with values
3362 // known at compile time, never SCEV objects.
3363 //
3364 // The implementation avoids pointless extra computations when the two
3365 // addrec's are of different length (mathematically, it's equivalent to
3366 // an infinite stream of zeros on the right).
3367 bool OpsModified = false;
3368 for (unsigned OtherIdx = Idx+1;
3369 OtherIdx != Ops.size() && isa<SCEVAddRecExpr>(Ops[OtherIdx]);
3370 ++OtherIdx) {
3371 const SCEVAddRecExpr *OtherAddRec =
3372 dyn_cast<SCEVAddRecExpr>(Ops[OtherIdx]);
3373 if (!OtherAddRec || OtherAddRec->getLoop() != AddRec->getLoop())
3374 continue;
3375
3376 // Limit max number of arguments to avoid creation of unreasonably big
3377 // SCEVAddRecs with very complex operands.
3378 if (AddRec->getNumOperands() + OtherAddRec->getNumOperands() - 1 >
3379 MaxAddRecSize || hasHugeExpression({AddRec, OtherAddRec}))
3380 continue;
3381
3382 bool Overflow = false;
3383 Type *Ty = AddRec->getType();
3384 bool LargerThan64Bits = getTypeSizeInBits(Ty) > 64;
3385 SmallVector<SCEVUse, 7> AddRecOps;
3386 for (int x = 0, xe = AddRec->getNumOperands() +
3387 OtherAddRec->getNumOperands() - 1; x != xe && !Overflow; ++x) {
3389 for (int y = x, ye = 2*x+1; y != ye && !Overflow; ++y) {
3390 uint64_t Coeff1 = Choose(x, 2*x - y, Overflow);
3391 for (int z = std::max(y-x, y-(int)AddRec->getNumOperands()+1),
3392 ze = std::min(x+1, (int)OtherAddRec->getNumOperands());
3393 z < ze && !Overflow; ++z) {
3394 uint64_t Coeff2 = Choose(2*x - y, x-z, Overflow);
3395 uint64_t Coeff;
3396 if (LargerThan64Bits)
3397 Coeff = umul_ov(Coeff1, Coeff2, Overflow);
3398 else
3399 Coeff = Coeff1*Coeff2;
3400 const SCEV *CoeffTerm = getConstant(Ty, Coeff);
3401 const SCEV *Term1 = AddRec->getOperand(y-z);
3402 const SCEV *Term2 = OtherAddRec->getOperand(z);
3403 SumOps.push_back(
3404 getMulExpr(CoeffTerm, Term1, Term2, SCEV::FlagNone, Depth + 1));
3405 }
3406 }
3407 if (SumOps.empty())
3408 SumOps.push_back(getZero(Ty));
3409 AddRecOps.push_back(getAddExpr(SumOps, SCEV::FlagNone, Depth + 1));
3410 }
3411 if (!Overflow) {
3412 const SCEV *NewAddRec =
3413 getAddRecExpr(AddRecOps, AddRec->getLoop(), SCEV::FlagNone);
3414 if (Ops.size() == 2) return NewAddRec;
3415 Ops[Idx] = NewAddRec;
3416 Ops.erase(Ops.begin() + OtherIdx); --OtherIdx;
3417 OpsModified = true;
3418 AddRec = dyn_cast<SCEVAddRecExpr>(NewAddRec);
3419 if (!AddRec)
3420 break;
3421 }
3422 }
3423 if (OpsModified)
3424 return getMulExpr(Ops, SCEV::FlagNone, Depth + 1);
3425
3426 // Otherwise couldn't fold anything into this recurrence. Move onto the
3427 // next one.
3428 }
3429
3430 // Okay, it looks like we really DO need an mul expr. Check to see if we
3431 // already have one, otherwise create a new one.
3432 assert((UseFlags == SCEV::FlagNone || equal(OrigOps, Ops)) &&
3433 "Tried to add SCEVUse flags after operands changed");
3434 return {getOrCreateMulExpr(Ops, ComputeFlags(Ops)), UseFlags};
3435}
3436
3437/// Represents an unsigned remainder expression based on unsigned division.
3439 assert(getEffectiveSCEVType(LHS->getType()) ==
3440 getEffectiveSCEVType(RHS->getType()) &&
3441 "SCEVURemExpr operand types don't match!");
3442
3443 // Short-circuit easy cases
3444 if (const SCEVConstant *RHSC = dyn_cast<SCEVConstant>(RHS)) {
3445 // If constant is one, the result is trivial
3446 if (RHSC->getValue()->isOne())
3447 return getZero(LHS->getType()); // X urem 1 --> 0
3448
3449 // If constant is a power of two, fold into a zext(trunc(LHS)).
3450 if (RHSC->getAPInt().isPowerOf2()) {
3451 Type *FullTy = LHS->getType();
3452 Type *TruncTy =
3453 IntegerType::get(getContext(), RHSC->getAPInt().logBase2());
3454 return getZeroExtendExpr(getTruncateExpr(LHS, TruncTy), FullTy);
3455 }
3456 }
3457
3458 // Fallback to %a == %x urem %y == %x -<nuw> ((%x udiv %y) *<nuw> %y)
3459 const SCEV *UDiv = getUDivExpr(LHS, RHS);
3460 const SCEV *Mult = getMulExpr(UDiv, RHS, SCEV::FlagNUW);
3461 return getMinusSCEV(LHS, Mult, SCEV::FlagNUW);
3462}
3463
3464/// Get a canonical unsigned division expression, or something simpler if
3465/// possible.
3467 assert(!LHS->getType()->isPointerTy() &&
3468 "SCEVUDivExpr operand can't be pointer!");
3469 assert(LHS->getType() == RHS->getType() &&
3470 "SCEVUDivExpr operand types don't match!");
3471
3472 if (SCEV *S = findExistingSCEVInCache(scUDivExpr, {LHS, RHS}))
3473 return S;
3474
3475 // 0 udiv Y == 0
3476 if (match(LHS, m_scev_Zero()))
3477 return LHS;
3478
3479 if (const SCEVConstant *RHSC = dyn_cast<SCEVConstant>(RHS)) {
3480 if (RHSC->getValue()->isOne())
3481 return LHS; // X udiv 1 --> x
3482 // If the denominator is zero, the result of the udiv is undefined. Don't
3483 // try to analyze it, because the resolution chosen here may differ from
3484 // the resolution chosen in other parts of the compiler.
3485 if (!RHSC->getValue()->isZero()) {
3486 // Determine if the division can be folded into the operands of
3487 // its operands.
3488 // TODO: Generalize this to non-constants by using known-bits information.
3489 Type *Ty = LHS->getType();
3490 unsigned LZ = RHSC->getAPInt().countl_zero();
3491 unsigned MaxShiftAmt = getTypeSizeInBits(Ty) - LZ - 1;
3492 // For non-power-of-two values, effectively round the value up to the
3493 // nearest power of two.
3494 if (!RHSC->getAPInt().isPowerOf2())
3495 ++MaxShiftAmt;
3496 IntegerType *ExtTy =
3497 IntegerType::get(getContext(), getTypeSizeInBits(Ty) + MaxShiftAmt);
3498 if (const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(LHS))
3499 if (const SCEVConstant *Step =
3500 dyn_cast<SCEVConstant>(AR->getStepRecurrence(*this))) {
3501 // {X,+,N}/C --> {X/C,+,N/C} if safe and N/C can be folded.
3502 const APInt &StepInt = Step->getAPInt();
3503 const APInt &DivInt = RHSC->getAPInt();
3504 if (!StepInt.urem(DivInt) &&
3505 getZeroExtendExpr(AR, ExtTy) ==
3506 getAddRecExpr(getZeroExtendExpr(AR->getStart(), ExtTy),
3507 getZeroExtendExpr(Step, ExtTy), AR->getLoop(),
3508 SCEV::FlagNone)) {
3510 for (const SCEV *Op : AR->operands())
3511 Operands.push_back(getUDivExpr(Op, RHS));
3512 return getAddRecExpr(Operands, AR->getLoop(), SCEV::FlagNW);
3513 }
3514 /// Get a canonical UDivExpr for a recurrence.
3515 /// {X,+,N}/C => {Y,+,N}/C where Y=X-(X%N). Safe when C%N=0.
3516 const APInt *StartRem;
3517 if (!DivInt.urem(StepInt) && match(getURemExpr(AR->getStart(), Step),
3518 m_scev_APInt(StartRem))) {
3519 bool NoWrap =
3520 getZeroExtendExpr(AR, ExtTy) ==
3521 getAddRecExpr(getZeroExtendExpr(AR->getStart(), ExtTy),
3522 getZeroExtendExpr(Step, ExtTy), AR->getLoop(),
3524
3525 // With N <= C and both N, C as powers-of-2, the transformation
3526 // {X,+,N}/C => {(X - X%N),+,N}/C preserves division results even
3527 // if wrapping occurs, as the division results remain equivalent for
3528 // all offsets in [[(X - X%N), X).
3529 bool CanFoldWithWrap = StepInt.ule(DivInt) && // N <= C
3530 StepInt.isPowerOf2() && DivInt.isPowerOf2();
3531 // Only fold if the subtraction can be folded in the start
3532 // expression.
3533 const SCEV *NewStart =
3534 getMinusSCEV(AR->getStart(), getConstant(*StartRem));
3535 if (*StartRem != 0 && (NoWrap || CanFoldWithWrap) &&
3536 !isa<SCEVAddExpr>(NewStart)) {
3537 const SCEV *NewLHS =
3538 getAddRecExpr(NewStart, Step, AR->getLoop(),
3539 NoWrap ? SCEV::FlagNW : SCEV::FlagNone);
3540 if (LHS != NewLHS)
3541 return getUDivExpr(NewLHS, RHS);
3542 }
3543 }
3544 }
3545 // (A*B)/C --> A*(B/C) if safe and B/C can be folded.
3546 if (const SCEVMulExpr *M = dyn_cast<SCEVMulExpr>(LHS)) {
3547 if (M->hasNoUnsignedWrap()) {
3548 // Find an operand that's safely divisible.
3549 for (unsigned i = 0, e = M->getNumOperands(); i != e; ++i) {
3550 const SCEV *Op = M->getOperand(i);
3551 const SCEV *Div = getUDivExpr(Op, RHSC);
3552 if (!isa<SCEVUDivExpr>(Div) && getMulExpr(Div, RHSC) == Op) {
3553 SmallVector<SCEVUse, 4> Operands(M->operands());
3554 Operands[i] = Div;
3555 return getMulExpr(Operands);
3556 }
3557 }
3558
3559 // Even if it's not divisible, try to remove a common factor.
3560 if (const auto *LHSC = dyn_cast<SCEVConstant>(M->getOperand(0))) {
3561 APInt Factor = APIntOps::GreatestCommonDivisor(LHSC->getAPInt(),
3562 RHSC->getAPInt());
3563 if (!Factor.isIntN(1)) {
3564 SmallVector<SCEVUse, 2> NewOperands;
3565 NewOperands.push_back(getConstant(LHSC->getAPInt().udiv(Factor)));
3566 append_range(NewOperands, M->operands().drop_front());
3567 const SCEV *NewMul = getMulExpr(NewOperands);
3568 return getUDivExpr(NewMul,
3569 getConstant(RHSC->getAPInt().udiv(Factor)));
3570 }
3571 }
3572 }
3573 }
3574
3575 // (A/B)/C --> A/(B*C) if safe and B*C can be folded.
3576 if (const SCEVUDivExpr *OtherDiv = dyn_cast<SCEVUDivExpr>(LHS)) {
3577 if (auto *DivisorConstant =
3578 dyn_cast<SCEVConstant>(OtherDiv->getRHS())) {
3579 bool Overflow = false;
3580 APInt NewRHS =
3581 DivisorConstant->getAPInt().umul_ov(RHSC->getAPInt(), Overflow);
3582 if (Overflow) {
3583 return getConstant(RHSC->getType(), 0, false);
3584 }
3585 return getUDivExpr(OtherDiv->getLHS(), getConstant(NewRHS));
3586 }
3587 }
3588
3589 // (A+B)/C --> (A/C + B/C) if the add does not unsigned wrap and A/C and
3590 // B/C can be folded.
3591 if (const SCEVAddExpr *A = dyn_cast<SCEVAddExpr>(LHS)) {
3592 if (A->hasNoUnsignedWrap()) {
3594 for (unsigned i = 0, e = A->getNumOperands(); i != e; ++i) {
3595 const SCEV *Op = getUDivExpr(A->getOperand(i), RHS);
3596 if (isa<SCEVUDivExpr>(Op) ||
3597 getMulExpr(Op, RHS) != A->getOperand(i))
3598 break;
3599 Operands.push_back(Op);
3600 }
3601 if (Operands.size() == A->getNumOperands())
3602 return getAddExpr(Operands);
3603 }
3604 }
3605
3606 // ((N - M) + (M * A)) / N --> ((N - 1) + (M * A)) / N
3607 // This is an idiom for rounding A up to the next multiple of N, where A
3608 // is aready known to be a multiple of M. In this case, instcombine can
3609 // see that some low bits of the added constant are unused, so can clear
3610 // them, but we want to canonicalise to set the low bits. This makes the
3611 // pattern easier to match, without needing to check for known bits in
3612 // A*M.
3613 const APInt &N = RHSC->getAPInt();
3614 const APInt *NMinusM, *M;
3615 const SCEV *A;
3616 if (match(LHS, m_scev_Add(m_scev_APInt(NMinusM),
3617 m_scev_Mul(m_scev_APInt(M), m_SCEV(A))))) {
3618 if (N.isPowerOf2() && M->isPowerOf2() && M->ult(N) &&
3619 *NMinusM == N - *M) {
3620 return getUDivExpr(
3622 RHS);
3623 }
3624 }
3625
3626 // Fold if both operands are constant.
3627 if (const SCEVConstant *LHSC = dyn_cast<SCEVConstant>(LHS))
3628 return getConstant(LHSC->getAPInt().udiv(RHSC->getAPInt()));
3629 }
3630 }
3631
3632 // ((-C + (C smax %x)) /u %x) evaluates to zero, for any positive constant C.
3633 const APInt *NegC, *C;
3634 if (match(LHS,
3637 NegC->isNegative() && !NegC->isMinSignedValue() && *C == -*NegC)
3638 return getZero(LHS->getType());
3639
3640 // (%a * %b)<nuw> / %b -> %a
3641 const auto *Mul = dyn_cast<SCEVMulExpr>(LHS);
3642 if (Mul && Mul->hasNoUnsignedWrap()) {
3643 for (int i = 0, e = Mul->getNumOperands(); i != e; ++i) {
3644 if (Mul->getOperand(i) == RHS) {
3646 append_range(Operands, Mul->operands().take_front(i));
3647 append_range(Operands, Mul->operands().drop_front(i + 1));
3648 return getMulExpr(Operands);
3649 }
3650 }
3651 }
3652
3653 // TODO: Generalize to handle any common factors.
3654 // udiv (mul nuw a, vscale), (mul nuw b, vscale) --> udiv a, b
3655 const SCEV *NewLHS, *NewRHS;
3656 if (match(LHS, m_scev_c_NUWMul(m_SCEV(NewLHS), m_SCEVVScale())) &&
3657 match(RHS, m_scev_c_NUWMul(m_SCEV(NewRHS), m_SCEVVScale())))
3658 return getUDivExpr(NewLHS, NewRHS);
3659
3660 return getOrCreateUDivExpr(LHS, RHS);
3661}
3662
3663/// Get a canonical unsigned division expression, or something simpler if
3664/// possible. There is no representation for an exact udiv in SCEV IR, but we
3665/// can attempt to optimize it prior to construction.
3667 // Currently there is no exact specific logic.
3668
3669 return getUDivExpr(LHS, RHS);
3670}
3671
3672/// Get an add recurrence expression for the specified loop. Simplify the
3673/// expression as much as possible.
3675 const Loop *L, SCEVFlagsPair Flags) {
3677 Operands.push_back(Start);
3678 if (const SCEVAddRecExpr *StepChrec = dyn_cast<SCEVAddRecExpr>(Step))
3679 if (StepChrec->getLoop() == L) {
3680 append_range(Operands, StepChrec->operands());
3681 // The use flags describe the two-operand recurrence, not the flattened
3682 // one built here, so drop them just like the expression's NUW/NSW.
3683 return getAddRecExpr(Operands, L,
3684 maskFlags(Flags.ExprFlags, SCEV::FlagNW));
3685 }
3686
3687 Operands.push_back(Step);
3688 return getAddRecExpr(Operands, L, Flags);
3689}
3690
3691/// Get an add recurrence expression for the specified loop. Simplify the
3692/// expression as much as possible.
3694 const Loop *L, SCEVFlagsPair NWFlags) {
3695 SCEVFlags ExprFlags = NWFlags.ExprFlags;
3696 SCEVFlags UseFlags = NWFlags.UseFlags;
3697 assert(!(UseFlags & ~(SCEV::FlagNUW | SCEV::FlagNSW)) &&
3698 "only nuw or nsw allowed");
3699 if (Operands.size() == 1) return Operands[0];
3700#ifndef NDEBUG
3702 for (const SCEV *Op : llvm::drop_begin(Operands)) {
3703 assert(getEffectiveSCEVType(Op->getType()) == ETy &&
3704 "SCEVAddRecExpr operand types don't match!");
3705 assert(!Op->getType()->isPointerTy() && "Step must be integer");
3706 }
3707 for (const SCEV *Op : Operands)
3709 "SCEVAddRecExpr operand is not available at loop entry!");
3710
3711 // Keep track of the original operands, for verification when adding
3712 // use-specific flags.
3713 const SmallVector<SCEVUse, 4> OrigOperands(Operands.begin(), Operands.end());
3714#endif
3715
3716 if (Operands.back()->isZero()) {
3717 Operands.pop_back();
3718 return getAddRecExpr(Operands, L, SCEV::FlagNone); // {X,+,0} --> X
3719 }
3720
3721 // It's tempting to want to call getConstantMaxBackedgeTakenCount count here and
3722 // use that information to infer NUW and NSW flags. However, computing a
3723 // BE count requires calling getAddRecExpr, so we may not yet have a
3724 // meaningful BE count at this point (and if we don't, we'd be stuck
3725 // with a SCEVCouldNotCompute as the cached BE count).
3726
3727 ExprFlags = StrengthenNoWrapFlags(this, scAddRecExpr, Operands, ExprFlags);
3728
3729 // Canonicalize nested AddRecs in by nesting them in order of loop depth.
3730 if (const SCEVAddRecExpr *NestedAR = dyn_cast<SCEVAddRecExpr>(Operands[0])) {
3731 const Loop *NestedLoop = NestedAR->getLoop();
3732 if (L->contains(NestedLoop)
3733 ? (L->getLoopDepth() < NestedLoop->getLoopDepth())
3734 : (!NestedLoop->contains(L) &&
3735 DT.dominates(L->getHeader(), NestedLoop->getHeader()))) {
3736 SmallVector<SCEVUse, 4> NestedOperands(NestedAR->operands());
3737 Operands[0] = NestedAR->getStart();
3738 // AddRecs require their operands be loop-invariant with respect to their
3739 // loops. Don't perform this transformation if it would break this
3740 // requirement.
3741 bool AllInvariant = all_of(
3742 Operands, [&](const SCEV *Op) { return isLoopInvariant(Op, L); });
3743
3744 if (AllInvariant) {
3745 // Create a recurrence for the outer loop with the same step size.
3746 //
3747 // The outer recurrence keeps its NW flag but only keeps NUW/NSW if the
3748 // inner recurrence has the same property.
3749 SCEVFlags OuterFlags =
3750 maskFlags(ExprFlags, SCEV::FlagNW | NestedAR->getNoWrapFlags());
3751
3752 NestedOperands[0] = getAddRecExpr(Operands, L, OuterFlags);
3753 AllInvariant = all_of(NestedOperands, [&](const SCEV *Op) {
3754 return isLoopInvariant(Op, NestedLoop);
3755 });
3756
3757 if (AllInvariant) {
3758 // Ok, both add recurrences are valid after the transformation.
3759 //
3760 // The inner recurrence keeps its NW flag but only keeps NUW/NSW if
3761 // the outer recurrence has the same property.
3762 SCEVFlags InnerFlags =
3763 maskFlags(NestedAR->getNoWrapFlags(), SCEV::FlagNW | ExprFlags);
3764 return getAddRecExpr(NestedOperands, NestedLoop, InnerFlags);
3765 }
3766 }
3767 // Reset Operands to its original state.
3768 Operands[0] = NestedAR;
3769 }
3770 }
3771
3772 // Okay, it looks like we really DO need an addrec expr. Check to see if we
3773 // already have one, otherwise create a new one.
3774 assert((UseFlags == SCEV::FlagNone || equal(OrigOperands, Operands)) &&
3775 "Tried to add SCEVUse flags after operands changed");
3776 return {getOrCreateAddRecExpr(Operands, L, ExprFlags), UseFlags};
3777}
3778
3780 ArrayRef<SCEVUse> IndexExprs) {
3781 const SCEV *BaseExpr = getSCEV(GEP->getPointerOperand());
3782 // getSCEV(Base)->getType() has the same address space as Base->getType()
3783 // because SCEV::getType() preserves the address space.
3784 GEPNoWrapFlags NW = GEP->getNoWrapFlags();
3785 if (NW != GEPNoWrapFlags::none()) {
3786 // We'd like to propagate flags from the IR to the corresponding SCEV nodes,
3787 // but to do that, we have to ensure that said flag is valid in the entire
3788 // defined scope of the SCEV.
3789 // TODO: non-instructions have global scope. We might be able to prove
3790 // some global scope cases
3791 auto *GEPI = dyn_cast<Instruction>(GEP);
3792 if (!GEPI || !isSCEVExprNeverPoison(GEPI))
3793 NW = GEPNoWrapFlags::none();
3794 }
3795
3796 return getGEPExpr(BaseExpr, IndexExprs, GEP->getSourceElementType(), NW);
3797}
3798
3800 ArrayRef<SCEVUse> IndexExprs,
3801 Type *SrcElementTy, GEPNoWrapFlags NW) {
3802 SCEVFlags OffsetWrap = SCEV::FlagNone;
3803 if (NW.hasNoUnsignedSignedWrap())
3804 OffsetWrap = setFlags(OffsetWrap, SCEV::FlagNSW);
3805 if (NW.hasNoUnsignedWrap())
3806 OffsetWrap = setFlags(OffsetWrap, SCEV::FlagNUW);
3807
3808 Type *CurTy = BaseExpr->getType();
3809 Type *IntIdxTy = getEffectiveSCEVType(BaseExpr->getType());
3810 bool FirstIter = true;
3812 for (SCEVUse IndexExpr : IndexExprs) {
3813 // Compute the (potentially symbolic) offset in bytes for this index.
3814 if (StructType *STy = dyn_cast<StructType>(CurTy)) {
3815 // For a struct, add the member offset.
3816 ConstantInt *Index = cast<SCEVConstant>(IndexExpr)->getValue();
3817 unsigned FieldNo = Index->getZExtValue();
3818 const SCEV *FieldOffset = getOffsetOfExpr(IntIdxTy, STy, FieldNo);
3819 Offsets.push_back(FieldOffset);
3820
3821 // Update CurTy to the type of the field at Index.
3822 CurTy = STy->getTypeAtIndex(Index);
3823 } else {
3824 // Update CurTy to its element type.
3825 if (FirstIter) {
3826 assert(isa<PointerType>(CurTy) &&
3827 "The first index of a GEP indexes a pointer");
3828 CurTy = SrcElementTy;
3829 FirstIter = false;
3830 } else {
3831 CurTy = GetElementPtrInst::getTypeAtIndex(CurTy, (uint64_t)0);
3832 }
3833 // For an array, add the element offset, explicitly scaled.
3834 const SCEV *ElementSize = getSizeOfExpr(IntIdxTy, CurTy);
3835 // Getelementptr indices are signed.
3836 IndexExpr = getTruncateOrSignExtend(IndexExpr, IntIdxTy);
3837
3838 // Multiply the index by the element size to compute the element offset.
3839 const SCEV *LocalOffset = getMulExpr(IndexExpr, ElementSize, OffsetWrap);
3840 Offsets.push_back(LocalOffset);
3841 }
3842 }
3843
3844 // Handle degenerate case of GEP without offsets.
3845 if (Offsets.empty())
3846 return BaseExpr;
3847
3848 // Add the offsets together, assuming nsw if inbounds.
3849 const SCEV *Offset = getAddExpr(Offsets, OffsetWrap);
3850 // Add the base address and the offset. We cannot use the nsw flag, as the
3851 // base address is unsigned. However, if we know that the offset is
3852 // non-negative, we can use nuw.
3853 bool NUW = NW.hasNoUnsignedWrap() ||
3855 SCEVFlags BaseWrap = NUW ? SCEV::FlagNUW : SCEV::FlagNone;
3856 const SCEV *GEPExpr = getAddExpr(BaseExpr, Offset, BaseWrap);
3857 assert(BaseExpr->getType() == GEPExpr->getType() &&
3858 "GEP should not change type mid-flight.");
3859 return GEPExpr;
3860}
3861
3862SCEV *ScalarEvolution::findExistingSCEVInCache(SCEVTypes SCEVType,
3864 const Loop *L) {
3865 assert((SCEVType != scAddRecExpr || L) &&
3866 "L must be passed to find existing AddRecs");
3868 ID.AddInteger(SCEVType);
3869 for (SCEVUse Op : Ops)
3870 ID.AddPointer(Op.getOpaqueValue());
3871 if (L)
3872 ID.AddPointer(L);
3874 return UniqueSCEVs.lookup(ID, Token);
3875}
3876
3877const SCEV *ScalarEvolution::getAbsExpr(const SCEV *Op, bool IsNSW) {
3878 SCEVFlags Flags = IsNSW ? SCEV::FlagNSW : SCEV::FlagNone;
3879 return getSMaxExpr(Op, getNegativeSCEV(Op, Flags));
3880}
3881
3884 assert(SCEVMinMaxExpr::isMinMaxType(Kind) && "Not a SCEVMinMaxExpr!");
3885 assert(!Ops.empty() && "Cannot get empty (u|s)(min|max)!");
3886 if (Ops.size() == 1) return Ops[0];
3887#ifndef NDEBUG
3888 Type *ETy = getEffectiveSCEVType(Ops[0]->getType());
3889 for (unsigned i = 1, e = Ops.size(); i != e; ++i) {
3890 assert(getEffectiveSCEVType(Ops[i]->getType()) == ETy &&
3891 "Operand types don't match!");
3892 assert(Ops[0]->getType()->isPointerTy() ==
3893 Ops[i]->getType()->isPointerTy() &&
3894 "min/max should be consistently pointerish");
3895 }
3896#endif
3897
3898 bool IsSigned = Kind == scSMaxExpr || Kind == scSMinExpr;
3899 bool IsMax = Kind == scSMaxExpr || Kind == scUMaxExpr;
3900
3901 const SCEV *Folded = constantFoldAndGroupOps(
3902 *this, LI, DT, Ops,
3903 [&](const APInt &C1, const APInt &C2) {
3904 switch (Kind) {
3905 case scSMaxExpr:
3906 return APIntOps::smax(C1, C2);
3907 case scSMinExpr:
3908 return APIntOps::smin(C1, C2);
3909 case scUMaxExpr:
3910 return APIntOps::umax(C1, C2);
3911 case scUMinExpr:
3912 return APIntOps::umin(C1, C2);
3913 default:
3914 llvm_unreachable("Unknown SCEV min/max opcode");
3915 }
3916 },
3917 [&](const APInt &C) {
3918 // identity
3919 if (IsMax)
3920 return IsSigned ? C.isMinSignedValue() : C.isMinValue();
3921 else
3922 return IsSigned ? C.isMaxSignedValue() : C.isMaxValue();
3923 },
3924 [&](const APInt &C) {
3925 // absorber
3926 if (IsMax)
3927 return IsSigned ? C.isMaxSignedValue() : C.isMaxValue();
3928 else
3929 return IsSigned ? C.isMinSignedValue() : C.isMinValue();
3930 });
3931 if (Folded)
3932 return Folded;
3933
3934 // Check if we have created the same expression before.
3935 if (const SCEV *S = findExistingSCEVInCache(Kind, Ops)) {
3936 return S;
3937 }
3938
3939 // Find the first operation of the same kind
3940 unsigned Idx = 0;
3941 while (Idx < Ops.size() && Ops[Idx]->getSCEVType() < Kind)
3942 ++Idx;
3943
3944 // Check to see if one of the operands is of the same kind. If so, expand its
3945 // operands onto our operand list, and recurse to simplify.
3946 if (Idx < Ops.size()) {
3947 bool DeletedAny = false;
3948 while (Ops[Idx]->getSCEVType() == Kind) {
3949 const SCEVMinMaxExpr *SMME = cast<SCEVMinMaxExpr>(Ops[Idx]);
3950 Ops.erase(Ops.begin()+Idx);
3951 append_range(Ops, SMME->operands());
3952 DeletedAny = true;
3953 }
3954
3955 if (DeletedAny)
3956 return getMinMaxExpr(Kind, Ops);
3957 }
3958
3959 // Okay, check to see if the same value occurs in the operand list twice. If
3960 // so, delete one. Since we sorted the list, these values are required to
3961 // be adjacent.
3966 llvm::CmpInst::Predicate FirstPred = IsMax ? GEPred : LEPred;
3967 llvm::CmpInst::Predicate SecondPred = IsMax ? LEPred : GEPred;
3968 for (unsigned i = 0, e = Ops.size() - 1; i != e; ++i) {
3969 if (Ops[i] == Ops[i + 1] ||
3970 isKnownViaNonRecursiveReasoning(FirstPred, Ops[i], Ops[i + 1])) {
3971 // X op Y op Y --> X op Y
3972 // X op Y --> X, if we know X, Y are ordered appropriately
3973 Ops.erase(Ops.begin() + i + 1, Ops.begin() + i + 2);
3974 --i;
3975 --e;
3976 } else if (isKnownViaNonRecursiveReasoning(SecondPred, Ops[i],
3977 Ops[i + 1])) {
3978 // X op Y --> Y, if we know X, Y are ordered appropriately
3979 Ops.erase(Ops.begin() + i, Ops.begin() + i + 1);
3980 --i;
3981 --e;
3982 }
3983 }
3984
3985 if (Ops.size() == 1) return Ops[0];
3986
3987 assert(!Ops.empty() && "Reduced smax down to nothing!");
3988
3989 // Okay, it looks like we really DO need an expr. Check to see if we
3990 // already have one, otherwise create a new one.
3992 ID.AddInteger(Kind);
3993 for (SCEVUse Op : Ops)
3994 ID.AddPointer(Op.getOpaqueValue());
3996 const SCEV *ExistingSCEV = UniqueSCEVs.lookup(ID, Token);
3997 if (ExistingSCEV)
3998 return ExistingSCEV;
3999 SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
4001 SCEV *S = new (SCEVAllocator)
4002 SCEVMinMaxExpr(ID.Intern(SCEVAllocator), Kind, O, Ops.size());
4003
4004 UniqueSCEVs.insert(S, Token);
4005 S->computeAndSetCanonical(*this);
4006 registerUser(S, Ops);
4007 return S;
4008}
4009
4010namespace {
4011
4012class SCEVSequentialMinMaxDeduplicatingVisitor final
4013 : public SCEVVisitor<SCEVSequentialMinMaxDeduplicatingVisitor,
4014 std::optional<const SCEV *>> {
4015 using RetVal = std::optional<const SCEV *>;
4016
4017 ScalarEvolution &SE;
4018 const SCEVTypes RootKind; // Must be a sequential min/max expression.
4019 const SCEVTypes NonSequentialRootKind; // Non-sequential variant of RootKind.
4021
4022 bool canRecurseInto(SCEVTypes Kind) const {
4023 // We can only recurse into the SCEV expression of the same effective type
4024 // as the type of our root SCEV expression.
4025 return RootKind == Kind || NonSequentialRootKind == Kind;
4026 };
4027
4028 RetVal visit(const SCEV *S) {
4029 // Has the whole operand been seen already?
4030 if (!SeenOps.insert(S).second)
4031 return std::nullopt;
4033 SCEVTypes Kind = S->getSCEVType();
4034
4035 if (!canRecurseInto(Kind))
4036 return S;
4037
4038 auto *NAry = cast<SCEVNAryExpr>(S);
4039 SmallVector<SCEVUse> NewOps;
4040 bool Changed = visit(Kind, NAry->operands(), NewOps);
4041
4042 if (!Changed)
4043 return S;
4044 if (NewOps.empty())
4045 return std::nullopt;
4046
4048 ? SE.getSequentialMinMaxExpr(Kind, NewOps)
4049 : SE.getMinMaxExpr(Kind, NewOps);
4050 }
4051 return S;
4052 }
4053
4054public:
4055 SCEVSequentialMinMaxDeduplicatingVisitor(ScalarEvolution &SE,
4056 SCEVTypes RootKind)
4057 : SE(SE), RootKind(RootKind),
4058 NonSequentialRootKind(
4059 SCEVSequentialMinMaxExpr::getEquivalentNonSequentialSCEVType(
4060 RootKind)) {}
4061
4062 bool /*Changed*/ visit(SCEVTypes Kind, ArrayRef<SCEVUse> OrigOps,
4063 SmallVectorImpl<SCEVUse> &NewOps) {
4064 bool Changed = false;
4066 Ops.reserve(OrigOps.size());
4067
4068 for (const SCEV *Op : OrigOps) {
4069 RetVal NewOp = visit(Op);
4070 if (NewOp != Op)
4071 Changed = true;
4072 if (NewOp)
4073 Ops.emplace_back(*NewOp);
4074 }
4075
4076 if (Changed)
4077 NewOps = std::move(Ops);
4078 return Changed;
4079 }
4080};
4081
4082} // namespace
4083
4085 switch (Kind) {
4086 case scConstant:
4087 case scVScale:
4088 case scTruncate:
4089 case scZeroExtend:
4090 case scSignExtend:
4091 case scPtrToAddr:
4092 case scAddExpr:
4093 case scMulExpr:
4094 case scUDivExpr:
4095 case scAddRecExpr:
4096 case scUMaxExpr:
4097 case scSMaxExpr:
4098 case scUMinExpr:
4099 case scSMinExpr:
4100 case scUnknown:
4101 // If any operand is poison, the whole expression is poison.
4102 return true;
4104 // FIXME: if the *first* operand is poison, the whole expression is poison.
4105 return false; // Pessimistically, say that it does not propagate poison.
4106 case scCouldNotCompute:
4107 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
4108 }
4109 llvm_unreachable("Unknown SCEV kind!");
4110}
4111
4112namespace {
4113// The only way poison may be introduced in a SCEV expression is from a
4114// poison SCEVUnknown (ConstantExprs are also represented as SCEVUnknown,
4115// not SCEVConstant). Notably, SCEVFlags on SCEV nodes can *not*
4116// introduce poison -- they encode guaranteed, non-speculated knowledge.
4117//
4118// Additionally, all SCEV nodes propagate poison from inputs to outputs,
4119// with the notable exception of umin_seq, where only poison from the first
4120// operand is (unconditionally) propagated.
4121struct SCEVPoisonCollector {
4122 bool LookThroughMaybePoisonBlocking;
4123 SmallPtrSet<const SCEVUnknown *, 4> MaybePoison;
4124 SCEVPoisonCollector(bool LookThroughMaybePoisonBlocking)
4125 : LookThroughMaybePoisonBlocking(LookThroughMaybePoisonBlocking) {}
4126
4127 bool follow(const SCEV *S) {
4128 if (!LookThroughMaybePoisonBlocking &&
4130 return false;
4131
4132 if (auto *SU = dyn_cast<SCEVUnknown>(S)) {
4133 if (!isGuaranteedNotToBePoison(SU->getValue()))
4134 MaybePoison.insert(SU);
4135 }
4136 return true;
4137 }
4138 bool isDone() const { return false; }
4139};
4140} // namespace
4141
4142/// Return true if V is poison given that AssumedPoison is already poison.
4143static bool impliesPoison(const SCEV *AssumedPoison, const SCEV *S) {
4144 // First collect all SCEVs that might result in AssumedPoison to be poison.
4145 // We need to look through potentially poison-blocking operations here,
4146 // because we want to find all SCEVs that *might* result in poison, not only
4147 // those that are *required* to.
4148 SCEVPoisonCollector PC1(/* LookThroughMaybePoisonBlocking */ true);
4149 visitAll(AssumedPoison, PC1);
4150
4151 // AssumedPoison is never poison. As the assumption is false, the implication
4152 // is true. Don't bother walking the other SCEV in this case.
4153 if (PC1.MaybePoison.empty())
4154 return true;
4155
4156 // Collect all SCEVs in S that, if poison, *will* result in S being poison
4157 // as well. We cannot look through potentially poison-blocking operations
4158 // here, as their arguments only *may* make the result poison.
4159 SCEVPoisonCollector PC2(/* LookThroughMaybePoisonBlocking */ false);
4160 visitAll(S, PC2);
4161
4162 // Make sure that no matter which SCEV in PC1.MaybePoison is actually poison,
4163 // it will also make S poison by being part of PC2.MaybePoison.
4164 return llvm::set_is_subset(PC1.MaybePoison, PC2.MaybePoison);
4165}
4166
4168 SmallPtrSetImpl<const Value *> &Result, const SCEV *S) {
4169 SCEVPoisonCollector PC(/* LookThroughMaybePoisonBlocking */ false);
4170 visitAll(S, PC);
4171 for (const SCEVUnknown *SU : PC.MaybePoison)
4172 Result.insert(SU->getValue());
4173}
4174
4176 const SCEV *S, Instruction *I,
4177 SmallVectorImpl<Instruction *> &DropPoisonGeneratingInsts) {
4178 // If the instruction cannot be poison, it's always safe to reuse.
4180 return true;
4181
4182 // Otherwise, it is possible that I is more poisonous that S. Collect the
4183 // poison-contributors of S, and then check whether I has any additional
4184 // poison-contributors. Poison that is contributed through poison-generating
4185 // flags is handled by dropping those flags instead.
4187 getPoisonGeneratingValues(PoisonVals, S);
4188
4189 SmallVector<Value *> Worklist;
4191 Worklist.push_back(I);
4192 while (!Worklist.empty()) {
4193 Value *V = Worklist.pop_back_val();
4194 if (!Visited.insert(V).second)
4195 continue;
4196
4197 // Avoid walking large instruction graphs.
4198 if (Visited.size() > 16)
4199 return false;
4200
4201 // Either the value can't be poison, or the S would also be poison if it
4202 // is.
4203 if (PoisonVals.contains(V) || ::isGuaranteedNotToBePoison(V))
4204 continue;
4205
4206 auto *I = dyn_cast<Instruction>(V);
4207 if (!I)
4208 return false;
4209
4210 // Disjoint or instructions are interpreted as adds by SCEV. However, we
4211 // can't replace an arbitrary add with disjoint or, even if we drop the
4212 // flag. We would need to convert the or into an add.
4213 if (auto *PDI = dyn_cast<PossiblyDisjointInst>(I))
4214 if (PDI->isDisjoint())
4215 return false;
4216
4217 // FIXME: Ignore vscale, even though it technically could be poison. Do this
4218 // because SCEV currently assumes it can't be poison. Remove this special
4219 // case once we proper model when vscale can be poison.
4220 if (auto *II = dyn_cast<IntrinsicInst>(I);
4221 II && II->getIntrinsicID() == Intrinsic::vscale)
4222 continue;
4223
4224 if (canCreatePoison(cast<Operator>(I), /*ConsiderFlagsAndMetadata*/ false))
4225 return false;
4226
4227 // If the instruction can't create poison, we can recurse to its operands.
4228 if (I->hasPoisonGeneratingAnnotations())
4229 DropPoisonGeneratingInsts.push_back(I);
4230
4231 llvm::append_range(Worklist, I->operands());
4232 }
4233 return true;
4234}
4235
4236const SCEV *
4239 assert(SCEVSequentialMinMaxExpr::isSequentialMinMaxType(Kind) &&
4240 "Not a SCEVSequentialMinMaxExpr!");
4241 assert(!Ops.empty() && "Cannot get empty (u|s)(min|max)!");
4242 if (Ops.size() == 1)
4243 return Ops[0];
4244#ifndef NDEBUG
4245 Type *ETy = getEffectiveSCEVType(Ops[0]->getType());
4246 for (unsigned i = 1, e = Ops.size(); i != e; ++i) {
4247 assert(getEffectiveSCEVType(Ops[i]->getType()) == ETy &&
4248 "Operand types don't match!");
4249 assert(Ops[0]->getType()->isPointerTy() ==
4250 Ops[i]->getType()->isPointerTy() &&
4251 "min/max should be consistently pointerish");
4252 }
4253#endif
4254
4255 // Note that SCEVSequentialMinMaxExpr is *NOT* commutative,
4256 // so we can *NOT* do any kind of sorting of the expressions!
4257
4258 // Check if we have created the same expression before.
4259 if (const SCEV *S = findExistingSCEVInCache(Kind, Ops))
4260 return S;
4261
4262 // FIXME: there are *some* simplifications that we can do here.
4263
4264 // Keep only the first instance of an operand.
4265 {
4266 SCEVSequentialMinMaxDeduplicatingVisitor Deduplicator(*this, Kind);
4267 bool Changed = Deduplicator.visit(Kind, Ops, Ops);
4268 if (Changed)
4269 return getSequentialMinMaxExpr(Kind, Ops);
4270 }
4271
4272 // Check to see if one of the operands is of the same kind. If so, expand its
4273 // operands onto our operand list, and recurse to simplify.
4274 {
4275 unsigned Idx = 0;
4276 bool DeletedAny = false;
4277 while (Idx < Ops.size()) {
4278 if (Ops[Idx]->getSCEVType() != Kind) {
4279 ++Idx;
4280 continue;
4281 }
4282 const auto *SMME = cast<SCEVSequentialMinMaxExpr>(Ops[Idx]);
4283 Ops.erase(Ops.begin() + Idx);
4284 Ops.insert(Ops.begin() + Idx, SMME->operands().begin(),
4285 SMME->operands().end());
4286 DeletedAny = true;
4287 }
4288
4289 if (DeletedAny)
4290 return getSequentialMinMaxExpr(Kind, Ops);
4291 }
4292
4293 const SCEV *SaturationPoint;
4295 switch (Kind) {
4297 SaturationPoint = getZero(Ops[0]->getType());
4298 Pred = ICmpInst::ICMP_ULE;
4299 break;
4300 default:
4301 llvm_unreachable("Not a sequential min/max type.");
4302 }
4303
4304 for (unsigned i = 1, e = Ops.size(); i != e; ++i) {
4305 if (!isGuaranteedNotToCauseUB(Ops[i]))
4306 continue;
4307 // We can replace %x umin_seq %y with %x umin %y if either:
4308 // * %y being poison implies %x is also poison.
4309 // * %x cannot be the saturating value (e.g. zero for umin).
4310 if (::impliesPoison(Ops[i], Ops[i - 1]) ||
4311 isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_NE, Ops[i - 1],
4312 SaturationPoint)) {
4313 SmallVector<SCEVUse, 2> SeqOps = {Ops[i - 1], Ops[i]};
4314 Ops[i - 1] = getMinMaxExpr(
4316 SeqOps);
4317 Ops.erase(Ops.begin() + i);
4318 return getSequentialMinMaxExpr(Kind, Ops);
4319 }
4320 // Fold %x umin_seq %y to %x if %x ule %y.
4321 // TODO: We might be able to prove the predicate for a later operand.
4322 if (isKnownViaNonRecursiveReasoning(Pred, Ops[i - 1], Ops[i])) {
4323 Ops.erase(Ops.begin() + i);
4324 return getSequentialMinMaxExpr(Kind, Ops);
4325 }
4326 }
4327
4328 // Okay, it looks like we really DO need an expr. Check to see if we
4329 // already have one, otherwise create a new one.
4331 ID.AddInteger(Kind);
4332 for (SCEVUse Op : Ops)
4333 ID.AddPointer(Op.getOpaqueValue());
4335 const SCEV *ExistingSCEV = UniqueSCEVs.lookup(ID, Token);
4336 if (ExistingSCEV)
4337 return ExistingSCEV;
4338
4339 SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
4341 SCEV *S = new (SCEVAllocator)
4342 SCEVSequentialMinMaxExpr(ID.Intern(SCEVAllocator), Kind, O, Ops.size());
4343
4344 UniqueSCEVs.insert(S, Token);
4345 S->computeAndSetCanonical(*this);
4346 registerUser(S, Ops);
4347 return S;
4348}
4349
4354
4358
4363
4367
4372
4376
4378 bool Sequential) {
4379 SmallVector<SCEVUse, 2> Ops = {LHS, RHS};
4380 return getUMinExpr(Ops, Sequential);
4381}
4382
4388
4389const SCEV *
4391 const SCEV *Res = getConstant(IntTy, Size.getKnownMinValue());
4392 if (Size.isScalable())
4393 Res = getMulExpr(Res, getVScale(IntTy));
4394 return Res;
4395}
4396
4398 return getSizeOfExpr(IntTy, getDataLayout().getTypeAllocSize(AllocTy));
4399}
4400
4402 return getSizeOfExpr(IntTy, getDataLayout().getTypeStoreSize(StoreTy));
4403}
4404
4406 StructType *STy,
4407 unsigned FieldNo) {
4408 // We can bypass creating a target-independent constant expression and then
4409 // folding it back into a ConstantInt. This is just a compile-time
4410 // optimization.
4411 const StructLayout *SL = getDataLayout().getStructLayout(STy);
4412 assert(!SL->getSizeInBits().isScalable() &&
4413 "Cannot get offset for structure containing scalable vector types");
4414 return getConstant(IntTy, SL->getElementOffset(FieldNo));
4415}
4416
4418 // Don't attempt to do anything other than create a SCEVUnknown object
4419 // here. createSCEV only calls getUnknown after checking for all other
4420 // interesting possibilities, and any other code that calls getUnknown
4421 // is doing so in order to hide a value from SCEV canonicalization.
4422
4425 ID.AddPointer(V);
4427 if (SCEV *S = UniqueSCEVs.lookup(ID, Token)) {
4428 assert(cast<SCEVUnknown>(S)->getValue() == V &&
4429 "Stale SCEVUnknown in uniquing map!");
4430 return S;
4431 }
4432 SCEV *S = new (SCEVAllocator) SCEVUnknown(ID.Intern(SCEVAllocator), V, this,
4433 FirstUnknown);
4434 FirstUnknown = cast<SCEVUnknown>(S);
4435 UniqueSCEVs.insert(S, Token);
4436 S->computeAndSetCanonical(*this);
4437 return S;
4438}
4439
4440//===----------------------------------------------------------------------===//
4441// Basic SCEV Analysis and PHI Idiom Recognition Code
4442//
4443
4444/// Test if values of the given type are analyzable within the SCEV
4445/// framework. This primarily includes integer types, and it can optionally
4446/// include pointer types if the ScalarEvolution class has access to
4447/// target-specific information.
4449 // Integers and pointers are always SCEVable.
4450 return Ty->isIntOrPtrTy();
4451}
4452
4453/// Return the size in bits of the specified type, for which isSCEVable must
4454/// return true.
4456 assert(isSCEVable(Ty) && "Type is not SCEVable!");
4457 if (Ty->isPointerTy())
4459 return getDataLayout().getTypeSizeInBits(Ty);
4460}
4461
4462/// Return a type with the same bitwidth as the given type and which represents
4463/// how SCEV will treat the given type, for which isSCEVable must return
4464/// true. For pointer types, this is the pointer index sized integer type.
4466 assert(isSCEVable(Ty) && "Type is not SCEVable!");
4467
4468 if (Ty->isIntegerTy())
4469 return Ty;
4470
4471 // The only other support type is pointer.
4472 assert(Ty->isPointerTy() && "Unexpected non-pointer non-integer type!");
4473 return getDataLayout().getIndexType(Ty);
4474}
4475
4477 return getTypeSizeInBits(T1) >= getTypeSizeInBits(T2) ? T1 : T2;
4478}
4479
4481 const SCEV *B) {
4482 /// For a valid use point to exist, the defining scope of one operand
4483 /// must dominate the other.
4484 bool PreciseA, PreciseB;
4485 auto *ScopeA = getDefiningScopeBound({A}, PreciseA);
4486 auto *ScopeB = getDefiningScopeBound({B}, PreciseB);
4487 if (!PreciseA || !PreciseB)
4488 // Can't tell.
4489 return false;
4490 return (ScopeA == ScopeB) || DT.dominates(ScopeA, ScopeB) ||
4491 DT.dominates(ScopeB, ScopeA);
4492}
4493
4495 return CouldNotCompute.get();
4496}
4497
4498bool ScalarEvolution::checkValidity(const SCEV *S) const {
4499 bool ContainsNulls = SCEVExprContains(S, [](const SCEV *S) {
4500 auto *SU = dyn_cast<SCEVUnknown>(S);
4501 return SU && SU->getValue() == nullptr;
4502 });
4503
4504 return !ContainsNulls;
4505}
4506
4508 HasRecMapType::iterator I = HasRecMap.find(S);
4509 if (I != HasRecMap.end())
4510 return I->second;
4511
4512 bool FoundAddRec =
4513 SCEVExprContains(S, [](const SCEV *S) { return isa<SCEVAddRecExpr>(S); });
4514 HasRecMap.insert({S, FoundAddRec});
4515 return FoundAddRec;
4516}
4517
4518/// Return the ValueOffsetPair set for \p S. \p S can be represented
4519/// by the value and offset from any ValueOffsetPair in the set.
4520ArrayRef<Value *> ScalarEvolution::getSCEVValues(const SCEV *S) {
4521 ExprValueMapType::iterator SI = ExprValueMap.find_as(S);
4522 if (SI == ExprValueMap.end())
4523 return {};
4524 return SI->second.getArrayRef();
4525}
4526
4527/// Erase Value from ValueExprMap and ExprValueMap. ValueExprMap.erase(V)
4528/// cannot be used separately. eraseValueFromMap should be used to remove
4529/// V from ValueExprMap and ExprValueMap at the same time.
4530void ScalarEvolution::eraseValueFromMap(Value *V) {
4531 ValueExprMapType::iterator I = ValueExprMap.find_as(V);
4532 if (I != ValueExprMap.end()) {
4533 auto EVIt = ExprValueMap.find(I->second);
4534 bool Removed = EVIt->second.remove(V);
4535 (void) Removed;
4536 assert(Removed && "Value not in ExprValueMap?");
4537 ValueExprMap.erase(I);
4538 }
4539}
4540
4541void ScalarEvolution::insertValueToMap(Value *V, const SCEV *S) {
4542 // A recursive query may have already computed the SCEV. It should be
4543 // equivalent, but may not necessarily be exactly the same, e.g. due to lazily
4544 // inferred nowrap flags.
4545 auto It = ValueExprMap.find_as(V);
4546 if (It == ValueExprMap.end()) {
4547 ValueExprMap.insert({SCEVCallbackVH(V, this), S});
4548 ExprValueMap[S].insert(V);
4549 }
4550}
4551
4552/// Return an existing SCEV if it exists, otherwise analyze the expression and
4553/// create a new one.
4555 assert(isSCEVable(V->getType()) && "Value is not SCEVable!");
4556
4557 if (const SCEV *S = getExistingSCEV(V))
4558 return S;
4559 return createSCEVIter(V);
4560}
4561
4563 assert(isSCEVable(V->getType()) && "Value is not SCEVable!");
4564
4565 ValueExprMapType::iterator I = ValueExprMap.find_as(V);
4566 if (I != ValueExprMap.end()) {
4567 const SCEV *S = I->second;
4568 assert(checkValidity(S) &&
4569 "existing SCEV has not been properly invalidated");
4570 return S;
4571 }
4572 return nullptr;
4573}
4574
4575/// Return a SCEV corresponding to -V = -1*V
4577 if (const SCEVConstant *VC = dyn_cast<SCEVConstant>(V))
4578 return getConstant(
4579 cast<ConstantInt>(ConstantExpr::getNeg(VC->getValue())));
4580
4581 Type *Ty = V->getType();
4582 Ty = getEffectiveSCEVType(Ty);
4583 return getMulExpr(V, getMinusOne(Ty), Flags);
4584}
4585
4586/// If Expr computes ~A, return A else return nullptr
4587static const SCEV *MatchNotExpr(const SCEV *Expr) {
4588 const SCEV *MulOp;
4589 if (match(Expr, m_scev_Add(m_scev_AllOnes(),
4590 m_scev_Mul(m_scev_AllOnes(), m_SCEV(MulOp)))))
4591 return MulOp;
4592 return nullptr;
4593}
4594
4595/// Return a SCEV corresponding to ~V = -1-V
4597 assert(!V->getType()->isPointerTy() && "Can't negate pointer");
4598
4599 if (const SCEVConstant *VC = dyn_cast<SCEVConstant>(V))
4600 return getConstant(
4601 cast<ConstantInt>(ConstantExpr::getNot(VC->getValue())));
4602
4603 // Fold ~(u|s)(min|max)(~x, ~y) to (u|s)(max|min)(x, y)
4604 if (const SCEVMinMaxExpr *MME = dyn_cast<SCEVMinMaxExpr>(V)) {
4605 auto MatchMinMaxNegation = [&](const SCEVMinMaxExpr *MME) {
4606 SmallVector<SCEVUse, 2> MatchedOperands;
4607 for (const SCEV *Operand : MME->operands()) {
4608 const SCEV *Matched = MatchNotExpr(Operand);
4609 if (!Matched)
4610 return (const SCEV *)nullptr;
4611 MatchedOperands.push_back(Matched);
4612 }
4613 return getMinMaxExpr(SCEVMinMaxExpr::negate(MME->getSCEVType()),
4614 MatchedOperands);
4615 };
4616 if (const SCEV *Replaced = MatchMinMaxNegation(MME))
4617 return Replaced;
4618 }
4619
4620 Type *Ty = V->getType();
4621 Ty = getEffectiveSCEVType(Ty);
4622 return getMinusSCEV(getMinusOne(Ty), V);
4623}
4624
4626 assert(P->getType()->isPointerTy());
4627
4628 if (auto *AddRec = dyn_cast<SCEVAddRecExpr>(P)) {
4629 // The base of an AddRec is the first operand.
4630 SmallVector<SCEVUse> Ops{AddRec->operands()};
4631 Ops[0] = removePointerBase(Ops[0]);
4632 // Don't try to transfer nowrap flags for now. We could in some cases
4633 // (for example, if pointer operand of the AddRec is a SCEVUnknown).
4634 return getAddRecExpr(Ops, AddRec->getLoop(), SCEV::FlagNone);
4635 }
4636 if (auto *Add = dyn_cast<SCEVAddExpr>(P)) {
4637 // The base of an Add is the pointer operand.
4638 SmallVector<SCEVUse> Ops{Add->operands()};
4639 SCEVUse *PtrOp = nullptr;
4640 for (SCEVUse &AddOp : Ops) {
4641 if (AddOp->getType()->isPointerTy()) {
4642 assert(!PtrOp && "Cannot have multiple pointer ops");
4643 PtrOp = &AddOp;
4644 }
4645 }
4646 *PtrOp = removePointerBase(*PtrOp);
4647 // Don't try to transfer nowrap flags for now. We could in some cases
4648 // (for example, if the pointer operand of the Add is a SCEVUnknown).
4649 return getAddExpr(Ops);
4650 }
4651 // Any other expression must be a pointer base.
4652 return getZero(P->getType());
4653}
4654
4656 SCEVFlags Flags, unsigned Depth) {
4657 // Fast path: X - X --> 0.
4658 if (LHS == RHS)
4659 return getZero(LHS->getType());
4660
4661 // If we subtract two pointers with different pointer bases, bail.
4662 // Eventually, we're going to add an assertion to getMulExpr that we
4663 // can't multiply by a pointer.
4664 if (RHS->getType()->isPointerTy()) {
4665 if (!LHS->getType()->isPointerTy() ||
4666 getPointerBase(LHS) != getPointerBase(RHS))
4667 return getCouldNotCompute();
4668 LHS = removePointerBase(LHS);
4669 RHS = removePointerBase(RHS);
4670 }
4671
4672 // We represent LHS - RHS as LHS + (-1)*RHS. This transformation
4673 // makes it so that we cannot make much use of NUW.
4674 auto AddFlags = SCEV::FlagNone;
4675 const bool RHSIsNotMinSigned =
4677 if (hasFlags(Flags, SCEV::FlagNSW)) {
4678 // Let M be the minimum representable signed value. Then (-1)*RHS
4679 // signed-wraps if and only if RHS is M. That can happen even for
4680 // a NSW subtraction because e.g. (-1)*M signed-wraps even though
4681 // -1 - M does not. So to transfer NSW from LHS - RHS to LHS +
4682 // (-1)*RHS, we need to prove that RHS != M.
4683 //
4684 // If LHS is non-negative and we know that LHS - RHS does not
4685 // signed-wrap, then RHS cannot be M. So we can rule out signed-wrap
4686 // either by proving that RHS > M or that LHS >= 0.
4687 if (RHSIsNotMinSigned || isKnownNonNegative(LHS)) {
4688 AddFlags = SCEV::FlagNSW;
4689 }
4690 }
4691
4692 // FIXME: Find a correct way to transfer NSW to (-1)*M when LHS -
4693 // RHS is NSW and LHS >= 0.
4694 //
4695 // The difficulty here is that the NSW flag may have been proven
4696 // relative to a loop that is to be found in a recurrence in LHS and
4697 // not in RHS. Applying NSW to (-1)*M may then let the NSW have a
4698 // larger scope than intended.
4699 auto NegFlags = RHSIsNotMinSigned ? SCEV::FlagNSW : SCEV::FlagNone;
4700
4701 return getAddExpr(LHS, getNegativeSCEV(RHS, NegFlags), AddFlags, Depth);
4702}
4703
4705 unsigned Depth) {
4706 Type *SrcTy = V->getType();
4707 assert(SrcTy->isIntOrPtrTy() && Ty->isIntOrPtrTy() &&
4708 "Cannot truncate or zero extend with non-integer arguments!");
4709 if (getTypeSizeInBits(SrcTy) == getTypeSizeInBits(Ty))
4710 return V; // No conversion
4711 if (getTypeSizeInBits(SrcTy) > getTypeSizeInBits(Ty))
4712 return getTruncateExpr(V, Ty, Depth);
4713 return getZeroExtendExpr(V, Ty, Depth);
4714}
4715
4717 unsigned Depth) {
4718 Type *SrcTy = V->getType();
4719 assert(SrcTy->isIntOrPtrTy() && Ty->isIntOrPtrTy() &&
4720 "Cannot truncate or zero extend with non-integer arguments!");
4721 if (getTypeSizeInBits(SrcTy) == getTypeSizeInBits(Ty))
4722 return V; // No conversion
4723 if (getTypeSizeInBits(SrcTy) > getTypeSizeInBits(Ty))
4724 return getTruncateExpr(V, Ty, Depth);
4725 return getSignExtendExpr(V, Ty, Depth);
4726}
4727
4729 Type *SrcTy = V->getType();
4730 assert(SrcTy->isIntOrPtrTy() && Ty->isIntOrPtrTy() &&
4731 "Cannot noop or zero extend with non-integer arguments!");
4733 "getNoopOrZeroExtend cannot truncate!");
4734 if (getTypeSizeInBits(SrcTy) == getTypeSizeInBits(Ty))
4735 return V; // No conversion
4736 return getZeroExtendExpr(V, Ty);
4737}
4738
4740 Type *SrcTy = V->getType();
4741 assert(SrcTy->isIntOrPtrTy() && Ty->isIntOrPtrTy() &&
4742 "Cannot noop or sign extend with non-integer arguments!");
4744 "getNoopOrSignExtend cannot truncate!");
4745 if (getTypeSizeInBits(SrcTy) == getTypeSizeInBits(Ty))
4746 return V; // No conversion
4747 return getSignExtendExpr(V, Ty);
4748}
4749
4751 Type *SrcTy = V->getType();
4752 assert(SrcTy->isIntOrPtrTy() && Ty->isIntOrPtrTy() &&
4753 "Cannot noop or any extend with non-integer arguments!");
4755 "getNoopOrAnyExtend cannot truncate!");
4756 if (getTypeSizeInBits(SrcTy) == getTypeSizeInBits(Ty))
4757 return V; // No conversion
4758 return getAnyExtendExpr(V, Ty);
4759}
4760
4762 Type *SrcTy = V->getType();
4763 assert(SrcTy->isIntOrPtrTy() && Ty->isIntOrPtrTy() &&
4764 "Cannot truncate or noop with non-integer arguments!");
4766 "getTruncateOrNoop cannot extend!");
4767 if (getTypeSizeInBits(SrcTy) == getTypeSizeInBits(Ty))
4768 return V; // No conversion
4769 return getTruncateExpr(V, Ty);
4770}
4771
4773 const SCEV *RHS) {
4774 const SCEV *PromotedLHS = LHS;
4775 const SCEV *PromotedRHS = RHS;
4776
4777 if (getTypeSizeInBits(LHS->getType()) > getTypeSizeInBits(RHS->getType()))
4778 PromotedRHS = getZeroExtendExpr(RHS, LHS->getType());
4779 else
4780 PromotedLHS = getNoopOrZeroExtend(LHS, RHS->getType());
4781
4782 return getUMaxExpr(PromotedLHS, PromotedRHS);
4783}
4784
4786 const SCEV *RHS,
4787 bool Sequential) {
4788 SmallVector<SCEVUse, 2> Ops = {LHS, RHS};
4789 return getUMinFromMismatchedTypes(Ops, Sequential);
4790}
4791
4792const SCEV *
4794 bool Sequential) {
4795 assert(!Ops.empty() && "At least one operand must be!");
4796 // Trivial case.
4797 if (Ops.size() == 1)
4798 return Ops[0];
4799
4800 // Find the max type first.
4801 Type *MaxType = nullptr;
4802 for (SCEVUse S : Ops)
4803 if (MaxType)
4804 MaxType = getWiderType(MaxType, S->getType());
4805 else
4806 MaxType = S->getType();
4807 assert(MaxType && "Failed to find maximum type!");
4808
4809 // Extend all ops to max type.
4810 SmallVector<SCEVUse, 2> PromotedOps;
4811 for (SCEVUse S : Ops)
4812 PromotedOps.push_back(getNoopOrZeroExtend(S, MaxType));
4813
4814 // Generate umin.
4815 return getUMinExpr(PromotedOps, Sequential);
4816}
4817
4819 // A pointer operand may evaluate to a nonpointer expression, such as null.
4820 if (!V->getType()->isPointerTy())
4821 return V;
4822
4823 while (true) {
4824 if (auto *AddRec = dyn_cast<SCEVAddRecExpr>(V)) {
4825 V = AddRec->getStart();
4826 } else if (auto *Add = dyn_cast<SCEVAddExpr>(V)) {
4827 const SCEV *PtrOp = nullptr;
4828 for (const SCEV *AddOp : Add->operands()) {
4829 if (AddOp->getType()->isPointerTy()) {
4830 assert(!PtrOp && "Cannot have multiple pointer ops");
4831 PtrOp = AddOp;
4832 }
4833 }
4834 assert(PtrOp && "Must have pointer op");
4835 V = PtrOp;
4836 } else // Not something we can look further into.
4837 return V;
4838 }
4839}
4840
4841/// Push users of the given Instruction onto the given Worklist.
4845 // Push the def-use children onto the Worklist stack.
4846 for (User *U : I->users()) {
4847 auto *UserInsn = cast<Instruction>(U);
4848 if (Visited.insert(UserInsn).second)
4849 Worklist.push_back(UserInsn);
4850 }
4851}
4852
4853namespace {
4854
4855/// Takes SCEV S and Loop L. For each AddRec sub-expression, use its start
4856/// expression in case its Loop is L. If it is not L then
4857/// if IgnoreOtherLoops is true then use AddRec itself
4858/// otherwise rewrite cannot be done.
4859/// If SCEV contains non-invariant unknown SCEV rewrite cannot be done.
4860class SCEVInitRewriter : public SCEVRewriteVisitor<SCEVInitRewriter> {
4861public:
4862 static const SCEV *rewrite(const SCEV *S, const Loop *L, ScalarEvolution &SE,
4863 bool IgnoreOtherLoops = true) {
4864 SCEVInitRewriter Rewriter(L, SE);
4865 const SCEV *Result = Rewriter.visit(S);
4866 if (Rewriter.hasSeenLoopVariantSCEVUnknown())
4867 return SE.getCouldNotCompute();
4868 return Rewriter.hasSeenOtherLoops() && !IgnoreOtherLoops
4869 ? SE.getCouldNotCompute()
4870 : Result;
4871 }
4872
4873 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
4874 if (!SE.isLoopInvariant(Expr, L))
4875 SeenLoopVariantSCEVUnknown = true;
4876 return Expr;
4877 }
4878
4879 const SCEV *visitAddRecExpr(const SCEVAddRecExpr *Expr) {
4880 // Only re-write AddRecExprs for this loop.
4881 if (Expr->getLoop() == L)
4882 return Expr->getStart();
4883 SeenOtherLoops = true;
4884 return Expr;
4885 }
4886
4887 bool hasSeenLoopVariantSCEVUnknown() { return SeenLoopVariantSCEVUnknown; }
4888
4889 bool hasSeenOtherLoops() { return SeenOtherLoops; }
4890
4891private:
4892 explicit SCEVInitRewriter(const Loop *L, ScalarEvolution &SE)
4893 : SCEVRewriteVisitor(SE), L(L) {}
4894
4895 const Loop *L;
4896 bool SeenLoopVariantSCEVUnknown = false;
4897 bool SeenOtherLoops = false;
4898};
4899
4900/// Takes SCEV S and Loop L. For each AddRec sub-expression, use its post
4901/// increment expression in case its Loop is L. If it is not L then
4902/// use AddRec itself.
4903/// If SCEV contains non-invariant unknown SCEV rewrite cannot be done.
4904class SCEVPostIncRewriter : public SCEVRewriteVisitor<SCEVPostIncRewriter> {
4905public:
4906 static const SCEV *rewrite(const SCEV *S, const Loop *L, ScalarEvolution &SE) {
4907 SCEVPostIncRewriter Rewriter(L, SE);
4908 const SCEV *Result = Rewriter.visit(S);
4909 return Rewriter.hasSeenLoopVariantSCEVUnknown()
4910 ? SE.getCouldNotCompute()
4911 : Result;
4912 }
4913
4914 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
4915 if (!SE.isLoopInvariant(Expr, L))
4916 SeenLoopVariantSCEVUnknown = true;
4917 return Expr;
4918 }
4919
4920 const SCEV *visitAddRecExpr(const SCEVAddRecExpr *Expr) {
4921 // Only re-write AddRecExprs for this loop.
4922 if (Expr->getLoop() == L)
4923 return Expr->getPostIncExpr(SE);
4924 SeenOtherLoops = true;
4925 return Expr;
4926 }
4927
4928 bool hasSeenLoopVariantSCEVUnknown() { return SeenLoopVariantSCEVUnknown; }
4929
4930 bool hasSeenOtherLoops() { return SeenOtherLoops; }
4931
4932private:
4933 explicit SCEVPostIncRewriter(const Loop *L, ScalarEvolution &SE)
4934 : SCEVRewriteVisitor(SE), L(L) {}
4935
4936 const Loop *L;
4937 bool SeenLoopVariantSCEVUnknown = false;
4938 bool SeenOtherLoops = false;
4939};
4940
4941/// This class evaluates the compare condition by matching it against the
4942/// condition of loop latch. If there is a match we assume a true value
4943/// for the condition while building SCEV nodes.
4944class SCEVBackedgeConditionFolder
4945 : public SCEVRewriteVisitor<SCEVBackedgeConditionFolder> {
4946public:
4947 static const SCEV *rewrite(const SCEV *S, const Loop *L,
4948 ScalarEvolution &SE) {
4949 bool IsPosBECond = false;
4950 Value *BECond = nullptr;
4951 if (BasicBlock *Latch = L->getLoopLatch()) {
4952 if (CondBrInst *BI = dyn_cast<CondBrInst>(Latch->getTerminator())) {
4953 assert(BI->getSuccessor(0) != BI->getSuccessor(1) &&
4954 "Both outgoing branches should not target same header!");
4955 BECond = BI->getCondition();
4956 IsPosBECond = BI->getSuccessor(0) == L->getHeader();
4957 } else {
4958 return S;
4959 }
4960 }
4961 SCEVBackedgeConditionFolder Rewriter(L, BECond, IsPosBECond, SE);
4962 return Rewriter.visit(S);
4963 }
4964
4965 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
4966 const SCEV *Result = Expr;
4967 bool InvariantF = SE.isLoopInvariant(Expr, L);
4968
4969 if (!InvariantF) {
4971 switch (I->getOpcode()) {
4972 case Instruction::Select: {
4973 SelectInst *SI = cast<SelectInst>(I);
4974 std::optional<const SCEV *> Res =
4975 compareWithBackedgeCondition(SI->getCondition());
4976 if (Res) {
4977 bool IsOne = cast<SCEVConstant>(*Res)->getValue()->isOne();
4978 Result = SE.getSCEV(IsOne ? SI->getTrueValue() : SI->getFalseValue());
4979 }
4980 break;
4981 }
4982 default: {
4983 std::optional<const SCEV *> Res = compareWithBackedgeCondition(I);
4984 if (Res)
4985 Result = *Res;
4986 break;
4987 }
4988 }
4989 }
4990 return Result;
4991 }
4992
4993private:
4994 explicit SCEVBackedgeConditionFolder(const Loop *L, Value *BECond,
4995 bool IsPosBECond, ScalarEvolution &SE)
4996 : SCEVRewriteVisitor(SE), L(L), BackedgeCond(BECond),
4997 IsPositiveBECond(IsPosBECond) {}
4998
4999 std::optional<const SCEV *> compareWithBackedgeCondition(Value *IC);
5000
5001 const Loop *L;
5002 /// Loop back condition.
5003 Value *BackedgeCond = nullptr;
5004 /// Set to true if loop back is on positive branch condition.
5005 bool IsPositiveBECond;
5006};
5007
5008std::optional<const SCEV *>
5009SCEVBackedgeConditionFolder::compareWithBackedgeCondition(Value *IC) {
5010
5011 // If value matches the backedge condition for loop latch,
5012 // then return a constant evolution node based on loopback
5013 // branch taken.
5014 if (BackedgeCond == IC)
5015 return IsPositiveBECond ? SE.getOne(Type::getInt1Ty(SE.getContext()))
5017 return std::nullopt;
5018}
5019
5020class SCEVShiftRewriter : public SCEVRewriteVisitor<SCEVShiftRewriter> {
5021public:
5022 static const SCEV *rewrite(const SCEV *S, const Loop *L,
5023 ScalarEvolution &SE) {
5024 SCEVShiftRewriter Rewriter(L, SE);
5025 const SCEV *Result = Rewriter.visit(S);
5026 return Rewriter.isValid() ? Result : SE.getCouldNotCompute();
5027 }
5028
5029 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
5030 // Only allow AddRecExprs for this loop.
5031 if (!SE.isLoopInvariant(Expr, L))
5032 Valid = false;
5033 return Expr;
5034 }
5035
5036 const SCEV *visitAddRecExpr(const SCEVAddRecExpr *Expr) {
5037 if (Expr->getLoop() == L && Expr->isAffine())
5038 return SE.getMinusSCEV(Expr, Expr->getStepRecurrence(SE));
5039 Valid = false;
5040 return Expr;
5041 }
5042
5043 bool isValid() { return Valid; }
5044
5045private:
5046 explicit SCEVShiftRewriter(const Loop *L, ScalarEvolution &SE)
5047 : SCEVRewriteVisitor(SE), L(L) {}
5048
5049 const Loop *L;
5050 bool Valid = true;
5051};
5052
5053} // end anonymous namespace
5054
5055void ScalarEvolution::inferNoWrapViaConstantRanges(const SCEVAddRecExpr *AR) {
5056 if (!AR->isAffine())
5057 return;
5058
5059 // Force computation of ranges, which will also perform range-based flag
5060 // inference.
5061 if (!AR->hasNoSignedWrap())
5062 (void)getSignedRange(AR);
5063
5064 if (!AR->hasNoUnsignedWrap())
5065 (void)getUnsignedRange(AR);
5066
5067 if (!AR->hasNoSelfWrap()) {
5068 const SCEV *BECount = getConstantMaxBackedgeTakenCount(AR->getLoop());
5069 if (const SCEVConstant *BECountMax = dyn_cast<SCEVConstant>(BECount)) {
5070 ConstantRange StepCR = getSignedRange(AR->getStepRecurrence(*this));
5071 const APInt &BECountAP = BECountMax->getAPInt();
5072 unsigned NoOverflowBitWidth =
5073 BECountAP.getActiveBits() + StepCR.getMinSignedBits();
5074 if (NoOverflowBitWidth <= getTypeSizeInBits(AR->getType()))
5075 const_cast<SCEVAddRecExpr *>(AR)->setNoWrapFlags(SCEV::FlagNW);
5076 }
5077 }
5078}
5079
5081ScalarEvolution::proveNoSignedWrapViaInduction(const SCEVAddRecExpr *AR) {
5083
5084 if (AR->hasNoSignedWrap())
5085 return Result;
5086
5087 if (!AR->isAffine())
5088 return Result;
5089
5090 // This function can be expensive, only try to prove NSW once per AddRec.
5091 if (!SignedWrapViaInductionTried.insert(AR).second)
5092 return Result;
5093
5094 const SCEV *Step = AR->getStepRecurrence(*this);
5095 const Loop *L = AR->getLoop();
5096
5097 // Check whether the backedge-taken count is SCEVCouldNotCompute.
5098 // Note that this serves two purposes: It filters out loops that are
5099 // simply not analyzable, and it covers the case where this code is
5100 // being called from within backedge-taken count analysis, such that
5101 // attempting to ask for the backedge-taken count would likely result
5102 // in infinite recursion. In the later case, the analysis code will
5103 // cope with a conservative value, and it will take care to purge
5104 // that value once it has finished.
5105 const SCEV *MaxBECount = getConstantMaxBackedgeTakenCount(L);
5106
5107 // Normally, in the cases we can prove no-overflow via a
5108 // backedge guarding condition, we can also compute a backedge
5109 // taken count for the loop. The exceptions are assumptions and
5110 // guards present in the loop -- SCEV is not great at exploiting
5111 // these to compute max backedge taken counts, but can still use
5112 // these to prove lack of overflow. Use this fact to avoid
5113 // doing extra work that may not pay off.
5114
5115 if (isa<SCEVCouldNotCompute>(MaxBECount) && !HasGuards &&
5116 AC.assumptions().empty())
5117 return Result;
5118
5119 // If the backedge is guarded by a comparison with the pre-inc value the
5120 // addrec is safe. Also, if the entry is guarded by a comparison with the
5121 // start value and the backedge is guarded by a comparison with the post-inc
5122 // value, the addrec is safe.
5124 const SCEV *OverflowLimit =
5125 getSignedOverflowLimitForStep(Step, &Pred, this);
5126 if (OverflowLimit &&
5127 (isLoopBackedgeGuardedByCond(L, Pred, AR, OverflowLimit) ||
5128 isKnownOnEveryIteration(Pred, AR, OverflowLimit))) {
5129 Result = setFlags(Result, SCEV::FlagNSW);
5130 }
5131 return Result;
5132}
5134ScalarEvolution::proveNoUnsignedWrapViaInduction(const SCEVAddRecExpr *AR) {
5136
5137 if (AR->hasNoUnsignedWrap())
5138 return Result;
5139
5140 if (!AR->isAffine())
5141 return Result;
5142
5143 // This function can be expensive, only try to prove NUW once per AddRec.
5144 if (!UnsignedWrapViaInductionTried.insert(AR).second)
5145 return Result;
5146
5147 const SCEV *Step = AR->getStepRecurrence(*this);
5148 const Loop *L = AR->getLoop();
5149
5150 // Check whether the backedge-taken count is SCEVCouldNotCompute.
5151 // Note that this serves two purposes: It filters out loops that are
5152 // simply not analyzable, and it covers the case where this code is
5153 // being called from within backedge-taken count analysis, such that
5154 // attempting to ask for the backedge-taken count would likely result
5155 // in infinite recursion. In the later case, the analysis code will
5156 // cope with a conservative value, and it will take care to purge
5157 // that value once it has finished.
5158 const SCEV *MaxBECount = getConstantMaxBackedgeTakenCount(L);
5159
5160 // Normally, in the cases we can prove no-overflow via a
5161 // backedge guarding condition, we can also compute a backedge
5162 // taken count for the loop. The exceptions are assumptions and
5163 // guards present in the loop -- SCEV is not great at exploiting
5164 // these to compute max backedge taken counts, but can still use
5165 // these to prove lack of overflow. Use this fact to avoid
5166 // doing extra work that may not pay off.
5167
5168 if (isa<SCEVCouldNotCompute>(MaxBECount) && !HasGuards &&
5169 AC.assumptions().empty())
5170 return Result;
5171
5172 // If the backedge is guarded by a comparison with the pre-inc value the
5173 // addrec is safe. Also, if the entry is guarded by a comparison with the
5174 // start value and the backedge is guarded by a comparison with the post-inc
5175 // value, the addrec is safe.
5176 if (isKnownPositive(Step)) {
5178 const SCEV *OverflowLimit =
5179 getUnsignedOverflowLimitForStep(Step, &Pred, this);
5180 if (isLoopBackedgeGuardedByCond(L, Pred, AR, OverflowLimit) ||
5181 isKnownOnEveryIteration(Pred, AR, OverflowLimit))
5182 Result = setFlags(Result, SCEV::FlagNUW);
5183 }
5184 return Result;
5185}
5186
5187namespace {
5188
5189/// Represents an abstract binary operation. This may exist as a
5190/// normal instruction or constant expression, or may have been
5191/// derived from an expression tree.
5192struct BinaryOp {
5193 unsigned Opcode;
5194 Value *LHS;
5195 Value *RHS;
5196 bool IsNSW = false;
5197 bool IsNUW = false;
5198
5199 /// Op is set if this BinaryOp corresponds to a concrete LLVM instruction or
5200 /// constant expression.
5201 Operator *Op = nullptr;
5202
5203 explicit BinaryOp(Operator *Op)
5204 : Opcode(Op->getOpcode()), LHS(Op->getOperand(0)), RHS(Op->getOperand(1)),
5205 Op(Op) {
5206 if (auto *OBO = dyn_cast<OverflowingBinaryOperator>(Op)) {
5207 IsNSW = OBO->hasNoSignedWrap();
5208 IsNUW = OBO->hasNoUnsignedWrap();
5209 }
5210 }
5211
5212 explicit BinaryOp(unsigned Opcode, Value *LHS, Value *RHS, bool IsNSW = false,
5213 bool IsNUW = false)
5214 : Opcode(Opcode), LHS(LHS), RHS(RHS), IsNSW(IsNSW), IsNUW(IsNUW) {}
5215};
5216
5217} // end anonymous namespace
5218
5219/// Try to map \p V into a BinaryOp, and return \c std::nullopt on failure.
5220static std::optional<BinaryOp> MatchBinaryOp(Value *V, const DataLayout &DL,
5221 AssumptionCache &AC,
5222 const DominatorTree &DT,
5223 const Instruction *CtxI) {
5224 auto *Op = dyn_cast<Operator>(V);
5225 if (!Op)
5226 return std::nullopt;
5227
5228 // Implementation detail: all the cleverness here should happen without
5229 // creating new SCEV expressions -- our caller knowns tricks to avoid creating
5230 // SCEV expressions when possible, and we should not break that.
5231
5232 switch (Op->getOpcode()) {
5233 case Instruction::Add:
5234 case Instruction::Sub:
5235 case Instruction::Mul:
5236 case Instruction::UDiv:
5237 case Instruction::URem:
5238 case Instruction::And:
5239 case Instruction::AShr:
5240 case Instruction::Shl:
5241 return BinaryOp(Op);
5242
5243 case Instruction::Or: {
5244 // Convert or disjoint into add nuw nsw.
5245 if (cast<PossiblyDisjointInst>(Op)->isDisjoint()) {
5246 BinaryOp BinOp(Instruction::Add, Op->getOperand(0), Op->getOperand(1),
5247 /*IsNSW=*/true, /*IsNUW=*/true);
5248 // Keep the reference to the original instruction so that we can later
5249 // check whether it can produce poison value or not.
5250 BinOp.Op = Op;
5251 return BinOp;
5252 }
5253 return BinaryOp(Op);
5254 }
5255
5256 case Instruction::Xor:
5257 if (auto *RHSC = dyn_cast<ConstantInt>(Op->getOperand(1)))
5258 // If the RHS of the xor is a signmask, then this is just an add.
5259 // Instcombine turns add of signmask into xor as a strength reduction step.
5260 if (RHSC->getValue().isSignMask())
5261 return BinaryOp(Instruction::Add, Op->getOperand(0), Op->getOperand(1));
5262 // Binary `xor` is a bit-wise `add`.
5263 if (V->getType()->isIntegerTy(1))
5264 return BinaryOp(Instruction::Add, Op->getOperand(0), Op->getOperand(1));
5265 return BinaryOp(Op);
5266
5267 case Instruction::LShr:
5268 // Turn logical shift right of a constant into a unsigned divide.
5269 if (ConstantInt *SA = dyn_cast<ConstantInt>(Op->getOperand(1))) {
5270 uint32_t BitWidth = cast<IntegerType>(Op->getType())->getBitWidth();
5271
5272 // If the shift count is not less than the bitwidth, the result of
5273 // the shift is undefined. Don't try to analyze it, because the
5274 // resolution chosen here may differ from the resolution chosen in
5275 // other parts of the compiler.
5276 if (SA->getValue().ult(BitWidth)) {
5277 Constant *X =
5278 ConstantInt::get(SA->getContext(),
5279 APInt::getOneBitSet(BitWidth, SA->getZExtValue()));
5280 return BinaryOp(Instruction::UDiv, Op->getOperand(0), X);
5281 }
5282 }
5283 return BinaryOp(Op);
5284
5285 case Instruction::ExtractValue: {
5286 auto *EVI = cast<ExtractValueInst>(Op);
5287 if (EVI->getNumIndices() != 1 || EVI->getIndices()[0] != 0)
5288 break;
5289
5290 auto *WO = dyn_cast<WithOverflowInst>(EVI->getAggregateOperand());
5291 if (!WO)
5292 break;
5293
5294 Instruction::BinaryOps BinOp = WO->getBinaryOp();
5295 bool Signed = WO->isSigned();
5296 // TODO: Should add nuw/nsw flags for mul as well.
5297 if (BinOp == Instruction::Mul || !isOverflowIntrinsicNoWrap(WO, DT))
5298 return BinaryOp(BinOp, WO->getLHS(), WO->getRHS());
5299
5300 // Now that we know that all uses of the arithmetic-result component of
5301 // CI are guarded by the overflow check, we can go ahead and pretend
5302 // that the arithmetic is non-overflowing.
5303 return BinaryOp(BinOp, WO->getLHS(), WO->getRHS(),
5304 /* IsNSW = */ Signed, /* IsNUW = */ !Signed);
5305 }
5306
5307 default:
5308 break;
5309 }
5310
5311 // Recognise intrinsic loop.decrement.reg, and as this has exactly the same
5312 // semantics as a Sub, return a binary sub expression.
5313 if (auto *II = dyn_cast<IntrinsicInst>(V))
5314 if (II->getIntrinsicID() == Intrinsic::loop_decrement_reg)
5315 return BinaryOp(Instruction::Sub, II->getOperand(0), II->getOperand(1));
5316
5317 return std::nullopt;
5318}
5319
5320/// Helper function to createAddRecFromPHIWithCasts. We have a phi
5321/// node whose symbolic (unknown) SCEV is \p SymbolicPHI, which is updated via
5322/// the loop backedge by a SCEVAddExpr, possibly also with a few casts on the
5323/// way. This function checks if \p Op, an operand of this SCEVAddExpr,
5324/// follows one of the following patterns:
5325/// Op == (SExt ix (Trunc iy (%SymbolicPHI) to ix) to iy)
5326/// Op == (ZExt ix (Trunc iy (%SymbolicPHI) to ix) to iy)
5327/// If the SCEV expression of \p Op conforms with one of the expected patterns
5328/// we return the type of the truncation operation, and indicate whether the
5329/// truncated type should be treated as signed/unsigned by setting
5330/// \p Signed to true/false, respectively.
5331static Type *isSimpleCastedPHI(const SCEV *Op, const SCEVUnknown *SymbolicPHI,
5332 bool &Signed, ScalarEvolution &SE) {
5333 // The case where Op == SymbolicPHI (that is, with no type conversions on
5334 // the way) is handled by the regular add recurrence creating logic and
5335 // would have already been triggered in createAddRecForPHI. Reaching it here
5336 // means that createAddRecFromPHI had failed for this PHI before (e.g.,
5337 // because one of the other operands of the SCEVAddExpr updating this PHI is
5338 // not invariant).
5339 //
5340 // Here we look for the case where Op = (ext(trunc(SymbolicPHI))), and in
5341 // this case predicates that allow us to prove that Op == SymbolicPHI will
5342 // be added.
5343 if (Op == SymbolicPHI)
5344 return nullptr;
5345
5346 unsigned SourceBits = SE.getTypeSizeInBits(SymbolicPHI->getType());
5347 unsigned NewBits = SE.getTypeSizeInBits(Op->getType());
5348 if (SourceBits != NewBits)
5349 return nullptr;
5350
5351 if (match(Op, m_scev_SExt(m_scev_Trunc(m_scev_Specific(SymbolicPHI))))) {
5352 Signed = true;
5353 return cast<SCEVCastExpr>(Op)->getOperand()->getType();
5354 }
5355 if (match(Op, m_scev_ZExt(m_scev_Trunc(m_scev_Specific(SymbolicPHI))))) {
5356 Signed = false;
5357 return cast<SCEVCastExpr>(Op)->getOperand()->getType();
5358 }
5359 return nullptr;
5360}
5361
5362static const Loop *isIntegerLoopHeaderPHI(const PHINode *PN, LoopInfo &LI) {
5363 if (!PN->getType()->isIntegerTy())
5364 return nullptr;
5365 const Loop *L = LI.getLoopFor(PN->getParent());
5366 if (!L || L->getHeader() != PN->getParent())
5367 return nullptr;
5368 return L;
5369}
5370
5371// Analyze \p SymbolicPHI, a SCEV expression of a phi node, and check if the
5372// computation that updates the phi follows the following pattern:
5373// (SExt/ZExt ix (Trunc iy (%SymbolicPHI) to ix) to iy) + InvariantAccum
5374// which correspond to a phi->trunc->sext/zext->add->phi update chain.
5375// If so, try to see if it can be rewritten as an AddRecExpr under some
5376// Predicates. If successful, return them as a pair. Also cache the results
5377// of the analysis.
5378//
5379// Example usage scenario:
5380// Say the Rewriter is called for the following SCEV:
5381// 8 * ((sext i32 (trunc i64 %X to i32) to i64) + %Step)
5382// where:
5383// %X = phi i64 (%Start, %BEValue)
5384// It will visitMul->visitAdd->visitSExt->visitTrunc->visitUnknown(%X),
5385// and call this function with %SymbolicPHI = %X.
5386//
5387// The analysis will find that the value coming around the backedge has
5388// the following SCEV:
5389// BEValue = ((sext i32 (trunc i64 %X to i32) to i64) + %Step)
5390// Upon concluding that this matches the desired pattern, the function
5391// will return the pair {NewAddRec, SmallPredsVec} where:
5392// NewAddRec = {%Start,+,%Step}
5393// SmallPredsVec = {P1, P2, P3} as follows:
5394// P1(WrapPred): AR: {trunc(%Start),+,(trunc %Step)}<nsw> Flags: <nssw>
5395// P2(EqualPred): %Start == (sext i32 (trunc i64 %Start to i32) to i64)
5396// P3(EqualPred): %Step == (sext i32 (trunc i64 %Step to i32) to i64)
5397// The returned pair means that SymbolicPHI can be rewritten into NewAddRec
5398// under the predicates {P1,P2,P3}.
5399// This predicated rewrite will be cached in PredicatedSCEVRewrites:
5400// PredicatedSCEVRewrites[{%X,L}] = {NewAddRec, {P1,P2,P3)}
5401//
5402// TODO's:
5403//
5404// 1) Extend the Induction descriptor to also support inductions that involve
5405// casts: When needed (namely, when we are called in the context of the
5406// vectorizer induction analysis), a Set of cast instructions will be
5407// populated by this method, and provided back to isInductionPHI. This is
5408// needed to allow the vectorizer to properly record them to be ignored by
5409// the cost model and to avoid vectorizing them (otherwise these casts,
5410// which are redundant under the runtime overflow checks, will be
5411// vectorized, which can be costly).
5412//
5413// 2) Support additional induction/PHISCEV patterns: We also want to support
5414// inductions where the sext-trunc / zext-trunc operations (partly) occur
5415// after the induction update operation (the induction increment):
5416//
5417// (Trunc iy (SExt/ZExt ix (%SymbolicPHI + InvariantAccum) to iy) to ix)
5418// which correspond to a phi->add->trunc->sext/zext->phi update chain.
5419//
5420// (Trunc iy ((SExt/ZExt ix (%SymbolicPhi) to iy) + InvariantAccum) to ix)
5421// which correspond to a phi->trunc->add->sext/zext->phi update chain.
5422//
5423// 3) Outline common code with createAddRecFromPHI to avoid duplication.
5424std::optional<std::pair<const SCEV *, SmallVector<const SCEVPredicate *, 3>>>
5425ScalarEvolution::createAddRecFromPHIWithCastsImpl(const SCEVUnknown *SymbolicPHI) {
5427
5428 // *** Part1: Analyze if we have a phi-with-cast pattern for which we can
5429 // return an AddRec expression under some predicate.
5430
5431 auto *PN = cast<PHINode>(SymbolicPHI->getValue());
5432 const Loop *L = isIntegerLoopHeaderPHI(PN, LI);
5433 assert(L && "Expecting an integer loop header phi");
5434
5435 // The loop may have multiple entrances or multiple exits; we can analyze
5436 // this phi as an addrec if it has a unique entry value and a unique
5437 // backedge value.
5438 Value *BEValueV = nullptr, *StartValueV = nullptr;
5439 for (unsigned i = 0, e = PN->getNumIncomingValues(); i != e; ++i) {
5440 Value *V = PN->getIncomingValue(i);
5441 if (L->contains(PN->getIncomingBlock(i))) {
5442 if (!BEValueV) {
5443 BEValueV = V;
5444 } else if (BEValueV != V) {
5445 BEValueV = nullptr;
5446 break;
5447 }
5448 } else if (!StartValueV) {
5449 StartValueV = V;
5450 } else if (StartValueV != V) {
5451 StartValueV = nullptr;
5452 break;
5453 }
5454 }
5455 if (!BEValueV || !StartValueV)
5456 return std::nullopt;
5457
5458 const SCEV *BEValue = getSCEV(BEValueV);
5459
5460 // If the value coming around the backedge is an add with the symbolic
5461 // value we just inserted, possibly with casts that we can ignore under
5462 // an appropriate runtime guard, then we found a simple induction variable!
5463 const auto *Add = dyn_cast<SCEVAddExpr>(BEValue);
5464 if (!Add)
5465 return std::nullopt;
5466
5467 // If there is a single occurrence of the symbolic value, possibly
5468 // casted, replace it with a recurrence.
5469 unsigned FoundIndex = Add->getNumOperands();
5470 Type *TruncTy = nullptr;
5471 bool Signed;
5472 for (unsigned i = 0, e = Add->getNumOperands(); i != e; ++i)
5473 if ((TruncTy =
5474 isSimpleCastedPHI(Add->getOperand(i), SymbolicPHI, Signed, *this)))
5475 if (FoundIndex == e) {
5476 FoundIndex = i;
5477 break;
5478 }
5479
5480 if (FoundIndex == Add->getNumOperands())
5481 return std::nullopt;
5482
5483 // Create an add with everything but the specified operand.
5485 for (unsigned i = 0, e = Add->getNumOperands(); i != e; ++i)
5486 if (i != FoundIndex)
5487 Ops.push_back(Add->getOperand(i));
5488 const SCEV *Accum = getAddExpr(Ops);
5489
5490 // The runtime checks will not be valid if the step amount is
5491 // varying inside the loop.
5492 if (!isLoopInvariant(Accum, L))
5493 return std::nullopt;
5494
5495 // *** Part2: Create the predicates
5496
5497 // Analysis was successful: we have a phi-with-cast pattern for which we
5498 // can return an AddRec expression under the following predicates:
5499 //
5500 // P1: A Wrap predicate that guarantees that Trunc(Start) + i*Trunc(Accum)
5501 // fits within the truncated type (does not overflow) for i = 0 to n-1.
5502 // P2: An Equal predicate that guarantees that
5503 // Start = (Ext ix (Trunc iy (Start) to ix) to iy)
5504 // P3: An Equal predicate that guarantees that
5505 // Accum = (Ext ix (Trunc iy (Accum) to ix) to iy)
5506 //
5507 // As we next prove, the above predicates guarantee that:
5508 // Start + i*Accum = (Ext ix (Trunc iy ( Start + i*Accum ) to ix) to iy)
5509 //
5510 //
5511 // More formally, we want to prove that:
5512 // Expr(i+1) = Start + (i+1) * Accum
5513 // = (Ext ix (Trunc iy (Expr(i)) to ix) to iy) + Accum
5514 //
5515 // Given that:
5516 // 1) Expr(0) = Start
5517 // 2) Expr(1) = Start + Accum
5518 // = (Ext ix (Trunc iy (Start) to ix) to iy) + Accum :: from P2
5519 // 3) Induction hypothesis (step i):
5520 // Expr(i) = (Ext ix (Trunc iy (Expr(i-1)) to ix) to iy) + Accum
5521 //
5522 // Proof:
5523 // Expr(i+1) =
5524 // = Start + (i+1)*Accum
5525 // = (Start + i*Accum) + Accum
5526 // = Expr(i) + Accum
5527 // = (Ext ix (Trunc iy (Expr(i-1)) to ix) to iy) + Accum + Accum
5528 // :: from step i
5529 //
5530 // = (Ext ix (Trunc iy (Start + (i-1)*Accum) to ix) to iy) + Accum + Accum
5531 //
5532 // = (Ext ix (Trunc iy (Start + (i-1)*Accum) to ix) to iy)
5533 // + (Ext ix (Trunc iy (Accum) to ix) to iy)
5534 // + Accum :: from P3
5535 //
5536 // = (Ext ix (Trunc iy ((Start + (i-1)*Accum) + Accum) to ix) to iy)
5537 // + Accum :: from P1: Ext(x)+Ext(y)=>Ext(x+y)
5538 //
5539 // = (Ext ix (Trunc iy (Start + i*Accum) to ix) to iy) + Accum
5540 // = (Ext ix (Trunc iy (Expr(i)) to ix) to iy) + Accum
5541 //
5542 // By induction, the same applies to all iterations 1<=i<n:
5543 //
5544
5545 // Create a truncated addrec for which we will add a no overflow check (P1).
5546 const SCEV *StartVal = getSCEV(StartValueV);
5547 const SCEV *PHISCEV =
5548 getAddRecExpr(getTruncateExpr(StartVal, TruncTy),
5549 getTruncateExpr(Accum, TruncTy), L, SCEV::FlagNone);
5550
5551 // PHISCEV can be either a SCEVConstant or a SCEVAddRecExpr.
5552 // ex: If truncated Accum is 0 and StartVal is a constant, then PHISCEV
5553 // will be constant.
5554 //
5555 // If PHISCEV is a constant, then P1 degenerates into P2 or P3, so we don't
5556 // add P1.
5557 if (const auto *AR = dyn_cast<SCEVAddRecExpr>(PHISCEV)) {
5561 const SCEVPredicate *AddRecPred = getWrapPredicate(AR, AddedFlags);
5562 Predicates.push_back(AddRecPred);
5563 }
5564
5565 // Create the Equal Predicates P2,P3:
5566
5567 // It is possible that the predicates P2 and/or P3 are computable at
5568 // compile time due to StartVal and/or Accum being constants.
5569 // If either one is, then we can check that now and escape if either P2
5570 // or P3 is false.
5571
5572 // Construct the extended SCEV: (Ext ix (Trunc iy (Expr) to ix) to iy)
5573 // for each of StartVal and Accum
5574 auto getExtendedExpr = [&](const SCEV *Expr,
5575 bool CreateSignExtend) -> const SCEV * {
5576 assert(isLoopInvariant(Expr, L) && "Expr is expected to be invariant");
5577 const SCEV *TruncatedExpr = getTruncateExpr(Expr, TruncTy);
5578 const SCEV *ExtendedExpr =
5579 CreateSignExtend ? getSignExtendExpr(TruncatedExpr, Expr->getType())
5580 : getZeroExtendExpr(TruncatedExpr, Expr->getType());
5581 return ExtendedExpr;
5582 };
5583
5584 // Given:
5585 // ExtendedExpr = (Ext ix (Trunc iy (Expr) to ix) to iy
5586 // = getExtendedExpr(Expr)
5587 // Determine whether the predicate P: Expr == ExtendedExpr
5588 // is known to be false at compile time
5589 auto PredIsKnownFalse = [&](const SCEV *Expr,
5590 const SCEV *ExtendedExpr) -> bool {
5591 return Expr != ExtendedExpr &&
5592 isKnownPredicate(ICmpInst::ICMP_NE, Expr, ExtendedExpr);
5593 };
5594
5595 const SCEV *StartExtended = getExtendedExpr(StartVal, Signed);
5596 if (PredIsKnownFalse(StartVal, StartExtended)) {
5597 LLVM_DEBUG(dbgs() << "P2 is compile-time false\n";);
5598 return std::nullopt;
5599 }
5600
5601 // The Step is always Signed (because the overflow checks are either
5602 // NSSW or NUSW)
5603 const SCEV *AccumExtended = getExtendedExpr(Accum, /*CreateSignExtend=*/true);
5604 if (PredIsKnownFalse(Accum, AccumExtended)) {
5605 LLVM_DEBUG(dbgs() << "P3 is compile-time false\n";);
5606 return std::nullopt;
5607 }
5608
5609 auto AppendPredicate = [&](const SCEV *Expr,
5610 const SCEV *ExtendedExpr) -> void {
5611 if (Expr != ExtendedExpr &&
5612 !isKnownPredicate(ICmpInst::ICMP_EQ, Expr, ExtendedExpr)) {
5613 const SCEVPredicate *Pred = getEqualPredicate(Expr, ExtendedExpr);
5614 LLVM_DEBUG(dbgs() << "Added Predicate: " << *Pred);
5615 Predicates.push_back(Pred);
5616 }
5617 };
5618
5619 AppendPredicate(StartVal, StartExtended);
5620 AppendPredicate(Accum, AccumExtended);
5621
5622 // *** Part3: Predicates are ready. Now go ahead and create the new addrec in
5623 // which the casts had been folded away. The caller can rewrite SymbolicPHI
5624 // into NewAR if it will also add the runtime overflow checks specified in
5625 // Predicates.
5626 const SCEV *NewAR = getAddRecExpr(StartVal, Accum, L, SCEV::FlagNone);
5627
5628 std::pair<const SCEV *, SmallVector<const SCEVPredicate *, 3>> PredRewrite =
5629 std::make_pair(NewAR, Predicates);
5630 // Remember the result of the analysis for this SCEV at this locayyytion.
5631 PredicatedSCEVRewrites[{SymbolicPHI, L}] = PredRewrite;
5632 return PredRewrite;
5633}
5634
5635std::optional<std::pair<const SCEV *, SmallVector<const SCEVPredicate *, 3>>>
5637 auto *PN = cast<PHINode>(SymbolicPHI->getValue());
5638 const Loop *L = isIntegerLoopHeaderPHI(PN, LI);
5639 if (!L)
5640 return std::nullopt;
5641
5642 // Check to see if we already analyzed this PHI.
5643 auto I = PredicatedSCEVRewrites.find({SymbolicPHI, L});
5644 if (I != PredicatedSCEVRewrites.end()) {
5645 std::pair<const SCEV *, SmallVector<const SCEVPredicate *, 3>> Rewrite =
5646 I->second;
5647 // Analysis was done before and failed to create an AddRec:
5648 if (Rewrite.first == SymbolicPHI)
5649 return std::nullopt;
5650 // Analysis was done before and succeeded to create an AddRec under
5651 // a predicate:
5652 assert(isa<SCEVAddRecExpr>(Rewrite.first) && "Expected an AddRec");
5653 assert(!(Rewrite.second).empty() && "Expected to find Predicates");
5654 return Rewrite;
5655 }
5656
5657 std::optional<std::pair<const SCEV *, SmallVector<const SCEVPredicate *, 3>>>
5658 Rewrite = createAddRecFromPHIWithCastsImpl(SymbolicPHI);
5659
5660 // Record in the cache that the analysis failed
5661 if (!Rewrite) {
5663 PredicatedSCEVRewrites[{SymbolicPHI, L}] = {SymbolicPHI, Predicates};
5664 return std::nullopt;
5665 }
5666
5667 return Rewrite;
5668}
5669
5670// FIXME: This utility is currently required because the Rewriter currently
5671// does not rewrite this expression:
5672// {0, +, (sext ix (trunc iy to ix) to iy)}
5673// into {0, +, %step},
5674// even when the following Equal predicate exists:
5675// "%step == (sext ix (trunc iy to ix) to iy)".
5677 const SCEVAddRecExpr *AR1, const SCEVAddRecExpr *AR2,
5678 ArrayRef<const SCEVPredicate *> NoWrapPreds) const {
5679 if (AR1 == AR2)
5680 return true;
5681
5682 SCEVUnionPredicate NoWrapUnionPred(NoWrapPreds, SE);
5683 SCEVUnionPredicate AllPreds = Preds->getUnionWith(&NoWrapUnionPred, SE);
5684 auto areExprsEqual = [&](const SCEV *Expr1, const SCEV *Expr2) -> bool {
5685 if (Expr1 != Expr2 &&
5686 !AllPreds.implies(SE.getEqualPredicate(Expr1, Expr2), SE) &&
5687 !AllPreds.implies(SE.getEqualPredicate(Expr2, Expr1), SE))
5688 return false;
5689 return true;
5690 };
5691
5692 if (!areExprsEqual(AR1->getStart(), AR2->getStart()) ||
5693 !areExprsEqual(AR1->getStepRecurrence(SE), AR2->getStepRecurrence(SE)))
5694 return false;
5695 return true;
5696}
5697
5699 ScalarEvolution &SE) {
5700 SCEVFlags Flags = SCEV::FlagNone;
5701 GEPNoWrapFlags NW = GEP->getNoWrapFlags();
5702 // If the increment has any nowrap flags, then we know the address
5703 // space cannot be wrapped around.
5704 if (NW != GEPNoWrapFlags::none())
5706 // If the GEP is nuw or nusw with non-negative offset, we know that
5707 // no unsigned wrap occurs. We cannot set the nsw flag as only the
5708 // offset is treated as signed, while the base is unsigned.
5709 if (NW.hasNoUnsignedWrap() ||
5710 (NW.hasNoUnsignedSignedWrap() && SE.isKnownNonNegative(Accum)))
5712
5713 return Flags;
5714}
5715
5716/// A helper function for createAddRecFromPHI to handle simple cases.
5717///
5718/// This function tries to find an AddRec expression for the simplest (yet most
5719/// common) cases: PN = PHI(Start, OP(Self, LoopInvariant)).
5720/// If it fails, createAddRecFromPHI will use a more general, but slow,
5721/// technique for finding the AddRec expression.
5722const SCEV *ScalarEvolution::createSimpleAffineAddRec(PHINode *PN,
5723 Value *BEValueV,
5724 Value *StartValueV) {
5725 const Loop *L = LI.getLoopFor(PN->getParent());
5726 assert(L && L->getHeader() == PN->getParent());
5727 assert(BEValueV && StartValueV);
5728
5729 const SCEV *Accum = nullptr;
5731 if (auto BO = MatchBinaryOp(BEValueV, getDataLayout(), AC, DT, PN)) {
5732 if (BO->Opcode != Instruction::Add)
5733 return nullptr;
5734
5735 if (BO->LHS == PN && L->isLoopInvariant(BO->RHS))
5736 Accum = getSCEV(BO->RHS);
5737 else if (BO->RHS == PN && L->isLoopInvariant(BO->LHS))
5738 Accum = getSCEV(BO->LHS);
5739
5740 if (!Accum)
5741 return nullptr;
5742
5743 if (BO->IsNUW)
5744 Flags = setFlags(Flags, SCEV::FlagNUW);
5745 if (BO->IsNSW)
5746 Flags = setFlags(Flags, SCEV::FlagNSW);
5747 } else {
5748 // Handle pointer induction variable: PN = PHI(Start, gep PN,
5749 // LoopInvariant).
5750 auto *GEP = dyn_cast<GEPOperator>(BEValueV);
5751 if (!GEP || GEP->getPointerOperand() != PN || GEP->getNumIndices() != 1)
5752 return nullptr;
5753 Value *Idx = *GEP->idx_begin();
5754 if (!L->isLoopInvariant(Idx))
5755 return nullptr;
5756
5757 Type *IntIdxTy = getEffectiveSCEVType(GEP->getType());
5758 Accum = getMulExpr(getTruncateOrSignExtend(getSCEV(Idx), IntIdxTy),
5759 getSizeOfExpr(IntIdxTy, GEP->getSourceElementType()));
5760 Flags = getNoWrapFlagsForGEP(GEP, Accum, *this);
5761 }
5762
5763 const SCEV *StartVal = getSCEV(StartValueV);
5764 const SCEV *PHISCEV = getAddRecExpr(StartVal, Accum, L, Flags);
5765 insertValueToMap(PN, PHISCEV);
5766
5767 if (auto *AR = dyn_cast<SCEVAddRecExpr>(PHISCEV))
5768 inferNoWrapViaConstantRanges(AR);
5769
5770 // We can add Flags to the post-inc expression only if we
5771 // know that it is *undefined behavior* for BEValueV to
5772 // overflow.
5773 if (auto *BEInst = dyn_cast<Instruction>(BEValueV)) {
5774 assert(isLoopInvariant(Accum, L) &&
5775 "Accum is defined outside L, but is not invariant?");
5776 if (isAddRecNeverPoison(BEInst, L))
5777 (void)getAddRecExpr(getAddExpr(StartVal, Accum), Accum, L, Flags);
5778 }
5779
5780 return PHISCEV;
5781}
5782
5783const SCEV *ScalarEvolution::createAddRecFromPHI(PHINode *PN) {
5784 const Loop *L = LI.getLoopFor(PN->getParent());
5785 if (!L || L->getHeader() != PN->getParent())
5786 return nullptr;
5787
5788 // The loop may have multiple entrances or multiple exits; we can analyze
5789 // this phi as an addrec if it has a unique entry value and a unique
5790 // backedge value.
5791 Value *BEValueV = nullptr, *StartValueV = nullptr;
5792 for (unsigned i = 0, e = PN->getNumIncomingValues(); i != e; ++i) {
5793 Value *V = PN->getIncomingValue(i);
5794 if (L->contains(PN->getIncomingBlock(i))) {
5795 if (!BEValueV) {
5796 BEValueV = V;
5797 } else if (BEValueV != V) {
5798 BEValueV = nullptr;
5799 break;
5800 }
5801 } else if (!StartValueV) {
5802 StartValueV = V;
5803 } else if (StartValueV != V) {
5804 StartValueV = nullptr;
5805 break;
5806 }
5807 }
5808 if (!BEValueV || !StartValueV)
5809 return nullptr;
5810
5811 assert(ValueExprMap.find_as(PN) == ValueExprMap.end() &&
5812 "PHI node already processed?");
5813
5814 // First, try to find AddRec expression without creating a fictituos symbolic
5815 // value for PN.
5816 if (auto *S = createSimpleAffineAddRec(PN, BEValueV, StartValueV))
5817 return S;
5818
5819 // Handle PHI node value symbolically.
5820 const SCEV *SymbolicName = getUnknown(PN);
5821 insertValueToMap(PN, SymbolicName);
5822
5823 // Using this symbolic name for the PHI, analyze the value coming around
5824 // the back-edge.
5825 const SCEV *BEValue = getSCEV(BEValueV);
5826
5827 // NOTE: If BEValue is loop invariant, we know that the PHI node just
5828 // has a special value for the first iteration of the loop.
5829
5830 // If the value coming around the backedge is an add with the symbolic
5831 // value we just inserted, then we found a simple induction variable!
5832 if (const SCEVAddExpr *Add = dyn_cast<SCEVAddExpr>(BEValue)) {
5833 // If there is a single occurrence of the symbolic value, replace it
5834 // with a recurrence.
5835 unsigned FoundIndex = Add->getNumOperands();
5836 for (unsigned i = 0, e = Add->getNumOperands(); i != e; ++i)
5837 if (Add->getOperand(i) == SymbolicName)
5838 if (FoundIndex == e) {
5839 FoundIndex = i;
5840 break;
5841 }
5842
5843 if (FoundIndex != Add->getNumOperands()) {
5844 // Create an add with everything but the specified operand.
5846 for (unsigned i = 0, e = Add->getNumOperands(); i != e; ++i)
5847 if (i != FoundIndex)
5848 Ops.push_back(SCEVBackedgeConditionFolder::rewrite(Add->getOperand(i),
5849 L, *this));
5850 const SCEV *Accum = getAddExpr(Ops);
5851
5852 // This is not a valid addrec if the step amount is varying each
5853 // loop iteration, but is not itself an addrec in this loop.
5854 if (isLoopInvariant(Accum, L) ||
5855 (isa<SCEVAddRecExpr>(Accum) &&
5856 cast<SCEVAddRecExpr>(Accum)->getLoop() == L)) {
5858
5859 if (auto BO = MatchBinaryOp(BEValueV, getDataLayout(), AC, DT, PN)) {
5860 if (BO->Opcode == Instruction::Add && BO->LHS == PN) {
5861 if (BO->IsNUW)
5862 Flags = setFlags(Flags, SCEV::FlagNUW);
5863 if (BO->IsNSW)
5864 Flags = setFlags(Flags, SCEV::FlagNSW);
5865 }
5866 } else if (GEPOperator *GEP = dyn_cast<GEPOperator>(BEValueV)) {
5867 if (GEP->getOperand(0) == PN)
5868 Flags = getNoWrapFlagsForGEP(GEP, Accum, *this);
5869
5870 // We cannot transfer nuw and nsw flags from subtraction
5871 // operations -- sub nuw X, Y is not the same as add nuw X, -Y
5872 // for instance.
5873 }
5874
5875 const SCEV *StartVal = getSCEV(StartValueV);
5876 const SCEV *PHISCEV = getAddRecExpr(StartVal, Accum, L, Flags);
5877
5878 // Okay, for the entire analysis of this edge we assumed the PHI
5879 // to be symbolic. We now need to go back and purge all of the
5880 // entries for the scalars that use the symbolic expression.
5881 forgetMemoizedResults({SymbolicName});
5882 insertValueToMap(PN, PHISCEV);
5883
5884 if (auto *AR = dyn_cast<SCEVAddRecExpr>(PHISCEV))
5885 inferNoWrapViaConstantRanges(AR);
5886
5887 // We can add Flags to the post-inc expression only if we
5888 // know that it is *undefined behavior* for BEValueV to
5889 // overflow.
5890 if (auto *BEInst = dyn_cast<Instruction>(BEValueV))
5891 if (isLoopInvariant(Accum, L) && isAddRecNeverPoison(BEInst, L))
5892 (void)getAddRecExpr(getAddExpr(StartVal, Accum), Accum, L, Flags);
5893
5894 return PHISCEV;
5895 }
5896 }
5897 } else {
5898 // Otherwise, this could be a loop like this:
5899 // i = 0; for (j = 1; ..; ++j) { .... i = j; }
5900 // In this case, j = {1,+,1} and BEValue is j.
5901 // Because the other in-value of i (0) fits the evolution of BEValue
5902 // i really is an addrec evolution.
5903 //
5904 // We can generalize this saying that i is the shifted value of BEValue
5905 // by one iteration:
5906 // PHI(f(0), f({1,+,1})) --> f({0,+,1})
5907
5908 // Do not allow refinement in rewriting of BEValue.
5909 const SCEV *Shifted = SCEVShiftRewriter::rewrite(BEValue, L, *this);
5910 const SCEV *Start = SCEVInitRewriter::rewrite(Shifted, L, *this, false);
5911 if (Shifted != getCouldNotCompute() && Start != getCouldNotCompute() &&
5912 isGuaranteedNotToCauseUB(Shifted) && ::impliesPoison(Shifted, Start)) {
5913 const SCEV *StartVal = getSCEV(StartValueV);
5914 if (Start == StartVal) {
5915 // Okay, for the entire analysis of this edge we assumed the PHI
5916 // to be symbolic. We now need to go back and purge all of the
5917 // entries for the scalars that use the symbolic expression.
5918 forgetMemoizedResults({SymbolicName});
5919 insertValueToMap(PN, Shifted);
5920 return Shifted;
5921 }
5922 }
5923 }
5924
5925 // Remove the temporary PHI node SCEV that has been inserted while intending
5926 // to create an AddRecExpr for this PHI node. We can not keep this temporary
5927 // as it will prevent later (possibly simpler) SCEV expressions to be added
5928 // to the ValueExprMap.
5929 eraseValueFromMap(PN);
5930
5931 return nullptr;
5932}
5933
5934// Try to match a control flow sequence that branches out at BI and merges back
5935// at Merge into a "C ? LHS : RHS" select pattern. Return true on a successful
5936// match.
5938 Value *&C, Value *&LHS, Value *&RHS) {
5939 C = BI->getCondition();
5940
5941 BasicBlockEdge LeftEdge(BI->getParent(), BI->getSuccessor(0));
5942 BasicBlockEdge RightEdge(BI->getParent(), BI->getSuccessor(1));
5943
5944 Use &LeftUse = Merge->getOperandUse(0);
5945 Use &RightUse = Merge->getOperandUse(1);
5946
5947 if (DT.dominates(LeftEdge, LeftUse) && DT.dominates(RightEdge, RightUse)) {
5948 LHS = LeftUse;
5949 RHS = RightUse;
5950 return true;
5951 }
5952
5953 if (DT.dominates(LeftEdge, RightUse) && DT.dominates(RightEdge, LeftUse)) {
5954 LHS = RightUse;
5955 RHS = LeftUse;
5956 return true;
5957 }
5958
5959 return false;
5960}
5961
5963 Value *&Cond, Value *&LHS,
5964 Value *&RHS) {
5965 auto IsReachable =
5966 [&](BasicBlock *BB) { return DT.isReachableFromEntry(BB); };
5967 if (PN->getNumIncomingValues() == 2 && all_of(PN->blocks(), IsReachable)) {
5968 // Try to match
5969 //
5970 // br %cond, label %left, label %right
5971 // left:
5972 // br label %merge
5973 // right:
5974 // br label %merge
5975 // merge:
5976 // V = phi [ %x, %left ], [ %y, %right ]
5977 //
5978 // as "select %cond, %x, %y"
5979
5980 BasicBlock *IDom = DT[PN->getParent()]->getIDom()->getBlock();
5981 assert(IDom && "At least the entry block should dominate PN");
5982
5983 auto *BI = dyn_cast<CondBrInst>(IDom->getTerminator());
5984 return BI && BrPHIToSelect(DT, BI, PN, Cond, LHS, RHS);
5985 }
5986 return false;
5987}
5988
5989const SCEV *ScalarEvolution::createNodeFromSelectLikePHI(PHINode *PN) {
5990 Value *Cond = nullptr, *LHS = nullptr, *RHS = nullptr;
5991 if (getOperandsForSelectLikePHI(DT, PN, Cond, LHS, RHS) &&
5994 return createNodeForSelectOrPHI(PN, Cond, LHS, RHS);
5995
5996 return nullptr;
5997}
5998
6000 BinaryOperator *CommonInst = nullptr;
6001 // Check if instructions are identical.
6002 for (Value *Incoming : PN->incoming_values()) {
6003 auto *IncomingInst = dyn_cast<BinaryOperator>(Incoming);
6004 if (!IncomingInst)
6005 return nullptr;
6006 if (CommonInst) {
6007 if (!CommonInst->isIdenticalToWhenDefined(IncomingInst))
6008 return nullptr; // Not identical, give up
6009 } else {
6010 // Remember binary operator
6011 CommonInst = IncomingInst;
6012 }
6013 }
6014 return CommonInst;
6015}
6016
6017/// Returns SCEV for the first operand of a phi if all phi operands have
6018/// identical opcodes and operands
6019/// eg.
6020/// a: %add = %a + %b
6021/// br %c
6022/// b: %add1 = %a + %b
6023/// br %c
6024/// c: %phi = phi [%add, a], [%add1, b]
6025/// scev(%phi) => scev(%add)
6026const SCEV *
6027ScalarEvolution::createNodeForPHIWithIdenticalOperands(PHINode *PN) {
6028 BinaryOperator *CommonInst = getCommonInstForPHI(PN);
6029 if (!CommonInst)
6030 return nullptr;
6031
6032 // Check if SCEV exprs for instructions are identical.
6033 const SCEV *CommonSCEV = getSCEV(CommonInst);
6034 bool SCEVExprsIdentical =
6036 [this, CommonSCEV](Value *V) { return CommonSCEV == getSCEV(V); });
6037 return SCEVExprsIdentical ? CommonSCEV : nullptr;
6038}
6039
6040const SCEV *ScalarEvolution::createNodeForPHI(PHINode *PN) {
6041 if (const SCEV *S = createAddRecFromPHI(PN))
6042 return S;
6043
6044 // We do not allow simplifying phi (undef, X) to X here, to avoid reusing the
6045 // phi node for X.
6046 if (Value *V = simplifyInstruction(
6047 PN, {getDataLayout(), &TLI, &DT, &AC, /*CtxI=*/nullptr,
6048 /*UseInstrInfo=*/true, /*CanUseUndef=*/false}))
6049 return getSCEV(V);
6050
6051 if (const SCEV *S = createNodeForPHIWithIdenticalOperands(PN))
6052 return S;
6053
6054 if (const SCEV *S = createNodeFromSelectLikePHI(PN))
6055 return S;
6056
6057 // If it's not a loop phi, we can't handle it yet.
6058 return getUnknown(PN);
6059}
6060
6061bool SCEVMinMaxExprContains(const SCEV *Root, const SCEV *OperandToFind,
6062 SCEVTypes RootKind) {
6063 struct FindClosure {
6064 const SCEV *OperandToFind;
6065 const SCEVTypes RootKind; // Must be a sequential min/max expression.
6066 const SCEVTypes NonSequentialRootKind; // Non-seq variant of RootKind.
6067
6068 bool Found = false;
6069
6070 bool canRecurseInto(SCEVTypes Kind) const {
6071 // We can only recurse into the SCEV expression of the same effective type
6072 // as the type of our root SCEV expression, and into zero-extensions.
6073 return RootKind == Kind || NonSequentialRootKind == Kind ||
6074 scZeroExtend == Kind;
6075 };
6076
6077 FindClosure(const SCEV *OperandToFind, SCEVTypes RootKind)
6078 : OperandToFind(OperandToFind), RootKind(RootKind),
6079 NonSequentialRootKind(
6081 RootKind)) {}
6082
6083 bool follow(const SCEV *S) {
6084 Found = S == OperandToFind;
6085
6086 return !isDone() && canRecurseInto(S->getSCEVType());
6087 }
6088
6089 bool isDone() const { return Found; }
6090 };
6091
6092 FindClosure FC(OperandToFind, RootKind);
6093 visitAll(Root, FC);
6094 return FC.Found;
6095}
6096
6097std::optional<const SCEV *>
6098ScalarEvolution::createNodeForSelectOrPHIInstWithICmpInstCond(Type *Ty,
6099 ICmpInst *Cond,
6100 Value *TrueVal,
6101 Value *FalseVal) {
6102 // Try to match some simple smax or umax patterns.
6103 auto *ICI = Cond;
6104
6105 Value *LHS = ICI->getOperand(0);
6106 Value *RHS = ICI->getOperand(1);
6107
6108 switch (ICI->getPredicate()) {
6109 case ICmpInst::ICMP_SLT:
6110 case ICmpInst::ICMP_SLE:
6111 case ICmpInst::ICMP_ULT:
6112 case ICmpInst::ICMP_ULE:
6113 std::swap(LHS, RHS);
6114 [[fallthrough]];
6115 case ICmpInst::ICMP_SGT:
6116 case ICmpInst::ICMP_SGE:
6117 case ICmpInst::ICMP_UGT:
6118 case ICmpInst::ICMP_UGE:
6119 // a > b ? a+x : b+x -> max(a, b)+x
6120 // a > b ? b+x : a+x -> min(a, b)+x
6122 bool Signed = ICI->isSigned();
6123 const SCEV *LA = getSCEV(TrueVal);
6124 const SCEV *RA = getSCEV(FalseVal);
6125 const SCEV *LS = getSCEV(LHS);
6126 const SCEV *RS = getSCEV(RHS);
6127 if (LA->getType()->isPointerTy()) {
6128 // FIXME: Handle cases where LS/RS are pointers not equal to LA/RA.
6129 // Need to make sure we can't produce weird expressions involving
6130 // negated pointers.
6131 if (LA == LS && RA == RS)
6132 return Signed ? getSMaxExpr(LS, RS) : getUMaxExpr(LS, RS);
6133 if (LA == RS && RA == LS)
6134 return Signed ? getSMinExpr(LS, RS) : getUMinExpr(LS, RS);
6135 }
6136 auto CoerceOperand = [&](const SCEV *Op) -> const SCEV * {
6137 if (Op->getType()->isPointerTy()) {
6140 return Op;
6141 }
6142 if (Signed)
6143 Op = getNoopOrSignExtend(Op, Ty);
6144 else
6145 Op = getNoopOrZeroExtend(Op, Ty);
6146 return Op;
6147 };
6148 LS = CoerceOperand(LS);
6149 RS = CoerceOperand(RS);
6151 break;
6152 const SCEV *LDiff = getMinusSCEV(LA, LS);
6153 const SCEV *RDiff = getMinusSCEV(RA, RS);
6154 if (LDiff == RDiff)
6155 return getAddExpr(Signed ? getSMaxExpr(LS, RS) : getUMaxExpr(LS, RS),
6156 LDiff);
6157 LDiff = getMinusSCEV(LA, RS);
6158 RDiff = getMinusSCEV(RA, LS);
6159 if (LDiff == RDiff)
6160 return getAddExpr(Signed ? getSMinExpr(LS, RS) : getUMinExpr(LS, RS),
6161 LDiff);
6162 }
6163 break;
6164 case ICmpInst::ICMP_NE:
6165 // x != 0 ? x+y : C+y -> x == 0 ? C+y : x+y
6166 std::swap(TrueVal, FalseVal);
6167 [[fallthrough]];
6168 case ICmpInst::ICMP_EQ:
6169 // x == 0 ? C+y : x+y -> umax(x, C)+y iff C u<= 1
6172 const SCEV *X = getNoopOrZeroExtend(getSCEV(LHS), Ty);
6173 const SCEV *TrueValExpr = getSCEV(TrueVal); // C+y
6174 const SCEV *FalseValExpr = getSCEV(FalseVal); // x+y
6175 const SCEV *Y = getMinusSCEV(FalseValExpr, X); // y = (x+y)-x
6176 const SCEV *C = getMinusSCEV(TrueValExpr, Y); // C = (C+y)-y
6177 if (isa<SCEVConstant>(C) && cast<SCEVConstant>(C)->getAPInt().ule(1))
6178 return getAddExpr(getUMaxExpr(X, C), Y);
6179 }
6180 // x == 0 ? 0 : umin (..., x, ...) -> umin_seq(x, umin (...))
6181 // x == 0 ? 0 : umin_seq(..., x, ...) -> umin_seq(x, umin_seq(...))
6182 // x == 0 ? 0 : umin (..., umin_seq(..., x, ...), ...)
6183 // -> umin_seq(x, umin (..., umin_seq(...), ...))
6185 isa<ConstantInt>(TrueVal) && cast<ConstantInt>(TrueVal)->isZero()) {
6186 const SCEV *X = getSCEV(LHS);
6187 while (auto *ZExt = dyn_cast<SCEVZeroExtendExpr>(X))
6188 X = ZExt->getOperand();
6189 if (getTypeSizeInBits(X->getType()) <= getTypeSizeInBits(Ty)) {
6190 const SCEV *FalseValExpr = getSCEV(FalseVal);
6191 if (SCEVMinMaxExprContains(FalseValExpr, X, scSequentialUMinExpr))
6192 return getUMinExpr(getNoopOrZeroExtend(X, Ty), FalseValExpr,
6193 /*Sequential=*/true);
6194 }
6195 }
6196 break;
6197 default:
6198 break;
6199 }
6200
6201 return std::nullopt;
6202}
6203
6204static std::optional<const SCEV *>
6206 const SCEV *TrueExpr, const SCEV *FalseExpr) {
6207 assert(CondExpr->getType()->isIntegerTy(1) &&
6208 TrueExpr->getType() == FalseExpr->getType() &&
6209 TrueExpr->getType()->isIntegerTy(1) &&
6210 "Unexpected operands of a select.");
6211
6212 // i1 cond ? i1 x : i1 C --> C + (i1 cond ? (i1 x - i1 C) : i1 0)
6213 // --> C + (umin_seq cond, x - C)
6214 //
6215 // i1 cond ? i1 C : i1 x --> C + (i1 cond ? i1 0 : (i1 x - i1 C))
6216 // --> C + (i1 ~cond ? (i1 x - i1 C) : i1 0)
6217 // --> C + (umin_seq ~cond, x - C)
6218
6219 // FIXME: while we can't legally model the case where both of the hands
6220 // are fully variable, we only require that the *difference* is constant.
6221 if (!isa<SCEVConstant>(TrueExpr) && !isa<SCEVConstant>(FalseExpr))
6222 return std::nullopt;
6223
6224 const SCEV *X, *C;
6225 if (isa<SCEVConstant>(TrueExpr)) {
6226 CondExpr = SE->getNotSCEV(CondExpr);
6227 X = FalseExpr;
6228 C = TrueExpr;
6229 } else {
6230 X = TrueExpr;
6231 C = FalseExpr;
6232 }
6233 return SE->getAddExpr(C, SE->getUMinExpr(CondExpr, SE->getMinusSCEV(X, C),
6234 /*Sequential=*/true));
6235}
6236
6237static std::optional<const SCEV *>
6239 Value *FalseVal) {
6240 if (!isa<ConstantInt>(TrueVal) && !isa<ConstantInt>(FalseVal))
6241 return std::nullopt;
6242
6243 const auto *SECond = SE->getSCEV(Cond);
6244 const auto *SETrue = SE->getSCEV(TrueVal);
6245 const auto *SEFalse = SE->getSCEV(FalseVal);
6246 return createNodeForSelectViaUMinSeq(SE, SECond, SETrue, SEFalse);
6247}
6248
6249const SCEV *ScalarEvolution::createNodeForSelectOrPHIViaUMinSeq(
6250 Value *V, Value *Cond, Value *TrueVal, Value *FalseVal) {
6251 assert(Cond->getType()->isIntegerTy(1) && "Select condition is not an i1?");
6252 assert(TrueVal->getType() == FalseVal->getType() &&
6253 V->getType() == TrueVal->getType() &&
6254 "Types of select hands and of the result must match.");
6255
6256 // For now, only deal with i1-typed `select`s.
6257 if (!V->getType()->isIntegerTy(1))
6258 return getUnknown(V);
6259
6260 if (std::optional<const SCEV *> S =
6261 createNodeForSelectViaUMinSeq(this, Cond, TrueVal, FalseVal))
6262 return *S;
6263
6264 return getUnknown(V);
6265}
6266
6267const SCEV *ScalarEvolution::createNodeForSelectOrPHI(Value *V, Value *Cond,
6268 Value *TrueVal,
6269 Value *FalseVal) {
6270 // Handle "constant" branch or select. This can occur for instance when a
6271 // loop pass transforms an inner loop and moves on to process the outer loop.
6272 if (auto *CI = dyn_cast<ConstantInt>(Cond))
6273 return getSCEV(CI->isOne() ? TrueVal : FalseVal);
6274
6275 if (auto *I = dyn_cast<Instruction>(V)) {
6276 if (auto *ICI = dyn_cast<ICmpInst>(Cond)) {
6277 if (std::optional<const SCEV *> S =
6278 createNodeForSelectOrPHIInstWithICmpInstCond(I->getType(), ICI,
6279 TrueVal, FalseVal))
6280 return *S;
6281 }
6282 }
6283
6284 return createNodeForSelectOrPHIViaUMinSeq(V, Cond, TrueVal, FalseVal);
6285}
6286
6287/// Expand GEP instructions into add and multiply operations. This allows them
6288/// to be analyzed by regular SCEV code.
6289const SCEV *ScalarEvolution::createNodeForGEP(GEPOperator *GEP) {
6290 assert(GEP->getSourceElementType()->isSized() &&
6291 "GEP source element type must be sized");
6292
6293 SmallVector<SCEVUse, 4> IndexExprs;
6294 for (Value *Index : GEP->indices())
6295 IndexExprs.push_back(getSCEV(Index));
6296 return getGEPExpr(GEP, IndexExprs);
6297}
6298
6299APInt ScalarEvolution::getConstantMultipleImpl(const SCEV *S,
6300 const Instruction *CtxI) {
6302 auto GetShiftedByZeros = [BitWidth](uint32_t TrailingZeros) {
6303 return TrailingZeros >= BitWidth
6305 : APInt::getOneBitSet(BitWidth, TrailingZeros);
6306 };
6307 auto GetGCDMultiple = [this, CtxI](const SCEVNAryExpr *N) {
6308 // The result is GCD of all operands results.
6309 APInt Res = getConstantMultiple(N->getOperand(0), CtxI);
6310 for (unsigned I = 1, E = N->getNumOperands(); I < E && Res != 1; ++I)
6312 Res, getConstantMultiple(N->getOperand(I), CtxI));
6313 return Res;
6314 };
6315
6316 switch (S->getSCEVType()) {
6317 case scConstant:
6318 return cast<SCEVConstant>(S)->getAPInt();
6319 case scPtrToAddr:
6320 return getConstantMultiple(cast<SCEVCastExpr>(S)->getOperand());
6321 case scUDivExpr:
6322 case scVScale:
6323 return APInt(BitWidth, 1);
6324 case scTruncate: {
6325 // Only multiples that are a power of 2 will hold after truncation.
6326 const SCEVTruncateExpr *T = cast<SCEVTruncateExpr>(S);
6327 uint32_t TZ = getMinTrailingZeros(T->getOperand(), CtxI);
6328 return GetShiftedByZeros(TZ);
6329 }
6330 case scZeroExtend: {
6331 const SCEVZeroExtendExpr *Z = cast<SCEVZeroExtendExpr>(S);
6332 return getConstantMultiple(Z->getOperand(), CtxI).zext(BitWidth);
6333 }
6334 case scSignExtend: {
6335 // Only multiples that are a power of 2 will hold after sext.
6336 const SCEVSignExtendExpr *E = cast<SCEVSignExtendExpr>(S);
6337 uint32_t TZ = getMinTrailingZeros(E->getOperand(), CtxI);
6338 return GetShiftedByZeros(TZ);
6339 }
6340 case scMulExpr: {
6341 const SCEVMulExpr *M = cast<SCEVMulExpr>(S);
6342 if (M->hasNoUnsignedWrap()) {
6343 // The result is the product of all operand results.
6344 APInt Res = getConstantMultiple(M->getOperand(0), CtxI);
6345 for (const SCEV *Operand : M->operands().drop_front())
6346 Res = Res * getConstantMultiple(Operand, CtxI);
6347 return Res;
6348 }
6349
6350 // If there are no wrap guarentees, find the trailing zeros, which is the
6351 // sum of trailing zeros for all its operands.
6352 uint32_t TZ = 0;
6353 for (const SCEV *Operand : M->operands())
6354 TZ += getMinTrailingZeros(Operand, CtxI);
6355 return GetShiftedByZeros(TZ);
6356 }
6357 case scAddExpr:
6358 case scAddRecExpr: {
6359 const SCEVNAryExpr *N = cast<SCEVNAryExpr>(S);
6360 if (N->hasNoUnsignedWrap())
6361 return GetGCDMultiple(N);
6362 // Find the trailing bits, which is the minimum of its operands.
6363 uint32_t TZ = getMinTrailingZeros(N->getOperand(0), CtxI);
6364 for (const SCEV *Operand : N->operands().drop_front())
6365 TZ = std::min(TZ, getMinTrailingZeros(Operand, CtxI));
6366 return GetShiftedByZeros(TZ);
6367 }
6368 case scUMaxExpr:
6369 case scSMaxExpr:
6370 case scUMinExpr:
6371 case scSMinExpr:
6373 return GetGCDMultiple(cast<SCEVNAryExpr>(S));
6374 case scUnknown: {
6375 // Ask ValueTracking for known bits. SCEVUnknown only become available at
6376 // the point their underlying IR instruction has been defined. If CtxI was
6377 // not provided, use:
6378 // * the first instruction in the entry block if it is an argument
6379 // * the instruction itself otherwise.
6380 const SCEVUnknown *U = cast<SCEVUnknown>(S);
6381 if (!CtxI) {
6382 if (isa<Argument>(U->getValue()))
6383 CtxI = &*F.getEntryBlock().begin();
6384 else if (auto *I = dyn_cast<Instruction>(U->getValue()))
6385 CtxI = I;
6386 }
6387 unsigned Known =
6388 computeKnownBits(U->getValue(),
6389 SimplifyQuery(getDataLayout(), &DT, &AC, CtxI)
6390 .allowEphemerals(true))
6391 .countMinTrailingZeros();
6392 return GetShiftedByZeros(Known);
6393 }
6394 case scCouldNotCompute:
6395 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
6396 }
6397 llvm_unreachable("Unknown SCEV kind!");
6398}
6399
6401 const Instruction *CtxI) {
6402 // Skip looking up and updating the cache if there is a context instruction,
6403 // as the result will only be valid in the specified context.
6404 if (CtxI)
6405 return getConstantMultipleImpl(S, CtxI);
6406
6407 auto I = ConstantMultipleCache.find(S);
6408 if (I != ConstantMultipleCache.end())
6409 return I->second;
6410
6411 APInt Result = getConstantMultipleImpl(S, CtxI);
6412 auto InsertPair = ConstantMultipleCache.insert({S, Result});
6413 assert(InsertPair.second && "Should insert a new key");
6414 return InsertPair.first->second;
6415}
6416
6418 APInt Multiple = getConstantMultiple(S);
6419 return Multiple == 0 ? APInt(Multiple.getBitWidth(), 1) : Multiple;
6420}
6421
6423 const Instruction *CtxI) {
6424 return std::min(getConstantMultiple(S, CtxI).countTrailingZeros(),
6425 (unsigned)getTypeSizeInBits(S->getType()));
6426}
6427
6428/// Helper method to assign a range to V from metadata present in the IR.
6429static std::optional<ConstantRange> GetRangeFromMetadata(Value *V) {
6431 if (MDNode *MD = I->getMetadata(LLVMContext::MD_range))
6432 return getConstantRangeFromMetadata(*MD);
6433 if (const auto *CB = dyn_cast<CallBase>(V))
6434 if (std::optional<ConstantRange> Range = CB->getRange())
6435 return Range;
6436 }
6437 if (auto *A = dyn_cast<Argument>(V))
6438 if (std::optional<ConstantRange> Range = A->getRange())
6439 return Range;
6440
6441 return std::nullopt;
6442}
6443
6445 SCEVFlags NWFlags = Flags & SCEV::FlagsNoWrapMask;
6446 if (AddRec->getNoWrapFlags(NWFlags) != NWFlags) {
6447 AddRec->setNoWrapFlags(NWFlags);
6448 UnsignedRanges.erase(AddRec);
6449 SignedRanges.erase(AddRec);
6450 ConstantMultipleCache.erase(AddRec);
6451 }
6452}
6453
6454ConstantRange ScalarEvolution::
6455getRangeForUnknownRecurrence(const SCEVUnknown *U) {
6456 const DataLayout &DL = getDataLayout();
6457
6458 unsigned BitWidth = getTypeSizeInBits(U->getType());
6459 const ConstantRange FullSet(BitWidth, /*isFullSet=*/true);
6460
6461 // Match a simple recurrence of the form: <start, ShiftOp, Step>, and then
6462 // use information about the trip count to improve our available range. Note
6463 // that the trip count independent cases are already handled by known bits.
6464 // WARNING: The definition of recurrence used here is subtly different than
6465 // the one used by AddRec (and thus most of this file). Step is allowed to
6466 // be arbitrarily loop varying here, where AddRec allows only loop invariant
6467 // and other addrecs in the same loop (for non-affine addrecs). The code
6468 // below intentionally handles the case where step is not loop invariant.
6469 auto *P = dyn_cast<PHINode>(U->getValue());
6470 if (!P)
6471 return FullSet;
6472
6473 // Make sure that no Phi input comes from an unreachable block. Otherwise,
6474 // even the values that are not available in these blocks may come from them,
6475 // and this leads to false-positive recurrence test.
6476 for (auto *Pred : predecessors(P->getParent()))
6477 if (!DT.isReachableFromEntry(Pred))
6478 return FullSet;
6479
6480 BinaryOperator *BO;
6481 Value *Start, *Step;
6482 if (!matchSimpleRecurrence(P, BO, Start, Step))
6483 return FullSet;
6484
6485 // If we found a recurrence in reachable code, we must be in a loop. Note
6486 // that BO might be in some subloop of L, and that's completely okay.
6487 auto *L = LI.getLoopFor(P->getParent());
6488 assert(L && L->getHeader() == P->getParent());
6489 if (!L->contains(BO->getParent()))
6490 // NOTE: This bailout should be an assert instead. However, asserting
6491 // the condition here exposes a case where LoopFusion is querying SCEV
6492 // with malformed loop information during the midst of the transform.
6493 // There doesn't appear to be an obvious fix, so for the moment bailout
6494 // until the caller issue can be fixed. PR49566 tracks the bug.
6495 return FullSet;
6496
6497 // TODO: Extend to other opcodes such as mul, and div
6498 switch (BO->getOpcode()) {
6499 default:
6500 return FullSet;
6501 case Instruction::AShr:
6502 case Instruction::LShr:
6503 case Instruction::Shl:
6504 break;
6505 };
6506
6507 if (BO->getOperand(0) != P)
6508 // TODO: Handle the power function forms some day.
6509 return FullSet;
6510
6511 unsigned TC = getSmallConstantMaxTripCount(L);
6512 if (!TC || TC >= BitWidth)
6513 return FullSet;
6514
6515 auto KnownStart = computeKnownBits(Start, DL, &AC, nullptr, &DT);
6516 auto KnownStep = computeKnownBits(Step, DL, &AC, nullptr, &DT);
6517 assert(KnownStart.getBitWidth() == BitWidth &&
6518 KnownStep.getBitWidth() == BitWidth);
6519
6520 // Compute total shift amount, being careful of overflow and bitwidths.
6521 auto MaxShiftAmt = KnownStep.getMaxValue();
6522 APInt TCAP(BitWidth, TC-1);
6523 bool Overflow = false;
6524 auto TotalShift = MaxShiftAmt.umul_ov(TCAP, Overflow);
6525 if (Overflow)
6526 return FullSet;
6527
6528 switch (BO->getOpcode()) {
6529 default:
6530 llvm_unreachable("filtered out above");
6531 case Instruction::AShr: {
6532 // For each ashr, three cases:
6533 // shift = 0 => unchanged value
6534 // saturation => 0 or -1
6535 // other => a value closer to zero (of the same sign)
6536 // Thus, the end value is closer to zero than the start.
6537 auto KnownEnd = KnownBits::ashr(KnownStart,
6538 KnownBits::makeConstant(TotalShift));
6539 if (KnownStart.isNonNegative())
6540 // Analogous to lshr (simply not yet canonicalized)
6541 return ConstantRange::getNonEmpty(KnownEnd.getMinValue(),
6542 KnownStart.getMaxValue() + 1);
6543 if (KnownStart.isNegative())
6544 // End >=u Start && End <=s Start
6545 return ConstantRange::getNonEmpty(KnownStart.getMinValue(),
6546 KnownEnd.getMaxValue() + 1);
6547 break;
6548 }
6549 case Instruction::LShr: {
6550 // For each lshr, three cases:
6551 // shift = 0 => unchanged value
6552 // saturation => 0
6553 // other => a smaller positive number
6554 // Thus, the low end of the unsigned range is the last value produced.
6555 auto KnownEnd = KnownBits::lshr(KnownStart,
6556 KnownBits::makeConstant(TotalShift));
6557 return ConstantRange::getNonEmpty(KnownEnd.getMinValue(),
6558 KnownStart.getMaxValue() + 1);
6559 }
6560 case Instruction::Shl: {
6561 // Iff no bits are shifted out, value increases on every shift.
6562 auto KnownEnd = KnownBits::shl(KnownStart,
6563 KnownBits::makeConstant(TotalShift));
6564 if (TotalShift.ult(KnownStart.countMinLeadingZeros()))
6565 return ConstantRange(KnownStart.getMinValue(),
6566 KnownEnd.getMaxValue() + 1);
6567 break;
6568 }
6569 };
6570 return FullSet;
6571}
6572
6573// The goal of this function is to check if recursively visiting the operands
6574// of this PHI might lead to an infinite loop. If we do see such a loop,
6575// there's no good way to break it, so we avoid analyzing such cases.
6576//
6577// getRangeRef previously used a visited set to avoid infinite loops, but this
6578// caused other issues: the result was dependent on the order of getRangeRef
6579// calls, and the interaction with createSCEVIter could cause a stack overflow
6580// in some cases (see issue #148253).
6581//
6582// FIXME: The way this is implemented is overly conservative; this checks
6583// for a few obviously safe patterns, but anything that doesn't lead to
6584// recursion is fine.
6586 Value *Cond = nullptr, *LHS = nullptr, *RHS = nullptr;
6588 return true;
6589
6590 if (all_of(PHI->operands(),
6591 [&](Value *Operand) { return DT.dominates(Operand, PHI); }))
6592 return true;
6593
6594 return false;
6595}
6596
6597const ConstantRange &
6598ScalarEvolution::getRangeRefIter(const SCEV *S,
6599 ScalarEvolution::RangeSignHint SignHint) {
6600 DenseMap<const SCEV *, ConstantRange> &Cache =
6601 SignHint == ScalarEvolution::HINT_RANGE_UNSIGNED ? UnsignedRanges
6602 : SignedRanges;
6603 SmallVector<SCEVUse> WorkList;
6604 SmallPtrSet<const SCEV *, 8> Seen;
6605
6606 // Add Expr to the worklist, if Expr is either an N-ary expression or a
6607 // SCEVUnknown PHI node.
6608 auto AddToWorklist = [&WorkList, &Seen, &Cache](const SCEV *Expr) {
6609 if (!Seen.insert(Expr).second)
6610 return;
6611 if (Cache.contains(Expr))
6612 return;
6613 switch (Expr->getSCEVType()) {
6614 case scUnknown:
6616 break;
6617 [[fallthrough]];
6618 case scConstant:
6619 case scVScale:
6620 case scTruncate:
6621 case scZeroExtend:
6622 case scSignExtend:
6623 case scPtrToAddr:
6624 case scAddExpr:
6625 case scMulExpr:
6626 case scUDivExpr:
6627 case scAddRecExpr:
6628 case scUMaxExpr:
6629 case scSMaxExpr:
6630 case scUMinExpr:
6631 case scSMinExpr:
6633 WorkList.push_back(Expr);
6634 break;
6635 case scCouldNotCompute:
6636 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
6637 }
6638 };
6639 AddToWorklist(S);
6640
6641 // Build worklist by queuing operands of N-ary expressions and phi nodes.
6642 for (unsigned I = 0; I != WorkList.size(); ++I) {
6643 const SCEV *P = WorkList[I];
6644 auto *UnknownS = dyn_cast<SCEVUnknown>(P);
6645 // If it is not a `SCEVUnknown`, just recurse into operands.
6646 if (!UnknownS) {
6647 for (const SCEV *Op : P->operands())
6648 AddToWorklist(Op);
6649 continue;
6650 }
6651 // `SCEVUnknown`'s require special treatment.
6652 if (PHINode *P = dyn_cast<PHINode>(UnknownS->getValue())) {
6653 if (!RangeRefPHIAllowedOperands(DT, P))
6654 continue;
6655 for (auto &Op : reverse(P->operands()))
6656 AddToWorklist(getSCEV(Op));
6657 }
6658 }
6659
6660 if (!WorkList.empty()) {
6661 // Use getRangeRef to compute ranges for items in the worklist in reverse
6662 // order. This will force ranges for earlier operands to be computed before
6663 // their users in most cases.
6664 for (const SCEV *P : reverse(drop_begin(WorkList))) {
6665 getRangeRef(P, SignHint);
6666 }
6667 }
6668
6669 return getRangeRef(S, SignHint, 0);
6670}
6671
6672const APInt *ScalarEvolution::getConstantAPIntOrNull(const SCEV *S) {
6673 if (const auto *C = dyn_cast<SCEVConstant>(S))
6674 return &C->getAPInt();
6675 return nullptr;
6676}
6677
6678/// Determine the range for a particular SCEV. If SignHint is
6679/// HINT_RANGE_UNSIGNED (resp. HINT_RANGE_SIGNED) then getRange prefers ranges
6680/// with a "cleaner" unsigned (resp. signed) representation.
6681const ConstantRange &ScalarEvolution::getRangeRef(
6682 const SCEV *S, ScalarEvolution::RangeSignHint SignHint, unsigned Depth) {
6683 DenseMap<const SCEV *, ConstantRange> &Cache =
6684 SignHint == ScalarEvolution::HINT_RANGE_UNSIGNED ? UnsignedRanges
6685 : SignedRanges;
6687 SignHint == ScalarEvolution::HINT_RANGE_UNSIGNED ? ConstantRange::Unsigned
6689
6690 // See if we've computed this range already.
6691 auto I = Cache.find(S);
6692 if (I != Cache.end())
6693 return I->second;
6694
6695 if (const SCEVConstant *C = dyn_cast<SCEVConstant>(S))
6696 return setRange(C, SignHint, ConstantRange(C->getAPInt()));
6697
6698 // Switch to iteratively computing the range for S, if it is part of a deeply
6699 // nested expression.
6701 return getRangeRefIter(S, SignHint);
6702
6703 unsigned BitWidth = getTypeSizeInBits(S->getType());
6704 ConstantRange ConservativeResult(BitWidth, /*isFullSet=*/true);
6705 using OBO = OverflowingBinaryOperator;
6706
6707 // If the value has known zeros, the maximum value will have those known zeros
6708 // as well.
6709 if (SignHint == ScalarEvolution::HINT_RANGE_UNSIGNED) {
6710 APInt Multiple = getNonZeroConstantMultiple(S);
6711 APInt Remainder = APInt::getMaxValue(BitWidth).urem(Multiple);
6712 if (!Remainder.isZero())
6713 ConservativeResult =
6714 ConstantRange(APInt::getMinValue(BitWidth),
6715 APInt::getMaxValue(BitWidth) - Remainder + 1);
6716 }
6717 else {
6718 uint32_t TZ = getMinTrailingZeros(S);
6719 if (TZ != 0) {
6720 ConservativeResult = ConstantRange(
6722 APInt::getSignedMaxValue(BitWidth).ashr(TZ).shl(TZ) + 1);
6723 }
6724 }
6725
6726 switch (S->getSCEVType()) {
6727 case scConstant:
6728 llvm_unreachable("Already handled above.");
6729 case scVScale:
6730 return setRange(S, SignHint, getVScaleRange(&F, BitWidth));
6731 case scTruncate: {
6732 const SCEVTruncateExpr *Trunc = cast<SCEVTruncateExpr>(S);
6733 ConstantRange X = getRangeRef(Trunc->getOperand(), SignHint, Depth + 1);
6734 return setRange(
6735 Trunc, SignHint,
6736 ConservativeResult.intersectWith(X.truncate(BitWidth), RangeType));
6737 }
6738 case scZeroExtend: {
6739 const SCEVZeroExtendExpr *ZExt = cast<SCEVZeroExtendExpr>(S);
6740 ConstantRange X = getRangeRef(ZExt->getOperand(), SignHint, Depth + 1);
6741 return setRange(
6742 ZExt, SignHint,
6743 ConservativeResult.intersectWith(X.zeroExtend(BitWidth), RangeType));
6744 }
6745 case scSignExtend: {
6746 const SCEVSignExtendExpr *SExt = cast<SCEVSignExtendExpr>(S);
6747 ConstantRange X = getRangeRef(SExt->getOperand(), SignHint, Depth + 1);
6748 return setRange(
6749 SExt, SignHint,
6750 ConservativeResult.intersectWith(X.signExtend(BitWidth), RangeType));
6751 }
6752 case scPtrToAddr: {
6753 const SCEVCastExpr *Cast = cast<SCEVCastExpr>(S);
6754 ConstantRange X = getRangeRef(Cast->getOperand(), SignHint, Depth + 1);
6755 return setRange(Cast, SignHint, X);
6756 }
6757 case scAddExpr: {
6758 const SCEVAddExpr *Add = cast<SCEVAddExpr>(S);
6759 // Check if this is a URem pattern: A - (A / B) * B, which is always < B.
6760 const SCEV *URemLHS = nullptr, *URemRHS = nullptr;
6761 if (SignHint == ScalarEvolution::HINT_RANGE_UNSIGNED &&
6762 match(S, m_scev_URem(m_SCEV(URemLHS), m_SCEV(URemRHS), *this))) {
6763 ConstantRange LHSRange = getRangeRef(URemLHS, SignHint, Depth + 1);
6764 ConstantRange RHSRange = getRangeRef(URemRHS, SignHint, Depth + 1);
6765 ConservativeResult =
6766 ConservativeResult.intersectWith(LHSRange.urem(RHSRange), RangeType);
6767 }
6768 ConstantRange X = getRangeRef(Add->getOperand(0), SignHint, Depth + 1);
6769 unsigned WrapType = OBO::AnyWrap;
6770 if (Add->hasNoSignedWrap())
6771 WrapType |= OBO::NoSignedWrap;
6772 if (Add->hasNoUnsignedWrap())
6773 WrapType |= OBO::NoUnsignedWrap;
6774 for (const SCEV *Op : drop_begin(Add->operands()))
6775 X = X.addWithNoWrap(getRangeRef(Op, SignHint, Depth + 1), WrapType,
6776 RangeType);
6777 return setRange(Add, SignHint,
6778 ConservativeResult.intersectWith(X, RangeType));
6779 }
6780 case scMulExpr: {
6781 const SCEVMulExpr *Mul = cast<SCEVMulExpr>(S);
6782 ConstantRange X = getRangeRef(Mul->getOperand(0), SignHint, Depth + 1);
6783 for (const SCEV *Op : drop_begin(Mul->operands()))
6784 X = X.multiply(getRangeRef(Op, SignHint, Depth + 1));
6785 return setRange(Mul, SignHint,
6786 ConservativeResult.intersectWith(X, RangeType));
6787 }
6788 case scUDivExpr: {
6789 const SCEVUDivExpr *UDiv = cast<SCEVUDivExpr>(S);
6790 ConstantRange X = getRangeRef(UDiv->getLHS(), SignHint, Depth + 1);
6791 ConstantRange Y = getRangeRef(UDiv->getRHS(), SignHint, Depth + 1);
6792 return setRange(UDiv, SignHint,
6793 ConservativeResult.intersectWith(X.udiv(Y), RangeType));
6794 }
6795 case scAddRecExpr: {
6796 const SCEVAddRecExpr *AddRec = cast<SCEVAddRecExpr>(S);
6797 // If there's no unsigned wrap, the value will never be less than its
6798 // initial value.
6799 if (AddRec->hasNoUnsignedWrap()) {
6800 APInt UnsignedMinValue = getUnsignedRangeMin(AddRec->getStart());
6801 if (!UnsignedMinValue.isZero())
6802 ConservativeResult = ConservativeResult.intersectWith(
6803 ConstantRange(UnsignedMinValue, APInt(BitWidth, 0)), RangeType);
6804 }
6805
6806 // If there's no signed wrap, and all the operands except initial value have
6807 // the same sign or zero, the value won't ever be:
6808 // 1: smaller than initial value if operands are non negative,
6809 // 2: bigger than initial value if operands are non positive.
6810 // For both cases, value can not cross signed min/max boundary.
6811 if (AddRec->hasNoSignedWrap()) {
6812 bool AllNonNeg = true;
6813 bool AllNonPos = true;
6814 for (unsigned i = 1, e = AddRec->getNumOperands(); i != e; ++i) {
6815 if (!isKnownNonNegative(AddRec->getOperand(i)))
6816 AllNonNeg = false;
6817 if (!isKnownNonPositive(AddRec->getOperand(i)))
6818 AllNonPos = false;
6819 }
6820 if (AllNonNeg)
6821 ConservativeResult = ConservativeResult.intersectWith(
6824 RangeType);
6825 else if (AllNonPos)
6826 ConservativeResult = ConservativeResult.intersectWith(
6828 getSignedRangeMax(AddRec->getStart()) +
6829 1),
6830 RangeType);
6831 }
6832
6833 // TODO: non-affine addrec
6834 if (AddRec->isAffine()) {
6835 const SCEV *MaxBEScev =
6837 if (!isa<SCEVCouldNotCompute>(MaxBEScev)) {
6838 APInt MaxBECount = cast<SCEVConstant>(MaxBEScev)->getAPInt();
6839
6840 // Adjust MaxBECount to the same bitwidth as AddRec. We can truncate if
6841 // MaxBECount's active bits are all <= AddRec's bit width.
6842 if (MaxBECount.getBitWidth() > BitWidth &&
6843 MaxBECount.getActiveBits() <= BitWidth)
6844 MaxBECount = MaxBECount.trunc(BitWidth);
6845 else if (MaxBECount.getBitWidth() < BitWidth)
6846 MaxBECount = MaxBECount.zext(BitWidth);
6847
6848 if (MaxBECount.getBitWidth() == BitWidth) {
6849 auto [RangeFromAffine, Flags] = getRangeForAffineAR(
6850 AddRec->getStart(), AddRec->getStepRecurrence(*this), MaxBECount);
6851 ConservativeResult =
6852 ConservativeResult.intersectWith(RangeFromAffine, RangeType);
6853 const_cast<SCEVAddRecExpr *>(AddRec)->setNoWrapFlags(Flags);
6854
6855 auto RangeFromFactoring = getRangeViaFactoring(
6856 AddRec->getStart(), AddRec->getStepRecurrence(*this), MaxBECount);
6857 ConservativeResult =
6858 ConservativeResult.intersectWith(RangeFromFactoring, RangeType);
6859 }
6860 }
6861
6862 // Now try symbolic BE count and more powerful methods.
6864 const SCEV *SymbolicMaxBECount =
6866 if (!isa<SCEVCouldNotCompute>(SymbolicMaxBECount) &&
6867 getTypeSizeInBits(MaxBEScev->getType()) <= BitWidth &&
6868 AddRec->hasNoSelfWrap()) {
6869 auto RangeFromAffineNew = getRangeForAffineNoSelfWrappingAR(
6870 AddRec, SymbolicMaxBECount, BitWidth, SignHint);
6871 ConservativeResult =
6872 ConservativeResult.intersectWith(RangeFromAffineNew, RangeType);
6873 }
6874 }
6875 }
6876
6877 return setRange(AddRec, SignHint, std::move(ConservativeResult));
6878 }
6879 case scUMaxExpr:
6880 case scSMaxExpr:
6881 case scUMinExpr:
6882 case scSMinExpr:
6883 case scSequentialUMinExpr: {
6885 switch (S->getSCEVType()) {
6886 case scUMaxExpr:
6887 ID = Intrinsic::umax;
6888 break;
6889 case scSMaxExpr:
6890 ID = Intrinsic::smax;
6891 break;
6892 case scUMinExpr:
6894 ID = Intrinsic::umin;
6895 break;
6896 case scSMinExpr:
6897 ID = Intrinsic::smin;
6898 break;
6899 default:
6900 llvm_unreachable("Unknown SCEVMinMaxExpr/SCEVSequentialMinMaxExpr.");
6901 }
6902
6903 const auto *NAry = cast<SCEVNAryExpr>(S);
6904 ConstantRange X = getRangeRef(NAry->getOperand(0), SignHint, Depth + 1);
6905 for (unsigned i = 1, e = NAry->getNumOperands(); i != e; ++i)
6906 X = X.intrinsic(
6907 ID, {X, getRangeRef(NAry->getOperand(i), SignHint, Depth + 1)});
6908 return setRange(S, SignHint,
6909 ConservativeResult.intersectWith(X, RangeType));
6910 }
6911 case scUnknown: {
6912 const SCEVUnknown *U = cast<SCEVUnknown>(S);
6913 Value *V = U->getValue();
6914
6915 // Check if the IR explicitly contains !range metadata.
6916 std::optional<ConstantRange> MDRange = GetRangeFromMetadata(V);
6917 if (MDRange)
6918 ConservativeResult =
6919 ConservativeResult.intersectWith(*MDRange, RangeType);
6920
6921 // Use facts about recurrences in the underlying IR. Note that add
6922 // recurrences are AddRecExprs and thus don't hit this path. This
6923 // primarily handles shift recurrences.
6924 auto CR = getRangeForUnknownRecurrence(U);
6925 ConservativeResult = ConservativeResult.intersectWith(CR);
6926
6927 // See if ValueTracking can give us a useful range.
6928 const DataLayout &DL = getDataLayout();
6929 KnownBits Known = computeKnownBits(V, DL, &AC, nullptr, &DT);
6930 if (Known.getBitWidth() != BitWidth)
6931 Known = Known.zextOrTrunc(BitWidth);
6932
6933 // ValueTracking may be able to compute a tighter result for the number of
6934 // sign bits than for the value of those sign bits.
6935 unsigned NS = ComputeNumSignBits(V, DL, &AC, nullptr, &DT);
6936 if (U->getType()->isPointerTy()) {
6937 // NS counts the sign bits of the whole pointer; drop those above the
6938 // index bits.
6939 unsigned PtrIdxDiff =
6940 DL.getPointerTypeSizeInBits(U->getType()) - BitWidth;
6941 NS = NS > PtrIdxDiff ? NS - PtrIdxDiff : 1;
6942 }
6943
6944 if (NS > 1) {
6945 // If we know any of the sign bits, we know all of the sign bits.
6946 if (!Known.Zero.getHiBits(NS).isZero())
6947 Known.Zero.setHighBits(NS);
6948 if (!Known.One.getHiBits(NS).isZero())
6949 Known.One.setHighBits(NS);
6950 }
6951
6952 if (Known.getMinValue() != Known.getMaxValue() + 1)
6953 ConservativeResult = ConservativeResult.intersectWith(
6954 ConstantRange(Known.getMinValue(), Known.getMaxValue() + 1),
6955 RangeType);
6956 if (NS > 1)
6957 ConservativeResult = ConservativeResult.intersectWith(
6958 ConstantRange(APInt::getSignedMinValue(BitWidth).ashr(NS - 1),
6959 APInt::getSignedMaxValue(BitWidth).ashr(NS - 1) + 1),
6960 RangeType);
6961
6962 if (U->getType()->isPointerTy() && SignHint == HINT_RANGE_UNSIGNED) {
6963 // Strengthen the range if the underlying IR value is a
6964 // global/alloca/heap allocation using the size of the object.
6965 bool CanBeNull;
6966 uint64_t DerefBytes = V->getPointerDereferenceableBytes(
6967 DL, CanBeNull, /*CanBeFreed=*/nullptr);
6968 if (DerefBytes > 1 && isUIntN(BitWidth, DerefBytes)) {
6969 // The highest address the object can start is DerefBytes bytes before
6970 // the end (unsigned max value). If this value is not a multiple of the
6971 // alignment, the last possible start value is the next lowest multiple
6972 // of the alignment. Note: The computations below cannot overflow,
6973 // because if they would there's no possible start address for the
6974 // object.
6975 APInt MaxVal =
6976 APInt::getMaxValue(BitWidth) - APInt(BitWidth, DerefBytes);
6977 uint64_t Align = U->getValue()->getPointerAlignment(DL).value();
6978 uint64_t Rem = MaxVal.urem(Align);
6979 MaxVal -= APInt(BitWidth, Rem);
6980 APInt MinVal = APInt::getZero(BitWidth);
6981 if (llvm::isKnownNonZero(V, DL))
6982 MinVal = Align;
6983 ConservativeResult = ConservativeResult.intersectWith(
6984 ConstantRange::getNonEmpty(MinVal, MaxVal + 1), RangeType);
6985 }
6986 }
6987
6988 // A range of Phi is a subset of union of all ranges of its input.
6989 if (PHINode *Phi = dyn_cast<PHINode>(V)) {
6990 // SCEVExpander sometimes creates SCEVUnknowns that are secretly
6991 // AddRecs; return the range for the corresponding AddRec.
6992 if (auto *AR = dyn_cast<SCEVAddRecExpr>(getSCEV(V)))
6993 return getRangeRef(AR, SignHint, Depth + 1);
6994
6995 // Make sure that we do not run over cycled Phis.
6996 if (RangeRefPHIAllowedOperands(DT, Phi)) {
6997 ConstantRange RangeFromOps(BitWidth, /*isFullSet=*/false);
6998
6999 for (const auto &Op : Phi->operands()) {
7000 auto OpRange = getRangeRef(getSCEV(Op), SignHint, Depth + 1);
7001 RangeFromOps = RangeFromOps.unionWith(OpRange);
7002 // No point to continue if we already have a full set.
7003 if (RangeFromOps.isFullSet())
7004 break;
7005 }
7006 ConservativeResult =
7007 ConservativeResult.intersectWith(RangeFromOps, RangeType);
7008 }
7009 }
7010
7011 // vscale can't be equal to zero
7012 if (const auto *II = dyn_cast<IntrinsicInst>(V))
7013 if (II->getIntrinsicID() == Intrinsic::vscale) {
7014 ConstantRange Disallowed = APInt::getZero(BitWidth);
7015 ConservativeResult = ConservativeResult.difference(Disallowed);
7016 }
7017
7018 return setRange(U, SignHint, std::move(ConservativeResult));
7019 }
7020 case scCouldNotCompute:
7021 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
7022 }
7023
7024 return setRange(S, SignHint, std::move(ConservativeResult));
7025}
7026
7027// Given a StartRange, Step and MaxBECount for an expression compute a range of
7028// values that the expression can take. Initially, the expression has a value
7029// from StartRange and then is changed by Step up to MaxBECount times. Signed
7030// argument defines if we treat Step as signed or unsigned. The second return
7031// value indicates that no wrapping occurred.
7032static std::pair<ConstantRange, bool>
7034 const APInt &MaxBECount, bool Signed) {
7035 unsigned BitWidth = Step.getBitWidth();
7036 assert(BitWidth == StartRange.getBitWidth() &&
7037 BitWidth == MaxBECount.getBitWidth() && "mismatched bit widths");
7038 // If either Step or MaxBECount is 0, then the expression won't change, and we
7039 // just need to return the initial range.
7040 if (Step == 0 || MaxBECount == 0)
7041 return {StartRange, true};
7042
7043 // If we don't know anything about the initial value (i.e. StartRange is
7044 // FullRange), then we don't know anything about the final range either.
7045 // Return FullRange.
7046 if (StartRange.isFullSet())
7047 return {ConstantRange::getFull(BitWidth), false};
7048
7049 // If Step is signed and negative, then we use its absolute value, but we also
7050 // note that we're moving in the opposite direction.
7051 bool Descending = Signed && Step.isNegative();
7052
7053 if (Signed)
7054 // This is correct even for INT_SMIN. Let's look at i8 to illustrate this:
7055 // abs(INT_SMIN) = abs(-128) = abs(0x80) = -0x80 = 0x80 = 128.
7056 // This equations hold true due to the well-defined wrap-around behavior of
7057 // APInt.
7058 Step = Step.abs();
7059
7060 // Check if Offset is more than full span of BitWidth. If it is, the
7061 // expression is guaranteed to overflow.
7062 if (APInt::getMaxValue(StartRange.getBitWidth()).udiv(Step).ult(MaxBECount))
7063 return {ConstantRange::getFull(BitWidth), false};
7064
7065 // Offset is by how much the expression can change. Checks above guarantee no
7066 // overflow here.
7067 APInt Offset = Step * MaxBECount;
7068
7069 // Minimum value of the final range will match the minimal value of StartRange
7070 // if the expression is increasing and will be decreased by Offset otherwise.
7071 // Maximum value of the final range will match the maximal value of StartRange
7072 // if the expression is decreasing and will be increased by Offset otherwise.
7073 APInt StartLower = StartRange.getLower();
7074 APInt StartUpper = StartRange.getUpper() - 1;
7075 bool Overflow;
7076 APInt MovedBoundary;
7077 if (Signed) {
7078 // This does not use sadd_ov, as we want to check overflow for a signed
7079 // start with an unsigned offset.
7080 if (Descending) {
7081 MovedBoundary = StartLower - std::move(Offset);
7082 Overflow = MovedBoundary.sgt(StartLower) || StartRange.isSignWrappedSet();
7083 } else {
7084 MovedBoundary = StartUpper + std::move(Offset);
7085 Overflow = MovedBoundary.slt(StartUpper) || StartRange.isSignWrappedSet();
7086 }
7087 } else {
7088 MovedBoundary = StartUpper.uadd_ov(std::move(Offset), Overflow);
7089 Overflow |= StartRange.isWrappedSet();
7090 }
7091
7092 // It's possible that the new minimum/maximum value will fall into the initial
7093 // range (due to wrap around). This means that the expression can take any
7094 // value in this bitwidth, and we have to return full range.
7095 if (StartRange.contains(MovedBoundary))
7096 return {ConstantRange::getFull(BitWidth), false};
7097
7098 APInt NewLower =
7099 Descending ? std::move(MovedBoundary) : std::move(StartLower);
7100 APInt NewUpper =
7101 Descending ? std::move(StartUpper) : std::move(MovedBoundary);
7102 NewUpper += 1;
7103
7104 // No overflow detected, return [StartLower, StartUpper + Offset + 1) range.
7105 return {ConstantRange::getNonEmpty(std::move(NewLower), std::move(NewUpper)),
7106 !Overflow};
7107}
7108
7109std::pair<ConstantRange, SCEVFlags>
7110ScalarEvolution::getRangeForAffineAR(const SCEV *Start, const SCEV *Step,
7111 const APInt &MaxBECount) {
7112 assert(getTypeSizeInBits(Start->getType()) ==
7113 getTypeSizeInBits(Step->getType()) &&
7114 getTypeSizeInBits(Start->getType()) == MaxBECount.getBitWidth() &&
7115 "mismatched bit widths");
7116
7117 // First, consider step signed.
7118 ConstantRange StartSRange = getSignedRange(Start);
7119 ConstantRange StepSRange = getSignedRange(Step);
7120
7121 // If Step can be both positive and negative, we need to find ranges for the
7122 // maximum absolute step values in both directions and union them.
7123 auto [SR1, NSW1] = getRangeForAffineARHelper(
7124 StepSRange.getSignedMin(), StartSRange, MaxBECount, /*Signed=*/true);
7125 auto [SR2, NSW2] = getRangeForAffineARHelper(StepSRange.getSignedMax(),
7126 StartSRange, MaxBECount,
7127 /*Signed=*/true);
7128 ConstantRange SR = SR1.unionWith(SR2);
7129
7130 // Next, consider step unsigned.
7131 auto [UR, NUW] = getRangeForAffineARHelper(
7132 getUnsignedRangeMax(Step), getUnsignedRange(Start), MaxBECount,
7133 /*Signed=*/false);
7134
7136 if (NUW)
7138 if (NSW1 && NSW2)
7140
7141 // Finally, intersect signed and unsigned ranges.
7143}
7144
7145ConstantRange ScalarEvolution::getRangeForAffineNoSelfWrappingAR(
7146 const SCEVAddRecExpr *AddRec, const SCEV *MaxBECount, unsigned BitWidth,
7147 ScalarEvolution::RangeSignHint SignHint) {
7148 assert(AddRec->isAffine() && "Non-affine AddRecs are not suppored!\n");
7149 assert(AddRec->hasNoSelfWrap() &&
7150 "This only works for non-self-wrapping AddRecs!");
7151 const bool IsSigned = SignHint == HINT_RANGE_SIGNED;
7152 const SCEV *Step = AddRec->getStepRecurrence(*this);
7153 // Only deal with constant step to save compile time.
7154 if (!isa<SCEVConstant>(Step))
7155 return ConstantRange::getFull(BitWidth);
7156 // Let's make sure that we can prove that we do not self-wrap during
7157 // MaxBECount iterations. We need this because MaxBECount is a maximum
7158 // iteration count estimate, and we might infer nw from some exit for which we
7159 // do not know max exit count (or any other side reasoning).
7160 // TODO: Turn into assert at some point.
7161 if (getTypeSizeInBits(MaxBECount->getType()) >
7162 getTypeSizeInBits(AddRec->getType()))
7163 return ConstantRange::getFull(BitWidth);
7164 MaxBECount = getNoopOrZeroExtend(MaxBECount, AddRec->getType());
7165 const SCEV *RangeWidth = getMinusOne(AddRec->getType());
7166 const SCEV *StepAbs = getUMinExpr(Step, getNegativeSCEV(Step));
7167 const SCEV *MaxItersWithoutWrap = getUDivExpr(RangeWidth, StepAbs);
7168 if (!isKnownPredicateViaConstantRanges(ICmpInst::ICMP_ULE, MaxBECount,
7169 MaxItersWithoutWrap))
7170 return ConstantRange::getFull(BitWidth);
7171
7172 ICmpInst::Predicate LEPred =
7174 ICmpInst::Predicate GEPred =
7176 const SCEV *End = AddRec->evaluateAtIteration(MaxBECount, *this);
7177
7178 // We know that there is no self-wrap. Let's take Start and End values and
7179 // look at all intermediate values V1, V2, ..., Vn that IndVar takes during
7180 // the iteration. They either lie inside the range [Min(Start, End),
7181 // Max(Start, End)] or outside it:
7182 //
7183 // Case 1: RangeMin ... Start V1 ... VN End ... RangeMax;
7184 // Case 2: RangeMin Vk ... V1 Start ... End Vn ... Vk + 1 RangeMax;
7185 //
7186 // No self wrap flag guarantees that the intermediate values cannot be BOTH
7187 // outside and inside the range [Min(Start, End), Max(Start, End)]. Using that
7188 // knowledge, let's try to prove that we are dealing with Case 1. It is so if
7189 // Start <= End and step is positive, or Start >= End and step is negative.
7190 const SCEV *Start = applyLoopGuards(AddRec->getStart(), AddRec->getLoop());
7191 ConstantRange StartRange = getRangeRef(Start, SignHint);
7192 ConstantRange EndRange = getRangeRef(End, SignHint);
7193 ConstantRange RangeBetween = StartRange.unionWith(EndRange);
7194 // If they already cover full iteration space, we will know nothing useful
7195 // even if we prove what we want to prove.
7196 if (RangeBetween.isFullSet())
7197 return RangeBetween;
7198 // Only deal with ranges that do not wrap (i.e. RangeMin < RangeMax).
7199 bool IsWrappedSet = IsSigned ? RangeBetween.isSignWrappedSet()
7200 : RangeBetween.isWrappedSet();
7201 if (IsWrappedSet)
7202 return ConstantRange::getFull(BitWidth);
7203
7204 if (isKnownPositive(Step) &&
7205 isKnownPredicateViaConstantRanges(LEPred, Start, End))
7206 return RangeBetween;
7207 if (isKnownNegative(Step) &&
7208 isKnownPredicateViaConstantRanges(GEPred, Start, End))
7209 return RangeBetween;
7210 return ConstantRange::getFull(BitWidth);
7211}
7212
7213ConstantRange ScalarEvolution::getRangeViaFactoring(const SCEV *Start,
7214 const SCEV *Step,
7215 const APInt &MaxBECount) {
7216 // RangeOf({C?A:B,+,C?P:Q}) == RangeOf(C?{A,+,P}:{B,+,Q})
7217 // == RangeOf({A,+,P}) union RangeOf({B,+,Q})
7218
7219 unsigned BitWidth = MaxBECount.getBitWidth();
7220 assert(getTypeSizeInBits(Start->getType()) == BitWidth &&
7221 getTypeSizeInBits(Step->getType()) == BitWidth &&
7222 "mismatched bit widths");
7223
7224 struct SelectPattern {
7225 Value *Condition = nullptr;
7226 APInt TrueValue;
7227 APInt FalseValue;
7228
7229 explicit SelectPattern(ScalarEvolution &SE, unsigned BitWidth,
7230 const SCEV *S) {
7231 std::optional<unsigned> CastOp;
7232 APInt Offset(BitWidth, 0);
7233
7235 "Should be!");
7236
7237 // Peel off a constant offset. In the future we could consider being
7238 // smarter here and handle {Start+Step,+,Step} too.
7239 const APInt *Off;
7240 if (match(S, m_scev_Add(m_scev_APInt(Off), m_SCEV(S))))
7241 Offset = *Off;
7242
7243 // Peel off a cast operation
7244 if (auto *SCast = dyn_cast<SCEVIntegralCastExpr>(S)) {
7245 CastOp = SCast->getSCEVType();
7246 S = SCast->getOperand();
7247 }
7248
7249 using namespace llvm::PatternMatch;
7250
7251 auto *SU = dyn_cast<SCEVUnknown>(S);
7252 const APInt *TrueVal, *FalseVal;
7253 if (!SU ||
7254 !match(SU->getValue(), m_Select(m_Value(Condition), m_APInt(TrueVal),
7255 m_APInt(FalseVal)))) {
7256 Condition = nullptr;
7257 return;
7258 }
7259
7260 TrueValue = *TrueVal;
7261 FalseValue = *FalseVal;
7262
7263 // Re-apply the cast we peeled off earlier
7264 if (CastOp)
7265 switch (*CastOp) {
7266 default:
7267 llvm_unreachable("Unknown SCEV cast type!");
7268
7269 case scTruncate:
7270 TrueValue = TrueValue.trunc(BitWidth);
7271 FalseValue = FalseValue.trunc(BitWidth);
7272 break;
7273 case scZeroExtend:
7274 TrueValue = TrueValue.zext(BitWidth);
7275 FalseValue = FalseValue.zext(BitWidth);
7276 break;
7277 case scSignExtend:
7278 TrueValue = TrueValue.sext(BitWidth);
7279 FalseValue = FalseValue.sext(BitWidth);
7280 break;
7281 }
7282
7283 // Re-apply the constant offset we peeled off earlier
7284 TrueValue += Offset;
7285 FalseValue += Offset;
7286 }
7287
7288 bool isRecognized() { return Condition != nullptr; }
7289 };
7290
7291 SelectPattern StartPattern(*this, BitWidth, Start);
7292 if (!StartPattern.isRecognized())
7293 return ConstantRange::getFull(BitWidth);
7294
7295 SelectPattern StepPattern(*this, BitWidth, Step);
7296 if (!StepPattern.isRecognized())
7297 return ConstantRange::getFull(BitWidth);
7298
7299 if (StartPattern.Condition != StepPattern.Condition) {
7300 // We don't handle this case today; but we could, by considering four
7301 // possibilities below instead of two. I'm not sure if there are cases where
7302 // that will help over what getRange already does, though.
7303 return ConstantRange::getFull(BitWidth);
7304 }
7305
7306 // NB! Calling ScalarEvolution::getConstant is fine, but we should not try to
7307 // construct arbitrary general SCEV expressions here. This function is called
7308 // from deep in the call stack, and calling getSCEV (on a sext instruction,
7309 // say) can end up caching a suboptimal value.
7310
7311 // FIXME: without the explicit `this` receiver below, MSVC errors out with
7312 // C2352 and C2512 (otherwise it isn't needed).
7313
7314 const SCEV *TrueStart = this->getConstant(StartPattern.TrueValue);
7315 const SCEV *TrueStep = this->getConstant(StepPattern.TrueValue);
7316 const SCEV *FalseStart = this->getConstant(StartPattern.FalseValue);
7317 const SCEV *FalseStep = this->getConstant(StepPattern.FalseValue);
7318
7319 ConstantRange TrueRange =
7320 this->getRangeForAffineAR(TrueStart, TrueStep, MaxBECount).first;
7321 ConstantRange FalseRange =
7322 this->getRangeForAffineAR(FalseStart, FalseStep, MaxBECount).first;
7323
7324 return TrueRange.unionWith(FalseRange);
7325}
7326
7327SCEVFlags ScalarEvolution::getNoWrapFlagsFromUB(const Value *V) {
7328 if (isa<ConstantExpr>(V))
7329 return SCEV::FlagNone;
7330 const BinaryOperator *BinOp = cast<BinaryOperator>(V);
7331
7332 // Return early if there are no flags to propagate to the SCEV.
7334 if (auto *PDI = dyn_cast<PossiblyDisjointInst>(BinOp);
7335 PDI && PDI->isDisjoint()) {
7337 } else {
7338 if (BinOp->hasNoUnsignedWrap())
7340 if (BinOp->hasNoSignedWrap())
7342 }
7343 if (Flags == SCEV::FlagNone)
7344 return SCEV::FlagNone;
7345
7346 return isSCEVExprNeverPoison(BinOp) ? Flags : SCEV::FlagNone;
7347}
7348
7349const Instruction *
7350ScalarEvolution::getNonTrivialDefiningScopeBound(const SCEV *S) {
7351 if (auto *AddRec = dyn_cast<SCEVAddRecExpr>(S))
7352 return &*AddRec->getLoop()->getHeader()->begin();
7353 if (auto *U = dyn_cast<SCEVUnknown>(S))
7354 if (auto *I = dyn_cast<Instruction>(U->getValue()))
7355 return I;
7356 return nullptr;
7357}
7358
7359const Instruction *ScalarEvolution::getDefiningScopeBound(ArrayRef<SCEVUse> Ops,
7360 bool &Precise) {
7361 Precise = true;
7362 // Do a bounded search of the def relation of the requested SCEVs.
7363 SmallPtrSet<const SCEV *, 16> Visited;
7364 SmallVector<SCEVUse> Worklist;
7365 auto pushOp = [&](const SCEV *S) {
7366 if (!Visited.insert(S).second)
7367 return;
7368 // Threshold of 30 here is arbitrary.
7369 if (Visited.size() > 30) {
7370 Precise = false;
7371 return;
7372 }
7373 Worklist.push_back(S);
7374 };
7375
7376 for (SCEVUse S : Ops)
7377 pushOp(S);
7378
7379 const Instruction *Bound = nullptr;
7380 while (!Worklist.empty()) {
7381 SCEVUse S = Worklist.pop_back_val();
7382 if (auto *DefI = getNonTrivialDefiningScopeBound(S)) {
7383 if (!Bound || DT.dominates(Bound, DefI))
7384 Bound = DefI;
7385 } else {
7386 for (SCEVUse Op : S->operands())
7387 pushOp(Op);
7388 }
7389 }
7390 return Bound ? Bound : &*F.getEntryBlock().begin();
7391}
7392
7393const Instruction *
7394ScalarEvolution::getDefiningScopeBound(ArrayRef<SCEVUse> Ops) {
7395 bool Discard;
7396 return getDefiningScopeBound(Ops, Discard);
7397}
7398
7399bool ScalarEvolution::isGuaranteedToTransferExecutionTo(const Instruction *A,
7400 const Instruction *B) {
7401 if (A->getParent() == B->getParent() &&
7403 B->getIterator()))
7404 return true;
7405
7406 auto *BLoop = LI.getLoopFor(B->getParent());
7407 if (BLoop && BLoop->getHeader() == B->getParent() &&
7408 BLoop->getLoopPreheader() == A->getParent() &&
7410 A->getParent()->end()) &&
7411 isGuaranteedToTransferExecutionToSuccessor(B->getParent()->begin(),
7412 B->getIterator()))
7413 return true;
7414 return false;
7415}
7416
7418 SCEVPoisonCollector PC(/* LookThroughMaybePoisonBlocking */ true);
7419 visitAll(Op, PC);
7420 return PC.MaybePoison.empty();
7421}
7422
7423bool ScalarEvolution::isGuaranteedNotToCauseUB(const SCEV *Op) {
7424 return !SCEVExprContains(Op, [this](const SCEV *S) {
7425 const SCEV *Op1;
7426 bool M = match(S, m_scev_UDiv(m_SCEV(), m_SCEV(Op1)));
7427 // The UDiv may be UB if the divisor is poison or zero. Unless the divisor
7428 // is a non-zero constant, we have to assume the UDiv may be UB.
7429 return M && (!isKnownNonZero(Op1) || !isGuaranteedNotToBePoison(Op1));
7430 });
7431}
7432
7433bool ScalarEvolution::isSCEVExprNeverPoison(const Instruction *I) {
7434 // Only proceed if we can prove that I does not yield poison.
7436 return false;
7437
7438 // At this point we know that if I is executed, then it does not wrap
7439 // according to at least one of NSW or NUW. If I is not executed, then we do
7440 // not know if the calculation that I represents would wrap. Multiple
7441 // instructions can map to the same SCEV. If we apply NSW or NUW from I to
7442 // the SCEV, we must guarantee no wrapping for that SCEV also when it is
7443 // derived from other instructions that map to the same SCEV. We cannot make
7444 // that guarantee for cases where I is not executed. So we need to find a
7445 // upper bound on the defining scope for the SCEV, and prove that I is
7446 // executed every time we enter that scope. When the bounding scope is a
7447 // loop (the common case), this is equivalent to proving I executes on every
7448 // iteration of that loop.
7449 SmallVector<SCEVUse> SCEVOps;
7450 for (const Use &Op : I->operands()) {
7451 // I could be an extractvalue from a call to an overflow intrinsic.
7452 // TODO: We can do better here in some cases.
7453 if (isSCEVable(Op->getType()))
7454 SCEVOps.push_back(getSCEV(Op));
7455 }
7456 auto *DefI = getDefiningScopeBound(SCEVOps);
7457 return isGuaranteedToTransferExecutionTo(DefI, I);
7458}
7459
7460bool ScalarEvolution::isAddRecNeverPoison(const Instruction *I, const Loop *L) {
7461 // If we know that \c I can never be poison period, then that's enough.
7462 if (isSCEVExprNeverPoison(I))
7463 return true;
7464
7465 // If the loop only has one exit, then we know that, if the loop is entered,
7466 // any instruction dominating that exit will be executed. If any such
7467 // instruction would result in UB, the addrec cannot be poison.
7468 //
7469 // This is basically the same reasoning as in isSCEVExprNeverPoison(), but
7470 // also handles uses outside the loop header (they just need to dominate the
7471 // single exit).
7472
7473 auto *ExitingBB = L->getExitingBlock();
7474 if (!ExitingBB || !loopHasNoAbnormalExits(L))
7475 return false;
7476
7477 SmallPtrSet<const Value *, 16> KnownPoison;
7479
7480 // We start by assuming \c I, the post-inc add recurrence, is poison. Only
7481 // things that are known to be poison under that assumption go on the
7482 // Worklist.
7483 KnownPoison.insert(I);
7484 Worklist.push_back(I);
7485
7486 while (!Worklist.empty()) {
7487 const Instruction *Poison = Worklist.pop_back_val();
7488
7489 for (const Use &U : Poison->uses()) {
7490 const Instruction *PoisonUser = cast<Instruction>(U.getUser());
7491 if (mustTriggerUB(PoisonUser, KnownPoison) &&
7492 DT.dominates(PoisonUser->getParent(), ExitingBB))
7493 return true;
7494
7495 if (propagatesPoison(U) && L->contains(PoisonUser))
7496 if (KnownPoison.insert(PoisonUser).second)
7497 Worklist.push_back(PoisonUser);
7498 }
7499 }
7500
7501 return false;
7502}
7503
7504ScalarEvolution::LoopProperties
7505ScalarEvolution::getLoopProperties(const Loop *L) {
7506 using LoopProperties = ScalarEvolution::LoopProperties;
7507
7508 auto Itr = LoopPropertiesCache.find(L);
7509 if (Itr == LoopPropertiesCache.end()) {
7510 auto HasSideEffects = [](Instruction *I) {
7511 if (auto *SI = dyn_cast<StoreInst>(I))
7512 return !SI->isSimple();
7513
7514 if (I->mayThrow())
7515 return true;
7516
7517 // Non-volatile memset / memcpy do not count as side-effect for forward
7518 // progress.
7519 if (isa<MemIntrinsic>(I) && !I->isVolatile())
7520 return false;
7521
7522 return I->mayWriteToMemory();
7523 };
7524
7525 LoopProperties LP = {/* HasNoAbnormalExits */ true,
7526 /*HasNoSideEffects*/ true};
7527
7528 for (auto *BB : L->getBlocks())
7529 for (auto &I : *BB) {
7531 LP.HasNoAbnormalExits = false;
7532 if (HasSideEffects(&I))
7533 LP.HasNoSideEffects = false;
7534 if (!LP.HasNoAbnormalExits && !LP.HasNoSideEffects)
7535 break; // We're already as pessimistic as we can get.
7536 }
7537
7538 auto InsertPair = LoopPropertiesCache.insert({L, LP});
7539 assert(InsertPair.second && "We just checked!");
7540 Itr = InsertPair.first;
7541 }
7542
7543 return Itr->second;
7544}
7545
7547 // A mustprogress loop without side effects must be finite.
7548 // TODO: The check used here is very conservative. It's only *specific*
7549 // side effects which are well defined in infinite loops.
7550 return isFinite(L) || (isMustProgress(L) && loopHasNoSideEffects(L));
7551}
7552
7553const SCEV *ScalarEvolution::createSCEVIter(Value *V) {
7554 // Worklist item with a Value and a bool indicating whether all operands have
7555 // been visited already.
7558
7559 Stack.emplace_back(V, false);
7560 while (!Stack.empty()) {
7561 auto E = Stack.back();
7562 Value *CurV = E.getPointer();
7563
7564 if (getExistingSCEV(CurV)) {
7565 Stack.pop_back();
7566 continue;
7567 }
7568
7570 const SCEV *CreatedSCEV = nullptr;
7571 // If all operands have been visited already, create the SCEV.
7572 if (E.getInt()) {
7573 CreatedSCEV = createSCEV(CurV);
7574 } else {
7575 // Otherwise get the operands we need to create SCEV's for before creating
7576 // the SCEV for CurV. If the SCEV for CurV can be constructed trivially,
7577 // just use it.
7578 CreatedSCEV = getOperandsToCreate(CurV, Ops);
7579 }
7580
7581 if (CreatedSCEV) {
7582 insertValueToMap(CurV, CreatedSCEV);
7583 Stack.pop_back();
7584 } else {
7585 Stack.back().setInt(true);
7586 // Queue its operands which need to be constructed.
7587 for (Value *Op : Ops)
7588 Stack.emplace_back(Op, false);
7589 }
7590 }
7591
7592 return getExistingSCEV(V);
7593}
7594
7595const SCEV *
7596ScalarEvolution::getOperandsToCreate(Value *V, SmallVectorImpl<Value *> &Ops) {
7597 if (!isSCEVable(V->getType()))
7598 return getUnknown(V);
7599
7600 if (Instruction *I = dyn_cast<Instruction>(V)) {
7601 // Don't attempt to analyze instructions in blocks that aren't
7602 // reachable. Such instructions don't matter, and they aren't required
7603 // to obey basic rules for definitions dominating uses which this
7604 // analysis depends on.
7605 if (!DT.isReachableFromEntry(I->getParent()))
7606 return getUnknown(PoisonValue::get(V->getType()));
7607 } else if (ConstantInt *CI = dyn_cast<ConstantInt>(V))
7608 return getConstant(CI);
7609 else if (isa<GlobalAlias>(V))
7610 return getUnknown(V);
7611 else if (!isa<ConstantExpr>(V))
7612 return getUnknown(V);
7613
7615 if (auto BO =
7617 bool IsConstArg = isa<ConstantInt>(BO->RHS);
7618 switch (BO->Opcode) {
7619 case Instruction::Add:
7620 case Instruction::Mul: {
7621 // For additions and multiplications, traverse add/mul chains for which we
7622 // can potentially create a single SCEV, to reduce the number of
7623 // get{Add,Mul}Expr calls.
7624 do {
7625 if (BO->Op) {
7626 if (BO->Op != V && getExistingSCEV(BO->Op)) {
7627 Ops.push_back(BO->Op);
7628 break;
7629 }
7630 }
7631 Ops.push_back(BO->RHS);
7632 auto NewBO = MatchBinaryOp(BO->LHS, getDataLayout(), AC, DT,
7634 if (!NewBO ||
7635 (BO->Opcode == Instruction::Add &&
7636 (NewBO->Opcode != Instruction::Add &&
7637 NewBO->Opcode != Instruction::Sub)) ||
7638 (BO->Opcode == Instruction::Mul &&
7639 NewBO->Opcode != Instruction::Mul)) {
7640 Ops.push_back(BO->LHS);
7641 break;
7642 }
7643 // CreateSCEV calls getNoWrapFlagsFromUB, which under certain conditions
7644 // requires a SCEV for the LHS.
7645 if (BO->Op && (BO->IsNSW || BO->IsNUW)) {
7646 auto *I = dyn_cast<Instruction>(BO->Op);
7647 if (I && programUndefinedIfPoison(I)) {
7648 Ops.push_back(BO->LHS);
7649 break;
7650 }
7651 }
7652 BO = NewBO;
7653 } while (true);
7654 return nullptr;
7655 }
7656 case Instruction::Sub:
7657 case Instruction::UDiv:
7658 case Instruction::URem:
7659 break;
7660 case Instruction::AShr:
7661 case Instruction::Shl:
7662 case Instruction::Xor:
7663 if (!IsConstArg)
7664 return nullptr;
7665 break;
7666 case Instruction::And:
7667 case Instruction::Or:
7668 if (!IsConstArg && !BO->LHS->getType()->isIntegerTy(1))
7669 return nullptr;
7670 break;
7671 case Instruction::LShr:
7672 return getUnknown(V);
7673 default:
7674 llvm_unreachable("Unhandled binop");
7675 break;
7676 }
7677
7678 Ops.push_back(BO->LHS);
7679 Ops.push_back(BO->RHS);
7680 return nullptr;
7681 }
7682
7683 switch (U->getOpcode()) {
7684 case Instruction::Trunc:
7685 case Instruction::ZExt:
7686 case Instruction::SExt:
7687 case Instruction::PtrToAddr:
7688 case Instruction::PtrToInt:
7689 Ops.push_back(U->getOperand(0));
7690 return nullptr;
7691
7692 case Instruction::BitCast:
7693 if (isSCEVable(U->getType()) && isSCEVable(U->getOperand(0)->getType())) {
7694 Ops.push_back(U->getOperand(0));
7695 return nullptr;
7696 }
7697 return getUnknown(V);
7698
7699 case Instruction::SDiv:
7700 case Instruction::SRem:
7701 Ops.push_back(U->getOperand(0));
7702 Ops.push_back(U->getOperand(1));
7703 return nullptr;
7704
7705 case Instruction::GetElementPtr:
7706 assert(cast<GEPOperator>(U)->getSourceElementType()->isSized() &&
7707 "GEP source element type must be sized");
7708 llvm::append_range(Ops, U->operands());
7709 return nullptr;
7710
7711 case Instruction::IntToPtr:
7712 return getUnknown(V);
7713
7714 case Instruction::PHI:
7715 // getNodeForPHI has four ways to turn a PHI into a SCEV; retrieve the
7716 // relevant nodes for each of them.
7717 //
7718 // The first is just to call simplifyInstruction, and get something back
7719 // that isn't a PHI.
7720 if (Value *V = simplifyInstruction(
7721 cast<PHINode>(U),
7722 {getDataLayout(), &TLI, &DT, &AC, /*CtxI=*/nullptr,
7723 /*UseInstrInfo=*/true, /*CanUseUndef=*/false})) {
7724 assert(V);
7725 Ops.push_back(V);
7726 return nullptr;
7727 }
7728 // The second is createNodeForPHIWithIdenticalOperands: this looks for
7729 // operands which all perform the same operation, but haven't been
7730 // CSE'ed for whatever reason.
7731 if (BinaryOperator *BO = getCommonInstForPHI(cast<PHINode>(U))) {
7732 assert(BO);
7733 Ops.push_back(BO);
7734 return nullptr;
7735 }
7736 // The third is createNodeFromSelectLikePHI; this takes a PHI which
7737 // is equivalent to a select, and analyzes it like a select.
7738 {
7739 Value *Cond = nullptr, *LHS = nullptr, *RHS = nullptr;
7741 assert(Cond);
7742 assert(LHS);
7743 assert(RHS);
7744 if (auto *CondICmp = dyn_cast<ICmpInst>(Cond)) {
7745 Ops.push_back(CondICmp->getOperand(0));
7746 Ops.push_back(CondICmp->getOperand(1));
7747 }
7748 Ops.push_back(Cond);
7749 Ops.push_back(LHS);
7750 Ops.push_back(RHS);
7751 return nullptr;
7752 }
7753 }
7754 // The fourth way is createAddRecFromPHI. It's complicated to handle here,
7755 // so just construct it recursively.
7756 //
7757 // In addition to getNodeForPHI, also construct nodes which might be needed
7758 // by getRangeRef.
7760 for (Value *V : cast<PHINode>(U)->operands())
7761 Ops.push_back(V);
7762 return nullptr;
7763 }
7764 return nullptr;
7765
7766 case Instruction::Select: {
7767 // Check if U is a select that can be simplified to a SCEVUnknown.
7768 auto CanSimplifyToUnknown = [this, U]() {
7769 if (U->getType()->isIntegerTy(1) || isa<ConstantInt>(U->getOperand(0)))
7770 return false;
7771
7772 auto *ICI = dyn_cast<ICmpInst>(U->getOperand(0));
7773 if (!ICI)
7774 return false;
7775 Value *LHS = ICI->getOperand(0);
7776 Value *RHS = ICI->getOperand(1);
7777 if (ICI->getPredicate() == CmpInst::ICMP_EQ ||
7778 ICI->getPredicate() == CmpInst::ICMP_NE) {
7780 return true;
7781 } else if (getTypeSizeInBits(LHS->getType()) >
7782 getTypeSizeInBits(U->getType()))
7783 return true;
7784 return false;
7785 };
7786 if (CanSimplifyToUnknown())
7787 return getUnknown(U);
7788
7789 llvm::append_range(Ops, U->operands());
7790 return nullptr;
7791 break;
7792 }
7793 case Instruction::Call:
7794 case Instruction::Invoke:
7795 if (Value *RV = cast<CallBase>(U)->getReturnedArgOperand()) {
7796 Ops.push_back(RV);
7797 return nullptr;
7798 }
7799
7800 if (auto *II = dyn_cast<IntrinsicInst>(U)) {
7801 switch (II->getIntrinsicID()) {
7802 case Intrinsic::abs:
7803 Ops.push_back(II->getArgOperand(0));
7804 return nullptr;
7805 case Intrinsic::umax:
7806 case Intrinsic::umin:
7807 case Intrinsic::smax:
7808 case Intrinsic::smin:
7809 case Intrinsic::usub_sat:
7810 case Intrinsic::uadd_sat:
7811 Ops.push_back(II->getArgOperand(0));
7812 Ops.push_back(II->getArgOperand(1));
7813 return nullptr;
7814 case Intrinsic::start_loop_iterations:
7815 case Intrinsic::annotation:
7816 case Intrinsic::ptr_annotation:
7817 Ops.push_back(II->getArgOperand(0));
7818 return nullptr;
7819 default:
7820 break;
7821 }
7822 }
7823 break;
7824 }
7825
7826 return nullptr;
7827}
7828
7829const SCEV *ScalarEvolution::createSCEV(Value *V) {
7830 if (!isSCEVable(V->getType()))
7831 return getUnknown(V);
7832
7833 if (Instruction *I = dyn_cast<Instruction>(V)) {
7834 // Don't attempt to analyze instructions in blocks that aren't
7835 // reachable. Such instructions don't matter, and they aren't required
7836 // to obey basic rules for definitions dominating uses which this
7837 // analysis depends on.
7838 if (!DT.isReachableFromEntry(I->getParent()))
7839 return getUnknown(PoisonValue::get(V->getType()));
7840 } else if (ConstantInt *CI = dyn_cast<ConstantInt>(V))
7841 return getConstant(CI);
7842 else if (isa<GlobalAlias>(V))
7843 return getUnknown(V);
7844 else if (!isa<ConstantExpr>(V))
7845 return getUnknown(V);
7846
7847 const SCEV *LHS;
7848 const SCEV *RHS;
7849
7851 if (auto BO =
7853 switch (BO->Opcode) {
7854 case Instruction::Add: {
7855 // The simple thing to do would be to just call getSCEV on both operands
7856 // and call getAddExpr with the result. However if we're looking at a
7857 // bunch of things all added together, this can be quite inefficient,
7858 // because it leads to N-1 getAddExpr calls for N ultimate operands.
7859 // Instead, gather up all the operands and make a single getAddExpr call.
7860 // LLVM IR canonical form means we need only traverse the left operands.
7862 do {
7863 if (BO->Op) {
7864 if (auto *OpSCEV = getExistingSCEV(BO->Op)) {
7865 AddOps.push_back(OpSCEV);
7866 break;
7867 }
7868
7869 // If a NUW or NSW flag can be applied to the SCEV for this
7870 // addition, then compute the SCEV for this addition by itself
7871 // with a separate call to getAddExpr. We need to do that
7872 // instead of pushing the operands of the addition onto AddOps,
7873 // since the flags are only known to apply to this particular
7874 // addition - they may not apply to other additions that can be
7875 // formed with operands from AddOps.
7876 const SCEV *RHS = getSCEV(BO->RHS);
7877 SCEVFlags Flags = getNoWrapFlagsFromUB(BO->Op);
7878 if (Flags != SCEV::FlagNone) {
7879 const SCEV *LHS = getSCEV(BO->LHS);
7880 if (BO->Opcode == Instruction::Sub)
7881 AddOps.push_back(getMinusSCEV(LHS, RHS, Flags));
7882 else
7883 AddOps.push_back(getAddExpr(LHS, RHS, Flags));
7884 break;
7885 }
7886 }
7887
7888 if (BO->Opcode == Instruction::Sub)
7889 AddOps.push_back(getNegativeSCEV(getSCEV(BO->RHS)));
7890 else
7891 AddOps.push_back(getSCEV(BO->RHS));
7892
7893 auto NewBO = MatchBinaryOp(BO->LHS, getDataLayout(), AC, DT,
7895 if (!NewBO || (NewBO->Opcode != Instruction::Add &&
7896 NewBO->Opcode != Instruction::Sub)) {
7897 AddOps.push_back(getSCEV(BO->LHS));
7898 break;
7899 }
7900 BO = NewBO;
7901 } while (true);
7902
7903 return getAddExpr(AddOps);
7904 }
7905
7906 case Instruction::Mul: {
7908 do {
7909 if (BO->Op) {
7910 if (auto *OpSCEV = getExistingSCEV(BO->Op)) {
7911 MulOps.push_back(OpSCEV);
7912 break;
7913 }
7914
7915 SCEVFlags Flags = getNoWrapFlagsFromUB(BO->Op);
7916 if (Flags != SCEV::FlagNone) {
7917 LHS = getSCEV(BO->LHS);
7918 RHS = getSCEV(BO->RHS);
7919 MulOps.push_back(getMulExpr(LHS, RHS, Flags));
7920 break;
7921 }
7922 }
7923
7924 MulOps.push_back(getSCEV(BO->RHS));
7925 auto NewBO = MatchBinaryOp(BO->LHS, getDataLayout(), AC, DT,
7927 if (!NewBO || NewBO->Opcode != Instruction::Mul) {
7928 MulOps.push_back(getSCEV(BO->LHS));
7929 break;
7930 }
7931 BO = NewBO;
7932 } while (true);
7933
7934 return getMulExpr(MulOps);
7935 }
7936 case Instruction::UDiv:
7937 LHS = getSCEV(BO->LHS);
7938 RHS = getSCEV(BO->RHS);
7939 return getUDivExpr(LHS, RHS);
7940 case Instruction::URem:
7941 LHS = getSCEV(BO->LHS);
7942 RHS = getSCEV(BO->RHS);
7943 return getURemExpr(LHS, RHS);
7944 case Instruction::Sub: {
7946 if (BO->Op)
7947 Flags = getNoWrapFlagsFromUB(BO->Op);
7948
7949 // Try to use ptrtoaddr for subtracts with at least one ptrtoint
7950 // operand. While we don't model ptrtoint directly in SCEV, the
7951 // difference between two pointer addresses is well-defined.
7952 Value *PtrLHS = nullptr, *PtrRHS = nullptr;
7953 bool HasPtrLHS = match(BO->LHS, m_PtrToInt(m_Value(PtrLHS)));
7954 bool HasPtrRHS = match(BO->RHS, m_PtrToInt(m_Value(PtrRHS)));
7955 if (HasPtrLHS || HasPtrRHS) {
7956 // Convert a ptrtoint operand (OrigOp) to ptrtoaddr of its pointer
7957 // PtrOp. When only one side is ptrtoint (BothPtr is false), skip
7958 // SCEVUnknown pointers since wrapping them in ptrtoaddr adds no
7959 // useful structure.
7960 auto GetOp = [&](bool HasPtr, Value *PtrOp, Value *OrigOp,
7961 bool BothPtr) -> const SCEV * {
7962 if (!HasPtr)
7963 return getSCEV(OrigOp);
7964 const SCEV *PtrSCEV = getSCEV(PtrOp);
7965 if (BothPtr || !isa<SCEVUnknown>(PtrSCEV)) {
7966 const SCEV *Addr = getPtrToAddrExpr(PtrSCEV);
7967 if (!isa<SCEVCouldNotCompute>(Addr) &&
7968 getTypeSizeInBits(OrigOp->getType()) <=
7969 getTypeSizeInBits(Addr->getType()))
7970 return getTruncateOrNoop(Addr, OrigOp->getType());
7971 }
7972 return getSCEV(OrigOp);
7973 };
7974 const SCEV *L = GetOp(HasPtrLHS, PtrLHS, BO->LHS, HasPtrRHS);
7975 const SCEV *R = GetOp(HasPtrRHS, PtrRHS, BO->RHS, HasPtrLHS);
7976 return getMinusSCEV(L, R, Flags);
7977 }
7978
7979 LHS = getSCEV(BO->LHS);
7980 RHS = getSCEV(BO->RHS);
7981 return getMinusSCEV(LHS, RHS, Flags);
7982 }
7983 case Instruction::And:
7984 // For an expression like x&255 that merely masks off the high bits,
7985 // use zext(trunc(x)) as the SCEV expression.
7986 if (ConstantInt *CI = dyn_cast<ConstantInt>(BO->RHS)) {
7987 if (CI->isZero())
7988 return getSCEV(BO->RHS);
7989 if (CI->isMinusOne())
7990 return getSCEV(BO->LHS);
7991 const APInt &A = CI->getValue();
7992
7993 // Instcombine's ShrinkDemandedConstant may strip bits out of
7994 // constants, obscuring what would otherwise be a low-bits mask.
7995 // Use computeKnownBits to compute what ShrinkDemandedConstant
7996 // knew about to reconstruct a low-bits mask value.
7997 unsigned LZ = A.countl_zero();
7998 unsigned TZ = A.countr_zero();
7999 unsigned BitWidth = A.getBitWidth();
8000 KnownBits Known(BitWidth);
8001 computeKnownBits(BO->LHS, Known, getDataLayout(), &AC, nullptr, &DT);
8002
8003 APInt EffectiveMask =
8004 APInt::getLowBitsSet(BitWidth, BitWidth - LZ - TZ).shl(TZ);
8005 if ((LZ != 0 || TZ != 0) && !((~A & ~Known.Zero) & EffectiveMask)) {
8006 const SCEV *MulCount = getConstant(APInt::getOneBitSet(BitWidth, TZ));
8007 const SCEV *LHS = getSCEV(BO->LHS);
8008 const SCEV *ShiftedLHS = nullptr;
8009 if (auto *LHSMul = dyn_cast<SCEVMulExpr>(LHS)) {
8010 if (auto *OpC = dyn_cast<SCEVConstant>(LHSMul->getOperand(0))) {
8011 // For an expression like (x * 8) & 8, simplify the multiply.
8012 unsigned MulZeros = OpC->getAPInt().countr_zero();
8013 unsigned GCD = std::min(MulZeros, TZ);
8014 APInt DivAmt = APInt::getOneBitSet(BitWidth, TZ - GCD);
8016 MulOps.push_back(getConstant(OpC->getAPInt().ashr(GCD)));
8017 append_range(MulOps, LHSMul->operands().drop_front());
8018 const SCEV *NewMul = getMulExpr(MulOps, LHSMul->getNoWrapFlags());
8019 ShiftedLHS = getUDivExpr(NewMul, getConstant(DivAmt));
8020 }
8021 }
8022 if (!ShiftedLHS)
8023 ShiftedLHS = getUDivExpr(LHS, MulCount);
8024 return getMulExpr(
8026 getTruncateExpr(ShiftedLHS,
8027 IntegerType::get(getContext(), BitWidth - LZ - TZ)),
8028 BO->LHS->getType()),
8029 MulCount);
8030 }
8031 }
8032 // Binary `and` is a bit-wise `umin`.
8033 if (BO->LHS->getType()->isIntegerTy(1)) {
8034 LHS = getSCEV(BO->LHS);
8035 RHS = getSCEV(BO->RHS);
8036 return getUMinExpr(LHS, RHS);
8037 }
8038 break;
8039
8040 case Instruction::Or:
8041 // Binary `or` is a bit-wise `umax`.
8042 if (BO->LHS->getType()->isIntegerTy(1)) {
8043 LHS = getSCEV(BO->LHS);
8044 RHS = getSCEV(BO->RHS);
8045 return getUMaxExpr(LHS, RHS);
8046 }
8047 break;
8048
8049 case Instruction::Xor:
8050 if (ConstantInt *CI = dyn_cast<ConstantInt>(BO->RHS)) {
8051 // If the RHS of xor is -1, then this is a not operation.
8052 if (CI->isMinusOne())
8053 return getNotSCEV(getSCEV(BO->LHS));
8054
8055 // Model xor(and(x, C), C) as and(~x, C), if C is a low-bits mask.
8056 // This is a variant of the check for xor with -1, and it handles
8057 // the case where instcombine has trimmed non-demanded bits out
8058 // of an xor with -1.
8059 if (auto *LBO = dyn_cast<BinaryOperator>(BO->LHS))
8060 if (ConstantInt *LCI = dyn_cast<ConstantInt>(LBO->getOperand(1)))
8061 if (LBO->getOpcode() == Instruction::And &&
8062 LCI->getValue() == CI->getValue())
8063 if (const SCEVZeroExtendExpr *Z =
8065 Type *UTy = BO->LHS->getType();
8066 const SCEV *Z0 = Z->getOperand();
8067 Type *Z0Ty = Z0->getType();
8068 unsigned Z0TySize = getTypeSizeInBits(Z0Ty);
8069
8070 // If C is a low-bits mask, the zero extend is serving to
8071 // mask off the high bits. Complement the operand and
8072 // re-apply the zext.
8073 if (CI->getValue().isMask(Z0TySize))
8074 return getZeroExtendExpr(getNotSCEV(Z0), UTy);
8075
8076 // If C is a single bit, it may be in the sign-bit position
8077 // before the zero-extend. In this case, represent the xor
8078 // using an add, which is equivalent, and re-apply the zext.
8079 APInt Trunc = CI->getValue().trunc(Z0TySize);
8080 if (Trunc.zext(getTypeSizeInBits(UTy)) == CI->getValue() &&
8081 Trunc.isSignMask())
8082 return getZeroExtendExpr(getAddExpr(Z0, getConstant(Trunc)),
8083 UTy);
8084 }
8085 }
8086 break;
8087
8088 case Instruction::Shl:
8089 // Turn shift left of a constant amount into a multiply.
8090 if (ConstantInt *SA = dyn_cast<ConstantInt>(BO->RHS)) {
8091 uint32_t BitWidth = cast<IntegerType>(SA->getType())->getBitWidth();
8092
8093 // If the shift count is not less than the bitwidth, the result of
8094 // the shift is undefined. Don't try to analyze it, because the
8095 // resolution chosen here may differ from the resolution chosen in
8096 // other parts of the compiler.
8097 if (SA->getValue().uge(BitWidth))
8098 break;
8099
8100 // We can safely preserve the nuw flag in all cases. It's also safe to
8101 // turn a nuw nsw shl into a nuw nsw mul. However, nsw in isolation
8102 // requires special handling. It can be preserved as long as we're not
8103 // left shifting by bitwidth - 1.
8104 auto Flags = SCEV::FlagNone;
8105 if (BO->Op) {
8106 auto MulFlags = getNoWrapFlagsFromUB(BO->Op);
8107 if (any(MulFlags & SCEV::FlagNSW) &&
8108 (any(MulFlags & SCEV::FlagNUW) ||
8109 SA->getValue().ult(BitWidth - 1)))
8111 if (any(MulFlags & SCEV::FlagNUW))
8113 }
8114
8115 ConstantInt *X = ConstantInt::get(
8116 getContext(), APInt::getOneBitSet(BitWidth, SA->getZExtValue()));
8117 return getMulExpr(getSCEV(BO->LHS), getConstant(X), Flags);
8118 }
8119 break;
8120
8121 case Instruction::AShr:
8122 // AShr X, C, where C is a constant.
8123 ConstantInt *CI = dyn_cast<ConstantInt>(BO->RHS);
8124 if (!CI)
8125 break;
8126
8127 Type *OuterTy = BO->LHS->getType();
8129 // If the shift count is not less than the bitwidth, the result of
8130 // the shift is undefined. Don't try to analyze it, because the
8131 // resolution chosen here may differ from the resolution chosen in
8132 // other parts of the compiler.
8133 if (CI->getValue().uge(BitWidth))
8134 break;
8135
8136 if (CI->isZero())
8137 return getSCEV(BO->LHS); // shift by zero --> noop
8138
8139 uint64_t AShrAmt = CI->getZExtValue();
8140 Type *TruncTy = IntegerType::get(getContext(), BitWidth - AShrAmt);
8141
8142 Operator *L = dyn_cast<Operator>(BO->LHS);
8143 const SCEV *AddTruncateExpr = nullptr;
8144 ConstantInt *ShlAmtCI = nullptr;
8145 const SCEV *AddConstant = nullptr;
8146
8147 if (L && L->getOpcode() == Instruction::Add) {
8148 // X = Shl A, n
8149 // Y = Add X, c
8150 // Z = AShr Y, m
8151 // n, c and m are constants.
8152
8153 Operator *LShift = dyn_cast<Operator>(L->getOperand(0));
8154 ConstantInt *AddOperandCI = dyn_cast<ConstantInt>(L->getOperand(1));
8155 if (LShift && LShift->getOpcode() == Instruction::Shl) {
8156 if (AddOperandCI) {
8157 const SCEV *ShlOp0SCEV = getSCEV(LShift->getOperand(0));
8158 ShlAmtCI = dyn_cast<ConstantInt>(LShift->getOperand(1));
8159 // since we truncate to TruncTy, the AddConstant should be of the
8160 // same type, so create a new Constant with type same as TruncTy.
8161 // Also, the Add constant should be shifted right by AShr amount.
8162 APInt AddOperand = AddOperandCI->getValue().ashr(AShrAmt);
8163 AddConstant = getConstant(AddOperand.trunc(BitWidth - AShrAmt));
8164 // we model the expression as sext(add(trunc(A), c << n)), since the
8165 // sext(trunc) part is already handled below, we create a
8166 // AddExpr(TruncExp) which will be used later.
8167 AddTruncateExpr = getTruncateExpr(ShlOp0SCEV, TruncTy);
8168 }
8169 }
8170 } else if (L && L->getOpcode() == Instruction::Shl) {
8171 // X = Shl A, n
8172 // Y = AShr X, m
8173 // Both n and m are constant.
8174
8175 const SCEV *ShlOp0SCEV = getSCEV(L->getOperand(0));
8176 ShlAmtCI = dyn_cast<ConstantInt>(L->getOperand(1));
8177 AddTruncateExpr = getTruncateExpr(ShlOp0SCEV, TruncTy);
8178 }
8179
8180 if (AddTruncateExpr && ShlAmtCI) {
8181 // We can merge the two given cases into a single SCEV statement,
8182 // incase n = m, the mul expression will be 2^0, so it gets resolved to
8183 // a simpler case. The following code handles the two cases:
8184 //
8185 // 1) For a two-shift sext-inreg, i.e. n = m,
8186 // use sext(trunc(x)) as the SCEV expression.
8187 //
8188 // 2) When n > m, use sext(mul(trunc(x), 2^(n-m)))) as the SCEV
8189 // expression. We already checked that ShlAmt < BitWidth, so
8190 // the multiplier, 1 << (ShlAmt - AShrAmt), fits into TruncTy as
8191 // ShlAmt - AShrAmt < Amt.
8192 const APInt &ShlAmt = ShlAmtCI->getValue();
8193 if (ShlAmt.ult(BitWidth) && ShlAmt.uge(AShrAmt)) {
8194 APInt Mul = APInt::getOneBitSet(BitWidth - AShrAmt,
8195 ShlAmtCI->getZExtValue() - AShrAmt);
8196 const SCEV *CompositeExpr =
8197 getMulExpr(AddTruncateExpr, getConstant(Mul));
8198 if (L->getOpcode() != Instruction::Shl)
8199 CompositeExpr = getAddExpr(CompositeExpr, AddConstant);
8200
8201 return getSignExtendExpr(CompositeExpr, OuterTy);
8202 }
8203 }
8204 break;
8205 }
8206 }
8207
8208 switch (U->getOpcode()) {
8209 case Instruction::Trunc:
8210 return getTruncateExpr(getSCEV(U->getOperand(0)), U->getType());
8211
8212 case Instruction::ZExt:
8213 return getZeroExtendExpr(getSCEV(U->getOperand(0)), U->getType());
8214
8215 case Instruction::SExt:
8216 if (auto BO = MatchBinaryOp(U->getOperand(0), getDataLayout(), AC, DT,
8218 // The NSW flag of a subtract does not always survive the conversion to
8219 // A + (-1)*B. By pushing sign extension onto its operands we are much
8220 // more likely to preserve NSW and allow later AddRec optimisations.
8221 //
8222 // NOTE: This is effectively duplicating this logic from getSignExtend:
8223 // sext((A + B + ...)<nsw>) --> (sext(A) + sext(B) + ...)<nsw>
8224 // but by that point the NSW information has potentially been lost.
8225 if (BO->Opcode == Instruction::Sub && BO->IsNSW) {
8226 Type *Ty = U->getType();
8227 auto *V1 = getSignExtendExpr(getSCEV(BO->LHS), Ty);
8228 auto *V2 = getSignExtendExpr(getSCEV(BO->RHS), Ty);
8229 return getMinusSCEV(V1, V2, SCEV::FlagNSW);
8230 }
8231 }
8232 return getSignExtendExpr(getSCEV(U->getOperand(0)), U->getType());
8233
8234 case Instruction::BitCast:
8235 // BitCasts are no-op casts so we just eliminate the cast.
8236 if (isSCEVable(U->getType()) && isSCEVable(U->getOperand(0)->getType()))
8237 return getSCEV(U->getOperand(0));
8238 break;
8239
8240 case Instruction::PtrToAddr: {
8241 const SCEV *IntOp = getPtrToAddrExpr(getSCEV(U->getOperand(0)));
8242 if (isa<SCEVCouldNotCompute>(IntOp))
8243 return getUnknown(V);
8244 return IntOp;
8245 }
8246
8247 case Instruction::PtrToInt:
8248 // SCEV only models ptrtoaddr.
8249 return getUnknown(V);
8250
8251 case Instruction::IntToPtr:
8252 // Just don't deal with inttoptr casts.
8253 return getUnknown(V);
8254
8255 case Instruction::SDiv:
8256 // If both operands are non-negative, this is just an udiv.
8257 if (isKnownNonNegative(getSCEV(U->getOperand(0))) &&
8258 isKnownNonNegative(getSCEV(U->getOperand(1))))
8259 return getUDivExpr(getSCEV(U->getOperand(0)), getSCEV(U->getOperand(1)));
8260 break;
8261
8262 case Instruction::SRem:
8263 // If both operands are non-negative, this is just an urem.
8264 if (isKnownNonNegative(getSCEV(U->getOperand(0))) &&
8265 isKnownNonNegative(getSCEV(U->getOperand(1))))
8266 return getURemExpr(getSCEV(U->getOperand(0)), getSCEV(U->getOperand(1)));
8267 break;
8268
8269 case Instruction::GetElementPtr:
8270 return createNodeForGEP(cast<GEPOperator>(U));
8271
8272 case Instruction::PHI:
8273 return createNodeForPHI(cast<PHINode>(U));
8274
8275 case Instruction::Select:
8276 return createNodeForSelectOrPHI(U, U->getOperand(0), U->getOperand(1),
8277 U->getOperand(2));
8278
8279 case Instruction::Call:
8280 case Instruction::Invoke:
8281 if (Value *RV = cast<CallBase>(U)->getReturnedArgOperand())
8282 return getSCEV(RV);
8283
8284 if (auto *II = dyn_cast<IntrinsicInst>(U)) {
8285 switch (II->getIntrinsicID()) {
8286 case Intrinsic::abs:
8287 return getAbsExpr(
8288 getSCEV(II->getArgOperand(0)),
8289 /*IsNSW=*/cast<ConstantInt>(II->getArgOperand(1))->isOne());
8290 case Intrinsic::umax:
8291 LHS = getSCEV(II->getArgOperand(0));
8292 RHS = getSCEV(II->getArgOperand(1));
8293 return getUMaxExpr(LHS, RHS);
8294 case Intrinsic::umin:
8295 LHS = getSCEV(II->getArgOperand(0));
8296 RHS = getSCEV(II->getArgOperand(1));
8297 return getUMinExpr(LHS, RHS);
8298 case Intrinsic::smax:
8299 LHS = getSCEV(II->getArgOperand(0));
8300 RHS = getSCEV(II->getArgOperand(1));
8301 return getSMaxExpr(LHS, RHS);
8302 case Intrinsic::smin:
8303 LHS = getSCEV(II->getArgOperand(0));
8304 RHS = getSCEV(II->getArgOperand(1));
8305 return getSMinExpr(LHS, RHS);
8306 case Intrinsic::usub_sat: {
8307 const SCEV *X = getSCEV(II->getArgOperand(0));
8308 const SCEV *Y = getSCEV(II->getArgOperand(1));
8309 const SCEV *ClampedY = getUMinExpr(X, Y);
8310 return getMinusSCEV(X, ClampedY, SCEV::FlagNUW);
8311 }
8312 case Intrinsic::uadd_sat: {
8313 const SCEV *X = getSCEV(II->getArgOperand(0));
8314 const SCEV *Y = getSCEV(II->getArgOperand(1));
8315 const SCEV *ClampedX = getUMinExpr(X, getNotSCEV(Y));
8316 return getAddExpr(ClampedX, Y, SCEV::FlagNUW);
8317 }
8318 case Intrinsic::start_loop_iterations:
8319 case Intrinsic::annotation:
8320 case Intrinsic::ptr_annotation:
8321 // A start_loop_iterations or llvm.annotation or llvm.prt.annotation is
8322 // just eqivalent to the first operand for SCEV purposes.
8323 return getSCEV(II->getArgOperand(0));
8324 case Intrinsic::vscale:
8325 return getVScale(II->getType());
8326 default:
8327 break;
8328 }
8329 }
8330 break;
8331 }
8332
8333 return getUnknown(V);
8334}
8335
8336//===----------------------------------------------------------------------===//
8337// Iteration Count Computation Code
8338//
8339
8341 if (isa<SCEVCouldNotCompute>(ExitCount))
8342 return getCouldNotCompute();
8343
8344 auto *ExitCountType = ExitCount->getType();
8345 assert(ExitCountType->isIntegerTy());
8346 auto *EvalTy = Type::getIntNTy(ExitCountType->getContext(),
8347 1 + ExitCountType->getScalarSizeInBits());
8348 return getTripCountFromExitCount(ExitCount, EvalTy, nullptr);
8349}
8350
8352 Type *EvalTy,
8353 const Loop *L) {
8354 if (isa<SCEVCouldNotCompute>(ExitCount))
8355 return getCouldNotCompute();
8356
8357 unsigned ExitCountSize = getTypeSizeInBits(ExitCount->getType());
8358 unsigned EvalSize = EvalTy->getPrimitiveSizeInBits();
8359
8360 auto CanAddOneWithoutOverflow = [&]() {
8361 ConstantRange ExitCountRange =
8362 getRangeRef(ExitCount, RangeSignHint::HINT_RANGE_UNSIGNED);
8363 if (!ExitCountRange.contains(APInt::getMaxValue(ExitCountSize)))
8364 return true;
8365
8366 return L && isLoopEntryGuardedByCond(L, ICmpInst::ICMP_NE, ExitCount,
8367 getMinusOne(ExitCount->getType()));
8368 };
8369
8370 // If we need to zero extend the backedge count, check if we can add one to
8371 // it prior to zero extending without overflow. Provided this is safe, it
8372 // allows better simplification of the +1.
8373 if (EvalSize > ExitCountSize && CanAddOneWithoutOverflow())
8374 return getZeroExtendExpr(
8375 getAddExpr(ExitCount, getOne(ExitCount->getType())), EvalTy);
8376
8377 // Get the total trip count from the count by adding 1. This may wrap.
8378 return getAddExpr(getTruncateOrZeroExtend(ExitCount, EvalTy), getOne(EvalTy));
8379}
8380
8381static unsigned getConstantTripCount(const SCEVConstant *ExitCount) {
8382 if (!ExitCount)
8383 return 0;
8384
8385 ConstantInt *ExitConst = ExitCount->getValue();
8386
8387 // Guard against huge trip counts.
8388 if (ExitConst->getValue().getActiveBits() > 32)
8389 return 0;
8390
8391 // In case of integer overflow, this returns 0, which is correct.
8392 return ((unsigned)ExitConst->getZExtValue()) + 1;
8393}
8394
8396 auto *ExitCount = dyn_cast<SCEVConstant>(getBackedgeTakenCount(L, Exact));
8397 return getConstantTripCount(ExitCount);
8398}
8399
8400unsigned
8402 const BasicBlock *ExitingBlock) {
8403 assert(ExitingBlock && "Must pass a non-null exiting block!");
8404 assert(L->isLoopExiting(ExitingBlock) &&
8405 "Exiting block must actually branch out of the loop!");
8406 const SCEVConstant *ExitCount =
8407 dyn_cast<SCEVConstant>(getExitCount(L, ExitingBlock));
8408 return getConstantTripCount(ExitCount);
8409}
8410
8412 const Loop *L, SmallVectorImpl<const SCEVPredicate *> *Predicates) {
8413
8414 const auto *MaxExitCount =
8415 Predicates ? getPredicatedConstantMaxBackedgeTakenCount(L, *Predicates)
8417 return getConstantTripCount(dyn_cast<SCEVConstant>(MaxExitCount));
8418}
8419
8421 SmallVector<BasicBlock *, 8> ExitingBlocks;
8422 L->getExitingBlocks(ExitingBlocks);
8423
8424 // An exit with an uncomputable exit count makes the result 1.
8425 if (ExitingBlocks.empty() ||
8426 any_of(ExitingBlocks, [this, L](BasicBlock *ExitingBB) {
8427 return isa<SCEVCouldNotCompute>(getExitCount(L, ExitingBB));
8428 }))
8429 return 1;
8430
8431 LoopGuards Guards = LoopGuards::collect(L, *this);
8432 unsigned Res = 0;
8433 for (BasicBlock *ExitingBB : ExitingBlocks)
8434 Res = std::gcd(
8435 Res, getSmallConstantTripMultiple(getExitCount(L, ExitingBB), Guards));
8436 return Res;
8437}
8438
8439unsigned
8441 const LoopGuards &Guards) {
8442 assert(!isa<SCEVCouldNotCompute>(ExitCount) && "Must be computable!");
8443
8444 // Get the trip count
8445 const SCEV *TCExpr =
8446 getTripCountFromExitCount(applyLoopGuards(ExitCount, Guards));
8447
8448 APInt Multiple = getNonZeroConstantMultiple(TCExpr);
8449 // If a trip multiple is huge (>=2^32), the trip count is still divisible by
8450 // the greatest power of 2 divisor less than 2^32.
8451 return Multiple.getActiveBits() > 32
8452 ? 1U << std::min(31U, Multiple.countTrailingZeros())
8453 : (unsigned)Multiple.getZExtValue();
8454}
8455
8457 const SCEV *ExitCount) {
8458 if (isa<SCEVCouldNotCompute>(ExitCount))
8459 return 1;
8460
8461 return getSmallConstantTripMultiple(ExitCount, LoopGuards::collect(L, *this));
8462}
8463
8464/// Returns the largest constant divisor of the trip count of this loop as a
8465/// normal unsigned value, if possible. This means that the actual trip count is
8466/// always a multiple of the returned value (don't forget the trip count could
8467/// very well be zero as well!).
8468///
8469/// Returns 1 if the trip count is unknown or not guaranteed to be the
8470/// multiple of a constant (which is also the case if the trip count is simply
8471/// constant, use getSmallConstantTripCount for that case), Will also return 1
8472/// if the trip count is very large (>= 2^32).
8473///
8474/// As explained in the comments for getSmallConstantTripCount, this assumes
8475/// that control exits the loop via ExitingBlock.
8476unsigned
8478 const BasicBlock *ExitingBlock) {
8479 assert(ExitingBlock && "Must pass a non-null exiting block!");
8480 assert(L->isLoopExiting(ExitingBlock) &&
8481 "Exiting block must actually branch out of the loop!");
8482 const SCEV *ExitCount = getExitCount(L, ExitingBlock);
8483 return getSmallConstantTripMultiple(L, ExitCount);
8484}
8485
8487 const BasicBlock *ExitingBlock,
8488 ExitCountKind Kind) {
8489 switch (Kind) {
8490 case Exact:
8491 return getBackedgeTakenInfo(L).getExact(ExitingBlock, this);
8492 case SymbolicMaximum:
8493 return getBackedgeTakenInfo(L).getSymbolicMax(ExitingBlock, this);
8494 case ConstantMaximum:
8495 return getBackedgeTakenInfo(L).getConstantMax(ExitingBlock, this);
8496 };
8497 llvm_unreachable("Invalid ExitCountKind!");
8498}
8499
8501 const Loop *L, const BasicBlock *ExitingBlock,
8503 switch (Kind) {
8504 case Exact:
8505 return getPredicatedBackedgeTakenInfo(L).getExact(ExitingBlock, this,
8506 Predicates);
8507 case SymbolicMaximum:
8508 return getPredicatedBackedgeTakenInfo(L).getSymbolicMax(ExitingBlock, this,
8509 Predicates);
8510 case ConstantMaximum:
8511 return getPredicatedBackedgeTakenInfo(L).getConstantMax(ExitingBlock, this,
8512 Predicates);
8513 };
8514 llvm_unreachable("Invalid ExitCountKind!");
8515}
8516
8519 return getPredicatedBackedgeTakenInfo(L).getExact(L, this, &Preds);
8520}
8521
8523 ExitCountKind Kind) {
8524 switch (Kind) {
8525 case Exact:
8526 return getBackedgeTakenInfo(L).getExact(L, this);
8527 case ConstantMaximum:
8528 return getBackedgeTakenInfo(L).getConstantMax(this);
8529 case SymbolicMaximum:
8530 return getBackedgeTakenInfo(L).getSymbolicMax(L, this);
8531 };
8532 llvm_unreachable("Invalid ExitCountKind!");
8533}
8534
8537 return getPredicatedBackedgeTakenInfo(L).getSymbolicMax(L, this, &Preds);
8538}
8539
8542 return getPredicatedBackedgeTakenInfo(L).getConstantMax(this, &Preds);
8543}
8544
8546 return getBackedgeTakenInfo(L).isConstantMaxOrZero(this);
8547}
8548
8549/// Push PHI nodes in the header of the given loop onto the given Worklist.
8550static void PushLoopPHIs(const Loop *L,
8553 BasicBlock *Header = L->getHeader();
8554
8555 // Push all Loop-header PHIs onto the Worklist stack.
8556 for (PHINode &PN : Header->phis())
8557 if (Visited.insert(&PN).second)
8558 Worklist.push_back(&PN);
8559}
8560
8561ScalarEvolution::BackedgeTakenInfo &
8562ScalarEvolution::getPredicatedBackedgeTakenInfo(const Loop *L) {
8563 auto &BTI = getBackedgeTakenInfo(L);
8564 if (BTI.hasFullInfo())
8565 return BTI;
8566
8567 auto Pair = PredicatedBackedgeTakenCounts.try_emplace(L);
8568
8569 if (!Pair.second)
8570 return Pair.first->second;
8571
8572 BackedgeTakenInfo Result =
8573 computeBackedgeTakenCount(L, /*AllowPredicates=*/true);
8574
8575 return PredicatedBackedgeTakenCounts.find(L)->second = std::move(Result);
8576}
8577
8578ScalarEvolution::BackedgeTakenInfo &
8579ScalarEvolution::getBackedgeTakenInfo(const Loop *L) {
8580 // Initially insert an invalid entry for this loop. If the insertion
8581 // succeeds, proceed to actually compute a backedge-taken count and
8582 // update the value. The temporary CouldNotCompute value tells SCEV
8583 // code elsewhere that it shouldn't attempt to request a new
8584 // backedge-taken count, which could result in infinite recursion.
8585 std::pair<DenseMap<const Loop *, BackedgeTakenInfo>::iterator, bool> Pair =
8586 BackedgeTakenCounts.try_emplace(L);
8587 if (!Pair.second)
8588 return Pair.first->second;
8589
8590 // computeBackedgeTakenCount may allocate memory for its result. Inserting it
8591 // into the BackedgeTakenCounts map transfers ownership. Otherwise, the result
8592 // must be cleared in this scope.
8593 BackedgeTakenInfo Result = computeBackedgeTakenCount(L);
8594
8595 // Now that we know more about the trip count for this loop, forget any
8596 // existing SCEV values for PHI nodes in this loop since they are only
8597 // conservative estimates made without the benefit of trip count
8598 // information. This invalidation is not necessary for correctness, and is
8599 // only done to produce more precise results.
8600 if (Result.hasAnyInfo()) {
8601 // Invalidate any expression using an addrec in this loop.
8602 SmallVector<SCEVUse, 8> ToForget;
8603 auto LoopUsersIt = LoopUsers.find(L);
8604 if (LoopUsersIt != LoopUsers.end())
8605 append_range(ToForget, LoopUsersIt->second);
8606 forgetMemoizedResults(ToForget);
8607
8608 // Invalidate constant-evolved loop header phis.
8609 for (PHINode &PN : L->getHeader()->phis())
8610 ConstantEvolutionLoopExitValue.erase(&PN);
8611 }
8612
8613 // Re-lookup the insert position, since the call to
8614 // computeBackedgeTakenCount above could result in a
8615 // recusive call to getBackedgeTakenInfo (on a different
8616 // loop), which would invalidate the iterator computed
8617 // earlier.
8618 return BackedgeTakenCounts.find(L)->second = std::move(Result);
8619}
8620
8622 // This method is intended to forget all info about loops. It should
8623 // invalidate caches as if the following happened:
8624 // - The trip counts of all loops have changed arbitrarily
8625 // - Every llvm::Value has been updated in place to produce a different
8626 // result.
8627 BackedgeTakenCounts.clear();
8628 PredicatedBackedgeTakenCounts.clear();
8629 BECountUsers.clear();
8630 LoopPropertiesCache.clear();
8631 ConstantEvolutionLoopExitValue.clear();
8632 ValueExprMap.clear();
8633 ValuesAtScopes.clear();
8634 ValuesAtScopesUsers.clear();
8635 LoopDispositions.clear();
8636 BlockDispositions.clear();
8637 UnsignedRanges.clear();
8638 SignedRanges.clear();
8639 ExprValueMap.clear();
8640 HasRecMap.clear();
8641 ConstantMultipleCache.clear();
8642 PredicatedSCEVRewrites.clear();
8643 FoldCache.clear();
8644 FoldCacheUser.clear();
8645}
8646void ScalarEvolution::visitAndClearUsers(
8649 SmallVectorImpl<SCEVUse> &ToForget) {
8650 // Nothing can be invalidated if no value has a SCEV yet.
8651 if (ValueExprMap.empty()) {
8652 Worklist.clear();
8653 return;
8654 }
8655 while (!Worklist.empty()) {
8656 Instruction *I = Worklist.pop_back_val();
8657 if (!isSCEVable(I->getType()) && !isa<WithOverflowInst>(I))
8658 continue;
8659
8661 ValueExprMap.find_as(static_cast<Value *>(I));
8662 if (It != ValueExprMap.end()) {
8663 ToForget.push_back(It->second);
8664 eraseValueFromMap(It->first);
8665 if (PHINode *PN = dyn_cast<PHINode>(I))
8666 ConstantEvolutionLoopExitValue.erase(PN);
8667 }
8668
8669 PushDefUseChildren(I, Worklist, Visited);
8670 }
8671}
8672
8674 SmallVector<const Loop *, 16> LoopWorklist(1, L);
8677 SmallVector<SCEVUse, 16> ToForget;
8678
8679 // Iterate over all the loops and sub-loops to drop SCEV information.
8680 while (!LoopWorklist.empty()) {
8681 auto *CurrL = LoopWorklist.pop_back_val();
8682
8683 // Drop any stored trip count value.
8684 forgetBackedgeTakenCounts(CurrL, /* Predicated */ false);
8685 forgetBackedgeTakenCounts(CurrL, /* Predicated */ true);
8686
8687 // Drop information about predicated SCEV rewrites for this loop.
8688 PredicatedSCEVRewrites.remove_if(
8689 [&](const auto &Entry) { return Entry.first.second == CurrL; });
8690
8691 auto LoopUsersItr = LoopUsers.find(CurrL);
8692 if (LoopUsersItr != LoopUsers.end())
8693 llvm::append_range(ToForget, LoopUsersItr->second);
8694
8695 // Drop information about expressions based on loop-header PHIs.
8696 PushLoopPHIs(CurrL, Worklist, Visited);
8697 visitAndClearUsers(Worklist, Visited, ToForget);
8698
8699 LoopPropertiesCache.erase(CurrL);
8700 // Forget all contained loops too, to avoid dangling entries in the
8701 // ValuesAtScopes map.
8702 LoopWorklist.append(CurrL->begin(), CurrL->end());
8703 }
8704 forgetMemoizedResults(ToForget);
8705}
8706
8708 forgetLoop(L->getOutermostLoop());
8709}
8710
8713 if (!I) return;
8714
8715 // Drop information about expressions based on loop-header PHIs.
8718 SmallVector<SCEVUse, 8> ToForget;
8719 Worklist.push_back(I);
8720 Visited.insert(I);
8721 visitAndClearUsers(Worklist, Visited, ToForget);
8722
8723 forgetMemoizedResults(ToForget);
8724}
8725
8729 SmallVector<SCEVUse, 8> ToForget;
8730 for (Value *V : Values)
8731 if (auto *I = dyn_cast<Instruction>(V))
8732 if (Visited.insert(I).second)
8733 Worklist.push_back(I);
8734 visitAndClearUsers(Worklist, Visited, ToForget);
8735
8736 forgetMemoizedResults(ToForget);
8737}
8738
8740 // If SCEV looked through a trivial LCSSA phi node, we might have SCEV's
8741 // directly using a SCEVUnknown/SCEVAddRec defined in the loop. After an
8742 // extra predecessor is added, this is no longer valid. Find all Unknowns and
8743 // AddRecs defined in the loop and invalidate any SCEV's making use of them.
8744 auto InvalidateValue = [&](Value *Val) {
8745 if (!isSCEVable(Val->getType()))
8746 return;
8747 if (const SCEV *S = getExistingSCEV(Val)) {
8748 struct InvalidationRootCollector {
8749 Loop *L;
8751
8752 InvalidationRootCollector(Loop *L) : L(L) {}
8753
8754 bool follow(const SCEV *S) {
8755 if (auto *SU = dyn_cast<SCEVUnknown>(S)) {
8756 if (auto *I = dyn_cast<Instruction>(SU->getValue()))
8757 if (L->contains(I))
8758 Roots.push_back(S);
8759 } else if (auto *AddRec = dyn_cast<SCEVAddRecExpr>(S)) {
8760 if (L->contains(AddRec->getLoop()))
8761 Roots.push_back(S);
8762 }
8763 return true;
8764 }
8765 bool isDone() const { return false; }
8766 };
8767
8768 InvalidationRootCollector C(L);
8769 visitAll(S, C);
8770 forgetMemoizedResults(C.Roots);
8771 }
8772 };
8773
8774 InvalidateValue(V);
8775
8776 // If V has a non-SCEV-able type (e.g. {i64, i1} from a with.overflow
8777 // intrinsic), its users (e.g. extractvalue) may have stale SCEV
8778 // expressions referencing loop-internal values.
8779 if (!isSCEVable(V->getType()) &&
8780 any_of(V->incoming_values(), IsaPred<WithOverflowInst>))
8781 for (User *U : V->users())
8782 InvalidateValue(U);
8783 // Also perform the normal invalidation.
8784 forgetValue(V);
8785}
8786
8787void ScalarEvolution::forgetLoopDispositions() { LoopDispositions.clear(); }
8788
8790 // Unless a specific value is passed to invalidation, completely clear both
8791 // caches.
8792 if (!V) {
8793 BlockDispositions.clear();
8794 LoopDispositions.clear();
8795 return;
8796 }
8797
8798 if (!isSCEVable(V->getType()))
8799 return;
8800
8801 const SCEV *S = getExistingSCEV(V);
8802 if (!S)
8803 return;
8804
8805 // Invalidate the block and loop dispositions cached for S. Dispositions of
8806 // S's users may change if S's disposition changes (i.e. a user may change to
8807 // loop-invariant, if S changes to loop invariant), so also invalidate
8808 // dispositions of S's users recursively.
8809 SmallVector<SCEVUse, 8> Worklist = {S};
8811 while (!Worklist.empty()) {
8812 const SCEV *Curr = Worklist.pop_back_val();
8813 bool LoopDispoRemoved = LoopDispositions.erase(Curr);
8814 bool BlockDispoRemoved = BlockDispositions.erase(Curr);
8815 if (!LoopDispoRemoved && !BlockDispoRemoved)
8816 continue;
8817 auto Users = SCEVUsers.find(Curr);
8818 if (Users != SCEVUsers.end())
8819 for (const auto *User : Users->second)
8820 if (Seen.insert(User).second)
8821 Worklist.push_back(User);
8822 }
8823}
8824
8825/// Get the exact loop backedge taken count considering all loop exits. A
8826/// computable result can only be returned for loops with all exiting blocks
8827/// dominating the latch. howFarToZero assumes that the limit of each loop test
8828/// is never skipped. This is a valid assumption as long as the loop exits via
8829/// that test. For precise results, it is the caller's responsibility to specify
8830/// the relevant loop exiting block using getExact(ExitingBlock, SE).
8831const SCEV *ScalarEvolution::BackedgeTakenInfo::getExact(
8832 const Loop *L, ScalarEvolution *SE,
8834 // If any exits were not computable, the loop is not computable.
8835 if (!isComplete() || ExitNotTaken.empty())
8836 return SE->getCouldNotCompute();
8837
8838 const BasicBlock *Latch = L->getLoopLatch();
8839 // All exiting blocks we have collected must dominate the only backedge.
8840 if (!Latch)
8841 return SE->getCouldNotCompute();
8842
8843 // All exiting blocks we have gathered dominate loop's latch, so exact trip
8844 // count is simply a minimum out of all these calculated exit counts.
8846 for (const auto &ENT : ExitNotTaken) {
8847 const SCEV *BECount = ENT.ExactNotTaken;
8848 assert(BECount != SE->getCouldNotCompute() && "Bad exit SCEV!");
8849 assert(SE->DT.dominates(ENT.ExitingBlock, Latch) &&
8850 "We should only have known counts for exiting blocks that dominate "
8851 "latch!");
8852
8853 Ops.push_back(BECount);
8854
8855 if (Preds)
8856 append_range(*Preds, ENT.Predicates);
8857
8858 assert((Preds || ENT.hasAlwaysTruePredicate()) &&
8859 "Predicate should be always true!");
8860 }
8861
8862 // If an earlier exit exits on the first iteration (exit count zero), then
8863 // a later poison exit count should not propagate into the result. This are
8864 // exactly the semantics provided by umin_seq.
8865 return SE->getUMinFromMismatchedTypes(Ops, /* Sequential */ true);
8866}
8867
8868const ScalarEvolution::ExitNotTakenInfo *
8869ScalarEvolution::BackedgeTakenInfo::getExitNotTaken(
8870 const BasicBlock *ExitingBlock,
8871 SmallVectorImpl<const SCEVPredicate *> *Predicates) const {
8872 for (const auto &ENT : ExitNotTaken)
8873 if (ENT.ExitingBlock == ExitingBlock) {
8874 if (ENT.hasAlwaysTruePredicate())
8875 return &ENT;
8876 else if (Predicates) {
8877 append_range(*Predicates, ENT.Predicates);
8878 return &ENT;
8879 }
8880 }
8881
8882 return nullptr;
8883}
8884
8885/// getConstantMax - Get the constant max backedge taken count for the loop.
8886const SCEV *ScalarEvolution::BackedgeTakenInfo::getConstantMax(
8887 ScalarEvolution *SE,
8888 SmallVectorImpl<const SCEVPredicate *> *Predicates) const {
8889 if (!getConstantMax())
8890 return SE->getCouldNotCompute();
8891
8892 for (const auto &ENT : ExitNotTaken)
8893 if (!ENT.hasAlwaysTruePredicate()) {
8894 if (!Predicates)
8895 return SE->getCouldNotCompute();
8896 append_range(*Predicates, ENT.Predicates);
8897 }
8898
8899 assert((isa<SCEVCouldNotCompute>(getConstantMax()) ||
8900 isa<SCEVConstant>(getConstantMax())) &&
8901 "No point in having a non-constant max backedge taken count!");
8902 return getConstantMax();
8903}
8904
8905const SCEV *ScalarEvolution::BackedgeTakenInfo::getSymbolicMax(
8906 const Loop *L, ScalarEvolution *SE,
8907 SmallVectorImpl<const SCEVPredicate *> *Predicates) {
8908 if (!SymbolicMax) {
8909 // Form an expression for the maximum exit count possible for this loop. We
8910 // merge the max and exact information to approximate a version of
8911 // getConstantMaxBackedgeTakenCount which isn't restricted to just
8912 // constants.
8913 SmallVector<SCEVUse, 4> ExitCounts;
8914
8915 for (const auto &ENT : ExitNotTaken) {
8916 const SCEV *ExitCount = ENT.SymbolicMaxNotTaken;
8917 if (!isa<SCEVCouldNotCompute>(ExitCount)) {
8918 assert(SE->DT.dominates(ENT.ExitingBlock, L->getLoopLatch()) &&
8919 "We should only have known counts for exiting blocks that "
8920 "dominate latch!");
8921 ExitCounts.push_back(ExitCount);
8922 if (Predicates)
8923 append_range(*Predicates, ENT.Predicates);
8924
8925 assert((Predicates || ENT.hasAlwaysTruePredicate()) &&
8926 "Predicate should be always true!");
8927 }
8928 }
8929 if (ExitCounts.empty())
8930 SymbolicMax = SE->getCouldNotCompute();
8931 else
8932 SymbolicMax =
8933 SE->getUMinFromMismatchedTypes(ExitCounts, /*Sequential*/ true);
8934 }
8935 return SymbolicMax;
8936}
8937
8938bool ScalarEvolution::BackedgeTakenInfo::isConstantMaxOrZero(
8939 ScalarEvolution *SE) const {
8940 auto PredicateNotAlwaysTrue = [](const ExitNotTakenInfo &ENT) {
8941 return !ENT.hasAlwaysTruePredicate();
8942 };
8943 return MaxOrZero && !any_of(ExitNotTaken, PredicateNotAlwaysTrue);
8944}
8945
8948
8950 const SCEV *E, const SCEV *ConstantMaxNotTaken,
8951 const SCEV *SymbolicMaxNotTaken, bool MaxOrZero,
8955 // If we prove the max count is zero, so is the symbolic bound. This happens
8956 // in practice due to differences in a) how context sensitive we've chosen
8957 // to be and b) how we reason about bounds implied by UB.
8958 if (ConstantMaxNotTaken->isZero()) {
8959 this->ExactNotTaken = E = ConstantMaxNotTaken;
8960 this->SymbolicMaxNotTaken = SymbolicMaxNotTaken = ConstantMaxNotTaken;
8961 }
8962
8965 "Exact is not allowed to be less precise than Constant Max");
8968 "Exact is not allowed to be less precise than Symbolic Max");
8971 "Symbolic Max is not allowed to be less precise than Constant Max");
8974 "No point in having a non-constant max backedge taken count!");
8976 for (const auto PredList : PredLists)
8977 for (const auto *P : PredList) {
8978 if (SeenPreds.contains(P))
8979 continue;
8980 assert(!isa<SCEVUnionPredicate>(P) && "Only add leaf predicates here!");
8981 SeenPreds.insert(P);
8982 Predicates.push_back(P);
8983 }
8984 assert((isa<SCEVCouldNotCompute>(E) || !E->getType()->isPointerTy()) &&
8985 "Backedge count should be int");
8987 !ConstantMaxNotTaken->getType()->isPointerTy()) &&
8988 "Max backedge count should be int");
8989}
8990
8998
8999/// Allocate memory for BackedgeTakenInfo and copy the not-taken count of each
9000/// computable exit into a persistent ExitNotTakenInfo array.
9001ScalarEvolution::BackedgeTakenInfo::BackedgeTakenInfo(
9003 bool IsComplete, const SCEV *ConstantMax, bool MaxOrZero)
9004 : ConstantMax(ConstantMax), IsComplete(IsComplete), MaxOrZero(MaxOrZero) {
9005 using EdgeExitInfo = ScalarEvolution::BackedgeTakenInfo::EdgeExitInfo;
9006
9007 ExitNotTaken.reserve(ExitCounts.size());
9008 std::transform(ExitCounts.begin(), ExitCounts.end(),
9009 std::back_inserter(ExitNotTaken),
9010 [&](const EdgeExitInfo &EEI) {
9011 BasicBlock *ExitBB = EEI.first;
9012 const ExitLimit &EL = EEI.second;
9013 return ExitNotTakenInfo(ExitBB, EL.ExactNotTaken,
9014 EL.ConstantMaxNotTaken, EL.SymbolicMaxNotTaken,
9015 EL.Predicates);
9016 });
9017 assert((isa<SCEVCouldNotCompute>(ConstantMax) ||
9018 isa<SCEVConstant>(ConstantMax)) &&
9019 "No point in having a non-constant max backedge taken count!");
9020}
9021
9022/// Compute the number of times the backedge of the specified loop will execute.
9023ScalarEvolution::BackedgeTakenInfo
9024ScalarEvolution::computeBackedgeTakenCount(const Loop *L,
9025 bool AllowPredicates) {
9026 SmallVector<BasicBlock *, 8> ExitingBlocks;
9027 L->getExitingBlocks(ExitingBlocks);
9028
9029 using EdgeExitInfo = ScalarEvolution::BackedgeTakenInfo::EdgeExitInfo;
9030
9032 bool CouldComputeBECount = true;
9033 BasicBlock *Latch = L->getLoopLatch(); // may be NULL.
9034 const SCEV *MustExitMaxBECount = nullptr;
9035 const SCEV *MayExitMaxBECount = nullptr;
9036 bool MustExitMaxOrZero = false;
9037 bool IsOnlyExit = ExitingBlocks.size() == 1;
9038
9039 // Compute the ExitLimit for each loop exit. Use this to populate ExitCounts
9040 // and compute maxBECount.
9041 // Do a union of all the predicates here.
9042 for (BasicBlock *ExitBB : ExitingBlocks) {
9043 // We canonicalize untaken exits to br (constant), ignore them so that
9044 // proving an exit untaken doesn't negatively impact our ability to reason
9045 // about the loop as whole.
9046 if (auto *BI = dyn_cast<CondBrInst>(ExitBB->getTerminator()))
9047 if (auto *CI = dyn_cast<ConstantInt>(BI->getCondition())) {
9048 bool ExitIfTrue = !L->contains(BI->getSuccessor(0));
9049 if (ExitIfTrue == CI->isZero())
9050 continue;
9051 }
9052
9053 ExitLimit EL = computeExitLimit(L, ExitBB, IsOnlyExit, AllowPredicates);
9054
9055 assert((AllowPredicates || EL.Predicates.empty()) &&
9056 "Predicated exit limit when predicates are not allowed!");
9057
9058 // 1. For each exit that can be computed, add an entry to ExitCounts.
9059 // CouldComputeBECount is true only if all exits can be computed.
9060 if (EL.ExactNotTaken != getCouldNotCompute())
9061 ++NumExitCountsComputed;
9062 else
9063 // We couldn't compute an exact value for this exit, so
9064 // we won't be able to compute an exact value for the loop.
9065 CouldComputeBECount = false;
9066 // Remember exit count if either exact or symbolic is known. Because
9067 // Exact always implies symbolic, only check symbolic.
9068 if (EL.SymbolicMaxNotTaken != getCouldNotCompute())
9069 ExitCounts.emplace_back(ExitBB, EL);
9070 else {
9071 assert(EL.ExactNotTaken == getCouldNotCompute() &&
9072 "Exact is known but symbolic isn't?");
9073 ++NumExitCountsNotComputed;
9074 }
9075
9076 // 2. Derive the loop's MaxBECount from each exit's max number of
9077 // non-exiting iterations. Partition the loop exits into two kinds:
9078 // LoopMustExits and LoopMayExits.
9079 //
9080 // If the exit dominates the loop latch, it is a LoopMustExit otherwise it
9081 // is a LoopMayExit. If any computable LoopMustExit is found, then
9082 // MaxBECount is the minimum EL.ConstantMaxNotTaken of computable
9083 // LoopMustExits. Otherwise, MaxBECount is conservatively the maximum
9084 // EL.ConstantMaxNotTaken, where CouldNotCompute is considered greater than
9085 // any
9086 // computable EL.ConstantMaxNotTaken.
9087 if (EL.ConstantMaxNotTaken != getCouldNotCompute() && Latch &&
9088 DT.dominates(ExitBB, Latch)) {
9089 if (!MustExitMaxBECount) {
9090 MustExitMaxBECount = EL.ConstantMaxNotTaken;
9091 MustExitMaxOrZero = EL.MaxOrZero;
9092 } else {
9093 MustExitMaxBECount = getUMinFromMismatchedTypes(MustExitMaxBECount,
9094 EL.ConstantMaxNotTaken);
9095 }
9096 } else if (MayExitMaxBECount != getCouldNotCompute()) {
9097 if (!MayExitMaxBECount || EL.ConstantMaxNotTaken == getCouldNotCompute())
9098 MayExitMaxBECount = EL.ConstantMaxNotTaken;
9099 else {
9100 MayExitMaxBECount = getUMaxFromMismatchedTypes(MayExitMaxBECount,
9101 EL.ConstantMaxNotTaken);
9102 }
9103 }
9104 }
9105 const SCEV *MaxBECount = MustExitMaxBECount ? MustExitMaxBECount :
9106 (MayExitMaxBECount ? MayExitMaxBECount : getCouldNotCompute());
9107 // The loop backedge will be taken the maximum or zero times if there's
9108 // a single exit that must be taken the maximum or zero times.
9109 bool MaxOrZero = (MustExitMaxOrZero && ExitingBlocks.size() == 1);
9110
9111 // Remember which SCEVs are used in exit limits for invalidation purposes.
9112 // We only care about non-constant SCEVs here, so we can ignore
9113 // EL.ConstantMaxNotTaken
9114 // and MaxBECount, which must be SCEVConstant.
9115 for (const auto &Pair : ExitCounts) {
9116 if (!isa<SCEVConstant>(Pair.second.ExactNotTaken))
9117 BECountUsers[Pair.second.ExactNotTaken].insert({L, AllowPredicates});
9118 if (!isa<SCEVConstant>(Pair.second.SymbolicMaxNotTaken))
9119 BECountUsers[Pair.second.SymbolicMaxNotTaken].insert(
9120 {L, AllowPredicates});
9121 }
9122 return BackedgeTakenInfo(std::move(ExitCounts), CouldComputeBECount,
9123 MaxBECount, MaxOrZero);
9124}
9125
9126ScalarEvolution::ExitLimit
9127ScalarEvolution::computeExitLimit(const Loop *L, BasicBlock *ExitingBlock,
9128 bool IsOnlyExit, bool AllowPredicates) {
9129 assert(L->contains(ExitingBlock) && "Exit count for non-loop block?");
9130 // If our exiting block does not dominate the latch, then its connection with
9131 // loop's exit limit may be far from trivial.
9132 const BasicBlock *Latch = L->getLoopLatch();
9133 if (!Latch || !DT.dominates(ExitingBlock, Latch))
9134 return getCouldNotCompute();
9135
9136 Instruction *Term = ExitingBlock->getTerminator();
9137 if (CondBrInst *BI = dyn_cast<CondBrInst>(Term)) {
9138 bool ExitIfTrue = !L->contains(BI->getSuccessor(0));
9139 assert(ExitIfTrue == L->contains(BI->getSuccessor(1)) &&
9140 "It should have one successor in loop and one exit block!");
9141 // Proceed to the next level to examine the exit condition expression.
9142 return computeExitLimitFromCond(L, BI->getCondition(), ExitIfTrue,
9143 /*ControlsOnlyExit=*/IsOnlyExit,
9144 AllowPredicates);
9145 }
9146
9147 if (SwitchInst *SI = dyn_cast<SwitchInst>(Term)) {
9148 // For switch, make sure that there is a single exit from the loop.
9149 BasicBlock *Exit = nullptr;
9150 for (auto *SBB : successors(ExitingBlock))
9151 if (!L->contains(SBB)) {
9152 if (Exit) // Multiple exit successors.
9153 return getCouldNotCompute();
9154 Exit = SBB;
9155 }
9156 assert(Exit && "Exiting block must have at least one exit");
9157 return computeExitLimitFromSingleExitSwitch(
9158 L, SI, Exit, /*ControlsOnlyExit=*/IsOnlyExit);
9159 }
9160
9161 return getCouldNotCompute();
9162}
9163
9165 const Loop *L, Value *ExitCond, bool ExitIfTrue, bool ControlsOnlyExit,
9166 bool AllowPredicates) {
9167 ScalarEvolution::ExitLimitCacheTy Cache(L, ExitIfTrue, AllowPredicates);
9168 return computeExitLimitFromCondCached(Cache, L, ExitCond, ExitIfTrue,
9169 ControlsOnlyExit, AllowPredicates);
9170}
9171
9172std::optional<ScalarEvolution::ExitLimit>
9173ScalarEvolution::ExitLimitCache::find(const Loop *L, Value *ExitCond,
9174 bool ExitIfTrue, bool ControlsOnlyExit,
9175 bool AllowPredicates) {
9176 (void)this->L;
9177 (void)this->ExitIfTrue;
9178 (void)this->AllowPredicates;
9179
9180 assert(this->L == L && this->ExitIfTrue == ExitIfTrue &&
9181 this->AllowPredicates == AllowPredicates &&
9182 "Variance in assumed invariant key components!");
9183 auto Itr = TripCountMap.find({ExitCond, ControlsOnlyExit});
9184 if (Itr == TripCountMap.end())
9185 return std::nullopt;
9186 return Itr->second;
9187}
9188
9189void ScalarEvolution::ExitLimitCache::insert(const Loop *L, Value *ExitCond,
9190 bool ExitIfTrue,
9191 bool ControlsOnlyExit,
9192 bool AllowPredicates,
9193 const ExitLimit &EL) {
9194 assert(this->L == L && this->ExitIfTrue == ExitIfTrue &&
9195 this->AllowPredicates == AllowPredicates &&
9196 "Variance in assumed invariant key components!");
9197
9198 auto InsertResult = TripCountMap.insert({{ExitCond, ControlsOnlyExit}, EL});
9199 assert(InsertResult.second && "Expected successful insertion!");
9200 (void)InsertResult;
9201 (void)ExitIfTrue;
9202}
9203
9204ScalarEvolution::ExitLimit ScalarEvolution::computeExitLimitFromCondCached(
9205 ExitLimitCacheTy &Cache, const Loop *L, Value *ExitCond, bool ExitIfTrue,
9206 bool ControlsOnlyExit, bool AllowPredicates) {
9207
9208 if (auto MaybeEL = Cache.find(L, ExitCond, ExitIfTrue, ControlsOnlyExit,
9209 AllowPredicates))
9210 return *MaybeEL;
9211
9212 ExitLimit EL = computeExitLimitFromCondImpl(
9213 Cache, L, ExitCond, ExitIfTrue, ControlsOnlyExit, AllowPredicates);
9214 Cache.insert(L, ExitCond, ExitIfTrue, ControlsOnlyExit, AllowPredicates, EL);
9215 return EL;
9216}
9217
9218ScalarEvolution::ExitLimit ScalarEvolution::computeExitLimitFromCondImpl(
9219 ExitLimitCacheTy &Cache, const Loop *L, Value *ExitCond, bool ExitIfTrue,
9220 bool ControlsOnlyExit, bool AllowPredicates) {
9221 // Handle BinOp conditions (And, Or).
9222 if (auto LimitFromBinOp = computeExitLimitFromCondFromBinOp(
9223 Cache, L, ExitCond, ExitIfTrue, AllowPredicates))
9224 return *LimitFromBinOp;
9225
9226 // With an icmp, it may be feasible to compute an exact backedge-taken count.
9227 // Proceed to the next level to examine the icmp.
9228 if (ICmpInst *ExitCondICmp = dyn_cast<ICmpInst>(ExitCond)) {
9229 ExitLimit EL =
9230 computeExitLimitFromICmp(L, ExitCondICmp, ExitIfTrue, ControlsOnlyExit);
9231 if (EL.hasFullInfo() || !AllowPredicates)
9232 return EL;
9233
9234 // Try again, but use SCEV predicates this time.
9235 return computeExitLimitFromICmp(L, ExitCondICmp, ExitIfTrue,
9236 ControlsOnlyExit,
9237 /*AllowPredicates=*/true);
9238 }
9239
9240 // Check for a constant condition. These are normally stripped out by
9241 // SimplifyCFG, but ScalarEvolution may be used by a pass which wishes to
9242 // preserve the CFG and is temporarily leaving constant conditions
9243 // in place.
9244 if (ConstantInt *CI = dyn_cast<ConstantInt>(ExitCond)) {
9245 if (ExitIfTrue == !CI->getZExtValue())
9246 // The backedge is always taken.
9247 return getCouldNotCompute();
9248 // The backedge is never taken.
9249 return getZero(CI->getType());
9250 }
9251
9252 // If we're exiting based on the overflow flag of an x.with.overflow intrinsic
9253 // with a constant step, we can form an equivalent icmp predicate and figure
9254 // out how many iterations will be taken before we exit.
9255 const WithOverflowInst *WO;
9256 const APInt *C;
9257 if (match(ExitCond, m_ExtractValue<1>(m_WithOverflowInst(WO))) &&
9258 match(WO->getRHS(), m_APInt(C))) {
9259 ConstantRange NWR =
9261 WO->getNoWrapKind());
9262 CmpInst::Predicate Pred;
9263 APInt NewRHSC, Offset;
9264 NWR.getEquivalentICmp(Pred, NewRHSC, Offset);
9265 if (!ExitIfTrue)
9266 Pred = ICmpInst::getInversePredicate(Pred);
9267 auto *LHS = getSCEV(WO->getLHS());
9268 if (Offset != 0)
9270 auto EL = computeExitLimitFromICmp(L, Pred, LHS, getConstant(NewRHSC),
9271 ControlsOnlyExit, AllowPredicates);
9272 if (EL.hasAnyInfo())
9273 return EL;
9274 }
9275
9276 // If it's not an integer or pointer comparison then compute it the hard way.
9277 return computeExitCountExhaustively(L, ExitCond, ExitIfTrue);
9278}
9279
9280std::optional<ScalarEvolution::ExitLimit>
9281ScalarEvolution::computeExitLimitFromCondFromBinOp(ExitLimitCacheTy &Cache,
9282 const Loop *L,
9283 Value *ExitCond,
9284 bool ExitIfTrue,
9285 bool AllowPredicates) {
9286 // Check if the controlling expression for this loop is an And or Or.
9287 Value *Op0, *Op1;
9288 bool IsAnd;
9289 if (match(ExitCond, m_LogicalAnd(m_Value(Op0), m_Value(Op1))))
9290 IsAnd = true;
9291 else if (match(ExitCond, m_LogicalOr(m_Value(Op0), m_Value(Op1))))
9292 IsAnd = false;
9293 else
9294 return std::nullopt;
9295
9296 // A sub-condition of a non-trivial binop never solely controls the exit,
9297 // whether we exit always depends on both conditions.
9298 ExitLimit EL0 = computeExitLimitFromCondCached(
9299 Cache, L, Op0, ExitIfTrue, /*ControlsOnlyExit=*/false, AllowPredicates);
9300 ExitLimit EL1 = computeExitLimitFromCondCached(
9301 Cache, L, Op1, ExitIfTrue, /*ControlsOnlyExit=*/false, AllowPredicates);
9302
9303 // EitherMayExit is true in these two cases:
9304 // br (and Op0 Op1), loop, exit
9305 // br (or Op0 Op1), exit, loop
9306 bool EitherMayExit = IsAnd ^ ExitIfTrue;
9307
9308 const SCEV *BECount = getCouldNotCompute();
9309 const SCEV *ConstantMaxBECount = getCouldNotCompute();
9310 const SCEV *SymbolicMaxBECount = getCouldNotCompute();
9311 if (EitherMayExit) {
9312 bool UseSequentialUMin = !isa<BinaryOperator>(ExitCond);
9313 // Both conditions must be same for the loop to continue executing.
9314 // Choose the less conservative count.
9315 if (EL0.ExactNotTaken != getCouldNotCompute() &&
9316 EL1.ExactNotTaken != getCouldNotCompute()) {
9317 BECount = getUMinFromMismatchedTypes(EL0.ExactNotTaken, EL1.ExactNotTaken,
9318 UseSequentialUMin);
9319 }
9320 if (EL0.ConstantMaxNotTaken == getCouldNotCompute())
9321 ConstantMaxBECount = EL1.ConstantMaxNotTaken;
9322 else if (EL1.ConstantMaxNotTaken == getCouldNotCompute())
9323 ConstantMaxBECount = EL0.ConstantMaxNotTaken;
9324 else
9325 ConstantMaxBECount = getUMinFromMismatchedTypes(EL0.ConstantMaxNotTaken,
9326 EL1.ConstantMaxNotTaken);
9327 if (EL0.SymbolicMaxNotTaken == getCouldNotCompute())
9328 SymbolicMaxBECount = EL1.SymbolicMaxNotTaken;
9329 else if (EL1.SymbolicMaxNotTaken == getCouldNotCompute())
9330 SymbolicMaxBECount = EL0.SymbolicMaxNotTaken;
9331 else
9332 SymbolicMaxBECount = getUMinFromMismatchedTypes(
9333 EL0.SymbolicMaxNotTaken, EL1.SymbolicMaxNotTaken, UseSequentialUMin);
9334 } else {
9335 // Both conditions must be same at the same time for the loop to exit.
9336 // For now, be conservative.
9337 if (EL0.ExactNotTaken == EL1.ExactNotTaken)
9338 BECount = EL0.ExactNotTaken;
9339 }
9340
9341 // There are cases (e.g. PR26207) where computeExitLimitFromCond is able
9342 // to be more aggressive when computing BECount than when computing
9343 // ConstantMaxBECount. In these cases it is possible for EL0.ExactNotTaken
9344 // and
9345 // EL1.ExactNotTaken to match, but for EL0.ConstantMaxNotTaken and
9346 // EL1.ConstantMaxNotTaken to not.
9347 if (isa<SCEVCouldNotCompute>(ConstantMaxBECount) &&
9348 !isa<SCEVCouldNotCompute>(BECount))
9349 ConstantMaxBECount = getConstant(getUnsignedRangeMax(BECount));
9350 if (isa<SCEVCouldNotCompute>(SymbolicMaxBECount))
9351 SymbolicMaxBECount =
9352 isa<SCEVCouldNotCompute>(BECount) ? ConstantMaxBECount : BECount;
9353 return ExitLimit(BECount, ConstantMaxBECount, SymbolicMaxBECount, false,
9354 {ArrayRef(EL0.Predicates), ArrayRef(EL1.Predicates)});
9355}
9356
9357ScalarEvolution::ExitLimit ScalarEvolution::computeExitLimitFromICmp(
9358 const Loop *L, ICmpInst *ExitCond, bool ExitIfTrue, bool ControlsOnlyExit,
9359 bool AllowPredicates) {
9360 // If the condition was exit on true, convert the condition to exit on false
9361 CmpPredicate Pred;
9362 if (!ExitIfTrue)
9363 Pred = ExitCond->getCmpPredicate();
9364 else
9365 Pred = ExitCond->getInverseCmpPredicate();
9366 const ICmpInst::Predicate OriginalPred = Pred;
9367
9368 const SCEV *LHS = getSCEV(ExitCond->getOperand(0));
9369 const SCEV *RHS = getSCEV(ExitCond->getOperand(1));
9370
9371 ExitLimit EL = computeExitLimitFromICmp(L, Pred, LHS, RHS, ControlsOnlyExit,
9372 AllowPredicates);
9373 if (EL.hasAnyInfo())
9374 return EL;
9375
9376 auto *ExhaustiveCount =
9377 computeExitCountExhaustively(L, ExitCond, ExitIfTrue);
9378
9379 if (!isa<SCEVCouldNotCompute>(ExhaustiveCount))
9380 return ExhaustiveCount;
9381
9382 return computeShiftCompareExitLimit(ExitCond->getOperand(0),
9383 ExitCond->getOperand(1), L, OriginalPred);
9384}
9385ScalarEvolution::ExitLimit ScalarEvolution::computeExitLimitFromICmp(
9386 const Loop *L, CmpPredicate Pred, SCEVUse LHS, SCEVUse RHS,
9387 bool ControlsOnlyExit, bool AllowPredicates) {
9388
9389 // Try to evaluate any dependencies out of the loop.
9390 LHS = getSCEVAtScope(LHS, L);
9391 RHS = getSCEVAtScope(RHS, L);
9392
9393 // At this point, we would like to compute how many iterations of the
9394 // loop the predicate will return true for these inputs.
9395 if (isLoopInvariant(LHS, L) && !isLoopInvariant(RHS, L)) {
9396 // If there is a loop-invariant, force it into the RHS.
9397 std::swap(LHS, RHS);
9399 }
9400
9401 bool ControllingFiniteLoop = ControlsOnlyExit && loopHasNoAbnormalExits(L) &&
9403 // Simplify the operands before analyzing them.
9404 (void)SimplifyICmpOperands(Pred, LHS, RHS, /*Depth=*/0);
9405
9406 // If we have a comparison of a chrec against a constant, try to use value
9407 // ranges to answer this query.
9408 if (const SCEVConstant *RHSC = dyn_cast<SCEVConstant>(RHS))
9409 if (const SCEVAddRecExpr *AddRec = dyn_cast<SCEVAddRecExpr>(LHS))
9410 if (AddRec->getLoop() == L) {
9411 // Form the constant range.
9412 ConstantRange CompRange =
9413 ConstantRange::makeExactICmpRegion(Pred, RHSC->getAPInt());
9414
9415 const SCEV *Ret = AddRec->getNumIterationsInRange(CompRange, *this);
9416 if (!isa<SCEVCouldNotCompute>(Ret)) return Ret;
9417 }
9418
9419 // If this loop must exit based on this condition (or execute undefined
9420 // behaviour), see if we can improve wrap flags. This is essentially
9421 // a must execute style proof.
9422 if (ControllingFiniteLoop && isLoopInvariant(RHS, L)) {
9423 // If we can prove the test sequence produced must repeat the same values
9424 // on self-wrap of the IV, then we can infer that IV doesn't self wrap
9425 // because if it did, we'd have an infinite (undefined) loop.
9426 // TODO: We can peel off any functions which are invertible *in L*. Loop
9427 // invariant terms are effectively constants for our purposes here.
9428 SCEVUse InnerLHS = LHS;
9429 if (auto *ZExt = dyn_cast<SCEVZeroExtendExpr>(LHS))
9430 InnerLHS = ZExt->getOperand();
9431 if (const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(InnerLHS);
9432 AR && !AR->hasNoSelfWrap() && AR->getLoop() == L && AR->isAffine() &&
9433 isKnownToBeAPowerOfTwo(AR->getStepRecurrence(*this), /*OrZero=*/true,
9434 /*OrNegative=*/true)) {
9435 auto Flags = AR->getNoWrapFlags();
9436 Flags = setFlags(Flags, SCEV::FlagNW);
9439 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), Flags);
9440 }
9441
9442 // For a slt/ult condition with a positive step, can we prove nsw/nuw?
9443 // From no-self-wrap, this follows trivially from the fact that every
9444 // (un)signed-wrapped, but not self-wrapped value must be LT than the
9445 // last value before (un)signed wrap. Since we know that last value
9446 // didn't exit, nor will any smaller one.
9447 if (Pred == ICmpInst::ICMP_SLT || Pred == ICmpInst::ICMP_ULT) {
9448 auto WrapType = Pred == ICmpInst::ICMP_SLT ? SCEV::FlagNSW : SCEV::FlagNUW;
9449 if (const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(LHS);
9450 AR && AR->getLoop() == L && AR->isAffine() &&
9451 !AR->getNoWrapFlags(WrapType) && AR->hasNoSelfWrap() &&
9452 isKnownPositive(AR->getStepRecurrence(*this))) {
9453 auto Flags = AR->getNoWrapFlags();
9454 Flags = setFlags(Flags, WrapType);
9457 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), Flags);
9458 }
9459 }
9460 }
9461
9462 switch (Pred) {
9463 case ICmpInst::ICMP_NE: { // while (X != Y)
9464 // Convert to: while (X-Y != 0)
9465 if (LHS->getType()->isPointerTy()) {
9468 return LHS;
9469 }
9470 if (RHS->getType()->isPointerTy()) {
9473 return RHS;
9474 }
9475 ExitLimit EL = howFarToZero(getMinusSCEV(LHS, RHS), L, ControlsOnlyExit,
9476 AllowPredicates);
9477 if (EL.hasAnyInfo())
9478 return EL;
9479 break;
9480 }
9481 case ICmpInst::ICMP_EQ: { // while (X == Y)
9482 // Convert to: while (X-Y == 0)
9483 if (LHS->getType()->isPointerTy()) {
9486 return LHS;
9487 }
9488 if (RHS->getType()->isPointerTy()) {
9491 return RHS;
9492 }
9493 ExitLimit EL = howFarToNonZero(getMinusSCEV(LHS, RHS), L);
9494 if (EL.hasAnyInfo()) return EL;
9495 break;
9496 }
9497 case ICmpInst::ICMP_SLE:
9498 case ICmpInst::ICMP_ULE:
9499 // Since the loop is finite, an invariant RHS cannot include the boundary
9500 // value, otherwise it would loop forever.
9501 if (!EnableFiniteLoopControl || !ControllingFiniteLoop ||
9502 !isLoopInvariant(RHS, L)) {
9503 // Otherwise, perform the addition in a wider type, to avoid overflow.
9504 // If the LHS is an addrec with the appropriate nowrap flag, the
9505 // extension will be sunk into it and the exit count can be analyzed.
9506 auto *OldType = dyn_cast<IntegerType>(LHS->getType());
9507 if (!OldType)
9508 break;
9509 // Prefer doubling the bitwidth over adding a single bit to make it more
9510 // likely that we use a legal type.
9511 auto *NewType =
9512 Type::getIntNTy(OldType->getContext(), OldType->getBitWidth() * 2);
9513 if (ICmpInst::isSigned(Pred)) {
9514 LHS = getSignExtendExpr(LHS, NewType);
9515 RHS = getSignExtendExpr(RHS, NewType);
9516 } else {
9517 LHS = getZeroExtendExpr(LHS, NewType);
9518 RHS = getZeroExtendExpr(RHS, NewType);
9519 }
9520 }
9522 [[fallthrough]];
9523 case ICmpInst::ICMP_SLT:
9524 case ICmpInst::ICMP_ULT: { // while (X < Y)
9525 bool IsSigned = ICmpInst::isSigned(Pred);
9526 ExitLimit EL = howManyLessThans(LHS, RHS, L, IsSigned, /*Invert=*/false,
9527 ControlsOnlyExit, AllowPredicates);
9528 if (EL.hasAnyInfo())
9529 return EL;
9530 break;
9531 }
9532 case ICmpInst::ICMP_SGE:
9533 case ICmpInst::ICMP_UGE:
9534 // Since the loop is finite, an invariant RHS cannot include the boundary
9535 // value, otherwise it would loop forever.
9536 if (!EnableFiniteLoopControl || !ControllingFiniteLoop ||
9537 !isLoopInvariant(RHS, L))
9538 break;
9540 [[fallthrough]];
9541 case ICmpInst::ICMP_SGT:
9542 case ICmpInst::ICMP_UGT: { // while (X > Y)
9543 // "X > Y" is analyzed as the equivalent "~X < ~Y".
9544 bool IsSigned = ICmpInst::isSigned(Pred);
9545 ExitLimit EL = howManyLessThans(LHS, RHS, L, IsSigned, /*Invert=*/true,
9546 ControlsOnlyExit, AllowPredicates);
9547 if (EL.hasAnyInfo())
9548 return EL;
9549 break;
9550 }
9551 default:
9552 break;
9553 }
9554
9555 return getCouldNotCompute();
9556}
9557
9558ScalarEvolution::ExitLimit
9559ScalarEvolution::computeExitLimitFromSingleExitSwitch(const Loop *L,
9560 SwitchInst *Switch,
9561 BasicBlock *ExitingBlock,
9562 bool ControlsOnlyExit) {
9563 assert(!L->contains(ExitingBlock) && "Not an exiting block!");
9564
9565 // Give up if the exit is the default dest of a switch.
9566 if (Switch->getDefaultDest() == ExitingBlock)
9567 return getCouldNotCompute();
9568
9569 assert(L->contains(Switch->getDefaultDest()) &&
9570 "Default case must not exit the loop!");
9571 const SCEV *LHS = getSCEVAtScope(Switch->getCondition(), L);
9572 const SCEV *RHS = getConstant(Switch->findCaseDest(ExitingBlock));
9573
9574 // while (X != Y) --> while (X-Y != 0)
9575 ExitLimit EL = howFarToZero(getMinusSCEV(LHS, RHS), L, ControlsOnlyExit);
9576 if (EL.hasAnyInfo())
9577 return EL;
9578
9579 return getCouldNotCompute();
9580}
9581
9582static ConstantInt *
9584 ScalarEvolution &SE) {
9585 const SCEV *InVal = SE.getConstant(C);
9586 const SCEV *Val = AddRec->evaluateAtIteration(InVal, SE);
9588 "Evaluation of SCEV at constant didn't fold correctly?");
9589 return cast<SCEVConstant>(Val)->getValue();
9590}
9591
9592ScalarEvolution::ExitLimit ScalarEvolution::computeShiftCompareExitLimit(
9593 Value *LHS, Value *RHSV, const Loop *L, ICmpInst::Predicate Pred) {
9594 ConstantInt *RHS = dyn_cast<ConstantInt>(RHSV);
9595 if (!RHS)
9596 return getCouldNotCompute();
9597
9598 const BasicBlock *Latch = L->getLoopLatch();
9599 if (!Latch)
9600 return getCouldNotCompute();
9601
9602 const BasicBlock *Predecessor = L->getLoopPredecessor();
9603 if (!Predecessor)
9604 return getCouldNotCompute();
9605
9606 // Return true if V is of the form "LHS `shift_op` <positive constant>".
9607 // Return LHS in OutLHS, shift_op in OutOpCode, and the shift amount in
9608 // OutShiftAmt.
9609 auto MatchPositiveShift = [](Value *V, Value *&OutLHS,
9610 Instruction::BinaryOps &OutOpCode,
9611 unsigned &OutShiftAmt) {
9612 using namespace PatternMatch;
9613
9614 ConstantInt *ShiftAmt;
9615 if (match(V, m_LShr(m_Value(OutLHS), m_ConstantInt(ShiftAmt))))
9616 OutOpCode = Instruction::LShr;
9617 else if (match(V, m_AShr(m_Value(OutLHS), m_ConstantInt(ShiftAmt))))
9618 OutOpCode = Instruction::AShr;
9619 else if (match(V, m_Shl(m_Value(OutLHS), m_ConstantInt(ShiftAmt))))
9620 OutOpCode = Instruction::Shl;
9621 else
9622 return false;
9623
9624 uint64_t Amt = ShiftAmt->getValue().getLimitedValue();
9625 if (Amt == 0 || Amt >= OutLHS->getType()->getScalarSizeInBits())
9626 return false;
9627 OutShiftAmt = Amt;
9628 return true;
9629 };
9630
9631 // Recognize a "shift recurrence" either of the form %iv or of %iv.shifted in
9632 //
9633 // loop:
9634 // %iv = phi i32 [ %iv.shifted, %loop ], [ %val, %preheader ]
9635 // %iv.shifted = lshr i32 %iv, <positive constant>
9636 //
9637 // Return true on a successful match. Return the corresponding PHI node (%iv
9638 // above) in PNOut, the opcode of the shift operation in OpCodeOut, and the
9639 // shift amount in ShiftAmtOut.
9640 auto MatchShiftRecurrence = [&](Value *V, PHINode *&PNOut,
9641 Instruction::BinaryOps &OpCodeOut,
9642 unsigned &ShiftAmtOut) {
9643 std::optional<Instruction::BinaryOps> PostShiftOpCode;
9644
9645 {
9647 Value *V;
9648 unsigned Amt;
9649
9650 // If we encounter a shift instruction, "peel off" the shift operation,
9651 // and remember that we did so. Later when we inspect %iv's backedge
9652 // value, we will make sure that the backedge value uses the same
9653 // operation.
9654 //
9655 // Note: the peeled shift operation does not have to be the same
9656 // instruction as the one feeding into the PHI's backedge value. We only
9657 // really care about it being the same *kind* of shift instruction --
9658 // that's all that is required for our later inferences to hold.
9659 if (MatchPositiveShift(LHS, V, OpC, Amt)) {
9660 PostShiftOpCode = OpC;
9661 LHS = V;
9662 }
9663 }
9664
9665 PNOut = dyn_cast<PHINode>(LHS);
9666 if (!PNOut || PNOut->getParent() != L->getHeader())
9667 return false;
9668
9669 Value *BEValue = PNOut->getIncomingValueForBlock(Latch);
9670 Value *OpLHS;
9671
9672 return
9673 // The backedge value for the PHI node must be a shift by a positive
9674 // amount
9675 MatchPositiveShift(BEValue, OpLHS, OpCodeOut, ShiftAmtOut) &&
9676
9677 // of the PHI node itself
9678 OpLHS == PNOut &&
9679
9680 // and the kind of shift should be match the kind of shift we peeled
9681 // off, if any.
9682 (!PostShiftOpCode || *PostShiftOpCode == OpCodeOut);
9683 };
9684
9685 PHINode *PN;
9687 unsigned ShiftAmt;
9688 if (!MatchShiftRecurrence(LHS, PN, OpCode, ShiftAmt))
9689 return getCouldNotCompute();
9690
9691 const DataLayout &DL = getDataLayout();
9692
9693 // The key rationale for this optimization is that for some kinds of shift
9694 // recurrences, the value of the recurrence "stabilizes" to either 0 or -1
9695 // within a finite number of iterations. If the condition guarding the
9696 // backedge (in the sense that the backedge is taken if the condition is true)
9697 // is false for the value the shift recurrence stabilizes to, then we know
9698 // that the backedge is taken only a finite number of times.
9699
9700 ConstantInt *StableValue = nullptr;
9701 switch (OpCode) {
9702 default:
9703 llvm_unreachable("Impossible case!");
9704
9705 case Instruction::AShr: {
9706 // {K,ashr,<positive-constant>} stabilizes to signum(K) in at most
9707 // bitwidth(K) iterations.
9708 Value *FirstValue = PN->getIncomingValueForBlock(Predecessor);
9709 KnownBits Known = computeKnownBits(FirstValue, DL, &AC,
9710 Predecessor->getTerminator(), &DT);
9711 auto *Ty = cast<IntegerType>(RHS->getType());
9712 if (Known.isNonNegative())
9713 StableValue = ConstantInt::get(Ty, 0);
9714 else if (Known.isNegative())
9715 StableValue = ConstantInt::get(Ty, -1, true);
9716 else
9717 return getCouldNotCompute();
9718
9719 break;
9720 }
9721 case Instruction::LShr:
9722 case Instruction::Shl:
9723 // Both {K,lshr,<positive-constant>} and {K,shl,<positive-constant>}
9724 // stabilize to 0 in at most bitwidth(K) iterations.
9725 StableValue = ConstantInt::get(cast<IntegerType>(RHS->getType()), 0);
9726 break;
9727 }
9728
9729 auto *Result =
9730 ConstantFoldCompareInstOperands(Pred, StableValue, RHS, DL, &TLI);
9731 assert(Result->getType()->isIntegerTy(1) &&
9732 "Otherwise cannot be an operand to a branch instruction");
9733
9734 if (Result->isNullValue()) {
9735 unsigned BitWidth = getTypeSizeInBits(RHS->getType());
9736 unsigned MaxBTC = BitWidth;
9737
9738 // For right-shift recurrences (lshr/ashr with non-negative start), we can
9739 // compute a tighter max backedge-taken count from the range of the start
9740 // value. After k shifts of ShiftAmt, value = start >> (k * ShiftAmt).
9741 // The value reaches 0 (the stable value) when k * ShiftAmt >=
9742 // activeBits(start), so max BTC = ceil(activeBits(maxStart) / ShiftAmt).
9743 if (OpCode == Instruction::LShr || OpCode == Instruction::AShr) {
9744 Value *StartValue = PN->getIncomingValueForBlock(Predecessor);
9745 const SCEV *StartSCEV = getSCEV(StartValue);
9746 APInt MaxStart = getUnsignedRangeMax(StartSCEV);
9747 if (MaxStart.isStrictlyPositive()) {
9748 unsigned ActiveBits = MaxStart.getActiveBits();
9749 unsigned RangeBTC = divideCeil(ActiveBits, ShiftAmt);
9750 MaxBTC = std::min(MaxBTC, RangeBTC);
9751 }
9752 }
9753
9754 const SCEV *UpperBound =
9756 return ExitLimit(getCouldNotCompute(), UpperBound, UpperBound, false);
9757 }
9758
9759 return getCouldNotCompute();
9760}
9761
9762/// Return true if we can constant fold an instruction of the specified type,
9763/// assuming that all operands were constants.
9764static bool canConstantFold(const Instruction *I,
9765 const TargetLibraryInfo *TLI) {
9769 return true;
9770
9771 if (const CallInst *CI = dyn_cast<CallInst>(I))
9772 if (const Function *F = CI->getCalledFunction())
9773 return canConstantFoldCallTo(CI, F, TLI);
9774 return false;
9775}
9776
9777/// Determine whether this instruction can constant evolve within this loop
9778/// assuming its operands can all constant evolve.
9779static bool canConstantEvolve(Instruction *I, const Loop *L,
9780 const TargetLibraryInfo *TLI) {
9781 // An instruction outside of the loop can't be derived from a loop PHI.
9782 if (!L->contains(I)) return false;
9783
9784 if (isa<PHINode>(I)) {
9785 // We don't currently keep track of the control flow needed to evaluate
9786 // PHIs, so we cannot handle PHIs inside of loops.
9787 return L->getHeader() == I->getParent();
9788 }
9789
9790 // If we won't be able to constant fold this expression even if the operands
9791 // are constants, bail early.
9792 return canConstantFold(I, TLI);
9793}
9794
9795/// getConstantEvolvingPHIOperands - Implement getConstantEvolvingPHI by
9796/// recursing through each instruction operand until reaching a loop header phi.
9797static PHINode *
9800 const TargetLibraryInfo *TLI, unsigned Depth) {
9802 return nullptr;
9803
9804 // Otherwise, we can evaluate this instruction if all of its operands are
9805 // constant or derived from a PHI node themselves.
9806 PHINode *PHI = nullptr;
9807 for (Value *Op : UseInst->operands()) {
9808 if (isa<Constant>(Op)) continue;
9809
9811 if (!OpInst || !canConstantEvolve(OpInst, L, TLI))
9812 return nullptr;
9813
9814 PHINode *P = dyn_cast<PHINode>(OpInst);
9815 if (!P)
9816 // If this operand is already visited, reuse the prior result.
9817 // We may have P != PHI if this is the deepest point at which the
9818 // inconsistent paths meet.
9819 P = PHIMap.lookup(OpInst);
9820 if (!P) {
9821 // Recurse and memoize the results, whether a phi is found or not.
9822 // This recursive call invalidates pointers into PHIMap.
9823 P = getConstantEvolvingPHIOperands(OpInst, L, PHIMap, TLI, Depth + 1);
9824 PHIMap[OpInst] = P;
9825 }
9826 if (!P)
9827 return nullptr; // Not evolving from PHI
9828 if (PHI && PHI != P)
9829 return nullptr; // Evolving from multiple different PHIs.
9830 PHI = P;
9831 }
9832 // This is a expression evolving from a constant PHI!
9833 return PHI;
9834}
9835
9836/// getConstantEvolvingPHI - Given an LLVM value and a loop, return a PHI node
9837/// in the loop that V is derived from. We allow arbitrary operations along the
9838/// way, but the operands of an operation must either be constants or a value
9839/// derived from a constant PHI. If this expression does not fit with these
9840/// constraints, return null.
9842 const TargetLibraryInfo *TLI) {
9844 if (!I || !canConstantEvolve(I, L, TLI))
9845 return nullptr;
9846
9847 if (PHINode *PN = dyn_cast<PHINode>(I))
9848 return PN;
9849
9850 // Record non-constant instructions contained by the loop.
9852 return getConstantEvolvingPHIOperands(I, L, PHIMap, TLI, 0);
9853}
9854
9855/// EvaluateExpression - Given an expression that passes the
9856/// getConstantEvolvingPHI predicate, evaluate its value assuming the PHI node
9857/// in the loop has the value PHIVal. If we can't fold this expression for some
9858/// reason, return null.
9861 const DataLayout &DL,
9862 const TargetLibraryInfo *TLI) {
9863 // Convenient constant check, but redundant for recursive calls.
9864 if (Constant *C = dyn_cast<Constant>(V)) return C;
9866 if (!I) return nullptr;
9867
9868 if (Constant *C = Vals.lookup(I)) return C;
9869
9870 // An instruction inside the loop depends on a value outside the loop that we
9871 // weren't given a mapping for, or a value such as a call inside the loop.
9872 if (!canConstantEvolve(I, L, TLI))
9873 return nullptr;
9874
9875 // An unmapped PHI can be due to a branch or another loop inside this loop,
9876 // or due to this not being the initial iteration through a loop where we
9877 // couldn't compute the evolution of this particular PHI last time.
9878 if (isa<PHINode>(I)) return nullptr;
9879
9880 std::vector<Constant*> Operands(I->getNumOperands());
9881
9882 for (unsigned i = 0, e = I->getNumOperands(); i != e; ++i) {
9883 Instruction *Operand = dyn_cast<Instruction>(I->getOperand(i));
9884 if (!Operand) {
9885 Operands[i] = dyn_cast<Constant>(I->getOperand(i));
9886 if (!Operands[i]) return nullptr;
9887 continue;
9888 }
9889 Constant *C = EvaluateExpression(Operand, L, Vals, DL, TLI);
9890 Vals[Operand] = C;
9891 if (!C) return nullptr;
9892 Operands[i] = C;
9893 }
9894
9895 return ConstantFoldInstOperands(I, Operands, DL, TLI,
9896 /*AllowNonDeterministic=*/false);
9897}
9898
9899
9900// If every incoming value to PN except the one for BB is a specific Constant,
9901// return that, else return nullptr.
9903 Constant *IncomingVal = nullptr;
9904
9905 for (unsigned i = 0, e = PN->getNumIncomingValues(); i != e; ++i) {
9906 if (PN->getIncomingBlock(i) == BB)
9907 continue;
9908
9909 auto *CurrentVal = dyn_cast<Constant>(PN->getIncomingValue(i));
9910 if (!CurrentVal)
9911 return nullptr;
9912
9913 if (IncomingVal != CurrentVal) {
9914 if (IncomingVal)
9915 return nullptr;
9916 IncomingVal = CurrentVal;
9917 }
9918 }
9919
9920 return IncomingVal;
9921}
9922
9923/// getConstantEvolutionLoopExitValue - If we know that the specified Phi is
9924/// in the header of its containing loop, we know the loop executes a
9925/// constant number of times, and the PHI node is just a recurrence
9926/// involving constants, fold it.
9927Constant *
9928ScalarEvolution::getConstantEvolutionLoopExitValue(PHINode *PN,
9929 const APInt &BEs,
9930 const Loop *L) {
9931 auto [I, Inserted] = ConstantEvolutionLoopExitValue.try_emplace(PN);
9932 if (!Inserted)
9933 return I->second;
9934
9936 return nullptr; // Not going to evaluate it.
9937
9938 Constant *&RetVal = I->second;
9939
9940 DenseMap<Instruction *, Constant *> CurrentIterVals;
9941 BasicBlock *Header = L->getHeader();
9942 assert(PN->getParent() == Header && "Can't evaluate PHI not in loop header!");
9943
9944 BasicBlock *Latch = L->getLoopLatch();
9945 if (!Latch)
9946 return nullptr;
9947
9948 for (PHINode &PHI : Header->phis()) {
9949 if (auto *StartCST = getOtherIncomingValue(&PHI, Latch))
9950 CurrentIterVals[&PHI] = StartCST;
9951 }
9952 if (!CurrentIterVals.count(PN))
9953 return RetVal = nullptr;
9954
9955 Value *BEValue = PN->getIncomingValueForBlock(Latch);
9956
9957 // Execute the loop symbolically to determine the exit value.
9958 assert(BEs.getActiveBits() < CHAR_BIT * sizeof(unsigned) &&
9959 "BEs is <= MaxBruteForceIterations which is an 'unsigned'!");
9960
9961 unsigned NumIterations = BEs.getZExtValue(); // must be in range
9962 unsigned IterationNum = 0;
9963 const DataLayout &DL = getDataLayout();
9964 for (; ; ++IterationNum) {
9965 if (IterationNum == NumIterations)
9966 return RetVal = CurrentIterVals[PN]; // Got exit value!
9967
9968 // Compute the value of the PHIs for the next iteration.
9969 // EvaluateExpression adds non-phi values to the CurrentIterVals map.
9970 DenseMap<Instruction *, Constant *> NextIterVals;
9971 Constant *NextPHI =
9972 EvaluateExpression(BEValue, L, CurrentIterVals, DL, &TLI);
9973 if (!NextPHI)
9974 return nullptr; // Couldn't evaluate!
9975 NextIterVals[PN] = NextPHI;
9976
9977 bool StoppedEvolving = NextPHI == CurrentIterVals[PN];
9978
9979 // Also evaluate the other PHI nodes. However, we don't get to stop if we
9980 // cease to be able to evaluate one of them or if they stop evolving,
9981 // because that doesn't necessarily prevent us from computing PN.
9983 for (const auto &I : CurrentIterVals) {
9984 PHINode *PHI = dyn_cast<PHINode>(I.first);
9985 if (!PHI || PHI == PN || PHI->getParent() != Header) continue;
9986 PHIsToCompute.emplace_back(PHI, I.second);
9987 }
9988 // We use two distinct loops because EvaluateExpression may invalidate any
9989 // iterators into CurrentIterVals.
9990 for (const auto &I : PHIsToCompute) {
9991 PHINode *PHI = I.first;
9992 Constant *&NextPHI = NextIterVals[PHI];
9993 if (!NextPHI) { // Not already computed.
9994 Value *BEValue = PHI->getIncomingValueForBlock(Latch);
9995 NextPHI = EvaluateExpression(BEValue, L, CurrentIterVals, DL, &TLI);
9996 }
9997 if (NextPHI != I.second)
9998 StoppedEvolving = false;
9999 }
10000
10001 // If all entries in CurrentIterVals == NextIterVals then we can stop
10002 // iterating, the loop can't continue to change.
10003 if (StoppedEvolving)
10004 return RetVal = CurrentIterVals[PN];
10005
10006 CurrentIterVals.swap(NextIterVals);
10007 }
10008}
10009
10010const SCEV *ScalarEvolution::computeExitCountExhaustively(const Loop *L,
10011 Value *Cond,
10012 bool ExitWhen) {
10013 PHINode *PN = getConstantEvolvingPHI(Cond, L, &TLI);
10014 if (!PN) return getCouldNotCompute();
10015
10016 // If the loop is canonicalized, the PHI will have exactly two entries.
10017 // That's the only form we support here.
10018 if (PN->getNumIncomingValues() != 2) return getCouldNotCompute();
10019
10020 DenseMap<Instruction *, Constant *> CurrentIterVals;
10021 BasicBlock *Header = L->getHeader();
10022 assert(PN->getParent() == Header && "Can't evaluate PHI not in loop header!");
10023
10024 BasicBlock *Latch = L->getLoopLatch();
10025 assert(Latch && "Should follow from NumIncomingValues == 2!");
10026
10027 for (PHINode &PHI : Header->phis()) {
10028 if (auto *StartCST = getOtherIncomingValue(&PHI, Latch))
10029 CurrentIterVals[&PHI] = StartCST;
10030 }
10031 if (!CurrentIterVals.count(PN))
10032 return getCouldNotCompute();
10033
10034 // Okay, we find a PHI node that defines the trip count of this loop. Execute
10035 // the loop symbolically to determine when the condition gets a value of
10036 // "ExitWhen".
10037 unsigned MaxIterations = MaxBruteForceIterations; // Limit analysis.
10038 const DataLayout &DL = getDataLayout();
10039 for (unsigned IterationNum = 0; IterationNum != MaxIterations;++IterationNum){
10040 auto *CondVal = dyn_cast_or_null<ConstantInt>(
10041 EvaluateExpression(Cond, L, CurrentIterVals, DL, &TLI));
10042
10043 // Couldn't symbolically evaluate.
10044 if (!CondVal) return getCouldNotCompute();
10045
10046 if (CondVal->getValue() == uint64_t(ExitWhen)) {
10047 ++NumBruteForceTripCountsComputed;
10048 return getConstant(Type::getInt32Ty(getContext()), IterationNum);
10049 }
10050
10051 // Update all the PHI nodes for the next iteration.
10052 DenseMap<Instruction *, Constant *> NextIterVals;
10053
10054 // Create a list of which PHIs we need to compute. We want to do this before
10055 // calling EvaluateExpression on them because that may invalidate iterators
10056 // into CurrentIterVals.
10057 SmallVector<PHINode *, 8> PHIsToCompute;
10058 for (const auto &I : CurrentIterVals) {
10059 PHINode *PHI = dyn_cast<PHINode>(I.first);
10060 if (!PHI || PHI->getParent() != Header) continue;
10061 PHIsToCompute.push_back(PHI);
10062 }
10063 for (PHINode *PHI : PHIsToCompute) {
10064 Constant *&NextPHI = NextIterVals[PHI];
10065 if (NextPHI) continue; // Already computed!
10066
10067 Value *BEValue = PHI->getIncomingValueForBlock(Latch);
10068 NextPHI = EvaluateExpression(BEValue, L, CurrentIterVals, DL, &TLI);
10069 }
10070 CurrentIterVals.swap(NextIterVals);
10071 }
10072
10073 // Too many iterations were needed to evaluate.
10074 return getCouldNotCompute();
10075}
10076
10078 auto &Values = ValuesAtScopes[V];
10079 // Check to see if we've folded this expression at this loop before.
10080 for (auto &LS : Values)
10081 if (LS.first == L)
10082 return LS.second ? LS.second : SCEVUse(V);
10083
10084 Values.emplace_back(L, nullptr);
10085
10086 // Otherwise compute it.
10087 SCEVUse C = computeSCEVAtScope(V, L);
10088 for (auto &LS : reverse(ValuesAtScopes[V]))
10089 if (LS.first == L) {
10090 LS.second = C;
10091 // Record the dependency under the bare expression: invalidation walks
10092 // expressions, and any use flags on C do not change which expression
10093 // this is the value at scope of.
10094 if (!isa<SCEVConstant>(C))
10095 ValuesAtScopesUsers[C.getPointer()].push_back({L, V});
10096 break;
10097 }
10098 return C;
10099}
10100
10102 const BasicBlock *ExitingBlock) {
10103 SCEVUse ExitValue = getSCEVAtScope(V, L->getParentLoop());
10104 if (!isLoopInvariant(ExitValue, L)) {
10105 // If we failed to evaluate it in the outer scope, try to evaluate an
10106 // addrec for the specific exit.
10107 // TODO: Generalize this to other expressions.
10108 const SCEV *ExitCount = getExitCount(L, ExitingBlock);
10109 if (!isa<SCEVCouldNotCompute>(ExitCount))
10110 if (auto *AddRec = dyn_cast<SCEVAddRecExpr>(V))
10111 if (AddRec->getLoop() == L)
10112 ExitValue = AddRec->evaluateAtIteration(ExitCount, *this);
10113 }
10114 return ExitValue;
10115}
10116
10117/// This builds up a Constant using the ConstantExpr interface. That way, we
10118/// will return Constants for objects which aren't represented by a
10119/// SCEVConstant, because SCEVConstant is restricted to ConstantInt.
10120/// Returns NULL if the SCEV isn't representable as a Constant.
10122 switch (V->getSCEVType()) {
10123 case scCouldNotCompute:
10124 case scAddRecExpr:
10125 case scVScale:
10126 return nullptr;
10127 case scConstant:
10128 return cast<SCEVConstant>(V)->getValue();
10129 case scUnknown:
10131 case scPtrToAddr: {
10133 if (Constant *CastOp = BuildConstantFromSCEV(P2I->getOperand()))
10134 return ConstantExpr::getPtrToAddr(CastOp, P2I->getType());
10135
10136 return nullptr;
10137 }
10138 case scTruncate: {
10140 if (Constant *CastOp = BuildConstantFromSCEV(ST->getOperand()))
10141 return ConstantExpr::getTrunc(CastOp, ST->getType());
10142 return nullptr;
10143 }
10144 case scAddExpr: {
10145 const SCEVAddExpr *SA = cast<SCEVAddExpr>(V);
10146 Constant *C = nullptr;
10147 for (const SCEV *Op : SA->operands()) {
10149 if (!OpC)
10150 return nullptr;
10151 if (!C) {
10152 C = OpC;
10153 continue;
10154 }
10155 assert(!C->getType()->isPointerTy() &&
10156 "Can only have one pointer, and it must be last");
10157 if (OpC->getType()->isPointerTy()) {
10158 // The offsets have been converted to bytes. We can add bytes using
10159 // an i8 GEP.
10160 C = ConstantExpr::getPtrAdd(OpC, C);
10161 } else {
10162 C = ConstantExpr::getAdd(C, OpC);
10163 }
10164 }
10165 return C;
10166 }
10167 case scMulExpr:
10168 case scSignExtend:
10169 case scZeroExtend:
10170 case scUDivExpr:
10171 case scSMaxExpr:
10172 case scUMaxExpr:
10173 case scSMinExpr:
10174 case scUMinExpr:
10176 return nullptr;
10177 }
10178 llvm_unreachable("Unknown SCEV kind!");
10179}
10180
10181const SCEV *ScalarEvolution::getWithOperands(const SCEV *S,
10182 SmallVectorImpl<SCEVUse> &NewOps) {
10183 switch (S->getSCEVType()) {
10184 case scTruncate:
10185 case scZeroExtend:
10186 case scSignExtend:
10187 case scPtrToAddr:
10188 return getCastExpr(S->getSCEVType(), NewOps[0], S->getType());
10189 case scAddRecExpr: {
10190 auto *AddRec = cast<SCEVAddRecExpr>(S);
10191 return getAddRecExpr(NewOps, AddRec->getLoop(), AddRec->getNoWrapFlags());
10192 }
10193 case scAddExpr:
10194 return getAddExpr(NewOps, cast<SCEVAddExpr>(S)->getNoWrapFlags());
10195 case scMulExpr:
10196 return getMulExpr(NewOps, cast<SCEVMulExpr>(S)->getNoWrapFlags());
10197 case scUDivExpr:
10198 return getUDivExpr(NewOps[0], NewOps[1]);
10199 case scUMaxExpr:
10200 case scSMaxExpr:
10201 case scUMinExpr:
10202 case scSMinExpr:
10203 return getMinMaxExpr(S->getSCEVType(), NewOps);
10205 return getSequentialMinMaxExpr(S->getSCEVType(), NewOps);
10206 case scConstant:
10207 case scVScale:
10208 case scUnknown:
10209 return S;
10210 case scCouldNotCompute:
10211 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
10212 }
10213 llvm_unreachable("Unknown SCEV kind!");
10214}
10215
10216SCEVUse ScalarEvolution::computeSCEVAtScope(const SCEV *V, const Loop *L) {
10217 switch (V->getSCEVType()) {
10218 case scConstant:
10219 case scVScale:
10220 return V;
10221 case scAddRecExpr: {
10222 // If this is a loop recurrence for a loop that does not contain L, then we
10223 // are dealing with the final value computed by the loop.
10224 const SCEVAddRecExpr *AddRec = cast<SCEVAddRecExpr>(V);
10225 // First, attempt to evaluate each operand.
10226 // Avoid performing the look-up in the common case where the specified
10227 // expression has no loop-variant portions.
10228 for (unsigned i = 0, e = AddRec->getNumOperands(); i != e; ++i) {
10229 SCEVUse OpAtScope = getSCEVAtScope(AddRec->getOperand(i), L);
10230 if (OpAtScope == AddRec->getOperand(i))
10231 continue;
10232
10233 // Okay, at least one of these operands is loop variant but might be
10234 // foldable. Build a new instance of the folded commutative expression.
10236 NewOps.reserve(AddRec->getNumOperands());
10237 append_range(NewOps, AddRec->operands().take_front(i));
10238 NewOps.push_back(OpAtScope);
10239 for (++i; i != e; ++i)
10240 NewOps.push_back(getSCEVAtScope(AddRec->getOperand(i), L));
10241
10242 const SCEV *FoldedRec = getAddRecExpr(
10243 NewOps, AddRec->getLoop(), AddRec->getNoWrapFlags(SCEV::FlagNW));
10244 AddRec = dyn_cast<SCEVAddRecExpr>(FoldedRec);
10245 // The addrec may be folded to a nonrecurrence, for example, if the
10246 // induction variable is multiplied by zero after constant folding. Go
10247 // ahead and return the folded value.
10248 if (!AddRec)
10249 return FoldedRec;
10250 break;
10251 }
10252
10253 // If the scope is outside the addrec's loop, evaluate it by using the
10254 // loop exit value of the addrec.
10255 if (!AddRec->getLoop()->contains(L)) {
10256 SCEVUse ExitValue = AddRec->getExitValue(*this);
10257 if (isa<SCEVCouldNotCompute>(ExitValue))
10258 return AddRec;
10259 return ExitValue;
10260 }
10261
10262 return AddRec;
10263 }
10264 case scTruncate:
10265 case scZeroExtend:
10266 case scSignExtend:
10267 case scPtrToAddr:
10268 case scAddExpr:
10269 case scMulExpr:
10270 case scUDivExpr:
10271 case scUMaxExpr:
10272 case scSMaxExpr:
10273 case scUMinExpr:
10274 case scSMinExpr:
10275 case scSequentialUMinExpr: {
10276 ArrayRef<SCEVUse> Ops = V->operands();
10277 // Avoid performing the look-up in the common case where the specified
10278 // expression has no loop-variant portions.
10279 for (unsigned i = 0, e = Ops.size(); i != e; ++i) {
10280 SCEVUse OpAtScope = getSCEVAtScope(Ops[i].getPointer(), L);
10281 if (OpAtScope != Ops[i].getPointer()) {
10282 // Okay, at least one of these operands is loop variant but might be
10283 // foldable. Build a new instance of the folded commutative expression.
10285 NewOps.reserve(Ops.size());
10286 append_range(NewOps, Ops.take_front(i));
10287 NewOps.push_back(OpAtScope);
10288
10289 for (++i; i != e; ++i) {
10290 OpAtScope = getSCEVAtScope(Ops[i].getPointer(), L);
10291 NewOps.push_back(OpAtScope);
10292 }
10293
10294 return getWithOperands(V, NewOps);
10295 }
10296 }
10297 // If we got here, all operands are loop invariant.
10298 return V;
10299 }
10300 case scUnknown: {
10301 // If this instruction is evolved from a constant-evolving PHI, compute the
10302 // exit value from the loop without using SCEVs.
10303 const SCEVUnknown *SU = cast<SCEVUnknown>(V);
10305 if (!I)
10306 return V; // This is some other type of SCEVUnknown, just return it.
10307
10308 if (PHINode *PN = dyn_cast<PHINode>(I)) {
10309 const Loop *CurrLoop = this->LI[I->getParent()];
10310 // Looking for loop exit value.
10311 if (CurrLoop && CurrLoop->getParentLoop() == L &&
10312 PN->getParent() == CurrLoop->getHeader()) {
10313 // Okay, there is no closed form solution for the PHI node. Check
10314 // to see if the loop that contains it has a known backedge-taken
10315 // count. If so, we may be able to force computation of the exit
10316 // value.
10317 const SCEV *BackedgeTakenCount = getBackedgeTakenCount(CurrLoop);
10318 // This trivial case can show up in some degenerate cases where
10319 // the incoming IR has not yet been fully simplified.
10320 if (BackedgeTakenCount->isZero()) {
10321 Value *InitValue = nullptr;
10322 bool MultipleInitValues = false;
10323 for (unsigned i = 0; i < PN->getNumIncomingValues(); i++) {
10324 if (!CurrLoop->contains(PN->getIncomingBlock(i))) {
10325 if (!InitValue)
10326 InitValue = PN->getIncomingValue(i);
10327 else if (InitValue != PN->getIncomingValue(i)) {
10328 MultipleInitValues = true;
10329 break;
10330 }
10331 }
10332 }
10333 if (!MultipleInitValues && InitValue)
10334 return getSCEV(InitValue);
10335 }
10336 // Do we have a loop invariant value flowing around the backedge
10337 // for a loop which must execute the backedge?
10338 if (!isa<SCEVCouldNotCompute>(BackedgeTakenCount) &&
10339 isKnownNonZero(BackedgeTakenCount) &&
10340 PN->getNumIncomingValues() == 2) {
10341
10342 unsigned InLoopPred =
10343 CurrLoop->contains(PN->getIncomingBlock(0)) ? 0 : 1;
10344 Value *BackedgeVal = PN->getIncomingValue(InLoopPred);
10345 if (CurrLoop->isLoopInvariant(BackedgeVal))
10346 return getSCEV(BackedgeVal);
10347 }
10348 if (auto *BTCC = dyn_cast<SCEVConstant>(BackedgeTakenCount)) {
10349 // Okay, we know how many times the containing loop executes. If
10350 // this is a constant evolving PHI node, get the final value at
10351 // the specified iteration number.
10352 Constant *RV =
10353 getConstantEvolutionLoopExitValue(PN, BTCC->getAPInt(), CurrLoop);
10354 if (RV)
10355 return getSCEV(RV);
10356 }
10357 }
10358 }
10359
10360 // Okay, this is an expression that we cannot symbolically evaluate
10361 // into a SCEV. Check to see if it's possible to symbolically evaluate
10362 // the arguments into constants, and if so, try to constant propagate the
10363 // result. This is particularly useful for computing loop exit values.
10364 if (!canConstantFold(I, &TLI))
10365 return V; // This is some other type of SCEVUnknown, just return it.
10366
10367 SmallVector<Constant *, 4> Operands;
10368 Operands.reserve(I->getNumOperands());
10369 bool MadeImprovement = false;
10370 for (Value *Op : I->operands()) {
10371 if (Constant *C = dyn_cast<Constant>(Op)) {
10372 Operands.push_back(C);
10373 continue;
10374 }
10375
10376 // If any of the operands is non-constant and if they are
10377 // non-integer and non-pointer, don't even try to analyze them
10378 // with scev techniques.
10379 if (!isSCEVable(Op->getType()))
10380 return V;
10381
10382 const SCEV *OrigV = getSCEV(Op);
10383 const SCEV *OpV = getSCEVAtScope(OrigV, L);
10384 MadeImprovement |= OrigV != OpV;
10385
10387 if (!C)
10388 return V;
10389 assert(C->getType() == Op->getType() && "Type mismatch");
10390 Operands.push_back(C);
10391 }
10392
10393 // Check to see if getSCEVAtScope actually made an improvement.
10394 if (!MadeImprovement)
10395 return V; // This is some other type of SCEVUnknown, just return it.
10396
10397 Constant *C = nullptr;
10398 const DataLayout &DL = getDataLayout();
10399 C = ConstantFoldInstOperands(I, Operands, DL, &TLI,
10400 /*AllowNonDeterministic=*/false);
10401 if (!C)
10402 return V;
10403 return getSCEV(C);
10404 }
10405 case scCouldNotCompute:
10406 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
10407 }
10408 llvm_unreachable("Unknown SCEV type!");
10409}
10410
10412 return getSCEVAtScope(getSCEV(V), L);
10413}
10414
10415const SCEV *ScalarEvolution::stripInjectiveFunctions(const SCEV *S) const {
10417 return stripInjectiveFunctions(ZExt->getOperand());
10419 return stripInjectiveFunctions(SExt->getOperand());
10420 return S;
10421}
10422
10423/// Finds the minimum unsigned root of the following equation:
10424///
10425/// A * X = B (mod N)
10426///
10427/// where N = 2^BW and BW is the common bit width of A and B. The signedness of
10428/// A and B isn't important.
10429///
10430/// If the equation does not have a solution, SCEVCouldNotCompute is returned.
10431static const SCEV *
10434 ScalarEvolution &SE, const Loop *L) {
10435 uint32_t BW = A.getBitWidth();
10436 assert(BW == SE.getTypeSizeInBits(B->getType()));
10437 assert(A != 0 && "A must be non-zero.");
10438
10439 // 1. D = gcd(A, N)
10440 //
10441 // The gcd of A and N may have only one prime factor: 2. The number of
10442 // trailing zeros in A is its multiplicity
10443 uint32_t Mult2 = A.countr_zero();
10444 // D = 2^Mult2
10445
10446 // 2. Check if B is divisible by D.
10447 //
10448 // B is divisible by D if and only if the multiplicity of prime factor 2 for B
10449 // is not less than multiplicity of this prime factor for D.
10450 unsigned MinTZ = SE.getMinTrailingZeros(B);
10451 // Try again with the terminator of the loop predecessor for context-specific
10452 // result, if MinTZ s too small.
10453 if (MinTZ < Mult2 && L->getLoopPredecessor())
10454 MinTZ = SE.getMinTrailingZeros(B, L->getLoopPredecessor()->getTerminator());
10455 if (MinTZ < Mult2) {
10456 // Check if we can prove there's no remainder using URem.
10457 const SCEV *URem =
10458 SE.getURemExpr(B, SE.getConstant(APInt::getOneBitSet(BW, Mult2)));
10459 const SCEV *Zero = SE.getZero(B->getType());
10460 if (!SE.isKnownPredicate(CmpInst::ICMP_EQ, URem, Zero)) {
10461 // Try to add a predicate ensuring B is a multiple of 1 << Mult2.
10462 if (!Predicates)
10463 return SE.getCouldNotCompute();
10464
10465 // Avoid adding a predicate that is known to be false.
10466 if (SE.isKnownPredicate(CmpInst::ICMP_NE, URem, Zero))
10467 return SE.getCouldNotCompute();
10468 Predicates->push_back(SE.getEqualPredicate(URem, Zero));
10469 }
10470 }
10471
10472 // 3. Compute I: the multiplicative inverse of (A / D) in arithmetic
10473 // modulo (N / D).
10474 //
10475 // If D == 1, (N / D) == N == 2^BW, so we need one extra bit to represent
10476 // (N / D) in general. The inverse itself always fits into BW bits, though,
10477 // so we immediately truncate it.
10478 APInt AD = A.lshr(Mult2).trunc(BW - Mult2); // AD = A / D
10479 APInt I = AD.multiplicativeInverse().zext(BW);
10480
10481 // 4. Compute the minimum unsigned root of the equation:
10482 // I * (B / D) mod (N / D)
10483 // To simplify the computation, we factor out the divide by D:
10484 // (I * B mod N) / D
10485 const SCEV *D = SE.getConstant(APInt::getOneBitSet(BW, Mult2));
10486 return SE.getUDivExactExpr(SE.getMulExpr(B, SE.getConstant(I)), D);
10487}
10488
10489/// For a given quadratic addrec, generate coefficients of the corresponding
10490/// quadratic equation, multiplied by a common value to ensure that they are
10491/// integers.
10492/// The returned value is a tuple { A, B, C, M, BitWidth }, where
10493/// Ax^2 + Bx + C is the quadratic function, M is the value that A, B and C
10494/// were multiplied by, and BitWidth is the bit width of the original addrec
10495/// coefficients.
10496/// This function returns std::nullopt if the addrec coefficients are not
10497/// compile- time constants.
10498static std::optional<std::tuple<APInt, APInt, APInt, APInt, unsigned>>
10500 assert(AddRec->getNumOperands() == 3 && "This is not a quadratic chrec!");
10501 const SCEVConstant *LC = dyn_cast<SCEVConstant>(AddRec->getOperand(0));
10502 const SCEVConstant *MC = dyn_cast<SCEVConstant>(AddRec->getOperand(1));
10503 const SCEVConstant *NC = dyn_cast<SCEVConstant>(AddRec->getOperand(2));
10504 LLVM_DEBUG(dbgs() << __func__ << ": analyzing quadratic addrec: "
10505 << *AddRec << '\n');
10506
10507 // We currently can only solve this if the coefficients are constants.
10508 if (!LC || !MC || !NC) {
10509 LLVM_DEBUG(dbgs() << __func__ << ": coefficients are not constant\n");
10510 return std::nullopt;
10511 }
10512
10513 APInt L = LC->getAPInt();
10514 APInt M = MC->getAPInt();
10515 APInt N = NC->getAPInt();
10516 assert(!N.isZero() && "This is not a quadratic addrec");
10517
10518 unsigned BitWidth = LC->getAPInt().getBitWidth();
10519 unsigned NewWidth = BitWidth + 1;
10520 LLVM_DEBUG(dbgs() << __func__ << ": addrec coeff bw: "
10521 << BitWidth << '\n');
10522 // The sign-extension (as opposed to a zero-extension) here matches the
10523 // extension used in SolveQuadraticEquationWrap (with the same motivation).
10524 N = N.sext(NewWidth);
10525 M = M.sext(NewWidth);
10526 L = L.sext(NewWidth);
10527
10528 // The increments are M, M+N, M+2N, ..., so the accumulated values are
10529 // L+M, (L+M)+(M+N), (L+M)+(M+N)+(M+2N), ..., that is,
10530 // L+M, L+2M+N, L+3M+3N, ...
10531 // After n iterations the accumulated value Acc is L + nM + n(n-1)/2 N.
10532 //
10533 // The equation Acc = 0 is then
10534 // L + nM + n(n-1)/2 N = 0, or 2L + 2M n + n(n-1) N = 0.
10535 // In a quadratic form it becomes:
10536 // N n^2 + (2M-N) n + 2L = 0.
10537
10538 APInt A = N;
10539 APInt B = 2 * M - A;
10540 APInt C = 2 * L;
10541 APInt T = APInt(NewWidth, 2);
10542 LLVM_DEBUG(dbgs() << __func__ << ": equation " << A << "x^2 + " << B
10543 << "x + " << C << ", coeff bw: " << NewWidth
10544 << ", multiplied by " << T << '\n');
10545 return std::make_tuple(A, B, C, T, BitWidth);
10546}
10547
10548/// Helper function to compare optional APInts:
10549/// (a) if X and Y both exist, return min(X, Y),
10550/// (b) if neither X nor Y exist, return std::nullopt,
10551/// (c) if exactly one of X and Y exists, return that value.
10552static std::optional<APInt> MinOptional(std::optional<APInt> X,
10553 std::optional<APInt> Y) {
10554 if (X && Y) {
10555 unsigned W = std::max(X->getBitWidth(), Y->getBitWidth());
10556 APInt XW = X->sext(W);
10557 APInt YW = Y->sext(W);
10558 return XW.slt(YW) ? *X : *Y;
10559 }
10560 if (!X && !Y)
10561 return std::nullopt;
10562 return X ? *X : *Y;
10563}
10564
10565/// Helper function to truncate an optional APInt to a given BitWidth.
10566/// When solving addrec-related equations, it is preferable to return a value
10567/// that has the same bit width as the original addrec's coefficients. If the
10568/// solution fits in the original bit width, truncate it (except for i1).
10569/// Returning a value of a different bit width may inhibit some optimizations.
10570///
10571/// In general, a solution to a quadratic equation generated from an addrec
10572/// may require BW+1 bits, where BW is the bit width of the addrec's
10573/// coefficients. The reason is that the coefficients of the quadratic
10574/// equation are BW+1 bits wide (to avoid truncation when converting from
10575/// the addrec to the equation).
10576static std::optional<APInt> TruncIfPossible(std::optional<APInt> X,
10577 unsigned BitWidth) {
10578 if (!X)
10579 return std::nullopt;
10580 unsigned W = X->getBitWidth();
10582 return X->trunc(BitWidth);
10583 return X;
10584}
10585
10586/// Let c(n) be the value of the quadratic chrec {L,+,M,+,N} after n
10587/// iterations. The values L, M, N are assumed to be signed, and they
10588/// should all have the same bit widths.
10589/// Find the least n >= 0 such that c(n) = 0 in the arithmetic modulo 2^BW,
10590/// where BW is the bit width of the addrec's coefficients.
10591/// If the calculated value is a BW-bit integer (for BW > 1), it will be
10592/// returned as such, otherwise the bit width of the returned value may
10593/// be greater than BW.
10594///
10595/// This function returns std::nullopt if
10596/// (a) the addrec coefficients are not constant, or
10597/// (b) SolveQuadraticEquationWrap was unable to find a solution. For cases
10598/// like x^2 = 5, no integer solutions exist, in other cases an integer
10599/// solution may exist, but SolveQuadraticEquationWrap may fail to find it.
10600static std::optional<APInt>
10602 APInt A, B, C, M;
10603 unsigned BitWidth;
10604 auto T = GetQuadraticEquation(AddRec);
10605 if (!T)
10606 return std::nullopt;
10607
10608 std::tie(A, B, C, M, BitWidth) = *T;
10609 LLVM_DEBUG(dbgs() << __func__ << ": solving for unsigned overflow\n");
10610 std::optional<APInt> X =
10612 if (!X)
10613 return std::nullopt;
10614
10615 ConstantInt *CX = ConstantInt::get(SE.getContext(), *X);
10616 ConstantInt *V = EvaluateConstantChrecAtConstant(AddRec, CX, SE);
10617 if (!V->isZero())
10618 return std::nullopt;
10619
10620 return TruncIfPossible(X, BitWidth);
10621}
10622
10623/// Let c(n) be the value of the quadratic chrec {0,+,M,+,N} after n
10624/// iterations. The values M, N are assumed to be signed, and they
10625/// should all have the same bit widths.
10626/// Find the least n such that c(n) does not belong to the given range,
10627/// while c(n-1) does.
10628///
10629/// This function returns std::nullopt if
10630/// (a) the addrec coefficients are not constant, or
10631/// (b) SolveQuadraticEquationWrap was unable to find a solution for the
10632/// bounds of the range.
10633static std::optional<APInt>
10635 const ConstantRange &Range, ScalarEvolution &SE) {
10636 assert(AddRec->getOperand(0)->isZero() &&
10637 "Starting value of addrec should be 0");
10638 LLVM_DEBUG(dbgs() << __func__ << ": solving boundary crossing for range "
10639 << Range << ", addrec " << *AddRec << '\n');
10640 // This case is handled in getNumIterationsInRange. Here we can assume that
10641 // we start in the range.
10642 assert(Range.contains(APInt(SE.getTypeSizeInBits(AddRec->getType()), 0)) &&
10643 "Addrec's initial value should be in range");
10644
10645 APInt A, B, C, M;
10646 unsigned BitWidth;
10647 auto T = GetQuadraticEquation(AddRec);
10648 if (!T)
10649 return std::nullopt;
10650
10651 // Be careful about the return value: there can be two reasons for not
10652 // returning an actual number. First, if no solutions to the equations
10653 // were found, and second, if the solutions don't leave the given range.
10654 // The first case means that the actual solution is "unknown", the second
10655 // means that it's known, but not valid. If the solution is unknown, we
10656 // cannot make any conclusions.
10657 // Return a pair: the optional solution and a flag indicating if the
10658 // solution was found.
10659 auto SolveForBoundary =
10660 [&](APInt Bound) -> std::pair<std::optional<APInt>, bool> {
10661 // Solve for signed overflow and unsigned overflow, pick the lower
10662 // solution.
10663 LLVM_DEBUG(dbgs() << "SolveQuadraticAddRecRange: checking boundary "
10664 << Bound << " (before multiplying by " << M << ")\n");
10665 Bound *= M; // The quadratic equation multiplier.
10666
10667 std::optional<APInt> SO;
10668 if (BitWidth > 1) {
10669 LLVM_DEBUG(dbgs() << "SolveQuadraticAddRecRange: solving for "
10670 "signed overflow\n");
10672 }
10673 LLVM_DEBUG(dbgs() << "SolveQuadraticAddRecRange: solving for "
10674 "unsigned overflow\n");
10675 std::optional<APInt> UO =
10677
10678 auto LeavesRange = [&] (const APInt &X) {
10679 ConstantInt *C0 = ConstantInt::get(SE.getContext(), X);
10680 ConstantInt *V0 = EvaluateConstantChrecAtConstant(AddRec, C0, SE);
10681 if (Range.contains(V0->getValue()))
10682 return false;
10683 // X should be at least 1, so X-1 is non-negative.
10684 ConstantInt *C1 = ConstantInt::get(SE.getContext(), X-1);
10686 if (Range.contains(V1->getValue()))
10687 return true;
10688 return false;
10689 };
10690
10691 // If SolveQuadraticEquationWrap returns std::nullopt, it means that there
10692 // can be a solution, but the function failed to find it. We cannot treat it
10693 // as "no solution".
10694 if (!SO || !UO)
10695 return {std::nullopt, false};
10696
10697 // Check the smaller value first to see if it leaves the range.
10698 // At this point, both SO and UO must have values.
10699 std::optional<APInt> Min = MinOptional(SO, UO);
10700 if (LeavesRange(*Min))
10701 return { Min, true };
10702 std::optional<APInt> Max = Min == SO ? UO : SO;
10703 if (LeavesRange(*Max))
10704 return { Max, true };
10705
10706 // Solutions were found, but were eliminated, hence the "true".
10707 return {std::nullopt, true};
10708 };
10709
10710 std::tie(A, B, C, M, BitWidth) = *T;
10711 // Lower bound is inclusive, subtract 1 to represent the exiting value.
10712 APInt Lower = Range.getLower().sext(A.getBitWidth()) - 1;
10713 APInt Upper = Range.getUpper().sext(A.getBitWidth());
10714 auto SL = SolveForBoundary(Lower);
10715 auto SU = SolveForBoundary(Upper);
10716 // If any of the solutions was unknown, no meaninigful conclusions can
10717 // be made.
10718 if (!SL.second || !SU.second)
10719 return std::nullopt;
10720
10721 // Claim: The correct solution is not some value between Min and Max.
10722 //
10723 // Justification: Assuming that Min and Max are different values, one of
10724 // them is when the first signed overflow happens, the other is when the
10725 // first unsigned overflow happens. Crossing the range boundary is only
10726 // possible via an overflow (treating 0 as a special case of it, modeling
10727 // an overflow as crossing k*2^W for some k).
10728 //
10729 // The interesting case here is when Min was eliminated as an invalid
10730 // solution, but Max was not. The argument is that if there was another
10731 // overflow between Min and Max, it would also have been eliminated if
10732 // it was considered.
10733 //
10734 // For a given boundary, it is possible to have two overflows of the same
10735 // type (signed/unsigned) without having the other type in between: this
10736 // can happen when the vertex of the parabola is between the iterations
10737 // corresponding to the overflows. This is only possible when the two
10738 // overflows cross k*2^W for the same k. In such case, if the second one
10739 // left the range (and was the first one to do so), the first overflow
10740 // would have to enter the range, which would mean that either we had left
10741 // the range before or that we started outside of it. Both of these cases
10742 // are contradictions.
10743 //
10744 // Claim: In the case where SolveForBoundary returns std::nullopt, the correct
10745 // solution is not some value between the Max for this boundary and the
10746 // Min of the other boundary.
10747 //
10748 // Justification: Assume that we had such Max_A and Min_B corresponding
10749 // to range boundaries A and B and such that Max_A < Min_B. If there was
10750 // a solution between Max_A and Min_B, it would have to be caused by an
10751 // overflow corresponding to either A or B. It cannot correspond to B,
10752 // since Min_B is the first occurrence of such an overflow. If it
10753 // corresponded to A, it would have to be either a signed or an unsigned
10754 // overflow that is larger than both eliminated overflows for A. But
10755 // between the eliminated overflows and this overflow, the values would
10756 // cover the entire value space, thus crossing the other boundary, which
10757 // is a contradiction.
10758
10759 return TruncIfPossible(MinOptional(SL.first, SU.first), BitWidth);
10760}
10761
10762ScalarEvolution::ExitLimit ScalarEvolution::howFarToZero(const SCEV *V,
10763 const Loop *L,
10764 bool ControlsOnlyExit,
10765 bool AllowPredicates) {
10766
10767 // This is only used for loops with a "x != y" exit test. The exit condition
10768 // is now expressed as a single expression, V = x-y. So the exit test is
10769 // effectively V != 0. We know and take advantage of the fact that this
10770 // expression only being used in a comparison by zero context.
10771
10773 // If the value is a constant
10774 if (const SCEVConstant *C = dyn_cast<SCEVConstant>(V)) {
10775 // If the value is already zero, the branch will execute zero times.
10776 if (C->getValue()->isZero()) return C;
10777 return getCouldNotCompute(); // Otherwise it will loop infinitely.
10778 }
10779
10780 const SCEVAddRecExpr *AddRec =
10781 dyn_cast<SCEVAddRecExpr>(stripInjectiveFunctions(V));
10782
10783 if (!AddRec && AllowPredicates)
10784 // Try to make this an AddRec using runtime tests, in the first X
10785 // iterations of this loop, where X is the SCEV expression found by the
10786 // algorithm below.
10787 AddRec = convertSCEVToAddRecWithPredicates(V, L, Predicates);
10788
10789 if (!AddRec || AddRec->getLoop() != L)
10790 return getCouldNotCompute();
10791
10792 // If this is a quadratic (3-term) AddRec {L,+,M,+,N}, find the roots of
10793 // the quadratic equation to solve it.
10794 if (AddRec->isQuadratic() && AddRec->getType()->isIntegerTy()) {
10795 // We can only use this value if the chrec ends up with an exact zero
10796 // value at this index. When solving for "X*X != 5", for example, we
10797 // should not accept a root of 2.
10798 if (auto S = SolveQuadraticAddRecExact(AddRec, *this)) {
10799 const auto *R = cast<SCEVConstant>(getConstant(*S));
10800 return ExitLimit(R, R, R, false, Predicates);
10801 }
10802 return getCouldNotCompute();
10803 }
10804
10805 // Otherwise we can only handle this if it is affine.
10806 if (!AddRec->isAffine())
10807 return getCouldNotCompute();
10808
10809 // If this is an affine expression, the execution count of this branch is
10810 // the minimum unsigned root of the following equation:
10811 //
10812 // Start + Step*N = 0 (mod 2^BW)
10813 //
10814 // equivalent to:
10815 //
10816 // Step*N = -Start (mod 2^BW)
10817 //
10818 // where BW is the common bit width of Start and Step.
10819
10820 // Get the initial value for the loop.
10821 const SCEV *Start = getSCEVAtScope(AddRec->getStart(), L->getParentLoop());
10822 const SCEV *Step = getSCEVAtScope(AddRec->getOperand(1), L->getParentLoop());
10823
10824 if (!isLoopInvariant(Step, L))
10825 return getCouldNotCompute();
10826
10827 LoopGuards Guards = LoopGuards::collect(L, *this);
10828 // Specialize step for this loop so we get context sensitive facts below.
10829 const SCEV *StepWLG = applyLoopGuards(Step, Guards);
10830
10831 // For positive steps (counting up until unsigned overflow):
10832 // N = -Start/Step (as unsigned)
10833 // For negative steps (counting down to zero):
10834 // N = Start/-Step
10835 // First compute the unsigned distance from zero in the direction of Step.
10836 bool CountDown = isKnownNegative(StepWLG);
10837 if (!CountDown && !isKnownNonNegative(StepWLG))
10838 return getCouldNotCompute();
10839
10840 const SCEV *Distance = CountDown ? Start : getNegativeSCEV(Start);
10841 // Handle unitary steps, which cannot wraparound.
10842 // 1*N = -Start; -1*N = Start (mod 2^BW), so:
10843 // N = Distance (as unsigned)
10844
10845 if (match(Step, m_CombineOr(m_scev_One(), m_scev_AllOnes()))) {
10846 APInt MaxBECount = getUnsignedRangeMax(applyLoopGuards(Distance, Guards));
10847 MaxBECount = APIntOps::umin(MaxBECount, getUnsignedRangeMax(Distance));
10848
10849 // When a loop like "for (int i = 0; i != n; ++i) { /* body */ }" is rotated,
10850 // we end up with a loop whose backedge-taken count is n - 1. Detect this
10851 // case, and see if we can improve the bound.
10852 //
10853 // Explicitly handling this here is necessary because getUnsignedRange
10854 // isn't context-sensitive; it doesn't know that we only care about the
10855 // range inside the loop.
10856 const SCEV *Zero = getZero(Distance->getType());
10857 const SCEV *One = getOne(Distance->getType());
10858 const SCEV *DistancePlusOne = getAddExpr(Distance, One);
10859 if (isLoopEntryGuardedByCond(L, ICmpInst::ICMP_NE, DistancePlusOne, Zero)) {
10860 // If Distance + 1 doesn't overflow, we can compute the maximum distance
10861 // as "unsigned_max(Distance + 1) - 1". Also apply the loop guards to
10862 // Distance + 1; the range of Distance itself may be a wrapped set even
10863 // when the guards bound Distance + 1 tightly.
10864 APInt Max = APIntOps::umin(
10865 getUnsignedRangeMax(applyLoopGuards(DistancePlusOne, Guards)),
10866 getUnsignedRangeMax(DistancePlusOne));
10867 MaxBECount = APIntOps::umin(MaxBECount, Max - 1);
10868 }
10869 return ExitLimit(Distance, getConstant(MaxBECount), Distance, false,
10870 Predicates);
10871 }
10872
10873 // If the condition controls loop exit (the loop exits only if the expression
10874 // is true) and the addition is no-wrap we can use unsigned divide to
10875 // compute the backedge count. In this case, the step may not divide the
10876 // distance, but we don't care because if the condition is "missed" the loop
10877 // will have undefined behavior due to wrapping.
10878 if (ControlsOnlyExit && AddRec->hasNoSelfWrap() &&
10879 loopHasNoAbnormalExits(AddRec->getLoop())) {
10880
10881 // If the stride is zero and the start is non-zero, the loop must be
10882 // infinite. In C++, most loops are finite by assumption, in which case the
10883 // step being zero implies UB must execute if the loop is entered.
10884 if (!(loopIsFiniteByAssumption(L) && isKnownNonZero(Start)) &&
10885 !isKnownNonZero(StepWLG))
10886 return getCouldNotCompute();
10887
10888 const SCEV *Exact =
10889 getUDivExpr(Distance, CountDown ? getNegativeSCEV(Step) : Step);
10890 const SCEV *ConstantMax = getCouldNotCompute();
10891 if (Exact != getCouldNotCompute()) {
10892 APInt MaxInt = getUnsignedRangeMax(applyLoopGuards(Exact, Guards));
10893 ConstantMax =
10895 }
10896 const SCEV *SymbolicMax =
10897 isa<SCEVCouldNotCompute>(Exact) ? ConstantMax : Exact;
10898 return ExitLimit(Exact, ConstantMax, SymbolicMax, false, Predicates);
10899 }
10900
10901 // Solve the general equation.
10902 const SCEVConstant *StepC = dyn_cast<SCEVConstant>(Step);
10903 if (!StepC || StepC->getValue()->isZero())
10904 return getCouldNotCompute();
10905 const SCEV *E = SolveLinEquationWithOverflow(
10906 StepC->getAPInt(), getNegativeSCEV(Start),
10907 AllowPredicates ? &Predicates : nullptr, *this, L);
10908
10909 const SCEV *M = E;
10910 if (E != getCouldNotCompute()) {
10911 APInt MaxWithGuards = getUnsignedRangeMax(applyLoopGuards(E, Guards));
10912 M = getConstant(APIntOps::umin(MaxWithGuards, getUnsignedRangeMax(E)));
10913 }
10914 auto *S = isa<SCEVCouldNotCompute>(E) ? M : E;
10915 return ExitLimit(E, M, S, false, Predicates);
10916}
10917
10918ScalarEvolution::ExitLimit
10919ScalarEvolution::howFarToNonZero(const SCEV *V, const Loop *L) {
10920 // Loops that look like: while (X == 0) are very strange indeed. We don't
10921 // handle them yet except for the trivial case. This could be expanded in the
10922 // future as needed.
10923
10924 // If the value is a constant, check to see if it is known to be non-zero
10925 // already. If so, the backedge will execute zero times.
10926 if (const SCEVConstant *C = dyn_cast<SCEVConstant>(V)) {
10927 if (!C->getValue()->isZero())
10928 return getZero(C->getType());
10929 return getCouldNotCompute(); // Otherwise it will loop infinitely.
10930 }
10931
10932 // We could implement others, but I really doubt anyone writes loops like
10933 // this, and if they did, they would already be constant folded.
10934 return getCouldNotCompute();
10935}
10936
10937std::pair<const BasicBlock *, const BasicBlock *>
10938ScalarEvolution::getPredecessorWithUniqueSuccessorForBB(const BasicBlock *BB)
10939 const {
10940 // If the block has a unique predecessor, then there is no path from the
10941 // predecessor to the block that does not go through the direct edge
10942 // from the predecessor to the block.
10943 if (const BasicBlock *Pred = BB->getSinglePredecessor())
10944 return {Pred, BB};
10945
10946 // A loop's header is defined to be a block that dominates the loop.
10947 // If the header has a unique predecessor outside the loop, it must be
10948 // a block that has exactly one successor that can reach the loop.
10949 if (const Loop *L = LI.getLoopFor(BB))
10950 return {L->getLoopPredecessor(), L->getHeader()};
10951
10952 return {nullptr, BB};
10953}
10954
10955/// SCEV structural equivalence is usually sufficient for testing whether two
10956/// expressions are equal, however for the purposes of looking for a condition
10957/// guarding a loop, it can be useful to be a little more general, since a
10958/// front-end may have replicated the controlling expression.
10959static bool HasSameValue(const SCEV *A, const SCEV *B) {
10960 // Quick check to see if they are the same SCEV.
10961 if (A == B) return true;
10962
10963 auto ComputesEqualValues = [](const Instruction *A, const Instruction *B) {
10964 // Not all instructions that are "identical" compute the same value. For
10965 // instance, two distinct alloca instructions allocating the same type are
10966 // identical and do not read memory; but compute distinct values.
10967 return A->isIdenticalTo(B) && (isa<BinaryOperator>(A) || isa<GetElementPtrInst>(A));
10968 };
10969
10970 // Otherwise, if they're both SCEVUnknown, it's possible that they hold
10971 // two different instructions with the same value. Check for this case.
10972 if (const SCEVUnknown *AU = dyn_cast<SCEVUnknown>(A))
10973 if (const SCEVUnknown *BU = dyn_cast<SCEVUnknown>(B))
10974 if (const Instruction *AI = dyn_cast<Instruction>(AU->getValue()))
10975 if (const Instruction *BI = dyn_cast<Instruction>(BU->getValue()))
10976 if (ComputesEqualValues(AI, BI))
10977 return true;
10978
10979 // Otherwise assume they may have a different value.
10980 return false;
10981}
10982
10983static bool MatchBinarySub(const SCEV *S, SCEVUse &LHS, SCEVUse &RHS) {
10984 const SCEV *Op0, *Op1;
10985 if (!match(S, m_scev_Add(m_SCEV(Op0), m_SCEV(Op1))))
10986 return false;
10987 if (match(Op0, m_scev_Mul(m_scev_AllOnes(), m_SCEV(RHS)))) {
10988 LHS = Op1;
10989 return true;
10990 }
10991 if (match(Op1, m_scev_Mul(m_scev_AllOnes(), m_SCEV(RHS)))) {
10992 LHS = Op0;
10993 return true;
10994 }
10995 return false;
10996}
10997
10999 SCEVUse &RHS, unsigned Depth) {
11000 bool Changed = false;
11001 // Simplifies ICMP to trivial true or false by turning it into '0 == 0' or
11002 // '0 != 0'.
11003 auto TrivialCase = [&](bool TriviallyTrue) {
11005 Pred = TriviallyTrue ? ICmpInst::ICMP_EQ : ICmpInst::ICMP_NE;
11006 return true;
11007 };
11008 // If we hit the max recursion limit bail out.
11009 if (Depth >= 3)
11010 return false;
11011
11012 const SCEV *NewLHS, *NewRHS;
11013 if (match(LHS, m_scev_c_Mul(m_SCEV(NewLHS), m_SCEVVScale())) &&
11014 match(RHS, m_scev_c_Mul(m_SCEV(NewRHS), m_SCEVVScale()))) {
11015 const SCEVMulExpr *LMul = cast<SCEVMulExpr>(LHS);
11016 const SCEVMulExpr *RMul = cast<SCEVMulExpr>(RHS);
11017
11018 // (X * vscale) pred (Y * vscale) ==> X pred Y
11019 // when both multiples are NSW.
11020 // (X * vscale) uicmp/eq/ne (Y * vscale) ==> X uicmp/eq/ne Y
11021 // when both multiples are NUW.
11022 if ((LMul->hasNoSignedWrap() && RMul->hasNoSignedWrap()) ||
11023 (LMul->hasNoUnsignedWrap() && RMul->hasNoUnsignedWrap() &&
11024 !ICmpInst::isSigned(Pred))) {
11025 LHS = NewLHS;
11026 RHS = NewRHS;
11027 Changed = true;
11028 }
11029 }
11030
11031 // Canonicalize a constant to the right side.
11032 if (const SCEVConstant *LHSC = dyn_cast<SCEVConstant>(LHS)) {
11033 // Check for both operands constant.
11034 if (const SCEVConstant *RHSC = dyn_cast<SCEVConstant>(RHS)) {
11035 if (!ICmpInst::compare(LHSC->getAPInt(), RHSC->getAPInt(), Pred))
11036 return TrivialCase(false);
11037 return TrivialCase(true);
11038 }
11039 // Otherwise swap the operands to put the constant on the right.
11040 std::swap(LHS, RHS);
11042 Changed = true;
11043 }
11044
11045 // (K + A) pred (K + B) --> A pred B
11046 // For equality, no flags are needed.
11047 // For signed, both adds must be NSW. For unsigned, both must be NUW.
11048 {
11049 const SCEVConstant *C = nullptr;
11050 if (match(LHS, m_scev_Add(m_SCEVConstant(C), m_SCEV(NewLHS))) &&
11051 match(RHS, m_scev_Add(m_scev_Specific(C), m_SCEV(NewRHS)))) {
11052 const auto *LAdd = cast<SCEVAddExpr>(LHS);
11053 const auto *RAdd = cast<SCEVAddExpr>(RHS);
11054 if (ICmpInst::isEquality(Pred) ||
11055 (ICmpInst::isSigned(Pred) && LAdd->hasNoSignedWrap() &&
11056 RAdd->hasNoSignedWrap()) ||
11057 (ICmpInst::isUnsigned(Pred) && LAdd->hasNoUnsignedWrap() &&
11058 RAdd->hasNoUnsignedWrap())) {
11059 LHS = NewLHS;
11060 RHS = NewRHS;
11061 Changed = true;
11062 }
11063 }
11064 }
11065
11066 // (C * A) pred (C * B) --> A pred B
11067 // For equality predicates, both muls must be NUW or both must be NSW
11068 // (either suffices to make multiplication by C injective; C == 0 is
11069 // impossible because SCEV folds 0 * X to 0).
11070 // For signed ordering, C must be positive and both muls must be NSW.
11071 // For unsigned ordering, both muls must be NUW.
11072 {
11073 const SCEVConstant *C = nullptr;
11074 if (match(LHS, m_scev_Mul(m_SCEVConstant(C), m_SCEV(NewLHS))) &&
11075 match(RHS, m_scev_Mul(m_scev_Specific(C), m_SCEV(NewRHS)))) {
11076 const auto *LMul = cast<SCEVMulExpr>(LHS);
11077 const auto *RMul = cast<SCEVMulExpr>(RHS);
11078 bool BothNUW = LMul->hasNoUnsignedWrap() && RMul->hasNoUnsignedWrap();
11079 bool BothNSW = LMul->hasNoSignedWrap() && RMul->hasNoSignedWrap();
11080 if ((ICmpInst::isEquality(Pred) && (BothNUW || BothNSW)) ||
11081 (ICmpInst::isSigned(Pred) && BothNSW &&
11082 C->getAPInt().isStrictlyPositive()) ||
11083 (ICmpInst::isUnsigned(Pred) && BothNUW)) {
11084 LHS = NewLHS;
11085 RHS = NewRHS;
11086 Changed = true;
11087 }
11088 }
11089 }
11090
11091 // If we're comparing an addrec with a value which is loop-invariant in the
11092 // addrec's loop, put the addrec on the left. Also make a dominance check,
11093 // as both operands could be addrecs loop-invariant in each other's loop.
11094 if (const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(RHS)) {
11095 const Loop *L = AR->getLoop();
11096 if (isLoopInvariant(LHS, L) && properlyDominates(LHS, L->getHeader())) {
11097 std::swap(LHS, RHS);
11099 Changed = true;
11100 }
11101 }
11102
11103 // If there's a constant operand, canonicalize comparisons with boundary
11104 // cases, and canonicalize *-or-equal comparisons to regular comparisons.
11105 if (const SCEVConstant *RC = dyn_cast<SCEVConstant>(RHS)) {
11106 const APInt &RA = RC->getAPInt();
11107
11108 bool SimplifiedByConstantRange = false;
11109
11110 if (!ICmpInst::isEquality(Pred)) {
11112 if (ExactCR.isFullSet())
11113 return TrivialCase(true);
11114 if (ExactCR.isEmptySet())
11115 return TrivialCase(false);
11116
11117 APInt NewRHS;
11118 CmpInst::Predicate NewPred;
11119 if (ExactCR.getEquivalentICmp(NewPred, NewRHS) &&
11120 ICmpInst::isEquality(NewPred)) {
11121 // We were able to convert an inequality to an equality.
11122 Pred = NewPred;
11123 RHS = getConstant(NewRHS);
11124 Changed = SimplifiedByConstantRange = true;
11125 }
11126 }
11127
11128 if (!SimplifiedByConstantRange) {
11129 switch (Pred) {
11130 default:
11131 break;
11132 case ICmpInst::ICMP_EQ:
11133 case ICmpInst::ICMP_NE:
11134 // Fold ((-1) * %a) + %b == 0 (equivalent to %b-%a == 0) into %a == %b.
11135 if (RA.isZero() && MatchBinarySub(LHS, LHS, RHS))
11136 Changed = true;
11137 break;
11138
11139 // The "Should have been caught earlier!" messages refer to the fact
11140 // that the ExactCR.isFullSet() or ExactCR.isEmptySet() check above
11141 // should have fired on the corresponding cases, and canonicalized the
11142 // check to trivial case.
11143
11144 case ICmpInst::ICMP_UGE:
11145 assert(!RA.isMinValue() && "Should have been caught earlier!");
11146 Pred = ICmpInst::ICMP_UGT;
11147 RHS = getConstant(RA - 1);
11148 Changed = true;
11149 break;
11150 case ICmpInst::ICMP_ULE:
11151 assert(!RA.isMaxValue() && "Should have been caught earlier!");
11152 Pred = ICmpInst::ICMP_ULT;
11153 RHS = getConstant(RA + 1);
11154 Changed = true;
11155 break;
11156 case ICmpInst::ICMP_SGE:
11157 assert(!RA.isMinSignedValue() && "Should have been caught earlier!");
11158 Pred = ICmpInst::ICMP_SGT;
11159 RHS = getConstant(RA - 1);
11160 Changed = true;
11161 break;
11162 case ICmpInst::ICMP_SLE:
11163 assert(!RA.isMaxSignedValue() && "Should have been caught earlier!");
11164 Pred = ICmpInst::ICMP_SLT;
11165 RHS = getConstant(RA + 1);
11166 Changed = true;
11167 break;
11168 }
11169 }
11170 }
11171
11172 // a /u b == 0 => a < b
11173 // a /u b != 0 => a >= b
11174 if (ICmpInst::isEquality(Pred) && RHS->isZero() &&
11175 match(LHS, m_scev_UDiv(m_SCEV(LHS), m_SCEV(RHS)))) {
11177 Changed = true;
11178 }
11179
11180 // Check for obvious equality.
11181 if (HasSameValue(LHS, RHS)) {
11182 if (ICmpInst::isTrueWhenEqual(Pred))
11183 return TrivialCase(true);
11185 return TrivialCase(false);
11186 }
11187
11188 // If possible, canonicalize GE/LE comparisons to GT/LT comparisons, by
11189 // adding or subtracting 1 from one of the operands.
11190 switch (Pred) {
11191 case ICmpInst::ICMP_SLE:
11192 if (!getSignedRangeMax(RHS).isMaxSignedValue()) {
11193 RHS = getAddExpr(getConstant(RHS->getType(), 1, true), RHS,
11195 Pred = ICmpInst::ICMP_SLT;
11196 Changed = true;
11197 } else if (!getSignedRangeMin(LHS).isMinSignedValue()) {
11198 LHS = getAddExpr(getConstant(RHS->getType(), (uint64_t)-1, true), LHS,
11200 Pred = ICmpInst::ICMP_SLT;
11201 Changed = true;
11202 }
11203 break;
11204 case ICmpInst::ICMP_SGE:
11205 if (!getSignedRangeMin(RHS).isMinSignedValue()) {
11206 RHS = getAddExpr(getConstant(RHS->getType(), (uint64_t)-1, true), RHS,
11208 Pred = ICmpInst::ICMP_SGT;
11209 Changed = true;
11210 } else if (!getSignedRangeMax(LHS).isMaxSignedValue()) {
11211 LHS = getAddExpr(getConstant(RHS->getType(), 1, true), LHS,
11213 Pred = ICmpInst::ICMP_SGT;
11214 Changed = true;
11215 }
11216 break;
11217 case ICmpInst::ICMP_ULE:
11218 if (!getUnsignedRangeMax(RHS).isMaxValue()) {
11219 RHS = getAddExpr(getConstant(RHS->getType(), 1, true), RHS,
11221 Pred = ICmpInst::ICMP_ULT;
11222 Changed = true;
11223 } else if (!getUnsignedRangeMin(LHS).isMinValue()) {
11224 LHS = getAddExpr(getConstant(RHS->getType(), (uint64_t)-1, true), LHS);
11225 Pred = ICmpInst::ICMP_ULT;
11226 Changed = true;
11227 }
11228 break;
11229 case ICmpInst::ICMP_UGE:
11230 // If RHS is an op we can fold the -1, try that first.
11231 // Otherwise prefer LHS to preserve the nuw flag.
11232 if ((isa<SCEVConstant>(RHS) ||
11234 isa<SCEVConstant>(cast<SCEVNAryExpr>(RHS)->getOperand(0)))) &&
11235 !getUnsignedRangeMin(RHS).isMinValue()) {
11236 RHS = getAddExpr(getConstant(RHS->getType(), (uint64_t)-1, true), RHS);
11237 Pred = ICmpInst::ICMP_UGT;
11238 Changed = true;
11239 } else if (!getUnsignedRangeMax(LHS).isMaxValue()) {
11240 LHS = getAddExpr(getConstant(RHS->getType(), 1, true), LHS,
11242 Pred = ICmpInst::ICMP_UGT;
11243 Changed = true;
11244 } else if (!getUnsignedRangeMin(RHS).isMinValue()) {
11245 RHS = getAddExpr(getConstant(RHS->getType(), (uint64_t)-1, true), RHS);
11246 Pred = ICmpInst::ICMP_UGT;
11247 Changed = true;
11248 }
11249 break;
11250 default:
11251 break;
11252 }
11253
11254 // TODO: More simplifications are possible here.
11255
11256 // Recursively simplify until we either hit a recursion limit or nothing
11257 // changes.
11258 if (Changed)
11259 (void)SimplifyICmpOperands(Pred, LHS, RHS, Depth + 1);
11260
11261 return Changed;
11262}
11263
11265 return getSignedRangeMax(S).isNegative();
11266}
11267
11271
11273 return !getSignedRangeMin(S).isNegative();
11274}
11275
11279
11281 // Query push down for cases where the unsigned range is
11282 // less than sufficient.
11283 if (const auto *SExt = dyn_cast<SCEVSignExtendExpr>(S))
11284 return isKnownNonZero(SExt->getOperand(0));
11285 return getUnsignedRangeMin(S) != 0;
11286}
11287
11289 bool OrNegative) {
11290 auto NonRecursive = [OrNegative](const SCEV *S) {
11291 if (auto *C = dyn_cast<SCEVConstant>(S))
11292 return C->getAPInt().isPowerOf2() ||
11293 (OrNegative && C->getAPInt().isNegatedPowerOf2());
11294
11295 // vscale is a power-of-two.
11296 return isa<SCEVVScale>(S);
11297 };
11298
11299 if (NonRecursive(S))
11300 return true;
11301
11302 auto *Mul = dyn_cast<SCEVMulExpr>(S);
11303 if (!Mul)
11304 return false;
11305 return all_of(Mul->operands(), NonRecursive) && (OrZero || isKnownNonZero(S));
11306}
11307
11309 const SCEV *S, uint64_t M,
11311 if (M == 0)
11312 return false;
11313 if (M == 1)
11314 return true;
11315
11316 // For a constant, check that "S % M == 0".
11317 if (auto *Cst = dyn_cast<SCEVConstant>(S)) {
11318 APInt C = Cst->getAPInt();
11319 return C.urem(M) == 0;
11320 }
11321
11322 // Basic tests have failed.
11323 // Check "S % M == 0" at compile time and record runtime Assumptions.
11324 auto *STy = dyn_cast<IntegerType>(S->getType());
11325 const SCEV *SmodM =
11326 getURemExpr(S, getConstant(ConstantInt::get(STy, M, false)));
11327 const SCEV *Zero = getZero(STy);
11328
11329 // Check whether "S % M == 0" is known at compile time.
11330 if (isKnownPredicate(ICmpInst::ICMP_EQ, SmodM, Zero))
11331 return true;
11332
11333 // Check whether "S % M != 0" is known at compile time.
11334 if (isKnownPredicate(ICmpInst::ICMP_NE, SmodM, Zero))
11335 return false;
11336
11337 if (!Predicates)
11338 return false;
11339
11340 // Look through Add and AddRec expressions with nuw to improve the
11341 // precision of added predicates. S is a multiple of M if S starts with a
11342 // multiple of M and at every iteration step S only adds multiples of M.
11345 all_of(S->operands(),
11346 [&](SCEVUse Op) { return isKnownMultipleOf(Op, M, Predicates); }))
11347 return true;
11348
11349 // Similarly, look through Mul with nuw, where any operand being a
11350 // known-multiple is sufficient.
11351 if (auto *Mul = dyn_cast<SCEVMulExpr>(S))
11352 if (Mul->hasNoUnsignedWrap() && any_of(S->operands(), [&](SCEVUse Op) {
11353 return isKnownMultipleOf(Op, M, Predicates);
11354 }))
11355 return true;
11356
11357 // Similarly, look through MinMax, with no wrapping arithmetic to consider.
11358 if (isa<SCEVMinMaxExpr>(S) && all_of(S->operands(), [&](SCEVUse Op) {
11359 return isKnownMultipleOf(Op, M, Predicates);
11360 }))
11361 return true;
11362
11364
11365 // Detect redundant predicates.
11366 for (auto *A : *Predicates)
11367 if (A->implies(P, *this))
11368 return true;
11369
11370 // Only record non-redundant predicates.
11371 Predicates->push_back(P);
11372 return true;
11373}
11374
11376 return ((isKnownNonNegative(S1) && isKnownNonNegative(S2)) ||
11378}
11379
11380std::pair<const SCEV *, const SCEV *>
11382 // Compute SCEV on entry of loop L.
11383 const SCEV *Start = SCEVInitRewriter::rewrite(S, L, *this);
11384 if (Start == getCouldNotCompute())
11385 return { Start, Start };
11386 // Compute post increment SCEV for loop L.
11387 const SCEV *PostInc = SCEVPostIncRewriter::rewrite(S, L, *this);
11388 assert(PostInc != getCouldNotCompute() && "Unexpected could not compute");
11389 return { Start, PostInc };
11390}
11391
11393 SCEVUse RHS) {
11394 // First collect all loops.
11396 getUsedLoops(LHS, LoopsUsed);
11397 getUsedLoops(RHS, LoopsUsed);
11398
11399 if (LoopsUsed.empty())
11400 return false;
11401
11402 // Domination relationship must be a linear order on collected loops.
11403#ifndef NDEBUG
11404 for (const auto *L1 : LoopsUsed)
11405 for (const auto *L2 : LoopsUsed)
11406 assert((DT.dominates(L1->getHeader(), L2->getHeader()) ||
11407 DT.dominates(L2->getHeader(), L1->getHeader())) &&
11408 "Domination relationship is not a linear order");
11409#endif
11410
11411 const Loop *MDL =
11412 *llvm::max_element(LoopsUsed, [&](const Loop *L1, const Loop *L2) {
11413 return DT.properlyDominates(L1->getHeader(), L2->getHeader());
11414 });
11415
11416 // Get init and post increment value for LHS.
11417 auto SplitLHS = SplitIntoInitAndPostInc(MDL, LHS);
11418 // if LHS contains unknown non-invariant SCEV then bail out.
11419 if (SplitLHS.first == getCouldNotCompute())
11420 return false;
11421 assert (SplitLHS.second != getCouldNotCompute() && "Unexpected CNC");
11422 // Get init and post increment value for RHS.
11423 auto SplitRHS = SplitIntoInitAndPostInc(MDL, RHS);
11424 // if RHS contains unknown non-invariant SCEV then bail out.
11425 if (SplitRHS.first == getCouldNotCompute())
11426 return false;
11427 assert (SplitRHS.second != getCouldNotCompute() && "Unexpected CNC");
11428 // It is possible that init SCEV contains an invariant load but it does
11429 // not dominate MDL and is not available at MDL loop entry, so we should
11430 // check it here.
11431 if (!isAvailableAtLoopEntry(SplitLHS.first, MDL) ||
11432 !isAvailableAtLoopEntry(SplitRHS.first, MDL))
11433 return false;
11434
11435 // It seems backedge guard check is faster than entry one so in some cases
11436 // it can speed up whole estimation by short circuit
11437 return isLoopBackedgeGuardedByCond(MDL, Pred, SplitLHS.second,
11438 SplitRHS.second) &&
11439 isLoopEntryGuardedByCond(MDL, Pred, SplitLHS.first, SplitRHS.first);
11440}
11441
11443 SCEVUse RHS) {
11444 // Canonicalize the inputs first.
11445 (void)SimplifyICmpOperands(Pred, LHS, RHS);
11446
11447 return isKnownViaInduction(Pred, LHS, RHS) ||
11448 isKnownPredicateViaSplitting(Pred, LHS, RHS) ||
11449 isKnownViaNonRecursiveReasoning(Pred, LHS, RHS);
11450}
11451
11453 const SCEV *LHS,
11454 const SCEV *RHS) {
11455 if (isKnownPredicate(Pred, LHS, RHS))
11456 return true;
11458 return false;
11459 return std::nullopt;
11460}
11461
11463 const SCEV *RHS,
11464 const Instruction *CtxI) {
11465 // TODO: Analyze guards and assumes from Context's block.
11466 return isKnownPredicate(Pred, LHS, RHS) ||
11467 isBasicBlockEntryGuardedByCond(CtxI->getParent(), Pred, LHS, RHS);
11468}
11469
11470std::optional<bool>
11472 const SCEV *RHS, const Instruction *CtxI) {
11473 std::optional<bool> KnownWithoutContext = evaluatePredicate(Pred, LHS, RHS);
11474 if (KnownWithoutContext)
11475 return KnownWithoutContext;
11476
11477 if (isBasicBlockEntryGuardedByCond(CtxI->getParent(), Pred, LHS, RHS))
11478 return true;
11480 CtxI->getParent(), ICmpInst::getInverseCmpPredicate(Pred), LHS, RHS))
11481 return false;
11482 return std::nullopt;
11483}
11484
11486 const SCEVAddRecExpr *LHS,
11487 const SCEV *RHS) {
11488 const Loop *L = LHS->getLoop();
11489 return isLoopEntryGuardedByCond(L, Pred, LHS->getStart(), RHS) &&
11490 isLoopBackedgeGuardedByCond(L, Pred, LHS->getPostIncExpr(*this), RHS);
11491}
11492
11493std::optional<ScalarEvolution::MonotonicPredicateType>
11495 ICmpInst::Predicate Pred) {
11496 auto Result = getMonotonicPredicateTypeImpl(LHS, Pred);
11497
11498#ifndef NDEBUG
11499 // Verify an invariant: inverting the predicate should turn a monotonically
11500 // increasing change to a monotonically decreasing one, and vice versa.
11501 if (Result) {
11502 auto ResultSwapped =
11503 getMonotonicPredicateTypeImpl(LHS, ICmpInst::getSwappedPredicate(Pred));
11504
11505 assert(*ResultSwapped != *Result &&
11506 "monotonicity should flip as we flip the predicate");
11507 }
11508#endif
11509
11510 return Result;
11511}
11512
11513std::optional<ScalarEvolution::MonotonicPredicateType>
11514ScalarEvolution::getMonotonicPredicateTypeImpl(const SCEVAddRecExpr *LHS,
11515 ICmpInst::Predicate Pred) {
11516 // A zero step value for LHS means the induction variable is essentially a
11517 // loop invariant value. We don't really depend on the predicate actually
11518 // flipping from false to true (for increasing predicates, and the other way
11519 // around for decreasing predicates), all we care about is that *if* the
11520 // predicate changes then it only changes from false to true.
11521 //
11522 // A zero step value in itself is not very useful, but there may be places
11523 // where SCEV can prove X >= 0 but not prove X > 0, so it is helpful to be
11524 // as general as possible.
11525
11526 // Only handle LE/LT/GE/GT predicates.
11527 if (!ICmpInst::isRelational(Pred))
11528 return std::nullopt;
11529
11530 bool IsGreater = ICmpInst::isGE(Pred) || ICmpInst::isGT(Pred);
11531 assert((IsGreater || ICmpInst::isLE(Pred) || ICmpInst::isLT(Pred)) &&
11532 "Should be greater or less!");
11533
11534 // Check that AR does not wrap.
11535 if (ICmpInst::isUnsigned(Pred)) {
11536 if (!LHS->hasNoUnsignedWrap())
11537 return std::nullopt;
11539 }
11540 assert(ICmpInst::isSigned(Pred) &&
11541 "Relational predicate is either signed or unsigned!");
11542 if (!LHS->hasNoSignedWrap())
11543 return std::nullopt;
11544
11545 const SCEV *Step = LHS->getStepRecurrence(*this);
11546
11547 if (isKnownNonNegative(Step))
11549
11550 if (isKnownNonPositive(Step))
11552
11553 return std::nullopt;
11554}
11555
11556std::optional<ScalarEvolution::LoopInvariantPredicate>
11558 const SCEV *RHS, const Loop *L,
11559 const Instruction *CtxI) {
11560 // If there is a loop-invariant, force it into the RHS, otherwise bail out.
11561 if (!isLoopInvariant(RHS, L)) {
11562 if (!isLoopInvariant(LHS, L))
11563 return std::nullopt;
11564
11565 std::swap(LHS, RHS);
11567 }
11568
11569 const SCEVAddRecExpr *ArLHS = dyn_cast<SCEVAddRecExpr>(LHS);
11570 if (!ArLHS || ArLHS->getLoop() != L)
11571 return std::nullopt;
11572
11573 auto MonotonicType = getMonotonicPredicateType(ArLHS, Pred);
11574 if (!MonotonicType)
11575 return std::nullopt;
11576 // If the predicate "ArLHS `Pred` RHS" monotonically increases from false to
11577 // true as the loop iterates, and the backedge is control dependent on
11578 // "ArLHS `Pred` RHS" == true then we can reason as follows:
11579 //
11580 // * if the predicate was false in the first iteration then the predicate
11581 // is never evaluated again, since the loop exits without taking the
11582 // backedge.
11583 // * if the predicate was true in the first iteration then it will
11584 // continue to be true for all future iterations since it is
11585 // monotonically increasing.
11586 //
11587 // For both the above possibilities, we can replace the loop varying
11588 // predicate with its value on the first iteration of the loop (which is
11589 // loop invariant).
11590 //
11591 // A similar reasoning applies for a monotonically decreasing predicate, by
11592 // replacing true with false and false with true in the above two bullets.
11594 auto P = Increasing ? Pred : ICmpInst::getInverseCmpPredicate(Pred);
11595
11596 if (isLoopBackedgeGuardedByCond(L, P, LHS, RHS))
11598 RHS);
11599
11600 if (!CtxI)
11601 return std::nullopt;
11602 // Try to prove via context.
11603 // TODO: Support other cases.
11604 switch (Pred) {
11605 default:
11606 break;
11607 case ICmpInst::ICMP_ULE:
11608 case ICmpInst::ICMP_ULT: {
11609 assert(ArLHS->hasNoUnsignedWrap() && "Is a requirement of monotonicity!");
11610 // Given preconditions
11611 // (1) ArLHS does not cross the border of positive and negative parts of
11612 // range because of:
11613 // - Positive step; (TODO: lift this limitation)
11614 // - nuw - does not cross zero boundary;
11615 // - nsw - does not cross SINT_MAX boundary;
11616 // (2) ArLHS <s RHS
11617 // (3) RHS >=s 0
11618 // we can replace the loop variant ArLHS <u RHS condition with loop
11619 // invariant Start(ArLHS) <u RHS.
11620 //
11621 // Because of (1) there are two options:
11622 // - ArLHS is always negative. It means that ArLHS <u RHS is always false;
11623 // - ArLHS is always non-negative. Because of (3) RHS is also non-negative.
11624 // It means that ArLHS <s RHS <=> ArLHS <u RHS.
11625 // Because of (2) ArLHS <u RHS is trivially true.
11626 // All together it means that ArLHS <u RHS <=> Start(ArLHS) >=s 0.
11627 // We can strengthen this to Start(ArLHS) <u RHS.
11628 auto SignFlippedPred = ICmpInst::getFlippedSignednessPredicate(Pred);
11629 if (ArLHS->hasNoSignedWrap() && ArLHS->isAffine() &&
11630 isKnownPositive(ArLHS->getStepRecurrence(*this)) &&
11631 isKnownNonNegative(RHS) &&
11632 isKnownPredicateAt(SignFlippedPred, ArLHS, RHS, CtxI))
11634 RHS);
11635 }
11636 }
11637
11638 return std::nullopt;
11639}
11640
11641std::optional<ScalarEvolution::LoopInvariantPredicate>
11643 CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, const Loop *L,
11644 const Instruction *CtxI, const SCEV *MaxIter) {
11646 Pred, LHS, RHS, L, CtxI, MaxIter))
11647 return LIP;
11648 if (auto *UMin = dyn_cast<SCEVUMinExpr>(MaxIter))
11649 // Number of iterations expressed as UMIN isn't always great for expressing
11650 // the value on the last iteration. If the straightforward approach didn't
11651 // work, try the following trick: if the a predicate is invariant for X, it
11652 // is also invariant for umin(X, ...). So try to find something that works
11653 // among subexpressions of MaxIter expressed as umin.
11654 for (SCEVUse Op : UMin->operands())
11656 Pred, LHS, RHS, L, CtxI, Op))
11657 return LIP;
11658 return std::nullopt;
11659}
11660
11661std::optional<ScalarEvolution::LoopInvariantPredicate>
11663 CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, const Loop *L,
11664 const Instruction *CtxI, const SCEV *MaxIter) {
11665 // Try to prove the following set of facts:
11666 // - The predicate is monotonic in the iteration space.
11667 // - If the check does not fail on the 1st iteration:
11668 // - No overflow will happen during first MaxIter iterations;
11669 // - It will not fail on the MaxIter'th iteration.
11670 // If the check does fail on the 1st iteration, we leave the loop and no
11671 // other checks matter.
11672
11673 // If there is a loop-invariant, force it into the RHS, otherwise bail out.
11674 if (!isLoopInvariant(RHS, L)) {
11675 if (!isLoopInvariant(LHS, L))
11676 return std::nullopt;
11677
11678 std::swap(LHS, RHS);
11680 }
11681
11682 auto *AR = dyn_cast<SCEVAddRecExpr>(LHS);
11683 if (!AR || AR->getLoop() != L)
11684 return std::nullopt;
11685
11686 // Even if both are valid, we need to consistently chose the unsigned or the
11687 // signed predicate below, not mixtures of both. For now, prefer the unsigned
11688 // predicate.
11689 Pred = Pred.dropSameSign();
11690
11691 // The predicate must be relational (i.e. <, <=, >=, >).
11692 if (!ICmpInst::isRelational(Pred))
11693 return std::nullopt;
11694
11695 // TODO: Support steps other than +/- 1.
11696 const SCEV *Step = AR->getStepRecurrence(*this);
11697 auto *One = getOne(Step->getType());
11698 auto *MinusOne = getNegativeSCEV(One);
11699 if (Step != One && Step != MinusOne)
11700 return std::nullopt;
11701
11702 // Type mismatch here means that MaxIter is potentially larger than max
11703 // unsigned value in start type, which mean we cannot prove no wrap for the
11704 // indvar.
11705 if (AR->getType() != MaxIter->getType())
11706 return std::nullopt;
11707
11708 // Value of IV on suggested last iteration.
11709 const SCEV *Last = AR->evaluateAtIteration(MaxIter, *this);
11710 // Does it still meet the requirement?
11711 if (!isLoopBackedgeGuardedByCond(L, Pred, Last, RHS))
11712 return std::nullopt;
11713 // Because step is +/- 1 and MaxIter has same type as Start (i.e. it does
11714 // not exceed max unsigned value of this type), this effectively proves
11715 // that there is no wrap during the iteration. To prove that there is no
11716 // signed/unsigned wrap, we need to check that
11717 // Start <= Last for step = 1 or Start >= Last for step = -1.
11718 ICmpInst::Predicate NoOverflowPred =
11720 if (Step == MinusOne)
11721 NoOverflowPred = ICmpInst::getSwappedPredicate(NoOverflowPred);
11722 const SCEV *Start = AR->getStart();
11723 if (!isKnownPredicateAt(NoOverflowPred, Start, Last, CtxI))
11724 return std::nullopt;
11725
11726 // Everything is fine.
11727 return ScalarEvolution::LoopInvariantPredicate(Pred, Start, RHS);
11728}
11729
11730bool ScalarEvolution::isKnownPredicateViaConstantRanges(CmpPredicate Pred,
11731 SCEVUse LHS,
11732 SCEVUse RHS) {
11733 if (HasSameValue(LHS, RHS))
11734 return ICmpInst::isTrueWhenEqual(Pred);
11735
11736 auto CheckRange = [&](bool IsSigned) {
11737 auto RangeLHS = IsSigned ? getSignedRange(LHS) : getUnsignedRange(LHS);
11738 auto RangeRHS = IsSigned ? getSignedRange(RHS) : getUnsignedRange(RHS);
11739 return RangeLHS.icmp(Pred, RangeRHS);
11740 };
11741
11742 // The check at the top of the function catches the case where the values are
11743 // known to be equal.
11744 if (Pred == CmpInst::ICMP_EQ)
11745 return false;
11746
11747 if (Pred == CmpInst::ICMP_NE) {
11748 if (CheckRange(true) || CheckRange(false))
11749 return true;
11750 auto *Diff = getMinusSCEV(LHS, RHS);
11751 return !isa<SCEVCouldNotCompute>(Diff) && isKnownNonZero(Diff);
11752 }
11753
11754 return CheckRange(CmpInst::isSigned(Pred));
11755}
11756
11757bool ScalarEvolution::isKnownPredicateViaNoOverflow(CmpPredicate Pred,
11759 // Match X to (A + C1)<ExpectedFlags> and Y to (A + C2)<ExpectedFlags>, where
11760 // C1 and C2 are constant integers. If either X or Y are not add expressions,
11761 // consider them as X + 0 and Y + 0 respectively. C1 and C2 are returned via
11762 // OutC1 and OutC2.
11763 auto MatchBinaryAddToConst = [this](SCEVUse X, SCEVUse Y, APInt &OutC1,
11764 APInt &OutC2, SCEVFlags ExpectedFlags) {
11765 SCEVUse XNonConstOp, XConstOp;
11766 SCEVUse YNonConstOp, YConstOp;
11767 SCEVFlags XFlagsPresent;
11768 SCEVFlags YFlagsPresent;
11769
11770 if (!splitBinaryAdd(X, XConstOp, XNonConstOp, XFlagsPresent)) {
11771 XConstOp = getZero(X->getType());
11772 XNonConstOp = X;
11773 XFlagsPresent = ExpectedFlags;
11774 }
11775 if (!isa<SCEVConstant>(XConstOp))
11776 return false;
11777
11778 if (!splitBinaryAdd(Y, YConstOp, YNonConstOp, YFlagsPresent)) {
11779 YConstOp = getZero(Y->getType());
11780 YNonConstOp = Y;
11781 YFlagsPresent = ExpectedFlags;
11782 }
11783
11784 if (YNonConstOp != XNonConstOp)
11785 return false;
11786
11787 if (!isa<SCEVConstant>(YConstOp))
11788 return false;
11789
11790 // When matching ADDs with NUW flags (and unsigned predicates), only the
11791 // second ADD (with the larger constant) requires NUW.
11792 if ((YFlagsPresent & ExpectedFlags) != ExpectedFlags)
11793 return false;
11794 if (ExpectedFlags != SCEV::FlagNUW &&
11795 (XFlagsPresent & ExpectedFlags) != ExpectedFlags) {
11796 return false;
11797 }
11798
11799 OutC1 = cast<SCEVConstant>(XConstOp)->getAPInt();
11800 OutC2 = cast<SCEVConstant>(YConstOp)->getAPInt();
11801
11802 return true;
11803 };
11804
11805 APInt C1;
11806 APInt C2;
11807
11808 switch (Pred) {
11809 default:
11810 break;
11811
11812 case ICmpInst::ICMP_SGE:
11813 std::swap(LHS, RHS);
11814 [[fallthrough]];
11815 case ICmpInst::ICMP_SLE:
11816 // (X + C1)<nsw> s<= (X + C2)<nsw> if C1 s<= C2.
11817 if (MatchBinaryAddToConst(LHS, RHS, C1, C2, SCEV::FlagNSW) && C1.sle(C2))
11818 return true;
11819
11820 break;
11821
11822 case ICmpInst::ICMP_SGT:
11823 std::swap(LHS, RHS);
11824 [[fallthrough]];
11825 case ICmpInst::ICMP_SLT:
11826 // (X + C1)<nsw> s< (X + C2)<nsw> if C1 s< C2.
11827 if (MatchBinaryAddToConst(LHS, RHS, C1, C2, SCEV::FlagNSW) && C1.slt(C2))
11828 return true;
11829
11830 break;
11831
11832 case ICmpInst::ICMP_UGE:
11833 std::swap(LHS, RHS);
11834 [[fallthrough]];
11835 case ICmpInst::ICMP_ULE:
11836 // (X + C1) u<= (X + C2)<nuw> for C1 u<= C2.
11837 if (MatchBinaryAddToConst(LHS, RHS, C1, C2, SCEV::FlagNUW) && C1.ule(C2))
11838 return true;
11839
11840 break;
11841
11842 case ICmpInst::ICMP_UGT:
11843 std::swap(LHS, RHS);
11844 [[fallthrough]];
11845 case ICmpInst::ICMP_ULT:
11846 // (X + C1) u< (X + C2)<nuw> if C1 u< C2.
11847 if (MatchBinaryAddToConst(LHS, RHS, C1, C2, SCEV::FlagNUW) && C1.ult(C2))
11848 return true;
11849 break;
11850 }
11851
11852 return false;
11853}
11854
11855bool ScalarEvolution::isKnownPredicateViaSplitting(CmpPredicate Pred,
11857 if (Pred != ICmpInst::ICMP_ULT || ProvingSplitPredicate)
11858 return false;
11859
11860 // Allowing arbitrary number of activations of isKnownPredicateViaSplitting on
11861 // the stack can result in exponential time complexity.
11862 SaveAndRestore Restore(ProvingSplitPredicate, true);
11863
11864 // If L >= 0 then I `ult` L <=> I >= 0 && I `slt` L
11865 //
11866 // To prove L >= 0 we use isKnownNonNegative whereas to prove I >= 0 we use
11867 // isKnownPredicate. isKnownPredicate is more powerful, but also more
11868 // expensive; and using isKnownNonNegative(RHS) is sufficient for most of the
11869 // interesting cases seen in practice. We can consider "upgrading" L >= 0 to
11870 // use isKnownPredicate later if needed.
11871 return isKnownNonNegative(RHS) &&
11874}
11875
11876bool ScalarEvolution::isImpliedViaGuard(const BasicBlock *BB, CmpPredicate Pred,
11877 const SCEV *LHS, const SCEV *RHS) {
11878 // No need to even try if we know the module has no guards.
11879 if (!HasGuards)
11880 return false;
11881
11882 return any_of(*BB, [&](const Instruction &I) {
11883 using namespace llvm::PatternMatch;
11884
11885 Value *Condition;
11887 m_Value(Condition))) &&
11888 isImpliedCond(Pred, LHS, RHS, Condition, false);
11889 });
11890}
11891
11892/// isLoopBackedgeGuardedByCond - Test whether the backedge of the loop is
11893/// protected by a conditional between LHS and RHS. This is used to
11894/// to eliminate casts.
11896 CmpPredicate Pred,
11897 const SCEV *LHS,
11898 const SCEV *RHS) {
11899 // Interpret a null as meaning no loop, where there is obviously no guard
11900 // (interprocedural conditions notwithstanding). Do not bother about
11901 // unreachable loops.
11902 if (!L || !DT.isReachableFromEntry(L->getHeader()))
11903 return true;
11904
11905 if (VerifyIR)
11906 assert(!verifyFunction(*L->getHeader()->getParent(), &dbgs()) &&
11907 "This cannot be done on broken IR!");
11908
11909
11910 if (isKnownViaNonRecursiveReasoning(Pred, LHS, RHS))
11911 return true;
11912
11913 BasicBlock *Latch = L->getLoopLatch();
11914 if (!Latch)
11915 return false;
11916
11917 CondBrInst *LoopContinuePredicate =
11919 if (LoopContinuePredicate &&
11920 isImpliedCond(Pred, LHS, RHS, LoopContinuePredicate->getCondition(),
11921 LoopContinuePredicate->getSuccessor(0) != L->getHeader()))
11922 return true;
11923
11924 // We don't want more than one activation of the following loops on the stack
11925 // -- that can lead to O(n!) time complexity.
11926 if (WalkingBEDominatingConds)
11927 return false;
11928
11929 SaveAndRestore ClearOnExit(WalkingBEDominatingConds, true);
11930
11931 // See if we can exploit a trip count to prove the predicate.
11932 const auto &BETakenInfo = getBackedgeTakenInfo(L);
11933 const SCEV *LatchBECount = BETakenInfo.getExact(Latch, this);
11934 if (LatchBECount != getCouldNotCompute()) {
11935 // We know that Latch branches back to the loop header exactly
11936 // LatchBECount times. This means the backdege condition at Latch is
11937 // equivalent to "{0,+,1} u< LatchBECount".
11938 Type *Ty = LatchBECount->getType();
11939 auto NoWrapFlags = SCEVFlags(SCEV::FlagNUW | SCEV::FlagNW);
11940 const SCEV *LoopCounter =
11941 getAddRecExpr(getZero(Ty), getOne(Ty), L, NoWrapFlags);
11942 if (isImpliedCond(Pred, LHS, RHS, ICmpInst::ICMP_ULT, LoopCounter,
11943 LatchBECount))
11944 return true;
11945 }
11946
11947 // Check conditions due to any @llvm.assume intrinsics.
11948 for (auto &AssumeVH : AC.assumptions()) {
11949 if (!AssumeVH)
11950 continue;
11951 auto *CI = cast<CallInst>(AssumeVH);
11952 if (!DT.dominates(CI, Latch->getTerminator()))
11953 continue;
11954
11955 if (isImpliedCond(Pred, LHS, RHS, CI->getArgOperand(0), false))
11956 return true;
11957 }
11958
11959 if (isImpliedViaGuard(Latch, Pred, LHS, RHS))
11960 return true;
11961
11962 for (DomTreeNode *DTN = DT[Latch], *HeaderDTN = DT[L->getHeader()];
11963 DTN != HeaderDTN; DTN = DTN->getIDom()) {
11964 assert(DTN && "should reach the loop header before reaching the root!");
11965
11966 BasicBlock *BB = DTN->getBlock();
11967 if (isImpliedViaGuard(BB, Pred, LHS, RHS))
11968 return true;
11969
11970 BasicBlock *PBB = BB->getSinglePredecessor();
11971 if (!PBB)
11972 continue;
11973
11975 if (!ContBr || ContBr->getSuccessor(0) == ContBr->getSuccessor(1))
11976 continue;
11977
11978 // If we have an edge `E` within the loop body that dominates the only
11979 // latch, the condition guarding `E` also guards the backedge. This
11980 // reasoning works only for loops with a single latch.
11981 // We're constructively (and conservatively) enumerating edges within the
11982 // loop body that dominate the latch. The dominator tree better agree
11983 // with us on this:
11984 assert(DT.dominates(BasicBlockEdge(PBB, BB), Latch) && "should be!");
11985 if (isImpliedCond(Pred, LHS, RHS, ContBr->getCondition(),
11986 BB != ContBr->getSuccessor(0)))
11987 return true;
11988 }
11989
11990 return false;
11991}
11992
11994 CmpPredicate Pred,
11995 const SCEV *LHS,
11996 const SCEV *RHS) {
11997 // Do not bother proving facts for unreachable code.
11998 if (!DT.isReachableFromEntry(BB))
11999 return true;
12000 if (VerifyIR)
12001 assert(!verifyFunction(*BB->getParent(), &dbgs()) &&
12002 "This cannot be done on broken IR!");
12003
12004 // If we cannot prove strict comparison (e.g. a > b), maybe we can prove
12005 // the facts (a >= b && a != b) separately. A typical situation is when the
12006 // non-strict comparison is known from ranges and non-equality is known from
12007 // dominating predicates. If we are proving strict comparison, we always try
12008 // to prove non-equality and non-strict comparison separately.
12009 CmpPredicate NonStrictPredicate = ICmpInst::getNonStrictCmpPredicate(Pred);
12010 const bool ProvingStrictComparison =
12011 Pred != NonStrictPredicate.dropSameSign();
12012 bool ProvedNonStrictComparison = false;
12013 bool ProvedNonEquality = false;
12014
12015 auto SplitAndProve = [&](std::function<bool(CmpPredicate)> Fn) -> bool {
12016 if (!ProvedNonStrictComparison)
12017 ProvedNonStrictComparison = Fn(NonStrictPredicate);
12018 if (!ProvedNonEquality)
12019 ProvedNonEquality = Fn(ICmpInst::ICMP_NE);
12020 if (ProvedNonStrictComparison && ProvedNonEquality)
12021 return true;
12022 return false;
12023 };
12024
12025 if (ProvingStrictComparison) {
12026 auto ProofFn = [&](CmpPredicate P) {
12027 return isKnownViaNonRecursiveReasoning(P, LHS, RHS);
12028 };
12029 if (SplitAndProve(ProofFn))
12030 return true;
12031 }
12032
12033 // Try to prove (Pred, LHS, RHS) using isImpliedCond.
12034 auto ProveViaCond = [&](const Value *Condition, bool Inverse) {
12035 const Instruction *CtxI = &BB->front();
12036 if (isImpliedCond(Pred, LHS, RHS, Condition, Inverse, CtxI))
12037 return true;
12038 if (ProvingStrictComparison) {
12039 auto ProofFn = [&](CmpPredicate P) {
12040 return isImpliedCond(P, LHS, RHS, Condition, Inverse, CtxI);
12041 };
12042 if (SplitAndProve(ProofFn))
12043 return true;
12044 }
12045 return false;
12046 };
12047
12048 // Starting at the block's predecessor, climb up the predecessor chain, as long
12049 // as there are predecessors that can be found that have unique successors
12050 // leading to the original block.
12051 const Loop *ContainingLoop = LI.getLoopFor(BB);
12052 const BasicBlock *PredBB;
12053 if (ContainingLoop && ContainingLoop->getHeader() == BB)
12054 PredBB = ContainingLoop->getLoopPredecessor();
12055 else
12056 PredBB = BB->getSinglePredecessor();
12057 for (std::pair<const BasicBlock *, const BasicBlock *> Pair(PredBB, BB);
12058 Pair.first; Pair = getPredecessorWithUniqueSuccessorForBB(Pair.first)) {
12059 const CondBrInst *BlockEntryPredicate =
12060 dyn_cast<CondBrInst>(Pair.first->getTerminator());
12061 if (!BlockEntryPredicate)
12062 continue;
12063
12064 if (ProveViaCond(BlockEntryPredicate->getCondition(),
12065 BlockEntryPredicate->getSuccessor(0) != Pair.second))
12066 return true;
12067 }
12068
12069 // Check conditions due to any @llvm.assume intrinsics.
12070 for (auto &AssumeVH : AC.assumptions()) {
12071 if (!AssumeVH)
12072 continue;
12073 auto *CI = cast<CallInst>(AssumeVH);
12074 if (!DT.dominates(CI, BB))
12075 continue;
12076
12077 if (ProveViaCond(CI->getArgOperand(0), false))
12078 return true;
12079 }
12080
12081 // Check conditions due to any @llvm.experimental.guard intrinsics.
12082 auto *GuardDecl = Intrinsic::getDeclarationIfExists(
12083 F.getParent(), Intrinsic::experimental_guard);
12084 if (GuardDecl)
12085 for (const auto *GU : GuardDecl->users())
12086 if (const auto *Guard = dyn_cast<IntrinsicInst>(GU))
12087 if (Guard->getFunction() == BB->getParent() && DT.dominates(Guard, BB))
12088 if (ProveViaCond(Guard->getArgOperand(0), false))
12089 return true;
12090 return false;
12091}
12092
12094 const SCEV *LHS,
12095 const SCEV *RHS) {
12096 // Interpret a null as meaning no loop, where there is obviously no guard
12097 // (interprocedural conditions notwithstanding).
12098 if (!L)
12099 return false;
12100
12101 // Both LHS and RHS must be available at loop entry.
12103 "LHS is not available at Loop Entry");
12105 "RHS is not available at Loop Entry");
12106
12107 if (isKnownViaNonRecursiveReasoning(Pred, LHS, RHS))
12108 return true;
12109
12110 return isBasicBlockEntryGuardedByCond(L->getHeader(), Pred, LHS, RHS);
12111}
12112
12113bool ScalarEvolution::isImpliedCond(CmpPredicate Pred, const SCEV *LHS,
12114 const SCEV *RHS,
12115 const Value *FoundCondValue, bool Inverse,
12116 const Instruction *CtxI) {
12117 // False conditions implies anything. Do not bother analyzing it further.
12118 if (FoundCondValue ==
12119 ConstantInt::getBool(FoundCondValue->getContext(), Inverse))
12120 return true;
12121
12122 if (!PendingLoopPredicates.insert(FoundCondValue).second)
12123 return false;
12124
12125 llvm::scope_exit ClearOnExit(
12126 [&]() { PendingLoopPredicates.erase(FoundCondValue); });
12127
12128 // Recursively handle And and Or conditions.
12129 const Value *Op0, *Op1;
12130 if (match(FoundCondValue, m_LogicalAnd(m_Value(Op0), m_Value(Op1)))) {
12131 if (!Inverse)
12132 return isImpliedCond(Pred, LHS, RHS, Op0, Inverse, CtxI) ||
12133 isImpliedCond(Pred, LHS, RHS, Op1, Inverse, CtxI);
12134 } else if (match(FoundCondValue, m_LogicalOr(m_Value(Op0), m_Value(Op1)))) {
12135 if (Inverse)
12136 return isImpliedCond(Pred, LHS, RHS, Op0, Inverse, CtxI) ||
12137 isImpliedCond(Pred, LHS, RHS, Op1, Inverse, CtxI);
12138 }
12139
12140 const ICmpInst *ICI = dyn_cast<ICmpInst>(FoundCondValue);
12141 if (!ICI) return false;
12142
12143 // Now that we found a conditional branch that dominates the loop or controls
12144 // the loop latch. Check to see if it is the comparison we are looking for.
12145 CmpPredicate FoundPred;
12146 if (Inverse)
12147 FoundPred = ICI->getInverseCmpPredicate();
12148 else
12149 FoundPred = ICI->getCmpPredicate();
12150
12151 const SCEV *FoundLHS = getSCEV(ICI->getOperand(0));
12152 const SCEV *FoundRHS = getSCEV(ICI->getOperand(1));
12153
12154 return isImpliedCond(Pred, LHS, RHS, FoundPred, FoundLHS, FoundRHS, CtxI);
12155}
12156
12157bool ScalarEvolution::isImpliedCond(CmpPredicate Pred, const SCEV *LHS,
12158 const SCEV *RHS, CmpPredicate FoundPred,
12159 const SCEV *FoundLHS, const SCEV *FoundRHS,
12160 const Instruction *CtxI) {
12161 // Balance the types.
12162 if (getTypeSizeInBits(LHS->getType()) <
12163 getTypeSizeInBits(FoundLHS->getType())) {
12164 // For unsigned and equality predicates, try to prove that both found
12165 // operands fit into narrow unsigned range. If so, try to prove facts in
12166 // narrow types.
12167 if (!CmpInst::isSigned(FoundPred) && !FoundLHS->getType()->isPointerTy() &&
12168 !FoundRHS->getType()->isPointerTy()) {
12169 auto *NarrowType = LHS->getType();
12170 auto *WideType = FoundLHS->getType();
12171 auto BitWidth = getTypeSizeInBits(NarrowType);
12172 const SCEV *MaxValue = getZeroExtendExpr(
12174 if (isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_ULE, FoundLHS,
12175 MaxValue) &&
12176 isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_ULE, FoundRHS,
12177 MaxValue)) {
12178 const SCEV *TruncFoundLHS = getTruncateExpr(FoundLHS, NarrowType);
12179 const SCEV *TruncFoundRHS = getTruncateExpr(FoundRHS, NarrowType);
12180 // We cannot preserve samesign after truncation.
12181 if (isImpliedCondBalancedTypes(Pred, LHS, RHS, FoundPred.dropSameSign(),
12182 TruncFoundLHS, TruncFoundRHS, CtxI))
12183 return true;
12184 }
12185 }
12186
12187 if (LHS->getType()->isPointerTy() || RHS->getType()->isPointerTy())
12188 return false;
12189 if (CmpInst::isSigned(Pred)) {
12190 LHS = getSignExtendExpr(LHS, FoundLHS->getType());
12191 RHS = getSignExtendExpr(RHS, FoundLHS->getType());
12192 } else {
12193 LHS = getZeroExtendExpr(LHS, FoundLHS->getType());
12194 RHS = getZeroExtendExpr(RHS, FoundLHS->getType());
12195 }
12196 } else if (getTypeSizeInBits(LHS->getType()) >
12197 getTypeSizeInBits(FoundLHS->getType())) {
12198 if (FoundLHS->getType()->isPointerTy() || FoundRHS->getType()->isPointerTy())
12199 return false;
12200 if (CmpInst::isSigned(FoundPred)) {
12201 FoundLHS = getSignExtendExpr(FoundLHS, LHS->getType());
12202 FoundRHS = getSignExtendExpr(FoundRHS, LHS->getType());
12203 } else {
12204 FoundLHS = getZeroExtendExpr(FoundLHS, LHS->getType());
12205 FoundRHS = getZeroExtendExpr(FoundRHS, LHS->getType());
12206 }
12207 }
12208 return isImpliedCondBalancedTypes(Pred, LHS, RHS, FoundPred, FoundLHS,
12209 FoundRHS, CtxI);
12210}
12211
12212bool ScalarEvolution::isImpliedCondBalancedTypes(
12213 CmpPredicate Pred, SCEVUse LHS, SCEVUse RHS, CmpPredicate FoundPred,
12214 SCEVUse FoundLHS, SCEVUse FoundRHS, const Instruction *CtxI) {
12216 getTypeSizeInBits(FoundLHS->getType()) &&
12217 "Types should be balanced!");
12218 // Canonicalize the query to match the way instcombine will have
12219 // canonicalized the comparison.
12220 if (SimplifyICmpOperands(Pred, LHS, RHS))
12221 if (LHS == RHS)
12222 return CmpInst::isTrueWhenEqual(Pred);
12223 if (SimplifyICmpOperands(FoundPred, FoundLHS, FoundRHS))
12224 if (FoundLHS == FoundRHS)
12225 return CmpInst::isFalseWhenEqual(FoundPred);
12226
12227 // Check to see if we can make the LHS or RHS match.
12228 if (LHS == FoundRHS || RHS == FoundLHS) {
12229 if (isa<SCEVConstant>(RHS)) {
12230 std::swap(FoundLHS, FoundRHS);
12231 FoundPred = ICmpInst::getSwappedCmpPredicate(FoundPred);
12232 } else {
12233 std::swap(LHS, RHS);
12235 }
12236 }
12237
12238 // Check whether the found predicate is the same as the desired predicate.
12239 if (auto P = CmpPredicate::getMatching(FoundPred, Pred))
12240 return isImpliedCondOperands(*P, LHS, RHS, FoundLHS, FoundRHS, CtxI);
12241
12242 // Check whether swapping the found predicate makes it the same as the
12243 // desired predicate.
12244 if (auto P = CmpPredicate::getMatching(
12245 ICmpInst::getSwappedCmpPredicate(FoundPred), Pred)) {
12246 // We can write the implication
12247 // 0. LHS Pred RHS <- FoundLHS SwapPred FoundRHS
12248 // using one of the following ways:
12249 // 1. LHS Pred RHS <- FoundRHS Pred FoundLHS
12250 // 2. RHS SwapPred LHS <- FoundLHS SwapPred FoundRHS
12251 // Both require swapping the operands of one condition. Don't do this if it
12252 // would break canonical constant/addrec ordering.
12254 return isImpliedCondOperands(ICmpInst::getSwappedCmpPredicate(*P), RHS,
12255 LHS, FoundLHS, FoundRHS, CtxI);
12256 if (!isa<SCEVConstant>(FoundRHS) && !isa<SCEVAddRecExpr>(FoundLHS))
12257 return isImpliedCondOperands(*P, LHS, RHS, FoundRHS, FoundLHS, CtxI);
12258
12259 return false;
12260 }
12261
12262 auto IsSignFlippedPredicate = [](CmpInst::Predicate P1,
12264 assert(P1 != P2 && "Handled earlier!");
12265 return CmpInst::isRelational(P2) &&
12267 };
12268 if (IsSignFlippedPredicate(Pred, FoundPred)) {
12269 // Unsigned comparison is the same as signed comparison when both the
12270 // operands are non-negative or negative.
12271 if (haveSameSign(FoundLHS, FoundRHS))
12272 return isImpliedCondOperands(Pred, LHS, RHS, FoundLHS, FoundRHS, CtxI);
12273 // Create local copies that we can freely swap and canonicalize our
12274 // conditions to "le/lt".
12275 CmpPredicate CanonicalPred = Pred, CanonicalFoundPred = FoundPred;
12276 const SCEV *CanonicalLHS = LHS, *CanonicalRHS = RHS,
12277 *CanonicalFoundLHS = FoundLHS, *CanonicalFoundRHS = FoundRHS;
12278 if (ICmpInst::isGT(CanonicalPred) || ICmpInst::isGE(CanonicalPred)) {
12279 CanonicalPred = ICmpInst::getSwappedCmpPredicate(CanonicalPred);
12280 CanonicalFoundPred = ICmpInst::getSwappedCmpPredicate(CanonicalFoundPred);
12281 std::swap(CanonicalLHS, CanonicalRHS);
12282 std::swap(CanonicalFoundLHS, CanonicalFoundRHS);
12283 }
12284 assert((ICmpInst::isLT(CanonicalPred) || ICmpInst::isLE(CanonicalPred)) &&
12285 "Must be!");
12286 assert((ICmpInst::isLT(CanonicalFoundPred) ||
12287 ICmpInst::isLE(CanonicalFoundPred)) &&
12288 "Must be!");
12289 if (ICmpInst::isSigned(CanonicalPred) && isKnownNonNegative(CanonicalRHS))
12290 // Use implication:
12291 // x <u y && y >=s 0 --> x <s y.
12292 // If we can prove the left part, the right part is also proven.
12293 return isImpliedCondOperands(CanonicalFoundPred, CanonicalLHS,
12294 CanonicalRHS, CanonicalFoundLHS,
12295 CanonicalFoundRHS);
12296 if (ICmpInst::isUnsigned(CanonicalPred) && isKnownNegative(CanonicalRHS))
12297 // Use implication:
12298 // x <s y && y <s 0 --> x <u y.
12299 // If we can prove the left part, the right part is also proven.
12300 return isImpliedCondOperands(CanonicalFoundPred, CanonicalLHS,
12301 CanonicalRHS, CanonicalFoundLHS,
12302 CanonicalFoundRHS);
12303 }
12304
12305 // Check if we can make progress by sharpening ranges.
12306 if (FoundPred == ICmpInst::ICMP_NE &&
12307 (isa<SCEVConstant>(FoundLHS) || isa<SCEVConstant>(FoundRHS))) {
12308
12309 const SCEVConstant *C = nullptr;
12310 const SCEV *V = nullptr;
12311
12312 if (isa<SCEVConstant>(FoundLHS)) {
12313 C = cast<SCEVConstant>(FoundLHS);
12314 V = FoundRHS;
12315 } else {
12316 C = cast<SCEVConstant>(FoundRHS);
12317 V = FoundLHS;
12318 }
12319
12320 // The guarding predicate tells us that C != V. If the known range
12321 // of V is [C, t), we can sharpen the range to [C + 1, t). The
12322 // range we consider has to correspond to same signedness as the
12323 // predicate we're interested in folding.
12324
12325 APInt Min = ICmpInst::isSigned(Pred) ?
12327
12328 if (Min == C->getAPInt()) {
12329 // Given (V >= Min && V != Min) we conclude V >= (Min + 1).
12330 // This is true even if (Min + 1) wraps around -- in case of
12331 // wraparound, (Min + 1) < Min, so (V >= Min => V >= (Min + 1)).
12332
12333 APInt SharperMin = Min + 1;
12334
12335 switch (Pred) {
12336 case ICmpInst::ICMP_SGE:
12337 case ICmpInst::ICMP_UGE:
12338 // We know V `Pred` SharperMin. If this implies LHS `Pred`
12339 // RHS, we're done.
12340 if (isImpliedCondOperands(Pred, LHS, RHS, V, getConstant(SharperMin),
12341 CtxI))
12342 return true;
12343 [[fallthrough]];
12344
12345 case ICmpInst::ICMP_SGT:
12346 case ICmpInst::ICMP_UGT:
12347 // We know from the range information that (V `Pred` Min ||
12348 // V == Min). We know from the guarding condition that !(V
12349 // == Min). This gives us
12350 //
12351 // V `Pred` Min || V == Min && !(V == Min)
12352 // => V `Pred` Min
12353 //
12354 // If V `Pred` Min implies LHS `Pred` RHS, we're done.
12355
12356 if (isImpliedCondOperands(Pred, LHS, RHS, V, getConstant(Min), CtxI))
12357 return true;
12358 break;
12359
12360 // `LHS < RHS` and `LHS <= RHS` are handled in the same way as `RHS > LHS` and `RHS >= LHS` respectively.
12361 case ICmpInst::ICMP_SLE:
12362 case ICmpInst::ICMP_ULE:
12363 if (isImpliedCondOperands(ICmpInst::getSwappedCmpPredicate(Pred), RHS,
12364 LHS, V, getConstant(SharperMin), CtxI))
12365 return true;
12366 [[fallthrough]];
12367
12368 case ICmpInst::ICMP_SLT:
12369 case ICmpInst::ICMP_ULT:
12370 if (isImpliedCondOperands(ICmpInst::getSwappedCmpPredicate(Pred), RHS,
12371 LHS, V, getConstant(Min), CtxI))
12372 return true;
12373 break;
12374
12375 default:
12376 // No change
12377 break;
12378 }
12379 }
12380 }
12381
12382 // Check whether the actual condition is beyond sufficient.
12383 if (FoundPred == ICmpInst::ICMP_EQ)
12384 if (ICmpInst::isTrueWhenEqual(Pred))
12385 if (isImpliedCondOperands(Pred, LHS, RHS, FoundLHS, FoundRHS, CtxI))
12386 return true;
12387 if (Pred == ICmpInst::ICMP_NE)
12388 if (!ICmpInst::isTrueWhenEqual(FoundPred))
12389 if (isImpliedCondOperands(FoundPred, LHS, RHS, FoundLHS, FoundRHS, CtxI))
12390 return true;
12391
12392 if (isImpliedCondOperandsViaRanges(Pred, LHS, RHS, FoundPred, FoundLHS, FoundRHS))
12393 return true;
12394
12395 // Otherwise assume the worst.
12396 return false;
12397}
12398
12399bool ScalarEvolution::splitBinaryAdd(SCEVUse Expr, SCEVUse &L, SCEVUse &R,
12400 SCEVFlags &Flags) {
12401 if (!match(Expr, m_scev_Add(m_SCEV(L), m_SCEV(R))))
12402 return false;
12403
12404 Flags = cast<SCEVAddExpr>(Expr)->getNoWrapFlags();
12405 return true;
12406}
12407
12408std::optional<APInt>
12410 // We avoid subtracting expressions here because this function is usually
12411 // fairly deep in the call stack (i.e. is called many times).
12412
12413 unsigned BW = getTypeSizeInBits(More->getType());
12414 APInt Diff(BW, 0);
12415 APInt DiffMul(BW, 1);
12416 // Try various simplifications to reduce the difference to a constant. Limit
12417 // the number of allowed simplifications to keep compile-time low.
12418 for (unsigned I = 0; I < 8; ++I) {
12419 if (More == Less)
12420 return Diff;
12421
12422 // Reduce addrecs with identical steps to their start value.
12424 const auto *LAR = cast<SCEVAddRecExpr>(Less);
12425 const auto *MAR = cast<SCEVAddRecExpr>(More);
12426
12427 if (LAR->getLoop() != MAR->getLoop())
12428 return std::nullopt;
12429
12430 // We look at affine expressions only; not for correctness but to keep
12431 // getStepRecurrence cheap.
12432 if (!LAR->isAffine() || !MAR->isAffine())
12433 return std::nullopt;
12434
12435 if (LAR->getStepRecurrence(*this) != MAR->getStepRecurrence(*this))
12436 return std::nullopt;
12437
12438 Less = LAR->getStart();
12439 More = MAR->getStart();
12440 continue;
12441 }
12442
12443 // Try to match a common constant multiply.
12444 auto MatchConstMul =
12445 [](const SCEV *S) -> std::optional<std::pair<const SCEV *, APInt>> {
12446 const APInt *C;
12447 const SCEV *Op;
12448 if (match(S, m_scev_Mul(m_scev_APInt(C), m_SCEV(Op))))
12449 return {{Op, *C}};
12450 return std::nullopt;
12451 };
12452 if (auto MatchedMore = MatchConstMul(More)) {
12453 if (auto MatchedLess = MatchConstMul(Less)) {
12454 if (MatchedMore->second == MatchedLess->second) {
12455 More = MatchedMore->first;
12456 Less = MatchedLess->first;
12457 DiffMul *= MatchedMore->second;
12458 continue;
12459 }
12460 }
12461 }
12462
12463 // Try to cancel out common factors in two add expressions.
12465 auto Add = [&](const SCEV *S, int Mul) {
12466 if (auto *C = dyn_cast<SCEVConstant>(S)) {
12467 if (Mul == 1) {
12468 Diff += C->getAPInt() * DiffMul;
12469 } else {
12470 assert(Mul == -1);
12471 Diff -= C->getAPInt() * DiffMul;
12472 }
12473 } else
12474 Multiplicity[S] += Mul;
12475 };
12476 auto Decompose = [&](const SCEV *S, int Mul) {
12477 if (isa<SCEVAddExpr>(S)) {
12478 for (const SCEV *Op : S->operands())
12479 Add(Op, Mul);
12480 } else
12481 Add(S, Mul);
12482 };
12483 Decompose(More, 1);
12484 Decompose(Less, -1);
12485
12486 // Check whether all the non-constants cancel out, or reduce to new
12487 // More/Less values.
12488 const SCEV *NewMore = nullptr, *NewLess = nullptr;
12489 for (const auto &[S, Mul] : Multiplicity) {
12490 if (Mul == 0)
12491 continue;
12492 if (Mul == 1) {
12493 if (NewMore)
12494 return std::nullopt;
12495 NewMore = S;
12496 } else if (Mul == -1) {
12497 if (NewLess)
12498 return std::nullopt;
12499 NewLess = S;
12500 } else
12501 return std::nullopt;
12502 }
12503
12504 // Values stayed the same, no point in trying further.
12505 if (NewMore == More || NewLess == Less)
12506 return std::nullopt;
12507
12508 More = NewMore;
12509 Less = NewLess;
12510
12511 // Reduced to constant.
12512 if (!More && !Less)
12513 return Diff;
12514
12515 // Left with variable on only one side, bail out.
12516 if (!More || !Less)
12517 return std::nullopt;
12518 }
12519
12520 // Did not reduce to constant.
12521 return std::nullopt;
12522}
12523
12524bool ScalarEvolution::isImpliedCondOperandsViaAddRecStart(
12525 CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, const SCEV *FoundLHS,
12526 const SCEV *FoundRHS, const Instruction *CtxI) {
12527 // Try to recognize the following pattern:
12528 //
12529 // FoundRHS = ...
12530 // ...
12531 // loop:
12532 // FoundLHS = {Start,+,W}
12533 // context_bb: // Basic block from the same loop
12534 // known(Pred, FoundLHS, FoundRHS)
12535 //
12536 // If some predicate is known in the context of a loop, it is also known on
12537 // each iteration of this loop, including the first iteration. Therefore, in
12538 // this case, `FoundLHS Pred FoundRHS` implies `Start Pred FoundRHS`. Try to
12539 // prove the original pred using this fact.
12540 if (!CtxI)
12541 return false;
12542 const BasicBlock *ContextBB = CtxI->getParent();
12543 // Make sure AR varies in the context block.
12544 if (auto *AR = dyn_cast<SCEVAddRecExpr>(FoundLHS)) {
12545 const Loop *L = AR->getLoop();
12546 const auto *Latch = L->getLoopLatch();
12547 // Make sure that context belongs to the loop and executes on 1st iteration
12548 // (if it ever executes at all).
12549 if (!L->contains(ContextBB) || !Latch || !DT.dominates(ContextBB, Latch))
12550 return false;
12551 if (!isAvailableAtLoopEntry(FoundRHS, AR->getLoop()))
12552 return false;
12553 return isImpliedCondOperands(Pred, LHS, RHS, AR->getStart(), FoundRHS);
12554 }
12555
12556 if (auto *AR = dyn_cast<SCEVAddRecExpr>(FoundRHS)) {
12557 const Loop *L = AR->getLoop();
12558 const auto *Latch = L->getLoopLatch();
12559 // Make sure that context belongs to the loop and executes on 1st iteration
12560 // (if it ever executes at all).
12561 if (!L->contains(ContextBB) || !Latch || !DT.dominates(ContextBB, Latch))
12562 return false;
12563 if (!isAvailableAtLoopEntry(FoundLHS, AR->getLoop()))
12564 return false;
12565 return isImpliedCondOperands(Pred, LHS, RHS, FoundLHS, AR->getStart());
12566 }
12567
12568 return false;
12569}
12570
12571bool ScalarEvolution::isImpliedCondOperandsViaNoOverflow(CmpPredicate Pred,
12572 const SCEV *LHS,
12573 const SCEV *RHS,
12574 const SCEV *FoundLHS,
12575 const SCEV *FoundRHS) {
12576 if (Pred != CmpInst::ICMP_SLT && Pred != CmpInst::ICMP_ULT)
12577 return false;
12578
12579 const auto *AddRecLHS = dyn_cast<SCEVAddRecExpr>(LHS);
12580 if (!AddRecLHS)
12581 return false;
12582
12583 const auto *AddRecFoundLHS = dyn_cast<SCEVAddRecExpr>(FoundLHS);
12584 if (!AddRecFoundLHS)
12585 return false;
12586
12587 // We'd like to let SCEV reason about control dependencies, so we constrain
12588 // both the inequalities to be about add recurrences on the same loop. This
12589 // way we can use isLoopEntryGuardedByCond later.
12590
12591 const Loop *L = AddRecFoundLHS->getLoop();
12592 if (L != AddRecLHS->getLoop())
12593 return false;
12594
12595 // FoundLHS u< FoundRHS u< -C => (FoundLHS + C) u< (FoundRHS + C) ... (1)
12596 //
12597 // FoundLHS s< FoundRHS s< INT_MIN - C => (FoundLHS + C) s< (FoundRHS + C)
12598 // ... (2)
12599 //
12600 // Informal proof for (2), assuming (1) [*]:
12601 //
12602 // We'll also assume (A s< B) <=> ((A + INT_MIN) u< (B + INT_MIN)) ... (3)[**]
12603 //
12604 // Then
12605 //
12606 // FoundLHS s< FoundRHS s< INT_MIN - C
12607 // <=> (FoundLHS + INT_MIN) u< (FoundRHS + INT_MIN) u< -C [ using (3) ]
12608 // <=> (FoundLHS + INT_MIN + C) u< (FoundRHS + INT_MIN + C) [ using (1) ]
12609 // <=> (FoundLHS + INT_MIN + C + INT_MIN) s<
12610 // (FoundRHS + INT_MIN + C + INT_MIN) [ using (3) ]
12611 // <=> FoundLHS + C s< FoundRHS + C
12612 //
12613 // [*]: (1) can be proved by ruling out overflow.
12614 //
12615 // [**]: This can be proved by analyzing all the four possibilities:
12616 // (A s< 0, B s< 0), (A s< 0, B s>= 0), (A s>= 0, B s< 0) and
12617 // (A s>= 0, B s>= 0).
12618 //
12619 // Note:
12620 // Despite (2), "FoundRHS s< INT_MIN - C" does not mean that "FoundRHS + C"
12621 // will not sign underflow. For instance, say FoundLHS = (i8 -128), FoundRHS
12622 // = (i8 -127) and C = (i8 -100). Then INT_MIN - C = (i8 -28), and FoundRHS
12623 // s< (INT_MIN - C). Lack of sign overflow / underflow in "FoundRHS + C" is
12624 // neither necessary nor sufficient to prove "(FoundLHS + C) s< (FoundRHS +
12625 // C)".
12626
12627 std::optional<APInt> LDiff = computeConstantDifference(LHS, FoundLHS);
12628 if (!LDiff)
12629 return false;
12630 std::optional<APInt> RDiff = computeConstantDifference(RHS, FoundRHS);
12631 if (!RDiff || *LDiff != *RDiff)
12632 return false;
12633
12634 if (LDiff->isMinValue())
12635 return true;
12636
12637 APInt FoundRHSLimit;
12638
12639 if (Pred == CmpInst::ICMP_ULT) {
12640 FoundRHSLimit = -(*RDiff);
12641 } else {
12642 assert(Pred == CmpInst::ICMP_SLT && "Checked above!");
12643 FoundRHSLimit = APInt::getSignedMinValue(getTypeSizeInBits(RHS->getType())) - *RDiff;
12644 }
12645
12646 // Try to prove (1) or (2), as needed.
12647 return isAvailableAtLoopEntry(FoundRHS, L) &&
12648 isLoopEntryGuardedByCond(L, Pred, FoundRHS,
12649 getConstant(FoundRHSLimit));
12650}
12651
12652bool ScalarEvolution::isImpliedViaMerge(CmpPredicate Pred, const SCEV *LHS,
12653 const SCEV *RHS, const SCEV *FoundLHS,
12654 const SCEV *FoundRHS, unsigned Depth) {
12655 const PHINode *LPhi = nullptr, *RPhi = nullptr;
12656
12657 llvm::scope_exit ClearOnExit([&]() {
12658 if (LPhi) {
12659 bool Erased = PendingMerges.erase(LPhi);
12660 assert(Erased && "Failed to erase LPhi!");
12661 (void)Erased;
12662 }
12663 if (RPhi) {
12664 bool Erased = PendingMerges.erase(RPhi);
12665 assert(Erased && "Failed to erase RPhi!");
12666 (void)Erased;
12667 }
12668 });
12669
12670 // Find respective Phis and check that they are not being pending.
12671 if (const SCEVUnknown *LU = dyn_cast<SCEVUnknown>(LHS))
12672 if (auto *Phi = dyn_cast<PHINode>(LU->getValue())) {
12673 if (!PendingMerges.insert(Phi).second)
12674 return false;
12675 LPhi = Phi;
12676 }
12677 if (const SCEVUnknown *RU = dyn_cast<SCEVUnknown>(RHS))
12678 if (auto *Phi = dyn_cast<PHINode>(RU->getValue())) {
12679 // If we detect a loop of Phi nodes being processed by this method, for
12680 // example:
12681 //
12682 // %a = phi i32 [ %some1, %preheader ], [ %b, %latch ]
12683 // %b = phi i32 [ %some2, %preheader ], [ %a, %latch ]
12684 //
12685 // we don't want to deal with a case that complex, so return conservative
12686 // answer false.
12687 if (!PendingMerges.insert(Phi).second)
12688 return false;
12689 RPhi = Phi;
12690 }
12691
12692 // If none of LHS, RHS is a Phi, nothing to do here.
12693 if (!LPhi && !RPhi)
12694 return false;
12695
12696 // If there is a SCEVUnknown Phi we are interested in, make it left.
12697 if (!LPhi) {
12698 std::swap(LHS, RHS);
12699 std::swap(FoundLHS, FoundRHS);
12700 std::swap(LPhi, RPhi);
12702 }
12703
12704 assert(LPhi && "LPhi should definitely be a SCEVUnknown Phi!");
12705 const BasicBlock *LBB = LPhi->getParent();
12706 const SCEVAddRecExpr *RAR = dyn_cast<SCEVAddRecExpr>(RHS);
12707
12708 auto ProvedEasily = [&](const SCEV *S1, const SCEV *S2) {
12709 return isKnownViaNonRecursiveReasoning(Pred, S1, S2) ||
12710 isImpliedCondOperandsViaRanges(Pred, S1, S2, Pred, FoundLHS, FoundRHS) ||
12711 isImpliedViaOperations(Pred, S1, S2, FoundLHS, FoundRHS, Depth);
12712 };
12713
12714 if (RPhi && RPhi->getParent() == LBB) {
12715 // Case one: RHS is also a SCEVUnknown Phi from the same basic block.
12716 // If we compare two Phis from the same block, and for each entry block
12717 // the predicate is true for incoming values from this block, then the
12718 // predicate is also true for the Phis.
12719 for (const BasicBlock *IncBB : predecessors(LBB)) {
12720 const SCEV *L = getSCEV(LPhi->getIncomingValueForBlock(IncBB));
12721 const SCEV *R = getSCEV(RPhi->getIncomingValueForBlock(IncBB));
12722 if (!ProvedEasily(L, R))
12723 return false;
12724 }
12725 } else if (RAR && RAR->getLoop()->getHeader() == LBB) {
12726 // Case two: RHS is also a Phi from the same basic block, and it is an
12727 // AddRec. It means that there is a loop which has both AddRec and Unknown
12728 // PHIs, for it we can compare incoming values of AddRec from above the loop
12729 // and latch with their respective incoming values of LPhi.
12730 // TODO: Generalize to handle loops with many inputs in a header.
12731 if (LPhi->getNumIncomingValues() != 2) return false;
12732
12733 auto *RLoop = RAR->getLoop();
12734 auto *Predecessor = RLoop->getLoopPredecessor();
12735 assert(Predecessor && "Loop with AddRec with no predecessor?");
12736 const SCEV *L1 = getSCEV(LPhi->getIncomingValueForBlock(Predecessor));
12737 if (!ProvedEasily(L1, RAR->getStart()))
12738 return false;
12739 auto *Latch = RLoop->getLoopLatch();
12740 assert(Latch && "Loop with AddRec with no latch?");
12741 const SCEV *L2 = getSCEV(LPhi->getIncomingValueForBlock(Latch));
12742 if (!ProvedEasily(L2, RAR->getPostIncExpr(*this)))
12743 return false;
12744 } else {
12745 // In all other cases go over inputs of LHS and compare each of them to RHS,
12746 // the predicate is true for (LHS, RHS) if it is true for all such pairs.
12747 // At this point RHS is either a non-Phi, or it is a Phi from some block
12748 // different from LBB.
12749 for (const BasicBlock *IncBB : predecessors(LBB)) {
12750 // Check that RHS is available in this block.
12751 if (!dominates(RHS, IncBB))
12752 return false;
12753 const SCEV *L = getSCEV(LPhi->getIncomingValueForBlock(IncBB));
12754 // Make sure L does not refer to a value from a potentially previous
12755 // iteration of a loop.
12756 if (!properlyDominates(L, LBB))
12757 return false;
12758 // Addrecs are considered to properly dominate their loop, so are missed
12759 // by the previous check. Discard any values that have computable
12760 // evolution in this loop.
12761 if (auto *Loop = LI.getLoopFor(LBB))
12763 return false;
12764 if (!ProvedEasily(L, RHS))
12765 return false;
12766 }
12767 }
12768 return true;
12769}
12770
12771bool ScalarEvolution::isImpliedCondOperandsViaShift(CmpPredicate Pred,
12772 const SCEV *LHS,
12773 const SCEV *RHS,
12774 const SCEV *FoundLHS,
12775 const SCEV *FoundRHS) {
12776 // We want to imply LHS < RHS from LHS < (RHS >> shiftvalue). First, make
12777 // sure that we are dealing with same LHS.
12778 if (RHS == FoundRHS) {
12779 std::swap(LHS, RHS);
12780 std::swap(FoundLHS, FoundRHS);
12782 }
12783 if (LHS != FoundLHS)
12784 return false;
12785
12786 auto *SUFoundRHS = dyn_cast<SCEVUnknown>(FoundRHS);
12787 if (!SUFoundRHS)
12788 return false;
12789
12790 Value *Shiftee, *ShiftValue;
12791
12792 using namespace PatternMatch;
12793 if (match(SUFoundRHS->getValue(),
12794 m_LShr(m_Value(Shiftee), m_Value(ShiftValue)))) {
12795 auto *ShifteeS = getSCEV(Shiftee);
12796 // Prove one of the following:
12797 // LHS <u (shiftee >> shiftvalue) && shiftee <=u RHS ---> LHS <u RHS
12798 // LHS <=u (shiftee >> shiftvalue) && shiftee <=u RHS ---> LHS <=u RHS
12799 // LHS <s (shiftee >> shiftvalue) && shiftee <=s RHS && shiftee >=s 0
12800 // ---> LHS <s RHS
12801 // LHS <=s (shiftee >> shiftvalue) && shiftee <=s RHS && shiftee >=s 0
12802 // ---> LHS <=s RHS
12803 if (Pred == ICmpInst::ICMP_ULT || Pred == ICmpInst::ICMP_ULE)
12804 return isKnownPredicate(ICmpInst::ICMP_ULE, ShifteeS, RHS);
12805 if (Pred == ICmpInst::ICMP_SLT || Pred == ICmpInst::ICMP_SLE)
12806 if (isKnownNonNegative(ShifteeS))
12807 return isKnownPredicate(ICmpInst::ICMP_SLE, ShifteeS, RHS);
12808 }
12809
12810 return false;
12811}
12812
12813bool ScalarEvolution::isImpliedCondOperandsViaMatchingDiff(
12814 CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, const SCEV *FoundLHS,
12815 const SCEV *FoundRHS) {
12816 // Only valid for equality predicates: (A == B) implies (C == D) when
12817 // the SCEV difference A - B equals C - D (they check the same
12818 // underlying relationship at every iteration).
12819 if (!ICmpInst::isEquality(Pred))
12820 return false;
12821
12822 // Restrict to cases involving loop recurrences - that's where this
12823 // pattern arises (correlated IV comparisons). This avoids calling
12824 // getMinusSCEV on arbitrary non-loop expressions.
12826 (!isa<SCEVAddRecExpr>(FoundLHS) && !isa<SCEVAddRecExpr>(FoundRHS)))
12827 return false;
12828
12829 // AddRecs from different loops can never produce matching differences.
12830 const SCEVAddRecExpr *QueryAddRec = dyn_cast<SCEVAddRecExpr>(LHS);
12831 if (!QueryAddRec)
12832 QueryAddRec = cast<SCEVAddRecExpr>(RHS);
12833 const SCEVAddRecExpr *FoundAddRec = dyn_cast<SCEVAddRecExpr>(FoundLHS);
12834 if (!FoundAddRec)
12835 FoundAddRec = cast<SCEVAddRecExpr>(FoundRHS);
12836 if (QueryAddRec->getLoop() != FoundAddRec->getLoop())
12837 return false;
12838
12839 // If the strides differ, the differences can never match.
12840 if (QueryAddRec->getStepRecurrence(*this) !=
12841 FoundAddRec->getStepRecurrence(*this))
12842 return false;
12843
12844 // Compute differences. For pointer-typed operands sharing the same base,
12845 // getMinusSCEV strips the common base and returns an integer SCEV.
12846 // For example, {base,+,8} - (base+8*n) = {-8n,+,8}
12847 const SCEV *FoundDiff = getMinusSCEV(FoundLHS, FoundRHS);
12848 if (isa<SCEVCouldNotCompute>(FoundDiff))
12849 return false;
12850
12851 const SCEV *Diff = getMinusSCEV(LHS, RHS);
12852 if (isa<SCEVCouldNotCompute>(Diff))
12853 return false;
12854
12855 return Diff == FoundDiff;
12856}
12857
12858bool ScalarEvolution::isImpliedCondOperands(CmpPredicate Pred, const SCEV *LHS,
12859 const SCEV *RHS,
12860 const SCEV *FoundLHS,
12861 const SCEV *FoundRHS,
12862 const Instruction *CtxI) {
12863 return isImpliedCondOperandsViaRanges(Pred, LHS, RHS, Pred, FoundLHS,
12864 FoundRHS) ||
12865 isImpliedCondOperandsViaNoOverflow(Pred, LHS, RHS, FoundLHS,
12866 FoundRHS) ||
12867 isImpliedCondOperandsViaShift(Pred, LHS, RHS, FoundLHS, FoundRHS) ||
12868 isImpliedCondOperandsViaAddRecStart(Pred, LHS, RHS, FoundLHS, FoundRHS,
12869 CtxI) ||
12870 isImpliedCondOperandsViaMatchingDiff(Pred, LHS, RHS, FoundLHS,
12871 FoundRHS) ||
12872 isImpliedCondOperandsHelper(Pred, LHS, RHS, FoundLHS, FoundRHS);
12873}
12874
12875/// Is MaybeMinMaxExpr an (U|S)(Min|Max) of Candidate and some other values?
12876template <typename MinMaxExprType>
12877static bool IsMinMaxConsistingOf(const SCEV *MaybeMinMaxExpr,
12878 const SCEV *Candidate) {
12879 const MinMaxExprType *MinMaxExpr = dyn_cast<MinMaxExprType>(MaybeMinMaxExpr);
12880 if (!MinMaxExpr)
12881 return false;
12882
12883 return is_contained(MinMaxExpr->operands(), Candidate);
12884}
12885
12887 CmpPredicate Pred, const SCEV *LHS,
12888 const SCEV *RHS) {
12889 // If both sides are affine addrecs for the same loop, with equal
12890 // steps, and we know the recurrences don't wrap, then we only
12891 // need to check the predicate on the starting values.
12892
12893 if (!ICmpInst::isRelational(Pred))
12894 return false;
12895
12896 const SCEV *LStart, *RStart, *Step;
12897 const Loop *L;
12898 if (!match(LHS,
12899 m_scev_AffineAddRec(m_SCEV(LStart), m_SCEV(Step), m_Loop(L))) ||
12901 m_SpecificLoop(L))))
12902 return false;
12906 if (!LAR->getNoWrapFlags(NW) || !RAR->getNoWrapFlags(NW))
12907 return false;
12908
12909 return SE.isKnownPredicate(Pred, LStart, RStart);
12910}
12911
12912/// Is LHS `Pred` RHS true because one of them is an AddRec that is known not to
12913/// go below its own start value?
12915 CmpPredicate Pred,
12916 const SCEV *LHS,
12917 const SCEV *RHS) {
12918 // Normalize to (AddRec Pred Start).
12921 std::swap(LHS, RHS);
12922 }
12923
12924 // The recurrence is equal to Start in the first iteration, so only the
12925 // non-strict predicate holds.
12926 if (Pred != ICmpInst::ICMP_UGE && Pred != ICmpInst::ICMP_SGE)
12927 return false;
12928
12929 const auto *AR = dyn_cast<SCEVAddRecExpr>(LHS);
12930 if (!AR || AR->getStart() != RHS)
12931 return false;
12932
12933 return SE.getMonotonicPredicateType(AR, Pred) ==
12935}
12936
12937/// Is LHS `Pred` RHS true on the virtue of LHS or RHS being a Min or Max
12938/// expression?
12940 const SCEV *LHS, const SCEV *RHS) {
12941 switch (Pred) {
12942 default:
12943 return false;
12944
12945 case ICmpInst::ICMP_SGE:
12946 std::swap(LHS, RHS);
12947 [[fallthrough]];
12948 case ICmpInst::ICMP_SLE:
12949 return
12950 // min(A, ...) <= A
12952 // A <= max(A, ...)
12954
12955 case ICmpInst::ICMP_UGE:
12956 std::swap(LHS, RHS);
12957 [[fallthrough]];
12958 case ICmpInst::ICMP_ULE:
12959 return
12960 // min(A, ...) <= A
12961 // FIXME: what about umin_seq?
12963 // A <= max(A, ...)
12965
12966 case ICmpInst::ICMP_UGT:
12967 std::swap(LHS, RHS);
12968 [[fallthrough]];
12969 case ICmpInst::ICMP_ULT:
12970 // umin(Ops) u<= each Op, so proving Op u< RHS for any Op proves
12971 // umin(Ops) u< RHS.
12972 //
12973 // Use computeConstantDifference instead of the more powerful
12974 // isKnownPredicate to keep this check cheap: isKnownPredicateViaMinOrMax
12975 // is called from isKnownViaNonRecursiveReasoning, so recursing into
12976 // the full predicate prover would be expensive.
12977 if (const auto *Min = dyn_cast<SCEVUMinExpr>(LHS)) {
12978 for (SCEVUse Op : Min->operands()) {
12979 std::optional<APInt> Diff = SE.computeConstantDifference(RHS, Op);
12980 // When Op and RHS share a common base differing by a
12981 // constant offset D (RHS - Op = D), Op u< RHS holds iff D != 0 and
12982 // RHS >= D (unsigned), i.e. the subtraction doesn't underflow.
12983 if (Diff && !Diff->isZero() && SE.getUnsignedRangeMin(RHS).uge(*Diff))
12984 return true;
12985 }
12986 }
12987 return false;
12988 }
12989
12990 llvm_unreachable("covered switch fell through?!");
12991}
12992
12993bool ScalarEvolution::isImpliedViaOperations(CmpPredicate Pred, const SCEV *LHS,
12994 const SCEV *RHS,
12995 const SCEV *FoundLHS,
12996 const SCEV *FoundRHS,
12997 unsigned Depth) {
13000 "LHS and RHS have different sizes?");
13001 assert(getTypeSizeInBits(FoundLHS->getType()) ==
13002 getTypeSizeInBits(FoundRHS->getType()) &&
13003 "FoundLHS and FoundRHS have different sizes?");
13004 // We want to avoid hurting the compile time with analysis of too big trees.
13006 return false;
13007
13008 // We only want to work with GT comparison so far.
13009 if (ICmpInst::isLT(Pred)) {
13011 std::swap(LHS, RHS);
13012 std::swap(FoundLHS, FoundRHS);
13013 }
13014
13016
13017 // For unsigned, try to reduce it to corresponding signed comparison.
13018 if (P == ICmpInst::ICMP_UGT)
13019 // We can replace unsigned predicate with its signed counterpart if all
13020 // involved values are non-negative.
13021 // TODO: We could have better support for unsigned.
13022 if (isKnownNonNegative(FoundLHS) && isKnownNonNegative(FoundRHS)) {
13023 // Knowing that both FoundLHS and FoundRHS are non-negative, and knowing
13024 // FoundLHS >u FoundRHS, we also know that FoundLHS >s FoundRHS. Let us
13025 // use this fact to prove that LHS and RHS are non-negative.
13026 const SCEV *MinusOne = getMinusOne(LHS->getType());
13027 if (isImpliedCondOperands(ICmpInst::ICMP_SGT, LHS, MinusOne, FoundLHS,
13028 FoundRHS) &&
13029 isImpliedCondOperands(ICmpInst::ICMP_SGT, RHS, MinusOne, FoundLHS,
13030 FoundRHS))
13032 }
13033
13034 if (P != ICmpInst::ICMP_SGT)
13035 return false;
13036
13037 auto GetOpFromSExt = [&](const SCEV *S) -> const SCEV * {
13038 if (auto *Ext = dyn_cast<SCEVSignExtendExpr>(S))
13039 return Ext->getOperand();
13040 // TODO: If S is a SCEVConstant then you can cheaply "strip" the sext off
13041 // the constant in some cases.
13042 return S;
13043 };
13044
13045 // Acquire values from extensions.
13046 auto *OrigLHS = LHS;
13047 auto *OrigFoundLHS = FoundLHS;
13048 LHS = GetOpFromSExt(LHS);
13049 FoundLHS = GetOpFromSExt(FoundLHS);
13050
13051 // Is the SGT predicate can be proved trivially or using the found context.
13052 auto IsSGTViaContext = [&](const SCEV *S1, const SCEV *S2) {
13053 return isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_SGT, S1, S2) ||
13054 isImpliedViaOperations(ICmpInst::ICMP_SGT, S1, S2, OrigFoundLHS,
13055 FoundRHS, Depth + 1);
13056 };
13057
13058 if (auto *LHSAddExpr = dyn_cast<SCEVAddExpr>(LHS)) {
13059 // We want to avoid creation of any new non-constant SCEV. Since we are
13060 // going to compare the operands to RHS, we should be certain that we don't
13061 // need any size extensions for this. So let's decline all cases when the
13062 // sizes of types of LHS and RHS do not match.
13063 // TODO: Maybe try to get RHS from sext to catch more cases?
13065 return false;
13066
13067 // Should not overflow.
13068 if (!LHSAddExpr->hasNoSignedWrap())
13069 return false;
13070
13071 SCEVUse LL = LHSAddExpr->getOperand(0);
13072 SCEVUse LR = LHSAddExpr->getOperand(1);
13073 auto *MinusOne = getMinusOne(RHS->getType());
13074
13075 // Checks that S1 >= 0 && S2 > RHS, trivially or using the found context.
13076 auto IsSumGreaterThanRHS = [&](const SCEV *S1, const SCEV *S2) {
13077 return IsSGTViaContext(S1, MinusOne) && IsSGTViaContext(S2, RHS);
13078 };
13079 // Try to prove the following rule:
13080 // (LHS = LL + LR) && (LL >= 0) && (LR > RHS) => (LHS > RHS).
13081 // (LHS = LL + LR) && (LR >= 0) && (LL > RHS) => (LHS > RHS).
13082 if (IsSumGreaterThanRHS(LL, LR) || IsSumGreaterThanRHS(LR, LL))
13083 return true;
13084 } else if (auto *LHSUnknownExpr = dyn_cast<SCEVUnknown>(LHS)) {
13085 Value *LL, *LR;
13086 // FIXME: Once we have SDiv implemented, we can get rid of this matching.
13087
13088 using namespace llvm::PatternMatch;
13089
13090 if (match(LHSUnknownExpr->getValue(), m_SDiv(m_Value(LL), m_Value(LR)))) {
13091 // Rules for division.
13092 // We are going to perform some comparisons with Denominator and its
13093 // derivative expressions. In general case, creating a SCEV for it may
13094 // lead to a complex analysis of the entire graph, and in particular it
13095 // can request trip count recalculation for the same loop. This would
13096 // cache as SCEVCouldNotCompute to avoid the infinite recursion. To avoid
13097 // this, we only want to create SCEVs that are constants in this section.
13098 // So we bail if Denominator is not a constant.
13099 if (!isa<ConstantInt>(LR))
13100 return false;
13101
13102 auto *Denominator = cast<SCEVConstant>(getSCEV(LR));
13103
13104 // We want to make sure that LHS = FoundLHS / Denominator. If it is so,
13105 // then a SCEV for the numerator already exists and matches with FoundLHS.
13106 auto *Numerator = getExistingSCEV(LL);
13107 if (!Numerator || Numerator->getType() != FoundLHS->getType())
13108 return false;
13109
13110 // Make sure that the numerator matches with FoundLHS and the denominator
13111 // is positive.
13112 if (!HasSameValue(Numerator, FoundLHS) || !isKnownPositive(Denominator))
13113 return false;
13114
13115 auto *DTy = Denominator->getType();
13116 auto *FRHSTy = FoundRHS->getType();
13117 if (DTy->isPointerTy() != FRHSTy->isPointerTy())
13118 // One of types is a pointer and another one is not. We cannot extend
13119 // them properly to a wider type, so let us just reject this case.
13120 // TODO: Usage of getEffectiveSCEVType for DTy, FRHSTy etc should help
13121 // to avoid this check.
13122 return false;
13123
13124 // Given that:
13125 // FoundLHS > FoundRHS, LHS = FoundLHS / Denominator, Denominator > 0.
13126 auto *WTy = getWiderType(DTy, FRHSTy);
13127 auto *DenominatorExt = getNoopOrSignExtend(Denominator, WTy);
13128 auto *FoundRHSExt = getNoopOrSignExtend(FoundRHS, WTy);
13129
13130 // Try to prove the following rule:
13131 // (FoundRHS > Denominator - 2) && (RHS <= 0) => (LHS > RHS).
13132 // For example, given that FoundLHS > 2. It means that FoundLHS is at
13133 // least 3. If we divide it by Denominator < 4, we will have at least 1.
13134 auto *DenomMinusTwo = getMinusSCEV(DenominatorExt, getConstant(WTy, 2));
13135 if (isKnownNonPositive(RHS) &&
13136 IsSGTViaContext(FoundRHSExt, DenomMinusTwo))
13137 return true;
13138
13139 // Try to prove the following rule:
13140 // (FoundRHS > -1 - Denominator) && (RHS < 0) => (LHS > RHS).
13141 // For example, given that FoundLHS > -3. Then FoundLHS is at least -2.
13142 // If we divide it by Denominator > 2, then:
13143 // 1. If FoundLHS is negative, then the result is 0.
13144 // 2. If FoundLHS is non-negative, then the result is non-negative.
13145 // Anyways, the result is non-negative.
13146 auto *MinusOne = getMinusOne(WTy);
13147 auto *NegDenomMinusOne = getMinusSCEV(MinusOne, DenominatorExt);
13148 if (isKnownNegative(RHS) &&
13149 IsSGTViaContext(FoundRHSExt, NegDenomMinusOne))
13150 return true;
13151 }
13152 }
13153
13154 // If our expression contained SCEVUnknown Phis, and we split it down and now
13155 // need to prove something for them, try to prove the predicate for every
13156 // possible incoming values of those Phis.
13157 if (isImpliedViaMerge(Pred, OrigLHS, RHS, OrigFoundLHS, FoundRHS, Depth + 1))
13158 return true;
13159
13160 return false;
13161}
13162
13164 const SCEV *RHS) {
13165 // zext x u<= sext x, sext x s<= zext x
13166 const SCEV *Op;
13167 switch (Pred) {
13168 case ICmpInst::ICMP_SGE:
13169 std::swap(LHS, RHS);
13170 [[fallthrough]];
13171 case ICmpInst::ICMP_SLE: {
13172 // If operand >=s 0 then ZExt == SExt. If operand <s 0 then SExt <s ZExt.
13173 return match(LHS, m_scev_SExt(m_SCEV(Op))) &&
13175 }
13176 case ICmpInst::ICMP_UGE:
13177 std::swap(LHS, RHS);
13178 [[fallthrough]];
13179 case ICmpInst::ICMP_ULE: {
13180 // If operand >=u 0 then ZExt == SExt. If operand <u 0 then ZExt <u SExt.
13181 return match(LHS, m_scev_ZExt(m_SCEV(Op))) &&
13183 }
13184 default:
13185 return false;
13186 };
13187 llvm_unreachable("unhandled case");
13188}
13189
13190bool ScalarEvolution::isKnownViaNonRecursiveReasoning(CmpPredicate Pred,
13191 SCEVUse LHS,
13192 SCEVUse RHS) {
13193 return isKnownPredicateExtendIdiom(Pred, LHS, RHS) ||
13194 isKnownPredicateViaConstantRanges(Pred, LHS, RHS) ||
13195 IsKnownPredicateViaMinOrMax(*this, Pred, LHS, RHS) ||
13196 IsKnownPredicateViaAddRecStart(*this, Pred, LHS, RHS) ||
13198 isKnownPredicateViaNoOverflow(Pred, LHS, RHS);
13199}
13200
13201bool ScalarEvolution::isImpliedCondOperandsHelper(CmpPredicate Pred,
13202 const SCEV *LHS,
13203 const SCEV *RHS,
13204 const SCEV *FoundLHS,
13205 const SCEV *FoundRHS) {
13206 switch (Pred) {
13207 default:
13208 llvm_unreachable("Unexpected CmpPredicate value!");
13209 case ICmpInst::ICMP_EQ:
13210 case ICmpInst::ICMP_NE:
13211 if (HasSameValue(LHS, FoundLHS) && HasSameValue(RHS, FoundRHS))
13212 return true;
13213 break;
13214 case ICmpInst::ICMP_SLT:
13215 case ICmpInst::ICMP_SLE:
13216 if (isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_SLE, LHS, FoundLHS) &&
13217 isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_SGE, RHS, FoundRHS))
13218 return true;
13219 break;
13220 case ICmpInst::ICMP_SGT:
13221 case ICmpInst::ICMP_SGE:
13222 if (isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_SGE, LHS, FoundLHS) &&
13223 isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_SLE, RHS, FoundRHS))
13224 return true;
13225 break;
13226 case ICmpInst::ICMP_ULT:
13227 case ICmpInst::ICMP_ULE:
13228 if (isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_ULE, LHS, FoundLHS) &&
13229 isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_UGE, RHS, FoundRHS))
13230 return true;
13231 break;
13232 case ICmpInst::ICMP_UGT:
13233 case ICmpInst::ICMP_UGE:
13234 if (isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_UGE, LHS, FoundLHS) &&
13235 isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_ULE, RHS, FoundRHS))
13236 return true;
13237 break;
13238 }
13239
13240 // Maybe it can be proved via operations?
13241 if (isImpliedViaOperations(Pred, LHS, RHS, FoundLHS, FoundRHS))
13242 return true;
13243
13244 return false;
13245}
13246
13247bool ScalarEvolution::isImpliedCondOperandsViaRanges(
13248 CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, CmpPredicate FoundPred,
13249 const SCEV *FoundLHS, const SCEV *FoundRHS) {
13250 if (!isa<SCEVConstant>(RHS) || !isa<SCEVConstant>(FoundRHS))
13251 // The restriction on `FoundRHS` be lifted easily -- it exists only to
13252 // reduce the compile time impact of this optimization.
13253 return false;
13254
13255 std::optional<APInt> Addend = computeConstantDifference(LHS, FoundLHS);
13256 if (!Addend)
13257 return false;
13258
13259 const APInt &ConstFoundRHS = cast<SCEVConstant>(FoundRHS)->getAPInt();
13260
13261 // `FoundLHSRange` is the range we know `FoundLHS` to be in by virtue of the
13262 // antecedent "`FoundLHS` `FoundPred` `FoundRHS`".
13263 ConstantRange FoundLHSRange =
13264 ConstantRange::makeExactICmpRegion(FoundPred, ConstFoundRHS);
13265
13266 // Since `LHS` is `FoundLHS` + `Addend`, we can compute a range for `LHS`:
13267 ConstantRange LHSRange = FoundLHSRange.add(ConstantRange(*Addend));
13268
13269 // We can also compute the range of values for `LHS` that satisfy the
13270 // consequent, "`LHS` `Pred` `RHS`":
13271 const APInt &ConstRHS = cast<SCEVConstant>(RHS)->getAPInt();
13272 // The antecedent implies the consequent if every value of `LHS` that
13273 // satisfies the antecedent also satisfies the consequent.
13274 return LHSRange.icmp(Pred, ConstRHS);
13275}
13276
13277bool ScalarEvolution::canIVOverflowOnLT(const SCEV *RHS, const SCEV *Stride,
13278 bool IsSigned, bool Invert) {
13279 assert(isKnownPositive(Stride) && "Positive stride expected!");
13280
13281 unsigned BitWidth = getTypeSizeInBits(RHS->getType());
13282 const SCEV *One = getOne(Stride->getType());
13283
13284 if (IsSigned) {
13285 APInt MaxRHS = getRangeMax(RHS, /*IsSigned=*/true, Invert);
13286 APInt MaxValue = APInt::getSignedMaxValue(BitWidth);
13287 APInt MaxStrideMinusOne = getSignedRangeMax(getMinusSCEV(Stride, One));
13288
13289 // SMaxRHS + SMaxStrideMinusOne > SMaxValue => overflow!
13290 return (std::move(MaxValue) - MaxStrideMinusOne).slt(MaxRHS);
13291 }
13292
13293 APInt MaxRHS = getRangeMax(RHS, /*IsSigned=*/false, Invert);
13294 APInt MaxValue = APInt::getMaxValue(BitWidth);
13295 APInt MaxStrideMinusOne = getUnsignedRangeMax(getMinusSCEV(Stride, One));
13296
13297 // UMaxRHS + UMaxStrideMinusOne > UMaxValue => overflow!
13298 return (std::move(MaxValue) - MaxStrideMinusOne).ult(MaxRHS);
13299}
13300
13302 // umin(N, 1) + floor((N - umin(N, 1)) / D)
13303 // This is equivalent to "1 + floor((N - 1) / D)" for N != 0. The umin
13304 // expression fixes the case of N=0.
13305 const SCEV *MinNOne = getUMinExpr(N, getOne(N->getType()));
13306 const SCEV *NMinusOne = getMinusSCEV(N, MinNOne);
13307 return getAddExpr(MinNOne, getUDivExpr(NMinusOne, D));
13308}
13309
13310const SCEV *
13311ScalarEvolution::computeMaxBECountForLT(const SCEV *Start, const SCEV *Stride,
13312 const SCEV *End, unsigned BitWidth,
13313 bool IsSigned, bool Invert) {
13314 // The logic in this function assumes we can represent a positive stride.
13315 // If we can't, the backedge-taken count must be zero.
13316 if (IsSigned && BitWidth == 1)
13317 return getZero(Stride->getType());
13318
13319 // This code below only been closely audited for negative strides in the
13320 // unsigned comparison case, it may be correct for signed comparison, but
13321 // that needs to be established.
13322 if (IsSigned && isKnownNegative(Stride))
13323 return getCouldNotCompute();
13324
13325 // Calculate the maximum backedge count based on the range of values
13326 // permitted by Start, End, and Stride. If Invert is true, both Start and End
13327 // need inverting. Stride was already negated by the caller.
13328 APInt MinStart = getRangeMin(Start, IsSigned, Invert);
13329
13330 APInt MinStride =
13331 IsSigned ? getSignedRangeMin(Stride) : getUnsignedRangeMin(Stride);
13332
13333 // We assume either the stride is positive, or the backedge-taken count
13334 // is zero. So force StrideForMaxBECount to be at least one.
13335 APInt One(BitWidth, 1);
13336 APInt StrideForMaxBECount = IsSigned ? APIntOps::smax(One, MinStride)
13337 : APIntOps::umax(One, MinStride);
13338
13339 APInt MaxValue = IsSigned ? APInt::getSignedMaxValue(BitWidth)
13340 : APInt::getMaxValue(BitWidth);
13341 APInt Limit = MaxValue - (StrideForMaxBECount - 1);
13342
13343 // Although End can be a MAX expression we estimate MaxEnd considering only
13344 // the case End = RHS of the loop termination condition. This is safe because
13345 // in the other case (End - Start) is zero, leading to a zero maximum backedge
13346 // taken count.
13347 APInt MaxEnd = getRangeMax(End, IsSigned, Invert);
13348 MaxEnd =
13349 IsSigned ? APIntOps::smin(MaxEnd, Limit) : APIntOps::umin(MaxEnd, Limit);
13350
13351 // MaxBECount = ceil((max(MaxEnd, MinStart) - MinStart) / Stride)
13352 MaxEnd = IsSigned ? APIntOps::smax(MaxEnd, MinStart)
13353 : APIntOps::umax(MaxEnd, MinStart);
13354
13355 APInt Delta = MaxEnd - MinStart;
13356
13357 // Try to refine Delta in case End - Start (or Start - End if Invert) gives a
13358 // tighter bound after folding.
13359 const SCEV *DeltaExpr =
13360 Invert ? getMinusSCEV(Start, End) : getMinusSCEV(End, Start);
13361 Delta = APIntOps::umin(Delta, getUnsignedRangeMax(DeltaExpr));
13362
13363 return getUDivCeilSCEV(getConstant(Delta), getConstant(StrideForMaxBECount));
13364}
13365
13367ScalarEvolution::howManyLessThans(const SCEV *LHS, const SCEV *RHS,
13368 const Loop *L, bool IsSigned, bool Invert,
13369 bool ControlsOnlyExit, bool AllowPredicates) {
13371
13372 // Loop guards for L, collected on demand.
13373 std::optional<LoopGuards> CachedGuards;
13374 auto getGuards = [&]() -> const LoopGuards & {
13375 if (!CachedGuards)
13376 CachedGuards.emplace(LoopGuards::collect(L, *this));
13377 return *CachedGuards;
13378 };
13379
13380 // FIXME: Extend the non-invariant RHS analysis to greater-than comparisons.
13381 if (Invert && !isLoopInvariant(RHS, L))
13382 return getCouldNotCompute();
13383
13384 const SCEVAddRecExpr *IV = dyn_cast<SCEVAddRecExpr>(LHS);
13385 bool PredicatedIV = false;
13386 // FIXME: Generalize the NUW inference below to decreasing IVs.
13387 if (!IV && !Invert) {
13388 if (auto *ZExt = dyn_cast<SCEVZeroExtendExpr>(LHS)) {
13389 const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(ZExt->getOperand());
13390 if (AR && AR->getLoop() == L && AR->isAffine()) {
13391 auto canProveNUW = [&]() {
13392 // We can use the comparison to infer no-wrap flags only if it fully
13393 // controls the loop exit.
13394 if (!ControlsOnlyExit)
13395 return false;
13396
13397 if (!isLoopInvariant(RHS, L))
13398 return false;
13399
13400 if (!isKnownNonZero(AR->getStepRecurrence(*this)))
13401 // We need the sequence defined by AR to strictly increase in the
13402 // unsigned integer domain for the logic below to hold.
13403 return false;
13404
13405 const unsigned InnerBitWidth = getTypeSizeInBits(AR->getType());
13406 const unsigned OuterBitWidth = getTypeSizeInBits(RHS->getType());
13407 // If RHS <=u Limit, then there must exist a value V in the sequence
13408 // defined by AR (e.g. {Start,+,Step}) such that V >u RHS, and
13409 // V <=u UINT_MAX. Thus, we must exit the loop before unsigned
13410 // overflow occurs. This limit also implies that a signed comparison
13411 // (in the wide bitwidth) is equivalent to an unsigned comparison as
13412 // the high bits on both sides must be zero.
13413 APInt StrideMax = getUnsignedRangeMax(AR->getStepRecurrence(*this));
13414 APInt Limit = APInt::getMaxValue(InnerBitWidth) - (StrideMax - 1);
13415 Limit = Limit.zext(OuterBitWidth);
13416 return getUnsignedRangeMax(applyLoopGuards(RHS, getGuards()))
13417 .ule(Limit);
13418 };
13419 auto Flags = AR->getNoWrapFlags();
13420 if (!hasFlags(Flags, SCEV::FlagNUW) && canProveNUW())
13421 Flags = setFlags(Flags, SCEV::FlagNUW);
13422
13423 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), Flags);
13424 if (AR->hasNoUnsignedWrap()) {
13425 // Emulate what getZeroExtendExpr would have done during construction
13426 // if we'd been able to infer the fact just above at that time.
13427 const SCEV *Step = AR->getStepRecurrence(*this);
13428 Type *Ty = ZExt->getType();
13429 const SCEV *S = getAddRecExpr(
13431 getZeroExtendExpr(Step, Ty, 0), L, AR->getNoWrapFlags());
13433 }
13434 }
13435 }
13436 }
13437
13438 if (!IV && AllowPredicates) {
13439 // Try to make this an AddRec using runtime tests, in the first X
13440 // iterations of this loop, where X is the SCEV expression found by the
13441 // algorithm below.
13442 IV = convertSCEVToAddRecWithPredicates(LHS, L, Predicates);
13443 PredicatedIV = true;
13444 }
13445
13446 // Avoid weird loops
13447 if (!IV || IV->getLoop() != L || !IV->isAffine())
13448 return getCouldNotCompute();
13449
13450 // A precondition of this method is that the condition being analyzed
13451 // reaches an exiting branch which dominates the latch. Given that, we can
13452 // assume that an increment which violates the nowrap specification and
13453 // produces poison must cause undefined behavior when the resulting poison
13454 // value is branched upon and thus we can conclude that the backedge is
13455 // taken no more often than would be required to produce that poison value.
13456 // Note that a well defined loop can exit on the iteration which violates
13457 // the nowrap specification if there is another exit (either explicit or
13458 // implicit/exceptional) which causes the loop to execute before the
13459 // exiting instruction we're analyzing would trigger UB.
13460 auto WrapType = IsSigned ? SCEV::FlagNSW : SCEV::FlagNUW;
13461 bool NoWrap = ControlsOnlyExit && any(IV->getNoWrapFlags(WrapType));
13462 // Reverse the ordering for greater-than comparisons.
13464 if (Invert)
13466
13467 // The step of ~IV is the negated step of IV.
13468 const SCEV *Stride = IV->getStepRecurrence(*this);
13469 if (Invert)
13470 Stride = getNegativeSCEV(Stride);
13471 const SCEV *GuardedStride = Stride;
13472
13473 // Whether the IV may reach the maximum (or minimum if inverted) value
13474 // before the exit is taken.
13475 bool IVMayOverflow = true;
13476
13477 bool PositiveStride = isKnownPositive(Stride);
13478 // A dominating guard may prove the stride positive.
13479 if (!PositiveStride) {
13480 const SCEV *LoopGuardedStride = applyLoopGuards(Stride, getGuards());
13481 if (isKnownPositive(LoopGuardedStride)) {
13482 GuardedStride = LoopGuardedStride;
13483 PositiveStride = true;
13484 // Encode the context-sensitive stride > 0 fact into the expression
13485 Stride = getUMaxExpr(Stride, getOne(Stride->getType()));
13486 }
13487 }
13488
13489 // Avoid negative or zero stride values.
13490 if (!PositiveStride) {
13491 // FIXME: Generalize the unknown-stride analysis to decreasing IVs.
13492 if (Invert)
13493 return getCouldNotCompute();
13494
13495 // We can compute the correct backedge taken count for loops with unknown
13496 // strides if we can prove that the loop is not an infinite loop with side
13497 // effects. Here's the loop structure we are trying to handle -
13498 //
13499 // i = start
13500 // do {
13501 // A[i] = i;
13502 // i += s;
13503 // } while (i < end);
13504 //
13505 // The backedge taken count for such loops is evaluated as -
13506 // (max(end, start + stride) - start - 1) /u stride
13507 //
13508 // The additional preconditions that we need to check to prove correctness
13509 // of the above formula is as follows -
13510 //
13511 // a) IV is either nuw or nsw depending upon signedness (indicated by the
13512 // NoWrap flag).
13513 // b) the loop is guaranteed to be finite (e.g. is mustprogress and has
13514 // b) the loop is guaranteed to be finite (e.g. is mustprogress and has
13515 // no side effects within the loop) or a predicate is added to ensure
13516 // stride is positive.
13517 // c) loop has a single static exit (with no abnormal exits)
13518 //
13519 // Precondition a) implies that if the stride is negative, this is a single
13520 // trip loop. The backedge taken count formula reduces to zero in this case.
13521 //
13522 // Precondition b) and c) combine to imply that if rhs is invariant in L,
13523 // then a zero stride means the backedge can't be taken without executing
13524 // undefined behavior.
13525 //
13526 // The positive stride case is the same as isKnownPositive(Stride) returning
13527 // true (original behavior of the function).
13528 //
13529 if (PredicatedIV || !NoWrap || !loopHasNoAbnormalExits(L))
13530 return getCouldNotCompute();
13531
13532 if (!loopIsFiniteByAssumption(L)) {
13533 // If the loop may be infinite, add a predicate ensuring Stride is
13534 // positive, to guarantee forward progress.
13535 if (!AllowPredicates || !isLoopInvariant(Stride, L))
13536 return getCouldNotCompute();
13537
13538 const SCEV *Zero = getZero(Stride->getType());
13539 const SCEVPredicate *P =
13541 Predicates.push_back(P);
13542 // When the predicate holds (Stride > 0), umax(Stride, 1) == Stride,
13543 // so the result is unchanged. To prevent div by zero.
13544 Stride = getUMaxExpr(Stride, getOne(Stride->getType()));
13545 } else if (!isKnownNonZero(Stride)) {
13546 // If we have a step of zero, and RHS isn't invariant in L, we don't know
13547 // if it might eventually be greater than start and if so, on which
13548 // iteration. We can't even produce a useful upper bound.
13549 if (!isLoopInvariant(RHS, L))
13550 return getCouldNotCompute();
13551
13552 // We allow a potentially zero stride, but we need to divide by stride
13553 // below. Since the loop can't be infinite and this check must control
13554 // the sole exit, we can infer the exit must be taken on the first
13555 // iteration (e.g. backedge count = 0) if the stride is zero. Given that,
13556 // we know the numerator in the divides below must be zero, so we can
13557 // pick an arbitrary non-zero value for the denominator (e.g. stride)
13558 // and produce the right result.
13559 // FIXME: Handle the case where Stride is poison?
13560 auto wouldZeroStrideBeUB = [&]() {
13561 // Proof by contradiction. Suppose the stride were zero. If we can
13562 // prove that the backedge *is* taken on the first iteration, then since
13563 // we know this condition controls the sole exit, we must have an
13564 // infinite loop. We can't have a (well defined) infinite loop per
13565 // check just above.
13566 // Note: The (Start - Stride) term is used to get the start' term from
13567 // (start' + stride,+,stride). Remember that we only care about the
13568 // result of this expression when stride == 0 at runtime.
13569 auto *StartIfZero = getMinusSCEV(IV->getStart(), Stride);
13570 return isLoopEntryGuardedByCond(L, Cond, StartIfZero, RHS);
13571 };
13572 if (!wouldZeroStrideBeUB()) {
13573 Stride = getUMaxExpr(Stride, getOne(Stride->getType()));
13574 }
13575 }
13576 } else {
13577 // Avoid proven overflow cases: this will ensure that the backedge taken
13578 // count will not generate any unsigned overflow.
13579 IVMayOverflow = canIVOverflowOnLT(RHS, GuardedStride, IsSigned, Invert);
13580 if (IVMayOverflow && !NoWrap)
13581 return getCouldNotCompute();
13582 }
13583
13584 // On all paths just preceeding, we established the following invariant:
13585 // IV can be assumed not to overflow up to and including the exiting
13586 // iteration. We proved this in one of two ways:
13587 // 1) We can show overflow doesn't occur before the exiting iteration
13588 // 1a) canIVOverflowOnLT, and b) step of one
13589 // 2) We can show that if overflow occurs, the loop must execute UB
13590 // before any possible exit.
13591 // Note that we have not yet proved RHS invariant (in general).
13592
13593 const SCEV *Start = IV->getStart();
13594
13595 // Preserve pointer-typed Start/RHS to pass to isLoopEntryGuardedByCond.
13596 // If we convert to integers, isLoopEntryGuardedByCond will miss some cases.
13597 // Use integer-typed versions for actual computation; we can't subtract
13598 // pointers in general.
13599 const SCEV *OrigStart = Start;
13600 const SCEV *OrigRHS = RHS;
13601 if (Start->getType()->isPointerTy()) {
13602 Start = getPtrToAddrExpr(Start);
13603 if (isa<SCEVCouldNotCompute>(Start))
13604 return Start;
13605 }
13606 if (RHS->getType()->isPointerTy()) {
13609 return RHS;
13610 }
13611
13612 const SCEV *End = nullptr, *BECount = getCouldNotCompute(),
13613 *BECountIfBackedgeTaken = getCouldNotCompute();
13614 if (!isLoopInvariant(RHS, L)) {
13615 assert(!Invert && "RHS must be loop-invariant for Invert");
13616 const auto *RHSAddRec = dyn_cast<SCEVAddRecExpr>(RHS);
13617 if (PositiveStride && RHSAddRec != nullptr && RHSAddRec->getLoop() == L &&
13618 any(RHSAddRec->getNoWrapFlags())) {
13619 // The structure of loop we are trying to calculate backedge count of:
13620 //
13621 // left = left_start
13622 // right = right_start
13623 //
13624 // while(left < right){
13625 // ... do something here ...
13626 // left += s1; // stride of left is s1 (s1 > 0)
13627 // right += s2; // stride of right is s2 (s2 < 0)
13628 // }
13629 //
13630
13631 const SCEV *RHSStart = RHSAddRec->getStart();
13632 const SCEV *RHSStride = RHSAddRec->getStepRecurrence(*this);
13633
13634 // If Stride - RHSStride is positive and does not overflow, we can write
13635 // backedge count as ->
13636 // ceil((End - Start) /u (Stride - RHSStride))
13637 // Where, End = max(RHSStart, Start)
13638
13639 // Check if RHSStride < 0 and Stride - RHSStride will not overflow.
13640 if (isKnownNegative(RHSStride) &&
13641 willNotOverflow(Instruction::Sub, /*Signed=*/true, Stride,
13642 RHSStride)) {
13643
13644 const SCEV *Denominator = getMinusSCEV(Stride, RHSStride);
13645 if (isKnownPositive(Denominator)) {
13646 End = IsSigned ? getSMaxExpr(RHSStart, Start)
13647 : getUMaxExpr(RHSStart, Start);
13648
13649 // We can do this because End >= Start, as End = max(RHSStart, Start)
13650 const SCEV *Delta = getMinusSCEV(End, Start);
13651
13652 BECount = getUDivCeilSCEV(Delta, Denominator);
13653 BECountIfBackedgeTaken =
13654 getUDivCeilSCEV(getMinusSCEV(RHSStart, Start), Denominator);
13655 }
13656 }
13657 }
13658 } else {
13659 // Let End = max(RHS,Start). We use the expression (End-Start)/Stride to
13660 // describe the backedge count: if the backedge is taken at least once then
13661 // End is RHS, and if not End is Start so we get a backedge count of zero.
13662 // Inverted, End is min(RHS, Start).
13663 //
13664 // AddingStrideMinusOneMayOverflow has the following preconditions:
13665 //
13666 // 1. Start <= End, signed if IsSigned (inverted: End <= Start)
13667 // 2. The index variable doesn't overflow.
13668 //
13669 // Therefore, we know N exists such that
13670 // (Start + Stride * N) >= End, and computing "(Start + Stride * N)"
13671 // doesn't overflow.
13672 //
13673 // Using this information, try to prove whether the addition in
13674 // "(End - Start) + (Stride - 1)" has unsigned overflow.
13675 //
13676 // If the IV cannot overflow, RHS is at least Stride - 1 below the maximum
13677 // value, so the distance End - Start is at most UMAX - (Stride - 1) and
13678 // the (Stride - 1) addition below cannot overflow.
13679 const SCEV *One = getOne(Stride->getType());
13680 bool AddingStrideMinusOneMayOverflow = IVMayOverflow && [&] {
13681 if (isKnownToBeAPowerOfTwo(Stride)) {
13682 // Suppose Stride is a power of two, and Start/End are unsigned
13683 // integers. Let UMAX be the largest representable unsigned
13684 // integer.
13685 //
13686 // By the preconditions of this function, we know
13687 // "(Start + Stride * N) >= End", and this doesn't overflow.
13688 // As a formula:
13689 //
13690 // End <= (Start + Stride * N) <= UMAX
13691 //
13692 // Subtracting Start from all the terms:
13693 //
13694 // End - Start <= Stride * N <= UMAX - Start
13695 //
13696 // Since Start is unsigned, UMAX - Start <= UMAX. Therefore:
13697 //
13698 // End - Start <= Stride * N <= UMAX
13699 //
13700 // Stride * N is a multiple of Stride. Therefore,
13701 //
13702 // End - Start <= Stride * N <= UMAX - (UMAX mod Stride)
13703 //
13704 // Since Stride is a power of two, UMAX + 1 is divisible by
13705 // Stride. Therefore, UMAX mod Stride == Stride - 1. So we can
13706 // write:
13707 //
13708 // End - Start <= Stride * N <= UMAX - Stride - 1
13709 //
13710 // Dropping the middle term:
13711 //
13712 // End - Start <= UMAX - Stride - 1
13713 //
13714 // Adding Stride - 1 to both sides:
13715 //
13716 // (End - Start) + (Stride - 1) <= UMAX
13717 //
13718 // In other words, the addition doesn't have unsigned overflow.
13719 //
13720 // A similar proof works if we treat Start/End as signed values.
13721 // Just rewrite steps before "End - Start <= Stride * N <= UMAX"
13722 // to use signed max instead of unsigned max. Note that we're
13723 // trying to prove a lack of unsigned overflow in either case.
13724 // Inverted: "Start - End <= Stride * N <= Start - MIN <= UMAX", same.
13725 return false;
13726 }
13727 if (!Invert && (Start == Stride || Start == getMinusSCEV(Stride, One))) {
13728 // If Start is equal to Stride, (End - Start) + (Stride - 1) == End
13729 // - 1. If !IsSigned, 0 <u Stride == Start <=u End; so 0 <u End - 1
13730 // <u End. If IsSigned, 0 <s Stride == Start <=s End; so 0 <s End -
13731 // 1 <s End.
13732 //
13733 // If Start is equal to Stride - 1, (End - Start) + Stride - 1 ==
13734 // End.
13735 //
13736 // Both need Start to be the smaller value, so neither applies inverted.
13737 return false;
13738 }
13739 return true;
13740 }();
13741
13742 // If inverted, the analyzed values are complements: "~V - Offset" is "~(V +
13743 // Offset)" and "~To - ~From" is "From - To".
13744 auto StepBack = [&](const SCEV *V, const SCEV *Offset) -> const SCEV * {
13745 if (Invert)
13746 return getAddExpr(V, Offset);
13747 return getMinusSCEV(V, Offset);
13748 };
13749 auto Distance = [&](const SCEV *From, const SCEV *To) {
13750 return Invert ? getMinusSCEV(From, To) : getMinusSCEV(To, From);
13751 };
13752
13753 const SCEV *OrigPrevStart = StepBack(OrigStart, Stride);
13754 assert(isAvailableAtLoopEntry(OrigPrevStart, L) && "Must be!");
13755 assert(isAvailableAtLoopEntry(OrigStart, L) && "Must be!");
13756 assert(isAvailableAtLoopEntry(OrigRHS, L) && "Must be!");
13757 // Can we prove Start - Stride < RHS, and either Start - Stride < Start or
13758 // (via !AddingStrideMinusOneMayOverflow) that (RHS - Start) + (Stride - 1)
13759 // does not overflow?
13760 if ((!AddingStrideMinusOneMayOverflow ||
13761 isLoopEntryGuardedByCond(L, Cond, OrigPrevStart, OrigStart)) &&
13762 isLoopEntryGuardedByCond(L, Cond, OrigPrevStart, OrigRHS)) {
13763 // In this case, we can use a refined formula for computing backedge
13764 // taken count. The general formula remains:
13765 // "End-Start /uceiling Stride"
13766 // We want to use the alternate formula:
13767 // "((RHS - 1) - (Start - Stride)) /u Stride"
13768 // Let's do a quick case analysis to show these are equivalent under
13769 // our preconditions. When inverted, the proof uses complemented Start,
13770 // RHS and End; Stride remains positive.
13771 // * For RHS <= Start (End is Start), the backedge-taken count must be
13772 // zero. Together with the precondition "Start - Stride < RHS", we have
13773 // "Start - Stride < RHS <= Start". Subtracting Start - Stride from
13774 // all sides we get "0 < RHS - (Start - Stride) <= Stride".
13775 // Subtracting 1 we get "0 <= (RHS - 1) - (Start - Stride) < Stride".
13776 // So dividing that by Stride gives zero.
13777 //
13778 // * For RHS > Start (End is RHS), the backedge count must be
13779 // "RHS-Start /uceil Stride", so it is sufficient to show that the
13780 // numerator "((RHS - 1) - (Start - Stride))" does not overflow.
13781 //
13782 // If "Start - Stride < Start" holds, we have
13783 // "RHS > Start > Start - Stride". As such
13784 // "RHS - (Start - Stride) - 1" does not overflow, which is the
13785 // reassociated numerator.
13786 //
13787 // Otherwise !AddingStrideMinusOneMayOverflow guarantees that
13788 // "(End - Start) + (Stride - 1)" does not overflow unsigned. Here
13789 // "End" is "RHS", as "RHS > Start", so this is the reassociated
13790 // numerator. Neither sub-term wraps unsigned: "RHS - Start"
13791 // due to "RHS > Start", and "Stride - 1", as Stride is non-zero.
13792 const SCEV *Numerator =
13793 getMinusSCEV(Distance(StepBack(Start, Stride), RHS), One);
13794 BECount = getUDivExpr(Numerator, Stride);
13795 }
13796
13797 if (isa<SCEVCouldNotCompute>(BECount)) {
13798 auto canProveRHSIsAtOrBeyondStart = [&]() {
13799 // Inverted, the claim is "Start >= RHS". Reverse the comparisons below
13800 // by swapping their operands rather than their predicates:
13801 // isLoopEntryGuardedByCond is sensitive to operand order and loses the
13802 // proof if the IV bound moves to the other side.
13803 auto SwapIfInverted = [&](const SCEV *A, const SCEV *B) {
13804 return Invert ? std::pair(B, A) : std::pair(A, B);
13805 };
13806
13807 auto CondGE = IsSigned ? ICmpInst::ICMP_SGE : ICmpInst::ICMP_UGE;
13808 const SCEV *GuardedRHS = applyLoopGuards(OrigRHS, getGuards());
13809 const SCEV *GuardedStart = applyLoopGuards(OrigStart, getGuards());
13810 if (Invert)
13811 std::swap(GuardedRHS, GuardedStart);
13812
13813 auto [GELHS, GERHS] = SwapIfInverted(OrigRHS, OrigStart);
13814 if (isLoopEntryGuardedByCond(L, CondGE, GELHS, GERHS) ||
13815 isKnownPredicate(CondGE, GuardedRHS, GuardedStart))
13816 return true;
13817
13818 // (RHS > Start - 1) implies RHS >= Start.
13819 // * "RHS >= Start" is trivially equivalent to "RHS > Start - 1" if
13820 // "Start - 1" doesn't overflow.
13821 // * For signed comparison, if Start - 1 does overflow, it's equal
13822 // to INT_MAX, and "RHS >s INT_MAX" is trivially false.
13823 // * For unsigned comparison, if Start - 1 does overflow, it's equal
13824 // to UINT_MAX, and "RHS >u UINT_MAX" is trivially false.
13825 //
13826 // FIXME: Should isLoopEntryGuardedByCond do this for us?
13827 auto CondGT = IsSigned ? ICmpInst::ICMP_SGT : ICmpInst::ICMP_UGT;
13828 auto [GTLHS, GTRHS] = SwapIfInverted(OrigRHS, StepBack(OrigStart, One));
13829 return isLoopEntryGuardedByCond(L, CondGT, GTLHS, GTRHS);
13830 };
13831
13832 // If we know that RHS >= Start in the context of loop, then we know
13833 // that max(RHS, Start) = RHS at this point.
13834 if (canProveRHSIsAtOrBeyondStart()) {
13835 End = RHS;
13836 } else {
13837 // If RHS < Start, the backedge will be taken zero times. So in
13838 // general, we can write the backedge-taken count as:
13839 //
13840 // RHS >= Start ? ceil(RHS - Start) / Stride : 0
13841 //
13842 // We convert it to the following to make it more convenient for SCEV:
13843 //
13844 // ceil(max(RHS, Start) - Start) / Stride
13845 //
13846 // Inverted, this is ceil(Start - min(RHS, Start)) / Stride.
13847 if (Invert)
13848 End = IsSigned ? getSMinExpr(RHS, Start) : getUMinExpr(RHS, Start);
13849 else
13850 End = IsSigned ? getSMaxExpr(RHS, Start) : getUMaxExpr(RHS, Start);
13851
13852 // See what would happen if we assume the backedge is taken. This is
13853 // used to compute MaxBECount.
13854 BECountIfBackedgeTaken = getUDivCeilSCEV(Distance(Start, RHS), Stride);
13855 }
13856
13857 const SCEV *Delta = Distance(Start, End);
13858 if (!AddingStrideMinusOneMayOverflow) {
13859 // floor((D + (S - 1)) / S)
13860 // We prefer this formulation if it's legal because it's fewer
13861 // operations.
13862 BECount =
13863 getUDivExpr(getAddExpr(Delta, getMinusSCEV(Stride, One)), Stride);
13864 } else {
13865 BECount = getUDivCeilSCEV(Delta, Stride);
13866 }
13867 }
13868 }
13869
13870 const SCEV *ConstantMaxBECount;
13871 bool MaxOrZero = false;
13872 if (isa<SCEVConstant>(BECount)) {
13873 ConstantMaxBECount = BECount;
13874 } else {
13875 ConstantMaxBECount = computeMaxBECountForLT(
13876 Start, Stride, RHS, getTypeSizeInBits(LHS->getType()), IsSigned,
13877 Invert);
13878 // If we know exactly how many times the backedge will be taken if it's
13879 // taken at least once, then the backedge count will either be that or
13880 // zero. If that count exceeds the range-based bound, the backedge can
13881 // never be taken.
13882 const APInt *IfTaken, *RangeMax;
13883 if (match(BECountIfBackedgeTaken, m_scev_APInt(IfTaken))) {
13884 if (match(ConstantMaxBECount, m_scev_APInt(RangeMax)) &&
13885 IfTaken->ugt(*RangeMax)) {
13886 ConstantMaxBECount = getZero(BECountIfBackedgeTaken->getType());
13887 } else {
13888 ConstantMaxBECount = BECountIfBackedgeTaken;
13889 MaxOrZero = true;
13890 }
13891 }
13892 }
13893
13894 if (isa<SCEVCouldNotCompute>(ConstantMaxBECount) &&
13895 !isa<SCEVCouldNotCompute>(BECount))
13896 ConstantMaxBECount = getConstant(getUnsignedRangeMax(BECount));
13897
13898 const SCEV *SymbolicMaxBECount =
13899 isa<SCEVCouldNotCompute>(BECount) ? ConstantMaxBECount : BECount;
13900 return ExitLimit(BECount, ConstantMaxBECount, SymbolicMaxBECount, MaxOrZero,
13901 Predicates);
13902}
13903
13905 ScalarEvolution &SE) const {
13906 if (Range.isFullSet()) // Infinite loop.
13907 return SE.getCouldNotCompute();
13908
13909 // If the start is a non-zero constant, shift the range to simplify things.
13910 if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(getStart()))
13911 if (!SC->getValue()->isZero()) {
13913 Operands[0] = SE.getZero(SC->getType());
13914 const SCEV *Shifted = SE.getAddRecExpr(Operands, getLoop(),
13916 if (const auto *ShiftedAddRec = dyn_cast<SCEVAddRecExpr>(Shifted))
13917 return ShiftedAddRec->getNumIterationsInRange(
13918 Range.subtract(SC->getAPInt()), SE);
13919 // This is strange and shouldn't happen.
13920 return SE.getCouldNotCompute();
13921 }
13922
13923 // The only time we can solve this is when we have all constant indices.
13924 // Otherwise, we cannot determine the overflow conditions.
13926 return SE.getCouldNotCompute();
13927
13928 // Okay at this point we know that all elements of the chrec are constants and
13929 // that the start element is zero.
13930
13931 // First check to see if the range contains zero. If not, the first
13932 // iteration exits.
13933 unsigned BitWidth = SE.getTypeSizeInBits(getType());
13934 if (!Range.contains(APInt(BitWidth, 0)))
13935 return SE.getZero(getType());
13936
13937 if (isAffine()) {
13938 // If this is an affine expression then we have this situation:
13939 // Solve {0,+,A} in Range === Ax in Range
13940
13941 // We know that zero is in the range. If A is positive then we know that
13942 // the upper value of the range must be the first possible exit value.
13943 // If A is negative then the lower of the range is the last possible loop
13944 // value. Also note that we already checked for a full range.
13945 APInt A = cast<SCEVConstant>(getOperand(1))->getAPInt();
13946 APInt End = A.sge(1) ? (Range.getUpper() - 1) : Range.getLower();
13947
13948 // The exit value should be (End+A)/A.
13949 APInt ExitVal = (End + A).udiv(A);
13950 ConstantInt *ExitValue = ConstantInt::get(SE.getContext(), ExitVal);
13951
13952 // Evaluate at the exit value. If we really did fall out of the valid
13953 // range, then we computed our trip count, otherwise wrap around or other
13954 // things must have happened.
13955 ConstantInt *Val = EvaluateConstantChrecAtConstant(this, ExitValue, SE);
13956 if (Range.contains(Val->getValue()))
13957 return SE.getCouldNotCompute(); // Something strange happened
13958
13959 // Ensure that the previous value is in the range.
13960 assert(Range.contains(
13962 ConstantInt::get(SE.getContext(), ExitVal - 1), SE)->getValue()) &&
13963 "Linear scev computation is off in a bad way!");
13964 return SE.getConstant(ExitValue);
13965 }
13966
13967 if (isQuadratic()) {
13968 if (auto S = SolveQuadraticAddRecRange(this, Range, SE))
13969 return SE.getConstant(*S);
13970 }
13971
13972 return SE.getCouldNotCompute();
13973}
13974
13975const SCEVAddRecExpr *
13977 assert(getNumOperands() > 1 && "AddRec with zero step?");
13978 // There is a temptation to just call getAddExpr(this, getStepRecurrence(SE)),
13979 // but in this case we cannot guarantee that the value returned will be an
13980 // AddRec because SCEV does not have a fixed point where it stops
13981 // simplification: it is legal to return ({rec1} + {rec2}). For example, it
13982 // may happen if we reach arithmetic depth limit while simplifying. So we
13983 // construct the returned value explicitly.
13985 // If this is {A,+,B,+,C,...,+,N}, then its step is {B,+,C,+,...,+,N}, and
13986 // (this + Step) is {A+B,+,B+C,+...,+,N}.
13987 for (unsigned i = 0, e = getNumOperands() - 1; i < e; ++i)
13988 Ops.push_back(SE.getAddExpr(getOperand(i), getOperand(i + 1)));
13989 // We know that the last operand is not a constant zero (otherwise it would
13990 // have been popped out earlier). This guarantees us that if the result has
13991 // the same last operand, then it will also not be popped out, meaning that
13992 // the returned value will be an AddRec.
13993 const SCEV *Last = getOperand(getNumOperands() - 1);
13994 assert(!Last->isZero() && "Recurrency with zero step?");
13995 Ops.push_back(Last);
13997}
13998
13999// Return true when S contains at least an undef value.
14001 return SCEVExprContains(
14002 S, [](const SCEV *S) { return match(S, m_scev_UndefOrPoison()); });
14003}
14004
14005// Return true when S contains a value that is a nullptr.
14007 return SCEVExprContains(S, [](const SCEV *S) {
14008 if (const auto *SU = dyn_cast<SCEVUnknown>(S))
14009 return SU->getValue() == nullptr;
14010 return false;
14011 });
14012}
14013
14014/// Return the size of an element read or written by Inst.
14016 if (!isa<LoadInst, StoreInst>(Inst))
14017 return nullptr;
14019 return getSizeOfExpr(ETy, getLoadStoreType(Inst));
14020}
14021
14022//===----------------------------------------------------------------------===//
14023// SCEVCallbackVH Class Implementation
14024//===----------------------------------------------------------------------===//
14025
14027 assert(SE && "SCEVCallbackVH called with a null ScalarEvolution!");
14028 if (PHINode *PN = dyn_cast<PHINode>(getValPtr()))
14029 SE->ConstantEvolutionLoopExitValue.erase(PN);
14030 SE->eraseValueFromMap(getValPtr());
14031 // this now dangles!
14032}
14033
14034void ScalarEvolution::SCEVCallbackVH::allUsesReplacedWith(Value *V) {
14035 assert(SE && "SCEVCallbackVH called with a null ScalarEvolution!");
14036
14037 // Forget all the expressions associated with users of the old value,
14038 // so that future queries will recompute the expressions using the new
14039 // value.
14040 SE->forgetValue(getValPtr());
14041 // this now dangles!
14042}
14043
14044ScalarEvolution::SCEVCallbackVH::SCEVCallbackVH(Value *V, ScalarEvolution *se)
14045 : CallbackVH(V), SE(se) {}
14046
14047//===----------------------------------------------------------------------===//
14048// ScalarEvolution Class Implementation
14049//===----------------------------------------------------------------------===//
14050
14053 LoopInfo &LI)
14054 : F(F), DL(F.getDataLayout()), TLI(TLI), AC(AC), DT(DT), LI(LI),
14055 CouldNotCompute(new SCEVCouldNotCompute()), ValuesAtScopes(64),
14056 LoopDispositions(64), BlockDispositions(64) {
14057 // To use guards for proving predicates, we need to scan every instruction in
14058 // relevant basic blocks, and not just terminators. Doing this is a waste of
14059 // time if the IR does not actually contain any calls to
14060 // @llvm.experimental.guard, so do a quick check and remember this beforehand.
14061 //
14062 // This pessimizes the case where a pass that preserves ScalarEvolution wants
14063 // to _add_ guards to the module when there weren't any before, and wants
14064 // ScalarEvolution to optimize based on those guards. For now we prefer to be
14065 // efficient in lieu of being smart in that rather obscure case.
14066
14067 auto *GuardDecl = Intrinsic::getDeclarationIfExists(
14068 F.getParent(), Intrinsic::experimental_guard);
14069 HasGuards = GuardDecl && !GuardDecl->use_empty();
14070}
14071
14073 : F(Arg.F), DL(Arg.DL), HasGuards(Arg.HasGuards), TLI(Arg.TLI), AC(Arg.AC),
14074 DT(Arg.DT), LI(Arg.LI), CouldNotCompute(std::move(Arg.CouldNotCompute)),
14075 ValueExprMap(std::move(Arg.ValueExprMap)),
14076 PendingLoopPredicates(std::move(Arg.PendingLoopPredicates)),
14077 PendingMerges(std::move(Arg.PendingMerges)),
14078 ConstantMultipleCache(std::move(Arg.ConstantMultipleCache)),
14079 BackedgeTakenCounts(std::move(Arg.BackedgeTakenCounts)),
14080 PredicatedBackedgeTakenCounts(
14081 std::move(Arg.PredicatedBackedgeTakenCounts)),
14082 BECountUsers(std::move(Arg.BECountUsers)),
14083 ConstantEvolutionLoopExitValue(
14084 std::move(Arg.ConstantEvolutionLoopExitValue)),
14085 ValuesAtScopes(std::move(Arg.ValuesAtScopes)),
14086 ValuesAtScopesUsers(std::move(Arg.ValuesAtScopesUsers)),
14087 LoopDispositions(std::move(Arg.LoopDispositions)),
14088 LoopPropertiesCache(std::move(Arg.LoopPropertiesCache)),
14089 BlockDispositions(std::move(Arg.BlockDispositions)),
14090 SCEVUsers(std::move(Arg.SCEVUsers)),
14091 UnsignedRanges(std::move(Arg.UnsignedRanges)),
14092 SignedRanges(std::move(Arg.SignedRanges)),
14093 UniqueSCEVs(std::move(Arg.UniqueSCEVs)),
14094 UniquePreds(std::move(Arg.UniquePreds)),
14095 SCEVAllocator(std::move(Arg.SCEVAllocator)),
14096 ConstantSCEVs(std::move(Arg.ConstantSCEVs)),
14097 LoopUsers(std::move(Arg.LoopUsers)),
14098 PredicatedSCEVRewrites(std::move(Arg.PredicatedSCEVRewrites)),
14099 FirstUnknown(Arg.FirstUnknown) {
14100 Arg.FirstUnknown = nullptr;
14101}
14102
14104 // Iterate through all the SCEVUnknown instances and call their
14105 // destructors, so that they release their references to their values.
14106 for (SCEVUnknown *U = FirstUnknown; U;) {
14107 SCEVUnknown *Tmp = U;
14108 U = U->Next;
14109 Tmp->~SCEVUnknown();
14110 }
14111 FirstUnknown = nullptr;
14112
14113 ExprValueMap.clear();
14114 ValueExprMap.clear();
14115 HasRecMap.clear();
14116 BackedgeTakenCounts.clear();
14117 PredicatedBackedgeTakenCounts.clear();
14118
14119 assert(PendingLoopPredicates.empty() && "isImpliedCond garbage");
14120 assert(PendingMerges.empty() && "isImpliedViaMerge garbage");
14121 assert(!WalkingBEDominatingConds && "isLoopBackedgeGuardedByCond garbage!");
14122 assert(!ProvingSplitPredicate && "ProvingSplitPredicate garbage!");
14123}
14124
14128
14129/// When printing a top-level SCEV for trip counts, it's helpful to include
14130/// a type for constants which are otherwise hard to disambiguate.
14131static void PrintSCEVWithTypeHint(raw_ostream &OS, const SCEV* S) {
14132 if (isa<SCEVConstant>(S))
14133 OS << *S->getType() << " ";
14134 OS << *S;
14135}
14136
14138 const Loop *L) {
14139 // Print all inner loops first
14140 for (Loop *I : *L)
14141 PrintLoopInfo(OS, SE, I);
14142
14143 OS << "Loop ";
14144 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14145 OS << ": ";
14146
14147 SmallVector<BasicBlock *, 8> ExitingBlocks;
14148 L->getExitingBlocks(ExitingBlocks);
14149 if (ExitingBlocks.size() != 1)
14150 OS << "<multiple exits> ";
14151
14152 auto *BTC = SE->getBackedgeTakenCount(L);
14153 if (!isa<SCEVCouldNotCompute>(BTC)) {
14154 OS << "backedge-taken count is ";
14155 PrintSCEVWithTypeHint(OS, BTC);
14156 } else
14157 OS << "Unpredictable backedge-taken count.";
14158 OS << "\n";
14159
14160 if (ExitingBlocks.size() > 1)
14161 for (BasicBlock *ExitingBlock : ExitingBlocks) {
14162 OS << " exit count for " << ExitingBlock->getName() << ": ";
14163 const SCEV *EC = SE->getExitCount(L, ExitingBlock);
14164 PrintSCEVWithTypeHint(OS, EC);
14165 if (isa<SCEVCouldNotCompute>(EC)) {
14166 // Retry with predicates.
14168 EC = SE->getPredicatedExitCount(L, ExitingBlock, &Predicates);
14169 if (!isa<SCEVCouldNotCompute>(EC)) {
14170 OS << "\n predicated exit count for " << ExitingBlock->getName()
14171 << ": ";
14172 PrintSCEVWithTypeHint(OS, EC);
14173 OS << "\n Predicates:\n";
14174 for (const auto *P : Predicates)
14175 P->print(OS, 4);
14176 }
14177 }
14178 OS << "\n";
14179 }
14180
14181 OS << "Loop ";
14182 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14183 OS << ": ";
14184
14185 auto *ConstantBTC = SE->getConstantMaxBackedgeTakenCount(L);
14186 if (!isa<SCEVCouldNotCompute>(ConstantBTC)) {
14187 OS << "constant max backedge-taken count is ";
14188 PrintSCEVWithTypeHint(OS, ConstantBTC);
14190 OS << ", actual taken count either this or zero.";
14191 } else {
14192 OS << "Unpredictable constant max backedge-taken count. ";
14193 }
14194
14195 OS << "\n"
14196 "Loop ";
14197 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14198 OS << ": ";
14199
14200 auto *SymbolicBTC = SE->getSymbolicMaxBackedgeTakenCount(L);
14201 if (!isa<SCEVCouldNotCompute>(SymbolicBTC)) {
14202 OS << "symbolic max backedge-taken count is ";
14203 PrintSCEVWithTypeHint(OS, SymbolicBTC);
14205 OS << ", actual taken count either this or zero.";
14206 } else {
14207 OS << "Unpredictable symbolic max backedge-taken count. ";
14208 }
14209 OS << "\n";
14210
14211 if (ExitingBlocks.size() > 1)
14212 for (BasicBlock *ExitingBlock : ExitingBlocks) {
14213 OS << " symbolic max exit count for " << ExitingBlock->getName() << ": ";
14214 auto *ExitBTC = SE->getExitCount(L, ExitingBlock,
14216 PrintSCEVWithTypeHint(OS, ExitBTC);
14217 if (isa<SCEVCouldNotCompute>(ExitBTC)) {
14218 // Retry with predicates.
14220 ExitBTC = SE->getPredicatedExitCount(L, ExitingBlock, &Predicates,
14222 if (!isa<SCEVCouldNotCompute>(ExitBTC)) {
14223 OS << "\n predicated symbolic max exit count for "
14224 << ExitingBlock->getName() << ": ";
14225 PrintSCEVWithTypeHint(OS, ExitBTC);
14226 OS << "\n Predicates:\n";
14227 for (const auto *P : Predicates)
14228 P->print(OS, 4);
14229 }
14230 }
14231 OS << "\n";
14232 }
14233
14235 auto *PBT = SE->getPredicatedBackedgeTakenCount(L, Preds);
14236 if (PBT != BTC) {
14237 OS << "Loop ";
14238 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14239 OS << ": ";
14240 if (!isa<SCEVCouldNotCompute>(PBT)) {
14241 OS << "Predicated backedge-taken count is ";
14242 PrintSCEVWithTypeHint(OS, PBT);
14243 } else
14244 OS << "Unpredictable predicated backedge-taken count.";
14245 OS << "\n";
14246 OS << " Predicates:\n";
14247 for (const auto *P : Preds)
14248 P->print(OS, 4);
14249 }
14250 Preds.clear();
14251
14252 auto *PredConstantMax =
14254 if (PredConstantMax != ConstantBTC) {
14255 OS << "Loop ";
14256 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14257 OS << ": ";
14258 if (!isa<SCEVCouldNotCompute>(PredConstantMax)) {
14259 OS << "Predicated constant max backedge-taken count is ";
14260 PrintSCEVWithTypeHint(OS, PredConstantMax);
14261 } else
14262 OS << "Unpredictable predicated constant max backedge-taken count.";
14263 OS << "\n";
14264 OS << " Predicates:\n";
14265 for (const auto *P : Preds)
14266 P->print(OS, 4);
14267 }
14268 Preds.clear();
14269
14270 auto *PredSymbolicMax =
14272 if (SymbolicBTC != PredSymbolicMax) {
14273 OS << "Loop ";
14274 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14275 OS << ": ";
14276 if (!isa<SCEVCouldNotCompute>(PredSymbolicMax)) {
14277 OS << "Predicated symbolic max backedge-taken count is ";
14278 PrintSCEVWithTypeHint(OS, PredSymbolicMax);
14279 } else
14280 OS << "Unpredictable predicated symbolic max backedge-taken count.";
14281 OS << "\n";
14282 OS << " Predicates:\n";
14283 for (const auto *P : Preds)
14284 P->print(OS, 4);
14285 }
14286
14288 OS << "Loop ";
14289 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14290 OS << ": ";
14291 OS << "Trip multiple is " << SE->getSmallConstantTripMultiple(L) << "\n";
14292 }
14293}
14294
14295namespace llvm {
14296// Note: these overloaded operators need to be in the llvm namespace for them
14297// to be resolved correctly. If we put them outside the llvm namespace, the
14298//
14299// OS << ": " << SE.getLoopDisposition(SV, InnerL);
14300//
14301// code below "breaks" and start printing raw enum values as opposed to the
14302// string values.
14305 switch (LD) {
14307 OS << "Variant";
14308 break;
14310 OS << "Invariant";
14311 break;
14313 OS << "Uniform";
14314 break;
14316 OS << "Computable";
14317 break;
14318 }
14319 return OS;
14320}
14321
14324 switch (BD) {
14326 OS << "DoesNotDominate";
14327 break;
14329 OS << "Dominates";
14330 break;
14332 OS << "ProperlyDominates";
14333 break;
14334 }
14335 return OS;
14336}
14337} // namespace llvm
14338
14340 // ScalarEvolution's implementation of the print method is to print
14341 // out SCEV values of all instructions that are interesting. Doing
14342 // this potentially causes it to create new SCEV objects though,
14343 // which technically conflicts with the const qualifier. This isn't
14344 // observable from outside the class though, so casting away the
14345 // const isn't dangerous.
14346 ScalarEvolution &SE = *const_cast<ScalarEvolution *>(this);
14347
14348 if (ClassifyExpressions) {
14349 OS << "Classifying expressions for: ";
14350 F.printAsOperand(OS, /*PrintType=*/false);
14351 OS << "\n";
14352 for (Instruction &I : instructions(F))
14353 if (isSCEVable(I.getType()) && !isa<CmpInst>(I)) {
14354 OS << I << '\n';
14355 OS << " --> ";
14356 const SCEV *SV = SE.getSCEV(&I);
14357 SV->print(OS);
14358 if (!isa<SCEVCouldNotCompute>(SV)) {
14359 OS << " U: ";
14360 SE.getUnsignedRange(SV).print(OS);
14361 OS << " S: ";
14362 SE.getSignedRange(SV).print(OS);
14363 }
14364
14365 const Loop *L = LI.getLoopFor(I.getParent());
14366
14367 SCEVUse AtUse = SE.getSCEVAtScope(SV, L);
14368 if (AtUse != SV) {
14369 OS << " --> ";
14370 OS << AtUse;
14371 if (!isa<SCEVCouldNotCompute>(AtUse)) {
14372 OS << " U: ";
14373 SE.getUnsignedRange(AtUse).print(OS);
14374 OS << " S: ";
14375 SE.getSignedRange(AtUse).print(OS);
14376 }
14377 }
14378
14379 if (L) {
14380 OS << "\t\t" "Exits: ";
14381 SCEVUse ExitValue = SE.getSCEVAtScope(SV, L->getParentLoop());
14382 if (!SE.isLoopInvariant(ExitValue, L)) {
14383 OS << "<<Unknown>>";
14384 } else {
14385 OS << ExitValue;
14386 }
14387
14388 ListSeparator LS(", ", "\t\tLoopDispositions: { ");
14389 for (const auto *Iter = L; Iter; Iter = Iter->getParentLoop()) {
14390 OS << LS;
14391 Iter->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14392 OS << ": " << SE.getLoopDisposition(SV, Iter);
14393 }
14394
14395 for (const auto *InnerL : depth_first(L)) {
14396 if (InnerL == L)
14397 continue;
14398 OS << LS;
14399 InnerL->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14400 OS << ": " << SE.getLoopDisposition(SV, InnerL);
14401 }
14402
14403 OS << " }";
14404 }
14405
14406 OS << "\n";
14407 }
14408 }
14409
14410 OS << "Determining loop execution counts for: ";
14411 F.printAsOperand(OS, /*PrintType=*/false);
14412 OS << "\n";
14413 for (Loop *I : LI)
14414 PrintLoopInfo(OS, &SE, I);
14415}
14416
14419 auto &Values = LoopDispositions[S];
14420 for (auto &V : Values) {
14421 if (V.getPointer() == L)
14422 return V.getInt();
14423 }
14424 Values.emplace_back(L, LoopVariant);
14425 LoopDisposition D = computeLoopDisposition(S, L);
14426 auto &Values2 = LoopDispositions[S];
14427 for (auto &V : llvm::reverse(Values2)) {
14428 if (V.getPointer() == L) {
14429 V.setInt(D);
14430 break;
14431 }
14432 }
14433 return D;
14434}
14435
14437ScalarEvolution::computeLoopDisposition(const SCEV *S, const Loop *L) {
14438 switch (S->getSCEVType()) {
14439 case scConstant:
14440 case scVScale:
14441 return LoopInvariant;
14442 case scAddRecExpr: {
14443 const SCEVAddRecExpr *AR = cast<SCEVAddRecExpr>(S);
14444
14445 // If L is the addrec's loop, it's computable.
14446 if (AR->getLoop() == L)
14447 return LoopComputable;
14448
14449 // Add recurrences are never invariant in the function-body (null loop).
14450 if (!L)
14451 return LoopVariant;
14452
14453 // Everything that is not defined at loop entry is variant.
14454 if (DT.dominates(L->getHeader(), AR->getLoop()->getHeader())) {
14455 if (L->contains(AR->getLoop()) &&
14456 llvm::all_of(AR->operands(),
14457 [&](const SCEV *Op) { return isLoopUniform(Op, L); }))
14458 return LoopUniform;
14459
14460 return LoopVariant;
14461 }
14462 assert(!L->contains(AR->getLoop()) && "Containing loop's header does not"
14463 " dominate the contained loop's header?");
14464
14465 // This recurrence is invariant w.r.t. L if AR's loop contains L.
14466 if (AR->getLoop()->contains(L))
14467 return LoopInvariant;
14468
14469 // This recurrence is variant w.r.t. L if any of its operands
14470 // are variant.
14471 for (SCEVUse Op : AR->operands())
14472 if (!isLoopInvariant(Op, L))
14473 return LoopVariant;
14474
14475 // Otherwise it's loop-invariant.
14476 return LoopInvariant;
14477 }
14478 case scTruncate:
14479 case scZeroExtend:
14480 case scSignExtend:
14481 case scPtrToAddr:
14482 case scAddExpr:
14483 case scMulExpr:
14484 case scUDivExpr:
14485 case scUMaxExpr:
14486 case scSMaxExpr:
14487 case scUMinExpr:
14488 case scSMinExpr:
14489 case scSequentialUMinExpr: {
14490 bool HasVarying = false;
14491 bool HasUniform = false;
14492 for (SCEVUse Op : S->operands()) {
14494 if (D == LoopVariant)
14495 return LoopVariant;
14496 if (D == LoopComputable)
14497 HasVarying = true;
14498 if (D == LoopUniform)
14499 HasUniform = true;
14500 }
14501 return HasVarying ? (HasUniform ? LoopVariant : LoopComputable)
14502 : (HasUniform ? LoopUniform : LoopInvariant);
14503 }
14504 case scUnknown:
14505 // All non-instruction values are loop invariant. All instructions are loop
14506 // invariant if they are not contained in the specified loop.
14507 // Instructions are never considered invariant in the function body
14508 // (null loop) because they are defined within the "loop".
14510 return (L && !L->contains(I)) ? LoopInvariant : LoopVariant;
14511 return LoopInvariant;
14512 case scCouldNotCompute:
14513 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
14514 }
14515 llvm_unreachable("Unknown SCEV kind!");
14516}
14517
14518bool ScalarEvolution::isLoopUniform(const SCEV *S, const Loop *L) {
14520 return D == LoopUniform || D == LoopInvariant;
14521}
14522
14524 return getLoopDisposition(S, L) == LoopInvariant;
14525}
14526
14528 return getLoopDisposition(S, L) == LoopComputable;
14529}
14530
14533 auto &Values = BlockDispositions[S];
14534 for (auto &V : Values) {
14535 if (V.getPointer() == BB)
14536 return V.getInt();
14537 }
14538 Values.emplace_back(BB, DoesNotDominateBlock);
14539 BlockDisposition D = computeBlockDisposition(S, BB);
14540 auto &Values2 = BlockDispositions[S];
14541 for (auto &V : llvm::reverse(Values2)) {
14542 if (V.getPointer() == BB) {
14543 V.setInt(D);
14544 break;
14545 }
14546 }
14547 return D;
14548}
14549
14551ScalarEvolution::computeBlockDisposition(const SCEV *S, const BasicBlock *BB) {
14552 switch (S->getSCEVType()) {
14553 case scConstant:
14554 case scVScale:
14556 case scAddRecExpr: {
14557 // This uses a "dominates" query instead of "properly dominates" query
14558 // to test for proper dominance too, because the instruction which
14559 // produces the addrec's value is a PHI, and a PHI effectively properly
14560 // dominates its entire containing block.
14561 const SCEVAddRecExpr *AR = cast<SCEVAddRecExpr>(S);
14562 if (!DT.dominates(AR->getLoop()->getHeader(), BB))
14563 return DoesNotDominateBlock;
14564
14565 // Fall through into SCEVNAryExpr handling.
14566 [[fallthrough]];
14567 }
14568 case scTruncate:
14569 case scZeroExtend:
14570 case scSignExtend:
14571 case scPtrToAddr:
14572 case scAddExpr:
14573 case scMulExpr:
14574 case scUDivExpr:
14575 case scUMaxExpr:
14576 case scSMaxExpr:
14577 case scUMinExpr:
14578 case scSMinExpr:
14579 case scSequentialUMinExpr: {
14580 bool Proper = true;
14581 for (const SCEV *NAryOp : S->operands()) {
14583 if (D == DoesNotDominateBlock)
14584 return DoesNotDominateBlock;
14585 if (D == DominatesBlock)
14586 Proper = false;
14587 }
14588 return Proper ? ProperlyDominatesBlock : DominatesBlock;
14589 }
14590 case scUnknown:
14591 if (Instruction *I =
14593 if (I->getParent() == BB)
14594 return DominatesBlock;
14595 if (DT.properlyDominates(I->getParent(), BB))
14597 return DoesNotDominateBlock;
14598 }
14600 case scCouldNotCompute:
14601 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
14602 }
14603 llvm_unreachable("Unknown SCEV kind!");
14604}
14605
14606bool ScalarEvolution::dominates(const SCEV *S, const BasicBlock *BB) {
14607 return getBlockDisposition(S, BB) >= DominatesBlock;
14608}
14609
14612}
14613
14614void ScalarEvolution::forgetBackedgeTakenCounts(const Loop *L,
14615 bool Predicated) {
14616 auto &BECounts =
14617 Predicated ? PredicatedBackedgeTakenCounts : BackedgeTakenCounts;
14618 auto It = BECounts.find(L);
14619 if (It != BECounts.end()) {
14620 for (const ExitNotTakenInfo &ENT : It->second.ExitNotTaken) {
14621 for (const SCEV *S : {ENT.ExactNotTaken, ENT.SymbolicMaxNotTaken}) {
14622 if (!isa<SCEVConstant>(S)) {
14623 auto UserIt = BECountUsers.find(S);
14624 assert(UserIt != BECountUsers.end());
14625 UserIt->second.erase({L, Predicated});
14626 }
14627 }
14628 }
14629 BECounts.erase(It);
14630 }
14631}
14632
14633void ScalarEvolution::forgetMemoizedResults(ArrayRef<SCEVUse> SCEVs) {
14634 SmallPtrSet<const SCEV *, 8> ToForget(llvm::from_range, SCEVs);
14635 SmallVector<SCEVUse, 8> Worklist(ToForget.begin(), ToForget.end());
14636
14637 while (!Worklist.empty()) {
14638 const SCEV *Curr = Worklist.pop_back_val();
14639 auto Users = SCEVUsers.find(Curr);
14640 if (Users != SCEVUsers.end())
14641 for (const auto *User : Users->second)
14642 if (ToForget.insert(User).second)
14643 Worklist.push_back(User);
14644 }
14645
14646 for (const auto *S : ToForget)
14647 forgetMemoizedResultsImpl(S);
14648
14649 PredicatedSCEVRewrites.remove_if(
14650 [&](const auto &Entry) { return ToForget.count(Entry.first.first); });
14651}
14652
14653void ScalarEvolution::forgetMemoizedResultsImpl(const SCEV *S) {
14654 LoopDispositions.erase(S);
14655 BlockDispositions.erase(S);
14656 UnsignedRanges.erase(S);
14657 SignedRanges.erase(S);
14658 HasRecMap.erase(S);
14659 ConstantMultipleCache.erase(S);
14660
14661 if (auto *AR = dyn_cast<SCEVAddRecExpr>(S)) {
14662 UnsignedWrapViaInductionTried.erase(AR);
14663 SignedWrapViaInductionTried.erase(AR);
14664 }
14665
14666 auto ExprIt = ExprValueMap.find(S);
14667 if (ExprIt != ExprValueMap.end()) {
14668 for (Value *V : ExprIt->second) {
14669 auto ValueIt = ValueExprMap.find_as(V);
14670 if (ValueIt != ValueExprMap.end())
14671 ValueExprMap.erase(ValueIt);
14672 }
14673 ExprValueMap.erase(ExprIt);
14674 }
14675
14676 auto ScopeIt = ValuesAtScopes.find(S);
14677 if (ScopeIt != ValuesAtScopes.end()) {
14678 for (const auto &Pair : ScopeIt->second)
14679 if (!isa_and_nonnull<SCEVConstant>(Pair.second))
14680 llvm::erase(ValuesAtScopesUsers[Pair.second.getPointer()],
14681 std::make_pair(Pair.first, S));
14682 ValuesAtScopes.erase(ScopeIt);
14683 }
14684
14685 auto ScopeUserIt = ValuesAtScopesUsers.find(S);
14686 if (ScopeUserIt != ValuesAtScopesUsers.end()) {
14687 for (const auto &Pair : ScopeUserIt->second)
14688 // The recorded value at scope is a use of S, which may carry no-wrap
14689 // flags that are not part of this key.
14690 llvm::erase_if(ValuesAtScopes[Pair.second], [&](const auto &LS) {
14691 return LS.first == Pair.first && LS.second.getPointer() == S;
14692 });
14693 ValuesAtScopesUsers.erase(ScopeUserIt);
14694 }
14695
14696 auto BEUsersIt = BECountUsers.find(S);
14697 if (BEUsersIt != BECountUsers.end()) {
14698 // Work on a copy, as forgetBackedgeTakenCounts() will modify the original.
14699 auto Copy = BEUsersIt->second;
14700 for (const auto &Pair : Copy)
14701 forgetBackedgeTakenCounts(Pair.getPointer(), Pair.getInt());
14702 BECountUsers.erase(BEUsersIt);
14703 }
14704
14705 auto FoldUser = FoldCacheUser.find(S);
14706 if (FoldUser != FoldCacheUser.end())
14707 for (auto &KV : FoldUser->second)
14708 FoldCache.erase(KV);
14709 FoldCacheUser.erase(S);
14710}
14711
14712void
14713ScalarEvolution::getUsedLoops(const SCEV *S,
14714 SmallPtrSetImpl<const Loop *> &LoopsUsed) {
14715 struct FindUsedLoops {
14716 FindUsedLoops(SmallPtrSetImpl<const Loop *> &LoopsUsed)
14717 : LoopsUsed(LoopsUsed) {}
14718 SmallPtrSetImpl<const Loop *> &LoopsUsed;
14719 bool follow(const SCEV *S) {
14720 if (auto *AR = dyn_cast<SCEVAddRecExpr>(S))
14721 LoopsUsed.insert(AR->getLoop());
14722 return true;
14723 }
14724
14725 bool isDone() const { return false; }
14726 };
14727
14728 FindUsedLoops F(LoopsUsed);
14729 SCEVTraversal<FindUsedLoops>(F).visitAll(S);
14730}
14731
14732void ScalarEvolution::getReachableBlocks(
14735 Worklist.push_back(&F.getEntryBlock());
14736 while (!Worklist.empty()) {
14737 BasicBlock *BB = Worklist.pop_back_val();
14738 if (!Reachable.insert(BB).second)
14739 continue;
14740
14741 Value *Cond;
14742 BasicBlock *TrueBB, *FalseBB;
14743 if (match(BB->getTerminator(), m_Br(m_Value(Cond), m_BasicBlock(TrueBB),
14744 m_BasicBlock(FalseBB)))) {
14745 if (auto *C = dyn_cast<ConstantInt>(Cond)) {
14746 Worklist.push_back(C->isOne() ? TrueBB : FalseBB);
14747 continue;
14748 }
14749
14750 if (auto *Cmp = dyn_cast<ICmpInst>(Cond)) {
14751 const SCEV *L = getSCEV(Cmp->getOperand(0));
14752 const SCEV *R = getSCEV(Cmp->getOperand(1));
14753 if (isKnownPredicateViaConstantRanges(Cmp->getCmpPredicate(), L, R)) {
14754 Worklist.push_back(TrueBB);
14755 continue;
14756 }
14757 if (isKnownPredicateViaConstantRanges(Cmp->getInverseCmpPredicate(), L,
14758 R)) {
14759 Worklist.push_back(FalseBB);
14760 continue;
14761 }
14762 }
14763 }
14764
14765 append_range(Worklist, successors(BB));
14766 }
14767}
14768
14770 ScalarEvolution &SE = *const_cast<ScalarEvolution *>(this);
14771 ScalarEvolution SE2(F, TLI, AC, DT, LI);
14772
14773 SmallVector<Loop *, 8> LoopStack(LI.begin(), LI.end());
14774
14775 // Map's SCEV expressions from one ScalarEvolution "universe" to another.
14776 struct SCEVMapper : public SCEVRewriteVisitor<SCEVMapper> {
14777 SCEVMapper(ScalarEvolution &SE) : SCEVRewriteVisitor<SCEVMapper>(SE) {}
14778
14779 const SCEV *visitConstant(const SCEVConstant *Constant) {
14780 return SE.getConstant(Constant->getAPInt());
14781 }
14782
14783 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
14784 return SE.getUnknown(Expr->getValue());
14785 }
14786
14787 const SCEV *visitCouldNotCompute(const SCEVCouldNotCompute *Expr) {
14788 return SE.getCouldNotCompute();
14789 }
14790 };
14791
14792 SCEVMapper SCM(SE2);
14793 SmallPtrSet<BasicBlock *, 16> ReachableBlocks;
14794 SE2.getReachableBlocks(ReachableBlocks, F);
14795
14796 auto GetDelta = [&](const SCEV *Old, const SCEV *New) -> const SCEV * {
14797 if (containsUndefs(Old) || containsUndefs(New)) {
14798 // SCEV treats "undef" as an unknown but consistent value (i.e. it does
14799 // not propagate undef aggressively). This means we can (and do) fail
14800 // verification in cases where a transform makes a value go from "undef"
14801 // to "undef+1" (say). The transform is fine, since in both cases the
14802 // result is "undef", but SCEV thinks the value increased by 1.
14803 return nullptr;
14804 }
14805
14806 // Unless VerifySCEVStrict is set, we only compare constant deltas.
14807 const SCEV *Delta = SE2.getMinusSCEV(Old, New);
14808 if (!VerifySCEVStrict && !isa<SCEVConstant>(Delta))
14809 return nullptr;
14810
14811 return Delta;
14812 };
14813
14814 while (!LoopStack.empty()) {
14815 auto *L = LoopStack.pop_back_val();
14816 llvm::append_range(LoopStack, *L);
14817
14818 // Only verify BECounts in reachable loops. For an unreachable loop,
14819 // any BECount is legal.
14820 if (!ReachableBlocks.contains(L->getHeader()))
14821 continue;
14822
14823 // Only verify cached BECounts. Computing new BECounts may change the
14824 // results of subsequent SCEV uses.
14825 auto It = BackedgeTakenCounts.find(L);
14826 if (It == BackedgeTakenCounts.end())
14827 continue;
14828
14829 auto *CurBECount =
14830 SCM.visit(It->second.getExact(L, const_cast<ScalarEvolution *>(this)));
14831 auto *NewBECount = SE2.getBackedgeTakenCount(L);
14832
14833 if (CurBECount == SE2.getCouldNotCompute() ||
14834 NewBECount == SE2.getCouldNotCompute()) {
14835 // NB! This situation is legal, but is very suspicious -- whatever pass
14836 // change the loop to make a trip count go from could not compute to
14837 // computable or vice-versa *should have* invalidated SCEV. However, we
14838 // choose not to assert here (for now) since we don't want false
14839 // positives.
14840 continue;
14841 }
14842
14843 if (SE.getTypeSizeInBits(CurBECount->getType()) >
14844 SE.getTypeSizeInBits(NewBECount->getType()))
14845 NewBECount = SE2.getZeroExtendExpr(NewBECount, CurBECount->getType());
14846 else if (SE.getTypeSizeInBits(CurBECount->getType()) <
14847 SE.getTypeSizeInBits(NewBECount->getType()))
14848 CurBECount = SE2.getZeroExtendExpr(CurBECount, NewBECount->getType());
14849
14850 const SCEV *Delta = GetDelta(CurBECount, NewBECount);
14851 if (Delta && !Delta->isZero()) {
14852 dbgs() << "Trip Count for " << *L << " Changed!\n";
14853 dbgs() << "Old: " << *CurBECount << "\n";
14854 dbgs() << "New: " << *NewBECount << "\n";
14855 dbgs() << "Delta: " << *Delta << "\n";
14856 std::abort();
14857 }
14858 }
14859
14860 // Collect all valid loops currently in LoopInfo.
14861 SmallPtrSet<Loop *, 32> ValidLoops;
14862 SmallVector<Loop *, 32> Worklist(LI.begin(), LI.end());
14863 while (!Worklist.empty()) {
14864 Loop *L = Worklist.pop_back_val();
14865 if (ValidLoops.insert(L).second)
14866 Worklist.append(L->begin(), L->end());
14867 }
14868 for (const auto &KV : ValueExprMap) {
14869#ifndef NDEBUG
14870 // Check for SCEV expressions referencing invalid/deleted loops.
14871 if (auto *AR = dyn_cast<SCEVAddRecExpr>(KV.second)) {
14872 assert(ValidLoops.contains(AR->getLoop()) &&
14873 "AddRec references invalid loop");
14874 }
14875#endif
14876
14877 // Check that the value is also part of the reverse map.
14878 auto It = ExprValueMap.find(KV.second);
14879 if (It == ExprValueMap.end() || !It->second.contains(KV.first)) {
14880 dbgs() << "Value " << *KV.first
14881 << " is in ValueExprMap but not in ExprValueMap\n";
14882 std::abort();
14883 }
14884
14885 if (auto *I = dyn_cast<Instruction>(&*KV.first)) {
14886 if (!ReachableBlocks.contains(I->getParent()))
14887 continue;
14888 const SCEV *OldSCEV = SCM.visit(KV.second);
14889 const SCEV *NewSCEV = SE2.getSCEV(I);
14890 const SCEV *Delta = GetDelta(OldSCEV, NewSCEV);
14891 if (Delta && !Delta->isZero()) {
14892 dbgs() << "SCEV for value " << *I << " changed!\n"
14893 << "Old: " << *OldSCEV << "\n"
14894 << "New: " << *NewSCEV << "\n"
14895 << "Delta: " << *Delta << "\n";
14896 std::abort();
14897 }
14898 }
14899 }
14900
14901 for (const auto &KV : ExprValueMap) {
14902 for (Value *V : KV.second) {
14903 const SCEV *S = ValueExprMap.lookup(V);
14904 if (!S) {
14905 dbgs() << "Value " << *V
14906 << " is in ExprValueMap but not in ValueExprMap\n";
14907 std::abort();
14908 }
14909 if (S != KV.first) {
14910 dbgs() << "Value " << *V << " mapped to " << *S << " rather than "
14911 << *KV.first << "\n";
14912 std::abort();
14913 }
14914 }
14915 }
14916
14917 // Verify integrity of SCEV users.
14918 for (const auto &S : UniqueSCEVs) {
14919 for (SCEVUse Op : S.operands()) {
14920 // We do not store dependencies of constants.
14921 if (isa<SCEVConstant>(Op))
14922 continue;
14923 auto It = SCEVUsers.find(Op);
14924 if (It != SCEVUsers.end() && It->second.count(&S))
14925 continue;
14926 dbgs() << "Use of operand " << *Op << " by user " << S
14927 << " is not being tracked!\n";
14928 std::abort();
14929 }
14930 }
14931
14932 // Verify integrity of ValuesAtScopes users.
14933 for (const auto &ValueAndVec : ValuesAtScopes) {
14934 const SCEV *Value = ValueAndVec.first;
14935 for (const auto &LoopAndValueAtScope : ValueAndVec.second) {
14936 const Loop *L = LoopAndValueAtScope.first;
14937 SCEVUse ValueAtScope = LoopAndValueAtScope.second;
14938 if (!isa<SCEVConstant>(ValueAtScope)) {
14939 auto It = ValuesAtScopesUsers.find(ValueAtScope.getPointer());
14940 if (It != ValuesAtScopesUsers.end() &&
14941 is_contained(It->second, std::make_pair(L, Value)))
14942 continue;
14943 dbgs() << "Value: " << *Value << ", Loop: " << *L << ", ValueAtScope: "
14944 << *ValueAtScope << " missing in ValuesAtScopesUsers\n";
14945 std::abort();
14946 }
14947 }
14948 }
14949
14950 for (const auto &ValueAtScopeAndVec : ValuesAtScopesUsers) {
14951 const SCEV *ValueAtScope = ValueAtScopeAndVec.first;
14952 for (const auto &LoopAndValue : ValueAtScopeAndVec.second) {
14953 const Loop *L = LoopAndValue.first;
14954 const SCEV *Value = LoopAndValue.second;
14956 auto It = ValuesAtScopes.find(Value);
14957 // The recorded value at scope may carry no-wrap flags that are not part
14958 // of the key it is recorded under.
14959 if (It != ValuesAtScopes.end() && any_of(It->second, [&](const auto &LS) {
14960 return LS.first == L && LS.second.getPointer() == ValueAtScope;
14961 }))
14962 continue;
14963 dbgs() << "Value: " << *Value << ", Loop: " << *L << ", ValueAtScope: "
14964 << *ValueAtScope << " missing in ValuesAtScopes\n";
14965 std::abort();
14966 }
14967 }
14968
14969 // Verify integrity of BECountUsers.
14970 auto VerifyBECountUsers = [&](bool Predicated) {
14971 auto &BECounts =
14972 Predicated ? PredicatedBackedgeTakenCounts : BackedgeTakenCounts;
14973 for (const auto &LoopAndBEInfo : BECounts) {
14974 for (const ExitNotTakenInfo &ENT : LoopAndBEInfo.second.ExitNotTaken) {
14975 for (const SCEV *S : {ENT.ExactNotTaken, ENT.SymbolicMaxNotTaken}) {
14976 if (!isa<SCEVConstant>(S)) {
14977 auto UserIt = BECountUsers.find(S);
14978 if (UserIt != BECountUsers.end() &&
14979 UserIt->second.contains({ LoopAndBEInfo.first, Predicated }))
14980 continue;
14981 dbgs() << "Value " << *S << " for loop " << *LoopAndBEInfo.first
14982 << " missing from BECountUsers\n";
14983 std::abort();
14984 }
14985 }
14986 }
14987 }
14988 };
14989 VerifyBECountUsers(/* Predicated */ false);
14990 VerifyBECountUsers(/* Predicated */ true);
14991
14992 // Verify intergity of loop disposition cache.
14993 for (auto &[S, Values] : LoopDispositions) {
14994 for (auto [Loop, CachedDisposition] : Values) {
14995 const auto RecomputedDisposition = SE2.getLoopDisposition(S, Loop);
14996 if (CachedDisposition != RecomputedDisposition) {
14997 dbgs() << "Cached disposition of " << *S << " for loop " << *Loop
14998 << " is incorrect: cached " << CachedDisposition << ", actual "
14999 << RecomputedDisposition << "\n";
15000 std::abort();
15001 }
15002 }
15003 }
15004
15005 // Verify integrity of the block disposition cache.
15006 for (auto &[S, Values] : BlockDispositions) {
15007 for (auto [BB, CachedDisposition] : Values) {
15008 const auto RecomputedDisposition = SE2.getBlockDisposition(S, BB);
15009 if (CachedDisposition != RecomputedDisposition) {
15010 dbgs() << "Cached disposition of " << *S << " for block %"
15011 << BB->getName() << " is incorrect: cached " << CachedDisposition
15012 << ", actual " << RecomputedDisposition << "\n";
15013 std::abort();
15014 }
15015 }
15016 }
15017
15018 // Verify FoldCache/FoldCacheUser caches.
15019 for (auto [FoldID, Expr] : FoldCache) {
15020 auto I = FoldCacheUser.find(Expr);
15021 if (I == FoldCacheUser.end()) {
15022 dbgs() << "Missing entry in FoldCacheUser for cached expression " << *Expr
15023 << "!\n";
15024 std::abort();
15025 }
15026 if (!is_contained(I->second, FoldID)) {
15027 dbgs() << "Missing FoldID in cached users of " << *Expr << "!\n";
15028 std::abort();
15029 }
15030 }
15031 for (auto [Expr, IDs] : FoldCacheUser) {
15032 for (auto &FoldID : IDs) {
15033 const SCEV *S = FoldCache.lookup(FoldID);
15034 if (!S) {
15035 dbgs() << "Missing entry in FoldCache for expression " << *Expr
15036 << "!\n";
15037 std::abort();
15038 }
15039 if (S != Expr) {
15040 dbgs() << "Entry in FoldCache doesn't match FoldCacheUser: " << *S
15041 << " != " << *Expr << "!\n";
15042 std::abort();
15043 }
15044 }
15045 }
15046
15047 // Verify that ConstantMultipleCache computations are correct. We check that
15048 // cached multiples and recomputed multiples are multiples of each other to
15049 // verify correctness. It is possible that a recomputed multiple is different
15050 // from the cached multiple due to strengthened no wrap flags or changes in
15051 // KnownBits computations.
15052 for (auto [S, Multiple] : ConstantMultipleCache) {
15053 APInt RecomputedMultiple = SE2.getConstantMultiple(S);
15054 if ((Multiple != 0 && RecomputedMultiple != 0 &&
15055 Multiple.urem(RecomputedMultiple) != 0 &&
15056 RecomputedMultiple.urem(Multiple) != 0)) {
15057 dbgs() << "Incorrect cached computation in ConstantMultipleCache for "
15058 << *S << " : Computed " << RecomputedMultiple
15059 << " but cache contains " << Multiple << "!\n";
15060 std::abort();
15061 }
15062 }
15063}
15064
15066 Function &F, const PreservedAnalyses &PA,
15067 FunctionAnalysisManager::Invalidator &Inv) {
15068 // Invalidate the ScalarEvolution object whenever it isn't preserved or one
15069 // of its dependencies is invalidated.
15070 auto PAC = PA.getChecker<ScalarEvolutionAnalysis>();
15071 return !(PAC.preserved() || PAC.preservedSet<AllAnalysesOn<Function>>()) ||
15072 Inv.invalidate<AssumptionAnalysis>(F, PA) ||
15073 Inv.invalidate<DominatorTreeAnalysis>(F, PA) ||
15074 Inv.invalidate<LoopAnalysis>(F, PA);
15075}
15076
15077AnalysisKey ScalarEvolutionAnalysis::Key;
15078
15081 auto &TLI = AM.getResult<TargetLibraryAnalysis>(F);
15082 auto &AC = AM.getResult<AssumptionAnalysis>(F);
15083 auto &DT = AM.getResult<DominatorTreeAnalysis>(F);
15084 auto &LI = AM.getResult<LoopAnalysis>(F);
15085 return ScalarEvolution(F, TLI, AC, DT, LI);
15086}
15087
15093
15096 // For compatibility with opt's -analyze feature under legacy pass manager
15097 // which was not ported to NPM. This keeps tests using
15098 // update_analyze_test_checks.py working.
15099 OS << "Printing analysis 'Scalar Evolution Analysis' for function '"
15100 << F.getName() << "':\n";
15102 return PreservedAnalyses::all();
15103}
15104
15106 "Scalar Evolution Analysis", false, true)
15112 "Scalar Evolution Analysis", false, true)
15113
15114char ScalarEvolutionWrapperPass::ID = 0;
15115
15117
15119 SE.reset(new ScalarEvolution(
15121 getAnalysis<AssumptionCacheTracker>().getAssumptionCache(F),
15123 getAnalysis<LoopInfoWrapperPass>().getLoopInfo()));
15124 return false;
15125}
15126
15128
15130 SE->print(OS);
15131}
15132
15134 if (!VerifySCEV)
15135 return;
15136
15137 SE->verify();
15138}
15139
15147
15149 const SCEV *RHS) {
15150 return getComparePredicate(ICmpInst::ICMP_EQ, LHS, RHS);
15151}
15152
15153const SCEVPredicate *
15155 const SCEV *LHS, const SCEV *RHS) {
15157 assert(LHS->getType() == RHS->getType() &&
15158 "Type mismatch between LHS and RHS");
15159 // Unique this node based on the arguments
15160 ID.AddInteger(SCEVPredicate::P_Compare);
15161 ID.AddInteger(Pred);
15162 ID.AddPointer(LHS);
15163 ID.AddPointer(RHS);
15165 if (const auto *S = UniquePreds.lookup(ID, Token))
15166 return S;
15167 SCEVComparePredicate *Eq = new (SCEVAllocator)
15168 SCEVComparePredicate(ID.Intern(SCEVAllocator), Pred, LHS, RHS);
15169 UniquePreds.insert(Eq, Token);
15170 return Eq;
15171}
15172
15174 const SCEVAddRecExpr *AR,
15177 // Unique this node based on the arguments
15179 ID.AddPointer(AR);
15180 ID.AddInteger(AddedFlags);
15182 if (const auto *S = UniquePreds.lookup(ID, Token))
15183 return S;
15184 auto *OF = new (SCEVAllocator)
15185 SCEVWrapPredicate(ID.Intern(SCEVAllocator), AR, AddedFlags);
15186 UniquePreds.insert(OF, Token);
15187 return OF;
15188}
15189
15190namespace {
15191
15192class SCEVPredicateRewriter : public SCEVRewriteVisitor<SCEVPredicateRewriter> {
15193public:
15194
15195 /// Rewrites \p S in the context of a loop L and the SCEV predication
15196 /// infrastructure.
15197 ///
15198 /// If \p Pred is non-null, the SCEV expression is rewritten to respect the
15199 /// equivalences present in \p Pred.
15200 ///
15201 /// If \p NewPreds is non-null, rewrite is free to add further predicates to
15202 /// \p NewPreds such that the result will be an AddRecExpr.
15203 static const SCEV *rewrite(const SCEV *S, const Loop *L, ScalarEvolution &SE,
15205 const SCEVPredicate *Pred) {
15206 SCEVPredicateRewriter Rewriter(L, SE, NewPreds, Pred);
15207 return Rewriter.visit(S);
15208 }
15209
15210 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
15211 if (Pred) {
15212 if (auto *U = dyn_cast<SCEVUnionPredicate>(Pred)) {
15213 for (const auto *Pred : U->getPredicates())
15214 if (const auto *IPred = dyn_cast<SCEVComparePredicate>(Pred))
15215 if (IPred->getLHS() == Expr &&
15216 IPred->getPredicate() == ICmpInst::ICMP_EQ)
15217 return IPred->getRHS();
15218 } else if (const auto *IPred = dyn_cast<SCEVComparePredicate>(Pred)) {
15219 if (IPred->getLHS() == Expr &&
15220 IPred->getPredicate() == ICmpInst::ICMP_EQ)
15221 return IPred->getRHS();
15222 }
15223 }
15224 return convertToAddRecWithPreds(Expr);
15225 }
15226
15227 const SCEV *visitZeroExtendExpr(const SCEVZeroExtendExpr *Expr) {
15228 const SCEV *Operand = visit(Expr->getOperand());
15229 const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(Operand);
15230 if (AR && AR->getLoop() == L && AR->isAffine()) {
15231 // This couldn't be folded because the operand didn't have the nuw
15232 // flag. Add the nusw flag as an assumption that we could make.
15233 const SCEV *Step = AR->getStepRecurrence(SE);
15234 Type *Ty = Expr->getType();
15235 if (addOverflowAssumption(AR, SCEVWrapPredicate::IncrementNUSW))
15236 return SE.getAddRecExpr(SE.getZeroExtendExpr(AR->getStart(), Ty),
15237 SE.getSignExtendExpr(Step, Ty), L,
15238 AR->getNoWrapFlags());
15239 }
15240 return SE.getZeroExtendExpr(Operand, Expr->getType());
15241 }
15242
15243 const SCEV *visitSignExtendExpr(const SCEVSignExtendExpr *Expr) {
15244 const SCEV *Operand = visit(Expr->getOperand());
15245 const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(Operand);
15246 if (AR && AR->getLoop() == L && AR->isAffine()) {
15247 // This couldn't be folded because the operand didn't have the nsw
15248 // flag. Add the nssw flag as an assumption that we could make.
15249 const SCEV *Step = AR->getStepRecurrence(SE);
15250 Type *Ty = Expr->getType();
15251 if (addOverflowAssumption(AR, SCEVWrapPredicate::IncrementNSSW))
15252 return SE.getAddRecExpr(SE.getSignExtendExpr(AR->getStart(), Ty),
15253 SE.getSignExtendExpr(Step, Ty), L,
15254 AR->getNoWrapFlags());
15255 }
15256 return SE.getSignExtendExpr(Operand, Expr->getType());
15257 }
15258
15259private:
15260 explicit SCEVPredicateRewriter(
15261 const Loop *L, ScalarEvolution &SE,
15262 SmallVectorImpl<const SCEVPredicate *> *NewPreds,
15263 const SCEVPredicate *Pred)
15264 : SCEVRewriteVisitor(SE), NewPreds(NewPreds), Pred(Pred), L(L) {}
15265
15266 bool addOverflowAssumption(const SCEVPredicate *P) {
15267 if (!NewPreds) {
15268 // Check if we've already made this assumption.
15269 return Pred && Pred->implies(P, SE);
15270 }
15271 NewPreds->push_back(P);
15272 return true;
15273 }
15274
15275 bool addOverflowAssumption(const SCEVAddRecExpr *AR,
15277 auto *A = SE.getWrapPredicate(AR, AddedFlags);
15278 return addOverflowAssumption(A);
15279 }
15280
15281 // If \p Expr represents a PHINode, we try to see if it can be represented
15282 // as an AddRec, possibly under a predicate (PHISCEVPred). If it is possible
15283 // to add this predicate as a runtime overflow check, we return the AddRec.
15284 // If \p Expr does not meet these conditions (is not a PHI node, or we
15285 // couldn't create an AddRec for it, or couldn't add the predicate), we just
15286 // return \p Expr.
15287 const SCEV *convertToAddRecWithPreds(const SCEVUnknown *Expr) {
15288 if (!isa<PHINode>(Expr->getValue()))
15289 return Expr;
15290 std::optional<
15291 std::pair<const SCEV *, SmallVector<const SCEVPredicate *, 3>>>
15292 PredicatedRewrite = SE.createAddRecFromPHIWithCasts(Expr);
15293 if (!PredicatedRewrite)
15294 return Expr;
15295 for (const auto *P : PredicatedRewrite->second){
15296 // Wrap predicates from outer loops are not supported.
15297 if (auto *WP = dyn_cast<const SCEVWrapPredicate>(P)) {
15298 if (L != WP->getExpr()->getLoop())
15299 return Expr;
15300 }
15301 if (!addOverflowAssumption(P))
15302 return Expr;
15303 }
15304 return PredicatedRewrite->first;
15305 }
15306
15307 SmallVectorImpl<const SCEVPredicate *> *NewPreds;
15308 const SCEVPredicate *Pred;
15309 const Loop *L;
15310};
15311
15312} // end anonymous namespace
15313
15314const SCEV *
15316 const SCEVPredicate &Preds) {
15317 return SCEVPredicateRewriter::rewrite(S, L, *this, nullptr, &Preds);
15318}
15319
15321 const SCEV *S, const Loop *L,
15324 S = SCEVPredicateRewriter::rewrite(S, L, *this, &TransformPreds, nullptr);
15325 auto *AddRec = dyn_cast<SCEVAddRecExpr>(S);
15326
15327 if (!AddRec)
15328 return nullptr;
15329
15330 // Check if any of the transformed predicates is known to be false. In that
15331 // case, it doesn't make sense to convert to a predicated AddRec, as the
15332 // versioned loop will never execute.
15333 for (const SCEVPredicate *Pred : TransformPreds) {
15334 auto *WrapPred = dyn_cast<SCEVWrapPredicate>(Pred);
15335 if (!WrapPred || WrapPred->getFlags() != SCEVWrapPredicate::IncrementNSSW)
15336 continue;
15337
15338 const SCEVAddRecExpr *AddRecToCheck = WrapPred->getExpr();
15339 const SCEV *ExitCount = getBackedgeTakenCount(AddRecToCheck->getLoop());
15340 if (isa<SCEVCouldNotCompute>(ExitCount))
15341 continue;
15342
15343 const SCEV *Step = AddRecToCheck->getStepRecurrence(*this);
15344 if (!Step->isOne())
15345 continue;
15346
15347 ExitCount = getTruncateOrSignExtend(ExitCount, Step->getType());
15348 const SCEV *Add = getAddExpr(AddRecToCheck->getStart(), ExitCount);
15349 if (isKnownPredicate(CmpInst::ICMP_SLT, Add, AddRecToCheck->getStart()))
15350 return nullptr;
15351 }
15352
15353 // Since the transformation was successful, we can now transfer the SCEV
15354 // predicates.
15355 Preds.append(TransformPreds.begin(), TransformPreds.end());
15356
15357 return AddRec;
15358}
15359
15360/// SCEV predicates
15364
15366 const ICmpInst::Predicate Pred,
15367 const SCEV *LHS, const SCEV *RHS)
15368 : SCEVPredicate(ID, P_Compare), Pred(Pred), LHS(LHS), RHS(RHS) {
15369 assert(LHS->getType() == RHS->getType() && "LHS and RHS types don't match");
15370 assert(LHS != RHS && "LHS and RHS are the same SCEV");
15371}
15372
15374 ScalarEvolution &SE) const {
15375 const auto *Op = dyn_cast<SCEVComparePredicate>(N);
15376
15377 if (!Op)
15378 return false;
15379
15380 if (Pred != ICmpInst::ICMP_EQ)
15381 return false;
15382
15383 return Op->LHS == LHS && Op->RHS == RHS;
15384}
15385
15386bool SCEVComparePredicate::isAlwaysTrue() const { return false; }
15387
15389 if (Pred == ICmpInst::ICMP_EQ)
15390 OS.indent(Depth) << "Equal predicate: " << *LHS << " == " << *RHS << "\n";
15391 else
15392 OS.indent(Depth) << "Compare predicate: " << *LHS << " " << Pred << ") "
15393 << *RHS << "\n";
15394
15395}
15396
15398 const SCEVAddRecExpr *AR,
15399 IncrementWrapFlags Flags)
15400 : SCEVPredicate(ID, P_Wrap), AR(AR), Flags(Flags) {}
15401
15402const SCEVAddRecExpr *SCEVWrapPredicate::getExpr() const { return AR; }
15403
15405 ScalarEvolution &SE) const {
15406 const auto *Op = dyn_cast<SCEVWrapPredicate>(N);
15407 if (!Op || setFlags(Flags, Op->Flags) != Flags)
15408 return false;
15409
15410 if (Op->AR == AR)
15411 return true;
15412
15413 if (Flags != SCEVWrapPredicate::IncrementNSSW &&
15415 return false;
15416
15417 const SCEV *Start = AR->getStart();
15418 const SCEV *OpStart = Op->AR->getStart();
15419 if (Start->getType()->isPointerTy() != OpStart->getType()->isPointerTy())
15420 return false;
15421
15422 // Reject pointers to different address spaces.
15423 if (Start->getType()->isPointerTy() && Start->getType() != OpStart->getType())
15424 return false;
15425
15426 // NUSW/NSSW on a wider-type AddRec does not imply the same on a
15427 // narrower-type AddRec.
15428 if (SE.getTypeSizeInBits(AR->getType()) >
15429 SE.getTypeSizeInBits(Op->AR->getType()))
15430 return false;
15431
15432 const SCEV *Step = AR->getStepRecurrence(SE);
15433 const SCEV *OpStep = Op->AR->getStepRecurrence(SE);
15434 if (!SE.isKnownPositive(Step) || !SE.isKnownPositive(OpStep))
15435 return false;
15436
15437 // If both steps are positive, this implies N, if N's start and step are
15438 // ULE/SLE (for NSUW/NSSW) than this'.
15439 Type *WiderTy = SE.getWiderType(Step->getType(), OpStep->getType());
15440 Step = SE.getNoopOrZeroExtend(Step, WiderTy);
15441 OpStep = SE.getNoopOrZeroExtend(OpStep, WiderTy);
15442
15443 bool IsNUW = Flags == SCEVWrapPredicate::IncrementNUSW;
15444 OpStart = IsNUW ? SE.getNoopOrZeroExtend(OpStart, WiderTy)
15445 : SE.getNoopOrSignExtend(OpStart, WiderTy);
15446 Start = IsNUW ? SE.getNoopOrZeroExtend(Start, WiderTy)
15447 : SE.getNoopOrSignExtend(Start, WiderTy);
15449 return SE.isKnownPredicate(Pred, OpStep, Step) &&
15450 SE.isKnownPredicate(Pred, OpStart, Start);
15451}
15452
15454 SCEVFlags ScevFlags = AR->getNoWrapFlags();
15455 IncrementWrapFlags IFlags = Flags;
15456
15457 if (ScalarEvolution::setFlags(ScevFlags, SCEV::FlagNSW) == ScevFlags)
15458 IFlags = clearFlags(IFlags, IncrementNSSW);
15459
15460 return IFlags == IncrementAnyWrap;
15461}
15462
15463void SCEVWrapPredicate::print(raw_ostream &OS, unsigned Depth) const {
15464 OS.indent(Depth) << *getExpr() << " Added Flags: ";
15466 OS << "<nusw>";
15468 OS << "<nssw>";
15469 OS << "\n";
15470}
15471
15472/// Union predicates don't get cached so create a dummy set ID for it.
15474 ScalarEvolution &SE)
15476 for (const auto *P : Preds)
15477 add(P, SE);
15478}
15479
15481 return all_of(Preds,
15482 [](const SCEVPredicate *I) { return I->isAlwaysTrue(); });
15483}
15484
15486 ScalarEvolution &SE) const {
15487 if (const auto *Set = dyn_cast<SCEVUnionPredicate>(N))
15488 return all_of(Set->Preds, [this, &SE](const SCEVPredicate *I) {
15489 return this->implies(I, SE);
15490 });
15491
15492 if (any_of(Preds,
15493 [N, &SE](const SCEVPredicate *I) { return I->implies(N, SE); }))
15494 return true;
15495
15496 // A wrap predicate may be implied by a wrap predicate in Preds after applying
15497 // equal predicates.
15498 const auto *NWrap = dyn_cast<SCEVWrapPredicate>(N);
15499 if (!NWrap)
15500 return false;
15501 const Loop *L = NWrap->getExpr()->getLoop();
15502 return any_of(Preds, [&](const SCEVPredicate *I) {
15503 const auto *IWrap = dyn_cast<SCEVWrapPredicate>(I);
15504 if (!IWrap)
15505 return false;
15506 const auto *RewrittenAR = dyn_cast<SCEVAddRecExpr>(
15507 SE.rewriteUsingPredicate(IWrap->getExpr(), L, *this));
15508 return RewrittenAR &&
15509 SE.getWrapPredicate(RewrittenAR, IWrap->getFlags())->implies(N, SE);
15510 });
15511}
15512
15514 for (const auto *Pred : Preds)
15515 Pred->print(OS, Depth);
15516}
15517
15518void SCEVUnionPredicate::add(const SCEVPredicate *N, ScalarEvolution &SE) {
15519 if (const auto *Set = dyn_cast<SCEVUnionPredicate>(N)) {
15520 for (const auto *Pred : Set->Preds)
15521 add(Pred, SE);
15522 return;
15523 }
15524
15525 // Implication checks are quadratic in the number of predicates. Stop doing
15526 // them if there are many predicates, as they should be too expensive to use
15527 // anyway at that point.
15528 bool CheckImplies = Preds.size() < 16;
15529
15530 // Only add predicate if it is not already implied by this union predicate.
15531 if (CheckImplies && implies(N, SE))
15532 return;
15533
15534 // Build a new vector containing the current predicates, except the ones that
15535 // are implied by the new predicate N.
15537 for (auto *P : Preds) {
15538 if (CheckImplies && N->implies(P, SE))
15539 continue;
15540 PrunedPreds.push_back(P);
15541 }
15542 Preds = std::move(PrunedPreds);
15543 Preds.push_back(N);
15544}
15545
15547 Loop &L)
15548 : SE(SE), L(L) {
15550 Preds = std::make_unique<SCEVUnionPredicate>(Empty, SE);
15551}
15552
15554 for (const SCEV *Op : Ops)
15555 // We do not expect that forgetting cached data for SCEVConstants will ever
15556 // open any prospects for sharpening or introduce any correctness issues,
15557 // so we don't bother storing their dependencies.
15558 if (!isa<SCEVConstant>(Op))
15559 SCEVUsers[Op].insert(User);
15560}
15561
15563 const SCEV *Expr = SE.getSCEV(V);
15564 return getPredicatedSCEV(Expr);
15565}
15566
15568 RewriteEntry &Entry = RewriteMap[Expr];
15569
15570 // If we already have an entry and the version matches, return it.
15571 if (Entry.second && Generation == Entry.first)
15572 return Entry.second;
15573
15574 // We found an entry but it's stale. Rewrite the stale entry
15575 // according to the current predicate.
15576 if (Entry.second)
15577 Expr = Entry.second;
15578
15579 const SCEV *NewSCEV = SE.rewriteUsingPredicate(Expr, &L, *Preds);
15580 Entry = {Generation, NewSCEV};
15581
15582 return NewSCEV;
15583}
15584
15586 if (!BackedgeCount) {
15588 BackedgeCount = SE.getPredicatedBackedgeTakenCount(&L, Preds);
15589 for (const auto *P : Preds)
15590 addPredicate(*P);
15591 }
15592 return BackedgeCount;
15593}
15594
15596 if (!SymbolicMaxBackedgeCount) {
15598 SymbolicMaxBackedgeCount =
15599 SE.getPredicatedSymbolicMaxBackedgeTakenCount(&L, Preds);
15600 for (const auto *P : Preds)
15601 addPredicate(*P);
15602 }
15603 return SymbolicMaxBackedgeCount;
15604}
15605
15607 if (!SmallConstantMaxTripCount) {
15609 SmallConstantMaxTripCount = SE.getSmallConstantMaxTripCount(&L, &Preds);
15610 for (const auto *P : Preds)
15611 addPredicate(*P);
15612 }
15613 return *SmallConstantMaxTripCount;
15614}
15615
15617 if (Preds->implies(&Pred, SE))
15618 return;
15619
15620 SmallVector<const SCEVPredicate *, 4> NewPreds(Preds->getPredicates());
15621 NewPreds.push_back(&Pred);
15622 Preds = std::make_unique<SCEVUnionPredicate>(NewPreds, SE);
15623 updateGeneration();
15624}
15625
15628 for (const SCEVPredicate *P : Preds)
15629 addPredicate(*P);
15630}
15631
15633 return *Preds;
15634}
15635
15636void PredicatedScalarEvolution::updateGeneration() {
15637 // If the generation number wrapped recompute everything.
15638 if (++Generation == 0) {
15639 for (auto &II : RewriteMap) {
15640 const SCEV *Rewritten = II.second.second;
15641 II.second = {Generation, SE.rewriteUsingPredicate(Rewritten, &L, *Preds)};
15642 }
15643 }
15644}
15645
15648 const SCEV *Expr = this->getSCEV(V);
15650 auto *New = SE.convertSCEVToAddRecWithPredicates(Expr, &L, NewPreds);
15651
15652 if (!New)
15653 return nullptr;
15654
15655 if (ExtraPreds) {
15656 ExtraPreds->append(NewPreds);
15657 return New;
15658 }
15659
15660 addPredicates(NewPreds);
15661
15662 RewriteMap[SE.getSCEV(V)] = {Generation, New};
15663 return New;
15664}
15665
15668 : RewriteMap(Init.RewriteMap), SE(Init.SE), L(Init.L),
15669 Preds(std::make_unique<SCEVUnionPredicate>(Init.Preds->getPredicates(),
15670 SE)),
15671 Generation(Init.Generation), BackedgeCount(Init.BackedgeCount) {}
15672
15674 // For each block.
15675 for (auto *BB : L.getBlocks())
15676 for (auto &I : *BB) {
15677 if (!SE.isSCEVable(I.getType()))
15678 continue;
15679
15680 auto *Expr = SE.getSCEV(&I);
15681 auto II = RewriteMap.find(Expr);
15682
15683 if (II == RewriteMap.end())
15684 continue;
15685
15686 // Don't print things that are not interesting.
15687 if (II->second.second == Expr)
15688 continue;
15689
15690 OS.indent(Depth) << "[PSE]" << I << ":\n";
15691 OS.indent(Depth + 2) << *Expr << "\n";
15692 OS.indent(Depth + 2) << "--> " << *II->second.second << "\n";
15693 }
15694}
15695
15698 BasicBlock *Header = L->getHeader();
15699 BasicBlock *Pred = L->getLoopPredecessor();
15700 LoopGuards Guards(SE);
15701 if (!Pred)
15702 return Guards;
15704 collectFromBlock(SE, Guards, Header, Pred, VisitedBlocks);
15705 return Guards;
15706}
15707
15708void ScalarEvolution::LoopGuards::collectFromPHI(
15712 unsigned Depth) {
15713 if (!SE.isSCEVable(Phi.getType()))
15714 return;
15715
15716 using MinMaxPattern = std::pair<const SCEVConstant *, SCEVTypes>;
15717 auto GetMinMaxConst = [&](unsigned IncomingIdx) -> MinMaxPattern {
15718 const BasicBlock *InBlock = Phi.getIncomingBlock(IncomingIdx);
15719 if (!VisitedBlocks.insert(InBlock).second)
15720 return {nullptr, scCouldNotCompute};
15721
15722 // Avoid analyzing unreachable blocks so that we don't get trapped
15723 // traversing cycles with ill-formed dominance or infinite cycles
15724 if (!SE.DT.isReachableFromEntry(InBlock))
15725 return {nullptr, scCouldNotCompute};
15726
15727 auto [G, Inserted] = IncomingGuards.try_emplace(InBlock, LoopGuards(SE));
15728 if (Inserted)
15729 collectFromBlock(SE, G->second, Phi.getParent(), InBlock, VisitedBlocks,
15730 Depth + 1);
15731 auto &RewriteMap = G->second.RewriteMap;
15732 if (RewriteMap.empty())
15733 return {nullptr, scCouldNotCompute};
15734 auto S = RewriteMap.find(SE.getSCEV(Phi.getIncomingValue(IncomingIdx)));
15735 if (S == RewriteMap.end())
15736 return {nullptr, scCouldNotCompute};
15737 auto *SM = dyn_cast_if_present<SCEVMinMaxExpr>(S->second);
15738 if (!SM)
15739 return {nullptr, scCouldNotCompute};
15740 if (const SCEVConstant *C0 = dyn_cast<SCEVConstant>(SM->getOperand(0)))
15741 return {C0, SM->getSCEVType()};
15742 return {nullptr, scCouldNotCompute};
15743 };
15744 auto MergeMinMaxConst = [](MinMaxPattern P1,
15745 MinMaxPattern P2) -> MinMaxPattern {
15746 auto [C1, T1] = P1;
15747 auto [C2, T2] = P2;
15748 if (!C1 || !C2 || T1 != T2)
15749 return {nullptr, scCouldNotCompute};
15750 switch (T1) {
15751 case scUMaxExpr:
15752 return {C1->getAPInt().ult(C2->getAPInt()) ? C1 : C2, T1};
15753 case scSMaxExpr:
15754 return {C1->getAPInt().slt(C2->getAPInt()) ? C1 : C2, T1};
15755 case scUMinExpr:
15756 return {C1->getAPInt().ugt(C2->getAPInt()) ? C1 : C2, T1};
15757 case scSMinExpr:
15758 return {C1->getAPInt().sgt(C2->getAPInt()) ? C1 : C2, T1};
15759 default:
15760 llvm_unreachable("Trying to merge non-MinMaxExpr SCEVs.");
15761 }
15762 };
15763 auto P = GetMinMaxConst(0);
15764 for (unsigned int In = 1; In < Phi.getNumIncomingValues(); In++) {
15765 if (!P.first)
15766 break;
15767 P = MergeMinMaxConst(P, GetMinMaxConst(In));
15768 }
15769 if (P.first) {
15770 const SCEV *LHS = SE.getSCEV(const_cast<PHINode *>(&Phi));
15771 SmallVector<SCEVUse, 2> Ops({P.first, LHS});
15772 const SCEV *RHS = SE.getMinMaxExpr(P.second, Ops);
15773 Guards.RewriteMap.insert({LHS, RHS});
15774 }
15775}
15776
15777// Return a new SCEV that modifies \p Expr to the closest number divides by
15778// \p Divisor and less or equal than Expr. For now, only handle constant
15779// Expr.
15781 const APInt &DivisorVal,
15782 ScalarEvolution &SE) {
15783 const APInt *ExprVal;
15784 if (!match(Expr, m_scev_APInt(ExprVal)) || ExprVal->isNegative() ||
15785 DivisorVal.isNonPositive())
15786 return Expr;
15787 APInt Rem = ExprVal->urem(DivisorVal);
15788 // return the SCEV: Expr - Expr % Divisor
15789 return SE.getConstant(*ExprVal - Rem);
15790}
15791
15792// Return a new SCEV that modifies \p Expr to the closest number divides by
15793// \p Divisor and greater or equal than Expr. For now, only handle constant
15794// Expr.
15795static const SCEV *getNextSCEVDivisibleByDivisor(const SCEV *Expr,
15796 const APInt &DivisorVal,
15797 ScalarEvolution &SE) {
15798 const APInt *ExprVal;
15799 if (!match(Expr, m_scev_APInt(ExprVal)) || ExprVal->isNegative() ||
15800 DivisorVal.isNonPositive())
15801 return Expr;
15802 APInt Rem = ExprVal->urem(DivisorVal);
15803 if (Rem.isZero())
15804 return Expr;
15805 // return the SCEV: Expr + Divisor - Expr % Divisor
15806 return SE.getConstant(*ExprVal + DivisorVal - Rem);
15807}
15808
15810 ICmpInst::Predicate Predicate, const SCEV *LHS, const SCEV *RHS,
15813 // If we have LHS == 0, check if LHS is computing a property of some unknown
15814 // SCEV %v which we can rewrite %v to express explicitly.
15816 return false;
15817 // If LHS is A % B, i.e. A % B == 0, rewrite A to (A /u B) * B to
15818 // explicitly express that.
15819 const SCEVUnknown *URemLHS = nullptr;
15820 const SCEV *URemRHS = nullptr;
15821 if (!match(LHS, m_scev_URem(m_SCEVUnknown(URemLHS), m_SCEV(URemRHS), SE)))
15822 return false;
15823
15824 const SCEV *Multiple =
15825 SE.getMulExpr(SE.getUDivExpr(URemLHS, URemRHS), URemRHS);
15826 DivInfo[URemLHS] = Multiple;
15827 if (auto *C = dyn_cast<SCEVConstant>(URemRHS))
15828 Multiples[URemLHS] = C->getAPInt();
15829 return true;
15830}
15831
15832// Check if the condition is a divisibility guard (A % B == 0).
15833static bool isDivisibilityGuard(const SCEV *LHS, const SCEV *RHS,
15834 ScalarEvolution &SE) {
15835 const SCEV *X, *Y;
15836 return match(LHS, m_scev_URem(m_SCEV(X), m_SCEV(Y), SE)) && RHS->isZero();
15837}
15838
15839// Apply divisibility by \p Divisor on MinMaxExpr with constant values,
15840// recursively. This is done by aligning up/down the constant value to the
15841// Divisor.
15842static const SCEV *applyDivisibilityOnMinMaxExpr(const SCEV *MinMaxExpr,
15843 APInt Divisor,
15844 ScalarEvolution &SE) {
15845 // Return true if \p Expr is a MinMax SCEV expression with a non-negative
15846 // constant operand. If so, return in \p SCTy the SCEV type and in \p RHS
15847 // the non-constant operand and in \p LHS the constant operand.
15848 auto IsMinMaxSCEVWithNonNegativeConstant =
15849 [&](const SCEV *Expr, SCEVTypes &SCTy, const SCEV *&LHS,
15850 const SCEV *&RHS) {
15851 if (auto *MinMax = dyn_cast<SCEVMinMaxExpr>(Expr)) {
15852 if (MinMax->getNumOperands() != 2)
15853 return false;
15854 if (auto *C = dyn_cast<SCEVConstant>(MinMax->getOperand(0))) {
15855 if (C->getAPInt().isNegative())
15856 return false;
15857 SCTy = MinMax->getSCEVType();
15858 LHS = MinMax->getOperand(0);
15859 RHS = MinMax->getOperand(1);
15860 return true;
15861 }
15862 }
15863 return false;
15864 };
15865
15866 const SCEV *MinMaxLHS = nullptr, *MinMaxRHS = nullptr;
15867 SCEVTypes SCTy;
15868 if (!IsMinMaxSCEVWithNonNegativeConstant(MinMaxExpr, SCTy, MinMaxLHS,
15869 MinMaxRHS))
15870 return MinMaxExpr;
15871 auto IsMin = isa<SCEVSMinExpr>(MinMaxExpr) || isa<SCEVUMinExpr>(MinMaxExpr);
15872 assert(SE.isKnownNonNegative(MinMaxLHS) && "Expected non-negative operand!");
15873 auto *DivisibleExpr =
15874 IsMin ? getPreviousSCEVDivisibleByDivisor(MinMaxLHS, Divisor, SE)
15875 : getNextSCEVDivisibleByDivisor(MinMaxLHS, Divisor, SE);
15877 applyDivisibilityOnMinMaxExpr(MinMaxRHS, Divisor, SE), DivisibleExpr};
15878 return SE.getMinMaxExpr(SCTy, Ops);
15879}
15880
15881void ScalarEvolution::LoopGuards::collectFromBlock(
15882 ScalarEvolution &SE, ScalarEvolution::LoopGuards &Guards,
15883 const BasicBlock *Block, const BasicBlock *Pred,
15884 SmallPtrSetImpl<const BasicBlock *> &VisitedBlocks, unsigned Depth) {
15885
15887
15888 SmallVector<SCEVUse> ExprsToRewrite;
15889 auto CollectCondition = [&](ICmpInst::Predicate Predicate, const SCEV *LHS,
15890 const SCEV *RHS,
15891 DenseMap<const SCEV *, const SCEV *> &RewriteMap,
15892 const LoopGuards &DivGuards) {
15893 // WARNING: It is generally unsound to apply any wrap flags to the proposed
15894 // replacement SCEV which isn't directly implied by the structure of that
15895 // SCEV. In particular, using contextual facts to imply flags is *NOT*
15896 // legal. See the scoping rules for flags in the header to understand why.
15897
15898 // Puts rewrite rule \p From -> \p To into the rewrite map. Also if \p From
15899 // and \p FromRewritten are the same (i.e. there has been no rewrite
15900 // registered for \p From), then puts this value in the list of rewritten
15901 // expressions.
15902 auto AddRewrite = [&](const SCEV *From, const SCEV *FromRewritten,
15903 const SCEV *To) {
15904 if (From == FromRewritten)
15905 ExprsToRewrite.push_back(From);
15906 RewriteMap[From] = To;
15907 };
15908
15909 // Checks whether \p S has already been rewritten. In that case returns the
15910 // existing rewrite because we want to chain further rewrites onto the
15911 // already rewritten value. Otherwise returns \p S.
15912 auto GetMaybeRewritten = [&](const SCEV *S) {
15913 return RewriteMap.lookup_or(S, S);
15914 };
15915
15916 // Check for a condition of the form (-C1 + X < C2). InstCombine will
15917 // create this form when combining two checks of the form (X u< C2 + C1) and
15918 // (X >=u C1).
15919 auto MatchRangeCheckIdiom = [&](ICmpInst::Predicate Pred,
15920 const SCEV *MatchLHS,
15921 const SCEV *MatchRHS) {
15922 const SCEVConstant *C1;
15923 const SCEVUnknown *LHSUnknown;
15924 auto *C2 = dyn_cast<SCEVConstant>(MatchRHS);
15925 if (!match(MatchLHS,
15926 m_scev_Add(m_SCEVConstant(C1), m_SCEVUnknown(LHSUnknown))) ||
15927 !C2)
15928 return false;
15929
15930 auto ExactRegion =
15931 ConstantRange::makeExactICmpRegion(Pred, C2->getAPInt())
15932 .sub(C1->getAPInt());
15933
15934 // Tighten the raw range with what we already know about LHSUnknown
15935 // from prior guards recorded in RewriteMap, or from SCEV's own range
15936 // analysis.
15937 const SCEV *RewrittenLHS = GetMaybeRewritten(LHSUnknown);
15938 ExactRegion = ExactRegion.intersectWith(SE.getUnsignedRange(RewrittenLHS),
15940
15941 // Bail if the guard is inconsistent with prior facts, or if the range
15942 // is still not a monotonic non-wrapping interval after tightening.
15943 if (ExactRegion.isEmptySet() || ExactRegion.isWrappedSet() ||
15944 ExactRegion.isFullSet())
15945 return false;
15946
15947 const SCEV *RegionMin = SE.getConstant(ExactRegion.getUnsignedMin());
15948 const SCEV *RegionMax = SE.getConstant(ExactRegion.getUnsignedMax());
15949 const SCEV *ClampedLHS =
15950 SE.getUMaxExpr(RegionMin, SE.getUMinExpr(RewrittenLHS, RegionMax));
15951 AddRewrite(LHSUnknown, RewrittenLHS, ClampedLHS);
15952 return true;
15953 };
15954 if (MatchRangeCheckIdiom(Predicate, LHS, RHS))
15955 return;
15956
15957 // Do not apply information for constants or if RHS contains an AddRec.
15959 return;
15960
15961 // If RHS is SCEVUnknown, make sure the information is applied to it.
15963 std::swap(LHS, RHS);
15965 }
15966
15967 const SCEV *RewrittenLHS = GetMaybeRewritten(LHS);
15968 // Apply divisibility information when computing the constant multiple.
15969 const APInt &DividesBy =
15970 SE.getConstantMultiple(DivGuards.rewrite(RewrittenLHS));
15971
15972 // Collect rewrites for LHS and its transitive operands based on the
15973 // condition.
15974 // For min/max expressions, also apply the guard to its operands:
15975 // 'min(a, b) >= c' -> '(a >= c) and (b >= c)',
15976 // 'min(a, b) > c' -> '(a > c) and (b > c)',
15977 // 'max(a, b) <= c' -> '(a <= c) and (b <= c)',
15978 // 'max(a, b) < c' -> '(a < c) and (b < c)'.
15979
15980 // We cannot express strict predicates in SCEV, so instead we replace them
15981 // with non-strict ones against plus or minus one of RHS depending on the
15982 // predicate.
15983 const SCEV *One = SE.getOne(RHS->getType());
15984 switch (Predicate) {
15985 case CmpInst::ICMP_ULT:
15986 if (RHS->getType()->isPointerTy())
15987 return;
15988 RHS = SE.getUMaxExpr(RHS, One);
15989 [[fallthrough]];
15990 case CmpInst::ICMP_SLT: {
15991 RHS = SE.getMinusSCEV(RHS, One);
15992 RHS = getPreviousSCEVDivisibleByDivisor(RHS, DividesBy, SE);
15993 break;
15994 }
15995 case CmpInst::ICMP_UGT:
15996 case CmpInst::ICMP_SGT:
15997 RHS = SE.getAddExpr(RHS, One);
15998 RHS = getNextSCEVDivisibleByDivisor(RHS, DividesBy, SE);
15999 break;
16000 case CmpInst::ICMP_ULE:
16001 case CmpInst::ICMP_SLE:
16002 RHS = getPreviousSCEVDivisibleByDivisor(RHS, DividesBy, SE);
16003 break;
16004 case CmpInst::ICMP_UGE:
16005 case CmpInst::ICMP_SGE:
16006 RHS = getNextSCEVDivisibleByDivisor(RHS, DividesBy, SE);
16007 break;
16008 default:
16009 break;
16010 }
16011
16012 SmallVector<SCEVUse, 16> Worklist(1, LHS);
16013 SmallPtrSet<const SCEV *, 16> Visited;
16014
16015 auto EnqueueOperands = [&Worklist](const SCEVNAryExpr *S) {
16016 append_range(Worklist, S->operands());
16017 };
16018
16019 while (!Worklist.empty()) {
16020 const SCEV *From = Worklist.pop_back_val();
16021 if (isa<SCEVConstant>(From))
16022 continue;
16023 if (!Visited.insert(From).second)
16024 continue;
16025 const SCEV *FromRewritten = GetMaybeRewritten(From);
16026 const SCEV *To = nullptr;
16027
16028 switch (Predicate) {
16029 case CmpInst::ICMP_ULT:
16030 case CmpInst::ICMP_ULE:
16031 To = SE.getUMinExpr(FromRewritten, RHS);
16032 if (auto *UMax = dyn_cast<SCEVUMaxExpr>(FromRewritten))
16033 EnqueueOperands(UMax);
16034 break;
16035 case CmpInst::ICMP_SLT:
16036 case CmpInst::ICMP_SLE:
16037 To = SE.getSMinExpr(FromRewritten, RHS);
16038 if (auto *SMax = dyn_cast<SCEVSMaxExpr>(FromRewritten))
16039 EnqueueOperands(SMax);
16040 break;
16041 case CmpInst::ICMP_UGT:
16042 case CmpInst::ICMP_UGE:
16043 To = SE.getUMaxExpr(FromRewritten, RHS);
16044 if (auto *UMin = dyn_cast<SCEVUMinExpr>(FromRewritten))
16045 EnqueueOperands(UMin);
16046 break;
16047 case CmpInst::ICMP_SGT:
16048 case CmpInst::ICMP_SGE:
16049 To = SE.getSMaxExpr(FromRewritten, RHS);
16050 if (auto *SMin = dyn_cast<SCEVSMinExpr>(FromRewritten))
16051 EnqueueOperands(SMin);
16052 break;
16053 case CmpInst::ICMP_EQ:
16055 To = RHS;
16056 break;
16057 case CmpInst::ICMP_NE:
16058 if (match(RHS, m_scev_Zero())) {
16059 const SCEV *OneAlignedUp =
16060 getNextSCEVDivisibleByDivisor(One, DividesBy, SE);
16061 To = SE.getUMaxExpr(FromRewritten, OneAlignedUp);
16062 } else {
16063 // LHS != RHS can be rewritten as (LHS - RHS) = UMax(1, LHS - RHS),
16064 // but creating the subtraction eagerly is expensive. Track the
16065 // inequalities in a separate map, and materialize the rewrite lazily
16066 // when encountering a suitable subtraction while re-writing.
16067 if (LHS->getType()->isPointerTy()) {
16068 LHS = SE.getPtrToAddrExpr(LHS);
16069 RHS = SE.getPtrToAddrExpr(RHS);
16071 break;
16072 }
16073 const SCEVConstant *C;
16074 const SCEV *A, *B;
16077 RHS = A;
16078 LHS = B;
16079 }
16080 if (LHS > RHS)
16081 std::swap(LHS, RHS);
16082 Guards.NotEqual.insert({LHS, RHS});
16083 continue;
16084 }
16085 break;
16086 default:
16087 break;
16088 }
16089
16090 if (To)
16091 AddRewrite(From, FromRewritten, To);
16092 }
16093 };
16094
16096 // First, collect information from assumptions dominating the loop.
16097 for (auto &AssumeVH : SE.AC.assumptions()) {
16098 if (!AssumeVH)
16099 continue;
16100 auto *AssumeI = cast<CallInst>(AssumeVH);
16101 if (!SE.DT.dominates(AssumeI, Block))
16102 continue;
16103 Terms.emplace_back(AssumeI->getOperand(0), true);
16104 }
16105
16106 // Second, collect information from llvm.experimental.guards dominating the loop.
16107 auto *GuardDecl = Intrinsic::getDeclarationIfExists(
16108 SE.F.getParent(), Intrinsic::experimental_guard);
16109 if (GuardDecl)
16110 for (const auto *GU : GuardDecl->users())
16111 if (const auto *Guard = dyn_cast<IntrinsicInst>(GU))
16112 if (Guard->getFunction() == Block->getParent() &&
16113 SE.DT.dominates(Guard, Block))
16114 Terms.emplace_back(Guard->getArgOperand(0), true);
16115
16116 // Third, collect conditions from dominating branches. Starting at the loop
16117 // predecessor, climb up the predecessor chain, as long as there are
16118 // predecessors that can be found that have unique successors leading to the
16119 // original header.
16120 // TODO: share this logic with isLoopEntryGuardedByCond.
16121 unsigned NumCollectedConditions = 0;
16123 std::pair<const BasicBlock *, const BasicBlock *> Pair(Pred, Block);
16124 for (; Pair.first;
16125 Pair = SE.getPredecessorWithUniqueSuccessorForBB(Pair.first)) {
16126 VisitedBlocks.insert(Pair.second);
16127 const CondBrInst *LoopEntryPredicate =
16128 dyn_cast<CondBrInst>(Pair.first->getTerminator());
16129 if (!LoopEntryPredicate)
16130 continue;
16131
16132 Terms.emplace_back(LoopEntryPredicate->getCondition(),
16133 LoopEntryPredicate->getSuccessor(0) == Pair.second);
16134 NumCollectedConditions++;
16135
16136 // If we are recursively collecting guards stop after 2
16137 // conditions to limit compile-time impact for now.
16138 if (Depth > 0 && NumCollectedConditions == 2)
16139 break;
16140 }
16141 // Finally, if we stopped climbing the predecessor chain because
16142 // there wasn't a unique one to continue, try to collect conditions
16143 // for PHINodes by recursively following all of their incoming
16144 // blocks and try to merge the found conditions to build a new one
16145 // for the Phi.
16146 if (Pair.second->hasNPredecessorsOrMore(2) &&
16148 SmallDenseMap<const BasicBlock *, LoopGuards> IncomingGuards;
16149 for (auto &Phi : Pair.second->phis())
16150 collectFromPHI(SE, Guards, Phi, VisitedBlocks, IncomingGuards, Depth);
16151 }
16152
16153 // Now apply the information from the collected conditions to
16154 // Guards.RewriteMap. Conditions are processed in reverse order, so the
16155 // earliest conditions is processed first, except guards with divisibility
16156 // information, which are moved to the back. This ensures the SCEVs with the
16157 // shortest dependency chains are constructed first.
16159 GuardsToProcess;
16160 for (auto [Term, EnterIfTrue] : reverse(Terms)) {
16161 SmallVector<Value *, 8> Worklist;
16162 SmallPtrSet<Value *, 8> Visited;
16163 Worklist.push_back(Term);
16164 while (!Worklist.empty()) {
16165 Value *Cond = Worklist.pop_back_val();
16166 if (!Visited.insert(Cond).second)
16167 continue;
16168
16169 if (auto *Cmp = dyn_cast<ICmpInst>(Cond)) {
16170 auto Predicate =
16171 EnterIfTrue ? Cmp->getPredicate() : Cmp->getInversePredicate();
16172 const auto *LHS = SE.getSCEV(Cmp->getOperand(0));
16173 const auto *RHS = SE.getSCEV(Cmp->getOperand(1));
16174 // If LHS is a constant, apply information to the other expression.
16175 // TODO: If LHS is not a constant, check if using CompareSCEVComplexity
16176 // can improve results.
16177 if (isa<SCEVConstant>(LHS)) {
16178 std::swap(LHS, RHS);
16180 }
16181 GuardsToProcess.emplace_back(Predicate, LHS, RHS);
16182 continue;
16183 }
16184
16185 Value *L, *R;
16186 if (EnterIfTrue ? match(Cond, m_LogicalAnd(m_Value(L), m_Value(R)))
16187 : match(Cond, m_LogicalOr(m_Value(L), m_Value(R)))) {
16188 Worklist.push_back(L);
16189 Worklist.push_back(R);
16190 }
16191 }
16192 }
16193
16194 // Process divisibility guards in reverse order to populate DivGuards early.
16195 DenseMap<const SCEV *, APInt> Multiples;
16196 LoopGuards DivGuards(SE);
16197 for (const auto &[Predicate, LHS, RHS] : GuardsToProcess) {
16198 if (!isDivisibilityGuard(LHS, RHS, SE))
16199 continue;
16200 collectDivisibilityInformation(Predicate, LHS, RHS, DivGuards.RewriteMap,
16201 Multiples, SE);
16202 }
16203
16204 for (const auto &[Predicate, LHS, RHS] : GuardsToProcess)
16205 CollectCondition(Predicate, LHS, RHS, Guards.RewriteMap, DivGuards);
16206
16207 // Apply divisibility information last. This ensures it is applied to the
16208 // outermost expression after other rewrites for the given value.
16209 for (const auto &[K, Divisor] : Multiples) {
16210 const SCEV *DivisorSCEV = SE.getConstant(Divisor);
16211 Guards.RewriteMap[K] =
16213 Guards.rewrite(K), Divisor, SE),
16214 DivisorSCEV),
16215 DivisorSCEV);
16216 ExprsToRewrite.push_back(K);
16217 }
16218
16219 // Let the rewriter preserve NUW/NSW flags if the unsigned/signed ranges of
16220 // the replacement expressions are contained in the ranges of the replaced
16221 // expressions.
16222 Guards.PreserveNUW = true;
16223 Guards.PreserveNSW = true;
16224 for (const SCEV *Expr : ExprsToRewrite) {
16225 const SCEV *RewriteTo = Guards.RewriteMap[Expr];
16226 Guards.PreserveNUW &=
16227 SE.getUnsignedRange(Expr).contains(SE.getUnsignedRange(RewriteTo));
16228 Guards.PreserveNSW &=
16229 SE.getSignedRange(Expr).contains(SE.getSignedRange(RewriteTo));
16230 }
16231
16232 // Now that all rewrite information is collect, rewrite the collected
16233 // expressions with the information in the map. This applies information to
16234 // sub-expressions.
16235 if (ExprsToRewrite.size() > 1) {
16236 for (const SCEV *Expr : ExprsToRewrite) {
16237 const SCEV *RewriteTo = Guards.RewriteMap[Expr];
16238 Guards.RewriteMap.erase(Expr);
16239 Guards.RewriteMap.insert({Expr, Guards.rewrite(RewriteTo)});
16240 }
16241 }
16242}
16243
16245 /// A rewriter to replace SCEV expressions in Map with the corresponding entry
16246 /// in the map. It skips AddRecExpr because we cannot guarantee that the
16247 /// replacement is loop invariant in the loop of the AddRec.
16248 class SCEVLoopGuardRewriter
16249 : public SCEVRewriteVisitor<SCEVLoopGuardRewriter> {
16252
16253 SCEVFlags FlagMask = SCEV::FlagNone;
16254
16255 public:
16256 SCEVLoopGuardRewriter(ScalarEvolution &SE,
16257 const ScalarEvolution::LoopGuards &Guards)
16258 : SCEVRewriteVisitor(SE), Map(Guards.RewriteMap),
16259 NotEqual(Guards.NotEqual) {
16260 if (Guards.PreserveNUW)
16261 FlagMask = ScalarEvolution::setFlags(FlagMask, SCEV::FlagNUW);
16262 if (Guards.PreserveNSW)
16263 FlagMask = ScalarEvolution::setFlags(FlagMask, SCEV::FlagNSW);
16264 }
16265
16266 const SCEV *visitAddRecExpr(const SCEVAddRecExpr *Expr) { return Expr; }
16267
16268 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
16269 return Map.lookup_or(Expr, Expr);
16270 }
16271
16272 const SCEV *visitPtrToAddrExpr(const SCEVPtrToAddrExpr *Expr) {
16273 if (const SCEV *S = Map.lookup(Expr))
16274 return S;
16276 Expr);
16277 }
16278
16279 const SCEV *visitZeroExtendExpr(const SCEVZeroExtendExpr *Expr) {
16280 if (const SCEV *S = Map.lookup(Expr))
16281 return S;
16282
16283 // If we didn't find the extact ZExt expr in the map, check if there's
16284 // an entry for a smaller ZExt we can use instead.
16285 Type *Ty = Expr->getType();
16286 const SCEV *Op = Expr->getOperand(0);
16287 unsigned Bitwidth = Ty->getScalarSizeInBits() / 2;
16288 while (Bitwidth % 8 == 0 && Bitwidth >= 8 &&
16289 Bitwidth > Op->getType()->getScalarSizeInBits()) {
16290 Type *NarrowTy = IntegerType::get(SE.getContext(), Bitwidth);
16291 auto *NarrowExt = SE.getZeroExtendExpr(Op, NarrowTy);
16292 if (const SCEV *S = Map.lookup(NarrowExt))
16293 return SE.getZeroExtendExpr(S, Ty);
16294 Bitwidth = Bitwidth / 2;
16295 }
16296
16298 Expr);
16299 }
16300
16301 const SCEV *visitSignExtendExpr(const SCEVSignExtendExpr *Expr) {
16302 if (const SCEV *S = Map.lookup(Expr))
16303 return S;
16305 Expr);
16306 }
16307
16308 const SCEV *visitUMinExpr(const SCEVUMinExpr *Expr) {
16309 if (const SCEV *S = Map.lookup(Expr))
16310 return S;
16312 }
16313
16314 const SCEV *visitSMinExpr(const SCEVSMinExpr *Expr) {
16315 if (const SCEV *S = Map.lookup(Expr))
16316 return S;
16318 }
16319
16320 const SCEV *visitAddExpr(const SCEVAddExpr *Expr) {
16321 if (const SCEV *S = Map.lookup(Expr))
16322 return S;
16323
16324 // Helper to check if S is a subtraction (A - B) where A != B, and if so,
16325 // return UMax(S, 1).
16326 auto RewriteSubtraction = [&](const SCEV *S) -> const SCEV * {
16327 SCEVUse LHS, RHS;
16328 if (MatchBinarySub(S, LHS, RHS)) {
16329 if (LHS > RHS)
16330 std::swap(LHS, RHS);
16331 if (NotEqual.contains({LHS, RHS})) {
16332 const SCEV *OneAlignedUp = getNextSCEVDivisibleByDivisor(
16333 SE.getOne(S->getType()), SE.getConstantMultiple(S), SE);
16334 return SE.getUMaxExpr(OneAlignedUp, S);
16335 }
16336 }
16337 return nullptr;
16338 };
16339
16340 // Check if Expr itself is a subtraction pattern with guard info.
16341 if (const SCEV *Rewritten = RewriteSubtraction(Expr))
16342 return Rewritten;
16343
16344 // Trip count expressions sometimes consist of adding 3 operands, i.e.
16345 // (Const + A + B). There may be guard info for A + B, and if so, apply
16346 // it.
16347 // TODO: Could more generally apply guards to Add sub-expressions.
16348 if (isa<SCEVConstant>(Expr->getOperand(0))) {
16349 if (Expr->getNumOperands() == 3) {
16350 const SCEV *Add =
16351 SE.getAddExpr(Expr->getOperand(1), Expr->getOperand(2));
16352 if (const SCEV *Rewritten = RewriteSubtraction(Add))
16353 return SE.getAddExpr(
16354 Expr->getOperand(0), Rewritten,
16355 ScalarEvolution::maskFlags(Expr->getNoWrapFlags(), FlagMask));
16356 if (const SCEV *S = Map.lookup(Add))
16357 return SE.getAddExpr(Expr->getOperand(0), S);
16358 }
16359
16360 // For expressions of the form (Const + A), check if we have guard info
16361 // for (Const + 1 + A), and rewrite to ((Const + 1 + A) - 1). This makes
16362 // sure we don't lose information when rewriting expressions based on
16363 // back-edge taken counts in some cases.
16364 if (Expr->getNumOperands() == 2) {
16365 const SCEV *S = nullptr;
16366 // Handle (-1 + 1 + A) without constructing SCEVs.
16367 if (match(Expr->getOperand(0), m_scev_AllOnes())) {
16368 S = Map.lookup(Expr->getOperand(1));
16369 } else {
16370 const SCEV *NewC =
16371 SE.getAddExpr(Expr->getOperand(0), SE.getOne(Expr->getType()));
16372 S = Map.lookup(SE.getAddExpr(NewC, Expr->getOperand(1)));
16373 }
16374 if (S)
16375 return SE.getAddExpr(S, SE.getMinusOne(Expr->getType()));
16376 }
16377 }
16379 bool Changed = false;
16380 for (SCEVUse Op : Expr->operands()) {
16381 Operands.push_back(
16383 Changed |= Op != Operands.back();
16384 }
16385 // We are only replacing operands with equivalent values, so transfer the
16386 // flags from the original expression.
16387 return !Changed ? Expr
16388 : SE.getAddExpr(Operands,
16390 Expr->getNoWrapFlags(), FlagMask));
16391 }
16392
16393 const SCEV *visitMulExpr(const SCEVMulExpr *Expr) {
16395 bool Changed = false;
16396 for (SCEVUse Op : Expr->operands()) {
16397 Operands.push_back(
16399 Changed |= Op != Operands.back();
16400 }
16401 // We are only replacing operands with equivalent values, so transfer the
16402 // flags from the original expression.
16403 return !Changed ? Expr
16404 : SE.getMulExpr(Operands,
16406 Expr->getNoWrapFlags(), FlagMask));
16407 }
16408 };
16409
16410 if (RewriteMap.empty() && NotEqual.empty())
16411 return Expr;
16412
16413 SCEVLoopGuardRewriter Rewriter(SE, *this);
16414 return Rewriter.visit(Expr);
16415}
16416
16417const SCEV *ScalarEvolution::applyLoopGuards(const SCEV *Expr, const Loop *L) {
16418 return applyLoopGuards(Expr, LoopGuards::collect(L, *this));
16419}
16420
16422 const LoopGuards &Guards) {
16423 return Guards.rewrite(Expr);
16424}
assert(UImm &&(UImm !=~static_cast< T >(0)) &&"Invalid immediate!")
unsigned uint64_t
constexpr LLT S1
Rewrite undef for PHI
This file implements a class to represent arbitrary precision integral constant values and operations...
@ PostInc
MachineBasicBlock MachineBasicBlock::iterator DebugLoc DL
Expand Atomic instructions
#define X(NUM, ENUM, NAME)
Definition ELF.h:857
static GCRegistry::Add< ShadowStackGC > C("shadow-stack", "Very portable GC for uncooperative code generators")
static GCRegistry::Add< ErlangGC > A("erlang", "erlang-compatible garbage collector")
static GCRegistry::Add< StatepointGC > D("statepoint-example", "an example strategy for statepoint")
static GCRegistry::Add< CoreCLRGC > E("coreclr", "CoreCLR-compatible GC")
static GCRegistry::Add< OcamlGC > B("ocaml", "ocaml 3.10-compatible GC")
#define LLVM_DUMP_METHOD
Mark debug helper function definitions like dump() that should not be stripped from debug builds.
Definition Compiler.h:686
This file contains the declarations for the subclasses of Constant, which represent the different fla...
SmallPtrSet< const BasicBlock *, 8 > VisitedBlocks
This file defines the DenseMap class.
This file builds on the ADT/GraphTraits.h file to build generic depth first graph iterator.
static bool isSigned(unsigned Opcode)
This file defines a hash set that can be used to remove duplication of nodes in a graph.
#define op(i)
Hexagon Common GEP
Value * getPointer(Value *Ptr)
This file provides various utilities for inspecting and working with the control flow graph in LLVM I...
This defines the Use class.
iv Induction Variable Users
Definition IVUsers.cpp:48
static bool hasNoUnsignedWrap(BinaryOperator &I)
static constexpr Value * getValue(Ty &ValueOrUse)
const AbstractManglingParser< Derived, Alloc >::OperatorInfo AbstractManglingParser< Derived, Alloc >::Ops[]
static bool isZero(Value *V, const DataLayout &DL, DominatorTree *DT, AssumptionCache *AC)
Definition Lint.cpp:540
#define F(x, y, z)
Definition MD5.cpp:54
#define I(x, y, z)
Definition MD5.cpp:57
#define G(x, y, z)
Definition MD5.cpp:55
#define T
#define T1
ConstantRange Range(APInt(BitWidth, Low), APInt(BitWidth, High))
uint64_t IntrinsicInst * II
#define P(N)
ppc ctr loops verify
PowerPC Reduce CR logical Operation
#define INITIALIZE_PASS_DEPENDENCY(depName)
Definition PassSupport.h:42
#define INITIALIZE_PASS_END(passName, arg, name, cfg, analysis)
Definition PassSupport.h:44
#define INITIALIZE_PASS_BEGIN(passName, arg, name, cfg, analysis)
Definition PassSupport.h:39
R600 Clause Merge
const SmallVectorImpl< MachineOperand > & Cond
static DominatorTree getDomTree(Function &F)
static bool isValid(const char C)
Returns true if C is a valid mangled character: <0-9a-zA-Z_>.
SI Fold Operands
SI optimize exec mask operations pre RA
static void visit(BasicBlock &Start, std::function< bool(BasicBlock *)> op)
This file contains some templates that are useful if you are working with the STL at all.
This file provides utility classes that use RAII to save and restore values.
bool SCEVMinMaxExprContains(const SCEV *Root, const SCEV *OperandToFind, SCEVTypes RootKind)
static cl::opt< unsigned > MaxAddRecSize("scalar-evolution-max-add-rec-size", cl::Hidden, cl::desc("Max coefficients in AddRec during evolving"), cl::init(8))
static cl::opt< unsigned > RangeIterThreshold("scev-range-iter-threshold", cl::Hidden, cl::desc("Threshold for switching to iteratively computing SCEV ranges"), cl::init(32))
static const Loop * isIntegerLoopHeaderPHI(const PHINode *PN, LoopInfo &LI)
static unsigned getConstantTripCount(const SCEVConstant *ExitCount)
static int CompareValueComplexity(const LoopInfo *const LI, Value *LV, Value *RV, unsigned Depth)
Compare the two values LV and RV in terms of their "complexity" where "complexity" is a partial (and ...
static const SCEV * getNextSCEVDivisibleByDivisor(const SCEV *Expr, const APInt &DivisorVal, ScalarEvolution &SE)
static void PushLoopPHIs(const Loop *L, SmallVectorImpl< Instruction * > &Worklist, SmallPtrSetImpl< Instruction * > &Visited)
Push PHI nodes in the header of the given loop onto the given Worklist.
static void insertFoldCacheEntry(const ScalarEvolution::FoldID &ID, const SCEV *S, DenseMap< ScalarEvolution::FoldID, const SCEV * > &FoldCache, DenseMap< const SCEV *, SmallVector< ScalarEvolution::FoldID, 2 > > &FoldCacheUser)
static cl::opt< bool > ClassifyExpressions("scalar-evolution-classify-expressions", cl::Hidden, cl::init(true), cl::desc("When printing analysis, include information on every instruction"))
static bool hasHugeExpression(ArrayRef< SCEVUse > Ops)
Returns true if Ops contains a huge SCEV (the subtree of S contains at least HugeExprThreshold nodes)...
static cl::opt< unsigned > AddOpsInlineThreshold("scev-addops-inline-threshold", cl::Hidden, cl::desc("Threshold for inlining addition operands into a SCEV"), cl::init(500))
static cl::opt< unsigned > MaxLoopGuardCollectionDepth("scalar-evolution-max-loop-guard-collection-depth", cl::Hidden, cl::desc("Maximum depth for recursive loop guard collection"), cl::init(1))
static cl::opt< bool > VerifyIR("scev-verify-ir", cl::Hidden, cl::desc("Verify IR correctness when making sensitive SCEV queries (slow)"), cl::init(false))
static bool RangeRefPHIAllowedOperands(DominatorTree &DT, PHINode *PHI)
static bool IsKnownPredicateViaAddRecMonotonicity(ScalarEvolution &SE, CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS)
Is LHS Pred RHS true because one of them is an AddRec that is known not to go below its own start val...
static std::optional< APInt > MinOptional(std::optional< APInt > X, std::optional< APInt > Y)
Helper function to compare optional APInts: (a) if X and Y both exist, return min(X,...
static PHINode * getConstantEvolvingPHI(Value *V, const Loop *L, const TargetLibraryInfo *TLI)
getConstantEvolvingPHI - Given an LLVM value and a loop, return a PHI node in the loop that V is deri...
static bool canConstantFold(const Instruction *I, const TargetLibraryInfo *TLI)
Return true if we can constant fold an instruction of the specified type, assuming that all operands ...
static cl::opt< unsigned > MulOpsInlineThreshold("scev-mulops-inline-threshold", cl::Hidden, cl::desc("Threshold for inlining multiplication operands into a SCEV"), cl::init(32))
static BinaryOperator * getCommonInstForPHI(PHINode *PN)
static PHINode * getConstantEvolvingPHIOperands(Instruction *UseInst, const Loop *L, DenseMap< Instruction *, PHINode * > &PHIMap, const TargetLibraryInfo *TLI, unsigned Depth)
getConstantEvolvingPHIOperands - Implement getConstantEvolvingPHI by recursing through each instructi...
static bool isDivisibilityGuard(const SCEV *LHS, const SCEV *RHS, ScalarEvolution &SE)
static std::optional< const SCEV * > createNodeForSelectViaUMinSeq(ScalarEvolution *SE, const SCEV *CondExpr, const SCEV *TrueExpr, const SCEV *FalseExpr)
static Constant * BuildConstantFromSCEV(const SCEV *V)
This builds up a Constant using the ConstantExpr interface.
static ConstantInt * EvaluateConstantChrecAtConstant(const SCEVAddRecExpr *AddRec, ConstantInt *C, ScalarEvolution &SE)
static const SCEV * BinomialCoefficient(const SCEV *It, unsigned K, ScalarEvolution &SE, Type *ResultTy)
Compute BC(It, K). The result has width W. Assume, K > 0.
static cl::opt< unsigned > MaxCastDepth("scalar-evolution-max-cast-depth", cl::Hidden, cl::desc("Maximum depth of recursive SExt/ZExt/Trunc"), cl::init(8))
static bool IsMinMaxConsistingOf(const SCEV *MaybeMinMaxExpr, const SCEV *Candidate)
Is MaybeMinMaxExpr an (U|S)(Min|Max) of Candidate and some other values?
static SCEVFlags getNoWrapFlagsForGEP(GEPOperator *GEP, const SCEV *Accum, ScalarEvolution &SE)
static const SCEV * SolveLinEquationWithOverflow(const APInt &A, const SCEV *B, SmallVectorImpl< const SCEVPredicate * > *Predicates, ScalarEvolution &SE, const Loop *L)
Finds the minimum unsigned root of the following equation:
static cl::opt< unsigned > MaxBruteForceIterations("scalar-evolution-max-iterations", cl::ReallyHidden, cl::desc("Maximum number of iterations SCEV will " "symbolically execute a constant " "derived loop"), cl::init(100))
static uint64_t umul_ov(uint64_t i, uint64_t j, bool &Overflow)
static void PrintSCEVWithTypeHint(raw_ostream &OS, const SCEV *S)
When printing a top-level SCEV for trip counts, it's helpful to include a type for constants which ar...
static void PrintLoopInfo(raw_ostream &OS, ScalarEvolution *SE, const Loop *L)
static bool containsConstantInAddMulChain(const SCEV *StartExpr)
Determine if any of the operands in this SCEV are a constant or if any of the add or multiply express...
static const SCEV * getExtendAddRecStart(const SCEVAddRecExpr *AR, Type *Ty, ScalarEvolution *SE, unsigned Depth)
static bool CollectAddOperandsWithScales(SmallDenseMap< SCEVUse, APInt, 16 > &M, SmallVectorImpl< SCEVUse > &NewOps, APInt &AccumulatedConstant, ArrayRef< SCEVUse > Ops, const APInt &Scale, ScalarEvolution &SE)
Process the given Ops list, which is a list of operands to be added under the given scale,...
static const SCEV * constantFoldAndGroupOps(ScalarEvolution &SE, LoopInfo &LI, DominatorTree &DT, SmallVectorImpl< SCEVUse > &Ops, FoldT Fold, IsIdentityT IsIdentity, IsAbsorberT IsAbsorber)
Performs a number of common optimizations on the passed Ops.
static bool IsKnownPredicateViaAddRecStart(ScalarEvolution &SE, CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS)
static const SCEV * getPreStartForExtend(const SCEVAddRecExpr *AR, ScalarEvolution *SE, unsigned Depth)
static void GroupByComplexity(SmallVectorImpl< SCEVUse > &Ops, LoopInfo *LI, DominatorTree &DT)
Given a list of SCEV objects, order them by their complexity, and group objects of the same complexit...
static bool collectDivisibilityInformation(ICmpInst::Predicate Predicate, const SCEV *LHS, const SCEV *RHS, DenseMap< const SCEV *, const SCEV * > &DivInfo, DenseMap< const SCEV *, APInt > &Multiples, ScalarEvolution &SE)
static cl::opt< unsigned > MaxSCEVOperationsImplicationDepth("scalar-evolution-max-scev-operations-implication-depth", cl::Hidden, cl::desc("Maximum depth of recursive SCEV operations implication analysis"), cl::init(2))
static void PushDefUseChildren(Instruction *I, SmallVectorImpl< Instruction * > &Worklist, SmallPtrSetImpl< Instruction * > &Visited)
Push users of the given Instruction onto the given Worklist.
static std::optional< APInt > SolveQuadraticAddRecRange(const SCEVAddRecExpr *AddRec, const ConstantRange &Range, ScalarEvolution &SE)
Let c(n) be the value of the quadratic chrec {0,+,M,+,N} after n iterations.
static cl::opt< bool > UseContextForNoWrapFlagInference("scalar-evolution-use-context-for-no-wrap-flag-strenghening", cl::Hidden, cl::desc("Infer nuw/nsw flags using context where suitable"), cl::init(true))
static cl::opt< bool > EnableFiniteLoopControl("scalar-evolution-finite-loop", cl::Hidden, cl::desc("Handle <= and >= in finite loops"), cl::init(true))
static bool getOperandsForSelectLikePHI(DominatorTree &DT, PHINode *PN, Value *&Cond, Value *&LHS, Value *&RHS)
static std::optional< std::tuple< APInt, APInt, APInt, APInt, unsigned > > GetQuadraticEquation(const SCEVAddRecExpr *AddRec)
For a given quadratic addrec, generate coefficients of the corresponding quadratic equation,...
static bool isKnownPredicateExtendIdiom(CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS)
static std::optional< APInt > SolveQuadraticAddRecExact(const SCEVAddRecExpr *AddRec, ScalarEvolution &SE)
Let c(n) be the value of the quadratic chrec {L,+,M,+,N} after n iterations.
static std::optional< APInt > TruncIfPossible(std::optional< APInt > X, unsigned BitWidth)
Helper function to truncate an optional APInt to a given BitWidth.
static cl::opt< unsigned > MaxSCEVCompareDepth("scalar-evolution-max-scev-compare-depth", cl::Hidden, cl::desc("Maximum depth of recursive SCEV complexity comparisons"), cl::init(32))
static SCEVFlags StrengthenNoWrapFlags(ScalarEvolution *SE, SCEVTypes Type, ArrayRef< SCEVUse > Ops, SCEVFlags Flags)
static APInt extractConstantWithoutWrapping(ScalarEvolution &SE, const SCEVConstant *ConstantTerm, const SCEVAddExpr *WholeAddExpr)
static cl::opt< unsigned > MaxConstantEvolvingDepth("scalar-evolution-max-constant-evolving-depth", cl::Hidden, cl::desc("Maximum depth of recursive constant evolving"), cl::init(32))
static bool canConstantEvolve(Instruction *I, const Loop *L, const TargetLibraryInfo *TLI)
Determine whether this instruction can constant evolve within this loop assuming its operands can all...
static bool MatchBinarySub(const SCEV *S, SCEVUse &LHS, SCEVUse &RHS)
static std::optional< ConstantRange > GetRangeFromMetadata(Value *V)
Helper method to assign a range to V from metadata present in the IR.
static cl::opt< unsigned > HugeExprThreshold("scalar-evolution-huge-expr-threshold", cl::Hidden, cl::desc("Size of the expression which is considered huge"), cl::init(4096))
static Type * isSimpleCastedPHI(const SCEV *Op, const SCEVUnknown *SymbolicPHI, bool &Signed, ScalarEvolution &SE)
Helper function to createAddRecFromPHIWithCasts.
static Constant * EvaluateExpression(Value *V, const Loop *L, DenseMap< Instruction *, Constant * > &Vals, const DataLayout &DL, const TargetLibraryInfo *TLI)
EvaluateExpression - Given an expression that passes the getConstantEvolvingPHI predicate,...
static const SCEV * getPreviousSCEVDivisibleByDivisor(const SCEV *Expr, const APInt &DivisorVal, ScalarEvolution &SE)
static const SCEV * MatchNotExpr(const SCEV *Expr)
If Expr computes ~A, return A else return nullptr.
static std::pair< ConstantRange, bool > getRangeForAffineARHelper(APInt Step, const ConstantRange &StartRange, const APInt &MaxBECount, bool Signed)
static cl::opt< unsigned > MaxValueCompareDepth("scalar-evolution-max-value-compare-depth", cl::Hidden, cl::desc("Maximum depth of recursive value complexity comparisons"), cl::init(2))
static const SCEV * applyDivisibilityOnMinMaxExpr(const SCEV *MinMaxExpr, APInt Divisor, ScalarEvolution &SE)
static cl::opt< bool, true > VerifySCEVOpt("verify-scev", cl::Hidden, cl::location(VerifySCEV), cl::desc("Verify ScalarEvolution's backedge taken counts (slow)"))
static const SCEV * getSignedOverflowLimitForStep(const SCEV *Step, ICmpInst::Predicate *Pred, ScalarEvolution *SE)
static cl::opt< unsigned > MaxArithDepth("scalar-evolution-max-arith-depth", cl::Hidden, cl::desc("Maximum depth of recursive arithmetics"), cl::init(32))
static bool HasSameValue(const SCEV *A, const SCEV *B)
SCEV structural equivalence is usually sufficient for testing whether two expressions are equal,...
static uint64_t Choose(uint64_t n, uint64_t k, bool &Overflow)
Compute the result of "n choose k", the binomial coefficient.
static std::optional< int > CompareSCEVComplexity(const LoopInfo *const LI, const SCEV *LHS, const SCEV *RHS, DominatorTree &DT, unsigned Depth=0)
static bool scevUnconditionallyPropagatesPoisonFromOperands(SCEVTypes Kind)
static cl::opt< bool > VerifySCEVStrict("verify-scev-strict", cl::Hidden, cl::desc("Enable stricter verification with -verify-scev is passed"))
static Constant * getOtherIncomingValue(PHINode *PN, BasicBlock *BB)
static std::optional< BinaryOp > MatchBinaryOp(Value *V, const DataLayout &DL, AssumptionCache &AC, const DominatorTree &DT, const Instruction *CtxI)
Try to map V into a BinaryOp, and return std::nullopt on failure.
static cl::opt< bool > UseExpensiveRangeSharpening("scalar-evolution-use-expensive-range-sharpening", cl::Hidden, cl::init(false), cl::desc("Use more powerful methods of sharpening expression ranges. May " "be costly in terms of compile time"))
static const SCEV * getUnsignedOverflowLimitForStep(const SCEV *Step, ICmpInst::Predicate *Pred, ScalarEvolution *SE)
static bool IsKnownPredicateViaMinOrMax(ScalarEvolution &SE, CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS)
Is LHS Pred RHS true on the virtue of LHS or RHS being a Min or Max expression?
static bool BrPHIToSelect(DominatorTree &DT, CondBrInst *BI, PHINode *Merge, Value *&C, Value *&LHS, Value *&RHS)
This file defines the scope_exit class, which executes user-defined cleanup logic at scope exit.
static bool InBlock(const Value *V, const BasicBlock *BB)
Provides some synthesis utilities to produce sequences of values.
This file defines the SmallPtrSet class.
This file defines the SmallVector class.
This file defines the 'Statistic' class, which is designed to be an easy way to expose various metric...
#define STATISTIC(VARNAME, DESC)
Definition Statistic.h:171
This file contains some functions that are useful when dealing with strings.
#define LLVM_DEBUG(...)
Definition Debug.h:119
static TableGen::Emitter::Opt Y("gen-skeleton-entry", EmitSkeleton, "Generate example skeleton entry")
static SymbolRef::Type getType(const Symbol *Sym)
Definition TapiFile.cpp:39
LocallyHashedType DenseMapInfo< LocallyHashedType >::Empty
static std::optional< bool > isImpliedCondOperands(CmpInst::Predicate Pred, const Value *ALHS, const Value *ARHS, const Value *BLHS, const Value *BRHS)
Return true if "icmp Pred BLHS BRHS" is true whenever "icmp PredALHS ARHS" is true.
Virtual Register Rewriter
Value * RHS
Value * LHS
BinaryOperator * Mul
static const uint32_t IV[8]
Definition blake3_impl.h:83
SCEVCastSinkingRewriter(ScalarEvolution &SE, Type *TargetTy, ConversionFn CreatePtrCast)
static const SCEV * rewrite(const SCEV *Scev, ScalarEvolution &SE, Type *TargetTy, ConversionFn CreatePtrCast)
const SCEV * visitUnknown(const SCEVUnknown *Expr)
const SCEV * visitAddExpr(const SCEVAddExpr *Expr)
const SCEV * visit(const SCEV *S)
Class for arbitrary precision integers.
Definition APInt.h:78
LLVM_ABI APInt umul_ov(const APInt &RHS, bool &Overflow) const
Definition APInt.cpp:2009
LLVM_ABI APInt zext(unsigned width) const
Zero extend to a new width.
Definition APInt.cpp:1057
bool isMinSignedValue() const
Determine if this is the smallest signed value.
Definition APInt.h:419
uint64_t getZExtValue() const
Get zero extended value.
Definition APInt.h:1560
unsigned getActiveBits() const
Compute the number of active bits in the value.
Definition APInt.h:1532
LLVM_ABI APInt trunc(unsigned width) const
Truncate to new width.
Definition APInt.cpp:970
static APInt getMaxValue(unsigned numBits)
Gets maximum unsigned value of APInt for specific bit width.
Definition APInt.h:202
APInt abs() const
Get the absolute value.
Definition APInt.h:1815
bool sgt(const APInt &RHS) const
Signed greater than comparison.
Definition APInt.h:1205
bool isAllOnes() const
Determine if all bits are set. This is true for zero-width values.
Definition APInt.h:367
bool ugt(const APInt &RHS) const
Unsigned greater than comparison.
Definition APInt.h:1186
bool isZero() const
Determine if this value is zero, i.e. all bits are clear.
Definition APInt.h:376
bool isSignMask() const
Check if the APInt's value is returned by getSignMask.
Definition APInt.h:462
LLVM_ABI APInt urem(const APInt &RHS) const
Unsigned remainder operation.
Definition APInt.cpp:1695
unsigned getBitWidth() const
Return the number of bits in the APInt.
Definition APInt.h:1508
bool ult(const APInt &RHS) const
Unsigned less than comparison.
Definition APInt.h:1115
static APInt getSignedMaxValue(unsigned numBits)
Gets maximum signed value of APInt for a specific bit width.
Definition APInt.h:205
static APInt getMinValue(unsigned numBits)
Gets minimum unsigned value of APInt for a specific bit width.
Definition APInt.h:212
bool isNegative() const
Determine sign of this APInt.
Definition APInt.h:325
bool sle(const APInt &RHS) const
Signed less or equal comparison.
Definition APInt.h:1170
LLVM_ABI APInt uadd_ov(const APInt &RHS, bool &Overflow) const
Definition APInt.cpp:1973
static APInt getSignedMinValue(unsigned numBits)
Gets minimum signed value of APInt for a specific bit width.
Definition APInt.h:215
bool isNonPositive() const
Determine if this APInt Value is non-positive (<= 0).
Definition APInt.h:357
unsigned countTrailingZeros() const
Definition APInt.h:1667
bool isStrictlyPositive() const
Determine if this APInt Value is positive.
Definition APInt.h:352
unsigned logBase2() const
Definition APInt.h:1781
uint64_t getLimitedValue(uint64_t Limit=UINT64_MAX) const
If this value is smaller than the specified limit, return it, otherwise return the limit value.
Definition APInt.h:471
APInt ashr(unsigned ShiftAmt) const
Arithmetic right-shift function.
Definition APInt.h:829
LLVM_ABI APInt multiplicativeInverse() const
Definition APInt.cpp:1303
bool ule(const APInt &RHS) const
Unsigned less or equal comparison.
Definition APInt.h:1154
LLVM_ABI APInt sext(unsigned width) const
Sign extend to a new width.
Definition APInt.cpp:1030
APInt shl(unsigned shiftAmt) const
Left-shift function.
Definition APInt.h:875
bool isPowerOf2() const
Check if this APInt's value is a power of two greater than zero.
Definition APInt.h:436
static APInt getLowBitsSet(unsigned numBits, unsigned loBitsSet)
Constructs an APInt value that has the bottom loBitsSet bits set.
Definition APInt.h:302
bool isSignBitSet() const
Determine if sign bit of this APInt is set.
Definition APInt.h:337
bool slt(const APInt &RHS) const
Signed less than comparison.
Definition APInt.h:1134
static APInt getZero(unsigned numBits)
Get the '0' value for the specified bit-width.
Definition APInt.h:196
bool isIntN(unsigned N) const
Check if this APInt has an N-bits unsigned integer value.
Definition APInt.h:428
bool sge(const APInt &RHS) const
Signed greater or equal comparison.
Definition APInt.h:1241
static APInt getOneBitSet(unsigned numBits, unsigned BitNo)
Return an APInt with exactly one bit set in the result.
Definition APInt.h:235
bool uge(const APInt &RHS) const
Unsigned greater or equal comparison.
Definition APInt.h:1225
This templated class represents "all analyses that operate over <aparticular IR unit>" (e....
Definition Analysis.h:50
PassT::Result & getResult(IRUnitT &IR, ExtraArgTs... ExtraArgs)
Get the result of an analysis pass for a given IR unit.
Represent the analysis usage information of a pass.
void setPreservesAll()
Set by analyses that do not transform their input at all.
AnalysisUsage & addRequiredTransitive()
Represent a constant reference to an array (0 or more elements consecutively in memory),...
Definition ArrayRef.h:40
iterator end() const
Definition ArrayRef.h:130
size_t size() const
Get the array size.
Definition ArrayRef.h:141
iterator begin() const
Definition ArrayRef.h:129
A function analysis which provides an AssumptionCache.
An immutable pass that tracks lazily created AssumptionCache objects.
A cache of @llvm.assume calls within a function.
MutableArrayRef< WeakVH > assumptions()
Access the list of assumption handles currently tracked for this function.
LLVM Basic Block Representation.
Definition BasicBlock.h:62
iterator begin()
Instruction iterator methods.
Definition BasicBlock.h:446
const Function * getParent() const
Return the enclosing method, or null if none.
Definition BasicBlock.h:213
LLVM_ABI const BasicBlock * getSinglePredecessor() const
Return the predecessor of this block if it has a single predecessor block.
const Instruction & front() const
Definition BasicBlock.h:469
const Instruction * getTerminator() const LLVM_READONLY
Returns the terminator instruction; assumes that the block is well-formed.
Definition BasicBlock.h:237
LLVM_ABI unsigned getNoWrapKind() const
Returns one of OBO::NoSignedWrap or OBO::NoUnsignedWrap.
LLVM_ABI Instruction::BinaryOps getBinaryOp() const
Returns the binary operation underlying the intrinsic.
BinaryOps getOpcode() const
Definition InstrTypes.h:409
This class represents a function call, abstracting a target machine's calling convention.
virtual void deleted()
Callback for Value destruction.
void setValPtr(Value *P)
This is the base class for all instructions that perform data casts.
Definition InstrTypes.h:512
This class is the base class for the comparison instructions.
Definition InstrTypes.h:728
bool isFalseWhenEqual() const
This is just a convenience.
Predicate
This enumeration lists the possible predicates for CmpInst subclasses.
Definition InstrTypes.h:740
@ ICMP_SLT
signed less than
Definition InstrTypes.h:769
@ ICMP_SLE
signed less or equal
Definition InstrTypes.h:770
@ ICMP_UGE
unsigned greater or equal
Definition InstrTypes.h:764
@ ICMP_UGT
unsigned greater than
Definition InstrTypes.h:763
@ ICMP_SGT
signed greater than
Definition InstrTypes.h:767
@ ICMP_ULT
unsigned less than
Definition InstrTypes.h:765
@ ICMP_NE
not equal
Definition InstrTypes.h:762
@ ICMP_SGE
signed greater or equal
Definition InstrTypes.h:768
@ ICMP_ULE
unsigned less or equal
Definition InstrTypes.h:766
bool isSigned() const
Definition InstrTypes.h:993
Predicate getSwappedPredicate() const
For example, EQ->EQ, SLE->SGE, ULT->UGT, OEQ->OEQ, ULE->UGE, OLT->OGT, etc.
Definition InstrTypes.h:890
bool isTrueWhenEqual() const
This is just a convenience.
Predicate getInversePredicate() const
For example, EQ -> NE, UGT -> ULE, SLT -> SGE, OEQ -> UNE, UGT -> OLE, OLT -> UGE,...
Definition InstrTypes.h:852
bool isUnsigned() const
Definition InstrTypes.h:999
bool isRelational() const
Return true if the predicate is relational (not EQ or NE).
Definition InstrTypes.h:989
An abstraction over a floating-point predicate, and a pack of an integer predicate with samesign info...
static LLVM_ABI std::optional< CmpPredicate > getMatching(CmpPredicate A, CmpPredicate B)
Compares two CmpPredicates taking samesign into account and returns the canonicalized CmpPredicate if...
LLVM_ABI CmpInst::Predicate getPreferredSignedPredicate() const
Attempts to return a signed CmpInst::Predicate from the CmpPredicate.
CmpInst::Predicate dropSameSign() const
Drops samesign information.
Conditional Branch instruction.
Value * getCondition() const
BasicBlock * getSuccessor(unsigned i) const
static LLVM_ABI Constant * getNot(Constant *C)
static Constant * getPtrAdd(Constant *Ptr, Constant *Offset, GEPNoWrapFlags NW=GEPNoWrapFlags::none(), std::optional< ConstantRange > InRange=std::nullopt, Type *OnlyIfReduced=nullptr)
Create a getelementptr i8, ptr, offset constant expression.
Definition Constants.h:1518
static LLVM_ABI Constant * getPtrToAddr(Constant *C, Type *Ty, bool OnlyIfReduced=false)
static LLVM_ABI Constant * getAdd(Constant *C1, Constant *C2, bool HasNUW=false, bool HasNSW=false)
static LLVM_ABI Constant * getNeg(Constant *C, bool HasNSW=false)
static LLVM_ABI Constant * getTrunc(Constant *C, Type *Ty, bool OnlyIfReduced=false)
This is the shared class of boolean and integer constants.
Definition Constants.h:87
bool isZero() const
This is just a convenience method to make client code smaller for a common code.
Definition Constants.h:219
static LLVM_ABI ConstantInt * getFalse(LLVMContext &Context)
uint64_t getZExtValue() const
Return the constant as a 64-bit unsigned integer value after it has been zero extended as appropriate...
Definition Constants.h:168
const APInt & getValue() const
Return the constant as an APInt value reference.
Definition Constants.h:159
static LLVM_ABI ConstantInt * getBool(LLVMContext &Context, bool V)
This class represents a range of values.
LLVM_ABI ConstantRange add(const ConstantRange &Other) const
Return a new range representing the possible values resulting from an addition of a value in this ran...
LLVM_ABI ConstantRange zextOrTrunc(uint32_t BitWidth) const
Make this range have the bit width given by BitWidth.
PreferredRangeType
If represented precisely, the result of some range operations may consist of multiple disjoint ranges...
LLVM_ABI bool getEquivalentICmp(CmpInst::Predicate &Pred, APInt &RHS) const
Set up Pred and RHS such that ConstantRange::makeExactICmpRegion(Pred, RHS) == *this.
const APInt & getLower() const
Return the lower value for this range.
LLVM_ABI ConstantRange urem(const ConstantRange &Other) const
Return a new range representing the possible values resulting from an unsigned remainder operation of...
LLVM_ABI bool isFullSet() const
Return true if this set contains all of the elements possible for this data-type.
LLVM_ABI bool icmp(CmpInst::Predicate Pred, const ConstantRange &Other) const
Does the predicate Pred hold between ranges this and Other?
LLVM_ABI bool isEmptySet() const
Return true if this set contains no members.
LLVM_ABI ConstantRange zeroExtend(uint32_t BitWidth) const
Return a new range in the specified integer type, which must be strictly larger than the current type...
LLVM_ABI bool isSignWrappedSet() const
Return true if this set wraps around the signed domain.
LLVM_ABI APInt getSignedMin() const
Return the smallest signed value contained in the ConstantRange.
LLVM_ABI bool isWrappedSet() const
Return true if this set wraps around the unsigned domain.
LLVM_ABI void print(raw_ostream &OS) const
Print out the bounds to a stream.
LLVM_ABI ConstantRange truncate(uint32_t BitWidth, unsigned NoWrapKind=0) const
Return a new range in the specified integer type, which must be strictly smaller than the current typ...
LLVM_ABI ConstantRange signExtend(uint32_t BitWidth) const
Return a new range in the specified integer type, which must be strictly larger than the current type...
const APInt & getUpper() const
Return the upper value for this range.
LLVM_ABI ConstantRange unionWith(const ConstantRange &CR, PreferredRangeType Type=Smallest) const
Return the range that results from the union of this range with another range.
static LLVM_ABI ConstantRange makeExactICmpRegion(CmpInst::Predicate Pred, const APInt &Other)
Produce the exact range such that all values in the returned range satisfy the given predicate with a...
LLVM_ABI bool contains(const APInt &Val) const
Return true if the specified value is in the set.
LLVM_ABI ConstantRange intersectWith(const ConstantRange &CR, PreferredRangeType Type=Smallest) const
Return the range that results from the intersection of this range with another range.
LLVM_ABI APInt getSignedMax() const
Return the largest signed value contained in the ConstantRange.
static ConstantRange getNonEmpty(APInt Lower, APInt Upper)
Create non-empty constant range with the given bounds.
static LLVM_ABI ConstantRange makeGuaranteedNoWrapRegion(Instruction::BinaryOps BinOp, const ConstantRange &Other, unsigned NoWrapKind)
Produce the largest range containing all X such that "X BinOp Y" is guaranteed not to wrap (overflow)...
LLVM_ABI unsigned getMinSignedBits() const
Compute the maximal number of bits needed to represent every value in this signed range.
uint32_t getBitWidth() const
Get the bit width of this ConstantRange.
LLVM_ABI ConstantRange sub(const ConstantRange &Other) const
Return a new range representing the possible values resulting from a subtraction of a value in this r...
LLVM_ABI ConstantRange sextOrTrunc(uint32_t BitWidth) const
Make this range have the bit width given by BitWidth.
static LLVM_ABI ConstantRange makeExactNoWrapRegion(Instruction::BinaryOps BinOp, const APInt &Other, unsigned NoWrapKind)
Produce the range that contains X if and only if "X BinOp Other" does not wrap.
This is an important base class in LLVM.
Definition Constant.h:43
A parsed version of the target data layout string in and methods for querying it.
Definition DataLayout.h:64
LLVM_ABI const StructLayout * getStructLayout(StructType *Ty) const
Returns a StructLayout object, indicating the alignment of the struct, its size, and the offsets of i...
LLVM_ABI unsigned getIndexTypeSizeInBits(Type *Ty) const
The size in bits of the index used in GEP calculation for this type.
LLVM_ABI IntegerType * getIndexType(LLVMContext &C, unsigned AddressSpace) const
Returns the type of a GEP index in AddressSpace.
TypeSize getTypeSizeInBits(Type *Ty) const
Size examples:
Definition DataLayout.h:791
bool contains(const_arg_type_t< KeyT > Val) const
Return true if the specified key is in the map, false otherwise.
Definition DenseMap.h:762
size_type count(const_arg_type_t< KeyT > Val) const
Return 1 if the specified key is in the map, 0 otherwise.
Definition DenseMap.h:767
bool empty() const
Definition DenseMap.h:721
iterator find(const_arg_type_t< KeyT > Val)
Definition DenseMap.h:771
iterator end()
Definition DenseMap.h:691
DenseMapIterator< KeyT, ValueT, KeyInfoT, BucketT > iterator
Definition DenseMap.h:683
ValueT lookup(const_arg_type_t< KeyT > Val) const
Return the entry for the specified key, or a default constructed value if no such entry exists.
Definition DenseMap.h:798
iterator find_as(const LookupKeyT &Val)
Alternate version of find() which allows a different, and possibly less expensive,...
Definition DenseMap.h:784
void swap(DenseMapBase &RHS)
Definition DenseMap.h:982
std::pair< iterator, bool > insert(const std::pair< KeyT, ValueT > &KV)
Definition DenseMap.h:832
std::pair< iterator, bool > try_emplace(KeyT &&Key, Ts &&...Args)
Definition DenseMap.h:861
Analysis pass which computes a DominatorTree.
Definition Dominators.h:241
Legacy analysis pass which computes a DominatorTree.
Definition Dominators.h:277
Concrete subclass of DominatorTreeBase that is used to compute a normal dominator tree.
Definition Dominators.h:122
LLVM_ABI bool isReachableFromEntry(const Use &U) const
Provide an overload for a Use.
LLVM_ABI bool dominates(const BasicBlock *BB, const Use &U) const
Return true if the (end of the) basic block BB dominates the use U.
This instruction extracts a single (scalar) element from a VectorType value.
This instruction extracts a struct member or array element value from an aggregate value.
Insertion token: a failed lookup fills it in, the matching insert consumes it.
Definition FoldingSet.h:284
This class describes a reference to an interned FoldingSetNodeID, which can be a useful to store node...
Definition FoldingSet.h:123
This class is used to gather all the unique data bits of a node.
Definition FoldingSet.h:162
void AddInteger(signed I)
Definition FoldingSet.h:190
This class represents a freeze function that returns random concrete value if an operand is either a ...
FunctionPass(char &pid)
Definition Pass.h:316
Represents flags for the getelementptr instruction/expression.
bool hasNoUnsignedSignedWrap() const
bool hasNoUnsignedWrap() const
static GEPNoWrapFlags none()
static LLVM_ABI Type * getTypeAtIndex(Type *Ty, Value *Idx)
Return the type of the element at the given index of an indexable type.
Module * getParent()
Get the module that this global value is contained inside of...
static bool isPrivateLinkage(LinkageTypes Linkage)
static bool isInternalLinkage(LinkageTypes Linkage)
This instruction compares its operands according to the predicate given to the constructor.
CmpPredicate getCmpPredicate() const
static bool isGE(Predicate P)
Return true if the predicate is SGE or UGE.
CmpPredicate getSwappedCmpPredicate() const
static LLVM_ABI bool compare(const APInt &LHS, const APInt &RHS, ICmpInst::Predicate Pred)
Return result of LHS Pred RHS comparison.
static bool isLT(Predicate P)
Return true if the predicate is SLT or ULT.
CmpPredicate getInverseCmpPredicate() const
Predicate getNonStrictCmpPredicate() const
For example, SGT -> SGE, SLT -> SLE, ULT -> ULE, UGT -> UGE.
static bool isGT(Predicate P)
Return true if the predicate is SGT or UGT.
Predicate getFlippedSignednessPredicate() const
For example, SLT->ULT, ULT->SLT, SLE->ULE, ULE->SLE, EQ->EQ.
static CmpPredicate getInverseCmpPredicate(CmpPredicate Pred)
bool isEquality() const
Return true if this predicate is either EQ or NE.
static bool isEquality(Predicate P)
Return true if this predicate is either EQ or NE.
bool isRelational() const
Return true if the predicate is relational (not EQ or NE).
static bool isLE(Predicate P)
Return true if the predicate is SLE or ULE.
This instruction inserts a single (scalar) element into a VectorType value.
This instruction inserts a struct field of array element value into an aggregate value.
LLVM_ABI bool hasNoUnsignedWrap() const LLVM_READONLY
Determine whether the no unsigned wrap flag is set.
LLVM_ABI bool hasNoSignedWrap() const LLVM_READONLY
Determine whether the no signed wrap flag is set.
LLVM_ABI bool isIdenticalToWhenDefined(const Instruction *I, bool IntersectAttrs=false) const LLVM_READONLY
This is like isIdenticalTo, except that it ignores the SubclassOptionalData flags,...
Class to represent integer types.
static LLVM_ABI IntegerType * get(LLVMContext &C, unsigned NumBits)
This static method is the primary way of constructing an IntegerType.
Definition Type.cpp:338
A helper class to return the specified delimiter string after the first invocation of operator String...
An instruction for reading from memory.
Analysis pass that exposes the LoopInfo for a function.
Definition LoopInfo.h:594
bool contains(const LoopT *L) const
Return true if the specified loop is contained within this loop.
BlockT * getHeader() const
unsigned getLoopDepth() const
Return the nesting level of this loop.
BlockT * getLoopPredecessor() const
If the given loop's header has exactly one unique predecessor outside the loop, return it.
LoopT * getParentLoop() const
Return the parent loop if it exists or nullptr for top level loops.
unsigned getLoopDepth(const BlockT *BB) const
Return the loop nesting level of the specified block.
LoopT * getLoopFor(const BlockT *BB) const
Return the inner most loop that BB lives in.
The legacy pass manager's analysis pass to compute loop information.
Definition LoopInfo.h:619
Represents a single loop in the control flow graph.
Definition LoopInfo.h:40
bool isLoopInvariant(const Value *V) const
Return true if the specified value is loop invariant.
Definition LoopInfo.cpp:67
Metadata node.
Definition Metadata.h:1081
A Module instance is used to store all the information related to an LLVM module.
Definition Module.h:68
unsigned getOpcode() const
Return the opcode for this Instruction or ConstantExpr.
Definition Operator.h:43
Utility class for integer operators which may exhibit overflow - Add, Sub, Mul, and Shl.
Definition Operator.h:78
bool hasNoSignedWrap() const
Test whether this operation is known to never undergo signed overflow, aka the nsw property.
Definition Operator.h:113
bool hasNoUnsignedWrap() const
Test whether this operation is known to never undergo unsigned overflow, aka the nuw property.
Definition Operator.h:107
iterator_range< const_block_iterator > blocks() const
op_range incoming_values()
Value * getIncomingValueForBlock(const BasicBlock *BB) const
BasicBlock * getIncomingBlock(unsigned i) const
Return incoming basic block number i.
Value * getIncomingValue(unsigned i) const
Return incoming value number x.
unsigned getNumIncomingValues() const
Return the number of incoming edges.
AnalysisType & getAnalysis() const
getAnalysis<AnalysisType>() - This function is used by subclasses to get to the analysis information ...
PointerIntPair - This class implements a pair of a pointer and small integer.
static LLVM_ABI PoisonValue * get(Type *T)
Static factory methods - Return an 'poison' object of the specified type.
LLVM_ABI void addPredicate(const SCEVPredicate &Pred)
Adds a new predicate.
LLVM_ABI const SCEVPredicate & getPredicate() const
LLVM_ABI const SCEV * getPredicatedSCEV(const SCEV *Expr)
Returns the rewritten SCEV for Expr in the context of the current SCEV predicate.
LLVM_ABI bool areAddRecsEqualWithPreds(const SCEVAddRecExpr *AR1, const SCEVAddRecExpr *AR2, ArrayRef< const SCEVPredicate * > ExtraPreds={}) const
Check if AR1 and AR2 are equal, while taking into account Equal predicates in Preds and ExtraPreds.
LLVM_ABI const SCEVAddRecExpr * getAsAddRec(Value *V, SmallVectorImpl< const SCEVPredicate * > *WrapPredsAdded=nullptr)
Attempts to produce an AddRecExpr for V by adding additional SCEV predicates.
LLVM_ABI void print(raw_ostream &OS, unsigned Depth) const
Print the SCEV mappings done by the Predicated Scalar Evolution.
LLVM_ABI PredicatedScalarEvolution(ScalarEvolution &SE, Loop &L)
LLVM_ABI unsigned getSmallConstantMaxTripCount()
Returns the upper bound of the loop trip count as a normal unsigned value, or 0 if the trip count is ...
LLVM_ABI void addPredicates(ArrayRef< const SCEVPredicate * > Preds)
Adds all predicates in Preds.
LLVM_ABI const SCEV * getBackedgeTakenCount()
Get the (predicated) backedge count for the analyzed loop.
LLVM_ABI const SCEV * getSymbolicMaxBackedgeTakenCount()
Get the (predicated) symbolic max backedge count for the analyzed loop.
LLVM_ABI const SCEV * getSCEV(Value *V)
Returns the SCEV expression of V, in the context of the current SCEV predicate.
A set of analyses that are preserved following a run of a transformation pass.
Definition Analysis.h:112
static PreservedAnalyses all()
Construct a special preserved set that preserves all passes.
Definition Analysis.h:118
PreservedAnalysisChecker getChecker() const
Build a checker for this PreservedAnalyses and the specified analysis type.
Definition Analysis.h:275
constexpr bool isValid() const
Definition Register.h:112
This node represents an addition of some number of SCEVs.
This node represents a polynomial recurrence on the trip count of the specified loop.
LLVM_ABI SCEVUse getExitValue(ScalarEvolution &SE) const
Return the value of this recurrences when its loop exits, i.e.
LLVM_ABI const SCEV * evaluateAtIteration(const SCEV *It, ScalarEvolution &SE) const
Return the value of this chain of recurrences at the specified iteration number.
void setNoWrapFlags(SCEVFlags Flags)
Set flags for a recurrence without clearing any previously set flags.
bool isAffine() const
Return true if this represents an expression A + B*x where A and B are loop invariant values.
bool isQuadratic() const
Return true if this represents an expression A + B*x + C*x^2 where A, B and C are loop invariant valu...
LLVM_ABI const SCEV * getNumIterationsInRange(const ConstantRange &Range, ScalarEvolution &SE) const
Return the number of iterations of this loop that produce values in the specified constant range.
LLVM_ABI const SCEVAddRecExpr * getPostIncExpr(ScalarEvolution &SE) const
Return an expression representing the value of this expression one iteration of the loop ahead.
SCEVUse getStepRecurrence(ScalarEvolution &SE) const
Constructs and returns the recurrence indicating how much this expression steps by.
This is the base class for unary cast operator classes.
LLVM_ABI SCEVCastExpr(const FoldingSetNodeIDRef ID, SCEVTypes SCEVTy, SCEVUse op, Type *ty)
This class represents an assumption that the expression LHS Pred RHS evaluates to true,...
SCEVComparePredicate(const FoldingSetNodeIDRef ID, const ICmpInst::Predicate Pred, const SCEV *LHS, const SCEV *RHS)
bool isAlwaysTrue() const override
Returns true if the predicate is always true.
void print(raw_ostream &OS, unsigned Depth=0) const override
Prints a textual representation of this predicate with an indentation of Depth.
bool implies(const SCEVPredicate *N, ScalarEvolution &SE) const override
Implementation of the SCEVPredicate interface.
This class represents a constant integer value.
ConstantInt * getValue() const
const APInt & getAPInt() const
This is the base class for unary integral cast operator classes.
LLVM_ABI SCEVIntegralCastExpr(const FoldingSetNodeIDRef ID, SCEVTypes SCEVTy, SCEVUse op, Type *ty)
This node is the base class min/max selections.
static enum SCEVTypes negate(enum SCEVTypes T)
This node represents multiplication of some number of SCEVs.
This node is a base class providing common functionality for n'ary operators.
ArrayRef< SCEVUse > operands() const
SCEVFlags getNoWrapFlags(SCEVFlags Mask=FlagsNoWrapMask) const
SCEVUse getOperand(unsigned i) const
This class represents an assumption made using SCEV expressions which can be checked at run-time.
SCEVPredicate(const SCEVPredicate &)=default
virtual bool implies(const SCEVPredicate *N, ScalarEvolution &SE) const =0
Returns true if this predicate implies N.
SCEVPredicateKind Kind
This class represents a cast from a pointer to a pointer-sized integer value, without capturing the p...
This visitor recursively visits a SCEV expression and re-writes it.
const SCEV * visitPtrToAddrExpr(const SCEVPtrToAddrExpr *Expr)
const SCEV * visitSignExtendExpr(const SCEVSignExtendExpr *Expr)
const SCEV * visitZeroExtendExpr(const SCEVZeroExtendExpr *Expr)
const SCEV * visitSMinExpr(const SCEVSMinExpr *Expr)
const SCEV * visitUMinExpr(const SCEVUMinExpr *Expr)
This class represents a signed minimum selection.
This node is the base class for sequential/in-order min/max selections.
static SCEVTypes getEquivalentNonSequentialSCEVType(SCEVTypes Ty)
This class represents a sign extension of a small integer value to a larger integer value.
Visit all nodes in the expression tree using worklist traversal.
This class represents a truncation of an integer value to a smaller integer value.
This class represents a binary unsigned division operation.
This class represents an unsigned minimum selection.
This class represents a composition of other SCEV predicates, and is the class that most clients will...
void print(raw_ostream &OS, unsigned Depth) const override
Prints a textual representation of this predicate with an indentation of Depth.
bool implies(const SCEVPredicate *N, ScalarEvolution &SE) const override
Returns true if this predicate implies N.
SCEVUnionPredicate(ArrayRef< const SCEVPredicate * > Preds, ScalarEvolution &SE)
Union predicates don't get cached so create a dummy set ID for it.
bool isAlwaysTrue() const override
Implementation of the SCEVPredicate interface.
SCEVUnionPredicate getUnionWith(const SCEVPredicate *N, ScalarEvolution &SE) const
Returns a new SCEVUnionPredicate that is the union of this predicate and the given predicate N.
This means that we are dealing with an entirely unknown SCEV value, and only represent it as its LLVM...
This class represents the value of vscale, as used when defining the length of a scalable vector or r...
This class represents an assumption made on an AddRec expression.
IncrementWrapFlags
Similar to SCEVFlags, but with slightly different semantics for FlagNUSW.
SCEVWrapPredicate(const FoldingSetNodeIDRef ID, const SCEVAddRecExpr *AR, IncrementWrapFlags Flags)
bool implies(const SCEVPredicate *N, ScalarEvolution &SE) const override
Returns true if this predicate implies N.
static SCEVWrapPredicate::IncrementWrapFlags setFlags(SCEVWrapPredicate::IncrementWrapFlags Flags, SCEVWrapPredicate::IncrementWrapFlags OnFlags)
void print(raw_ostream &OS, unsigned Depth=0) const override
Prints a textual representation of this predicate with an indentation of Depth.
bool isAlwaysTrue() const override
Returns true if the predicate is always true.
const SCEVAddRecExpr * getExpr() const
Implementation of the SCEVPredicate interface.
static SCEVWrapPredicate::IncrementWrapFlags clearFlags(SCEVWrapPredicate::IncrementWrapFlags Flags, SCEVWrapPredicate::IncrementWrapFlags OffFlags)
Convenient IncrementWrapFlags manipulation methods.
IncrementWrapFlags getFlags() const
Returns the set assumed no overflow flags.
This class represents a zero extension of a small integer value to a larger integer value.
This class represents an analyzed expression in the program.
unsigned short getExpressionSize() const
static constexpr auto FlagsNoWrapMask
LLVM_ABI bool isOne() const
Return true if the expression is a constant one.
SCEV(const FoldingSetNodeIDRef ID, SCEVTypes SCEVTy, unsigned short ExpressionSize, Type *Ty)
static constexpr auto FlagNUW
LLVM_ABI void computeAndSetCanonical(ScalarEvolution &SE)
Compute and set the canonical SCEV, by constructing a SCEV with the same operands,...
LLVM_ABI bool isZero() const
Return true if the expression is a constant zero.
const SCEV * CanonicalSCEV
Pointer to the canonical version of the SCEV, i.e.
LLVM_ABI void dump() const
This method is used for debugging.
LLVM_ABI bool isAllOnesValue() const
Return true if the expression is a constant all-ones value.
LLVM_ABI bool isNonConstantNegative() const
Return true if the specified scev is negated, but not a constant.
static constexpr auto FlagNSW
LLVM_ABI ArrayRef< SCEVUse > operands() const
Return operands of this SCEV expression.
Type * getType() const
Return the LLVM type of this SCEV expression.
static constexpr auto FlagNone
LLVM_ABI void print(raw_ostream &OS) const
Print out the internal representation of this scalar to the specified stream.
SCEVTypes getSCEVType() const
static constexpr auto FlagNW
Analysis pass that exposes the ScalarEvolution for a function.
LLVM_ABI ScalarEvolution run(Function &F, FunctionAnalysisManager &AM)
LLVM_ABI PreservedAnalyses run(Function &F, FunctionAnalysisManager &AM)
LLVM_ABI PreservedAnalyses run(Function &F, FunctionAnalysisManager &AM)
void getAnalysisUsage(AnalysisUsage &AU) const override
getAnalysisUsage - This function should be overriden by passes that need analysis information to do t...
void print(raw_ostream &OS, const Module *=nullptr) const override
print - Print out the internal state of the pass.
bool runOnFunction(Function &F) override
runOnFunction - Virtual method overriden by subclasses to do the per-function processing of the pass.
void releaseMemory() override
releaseMemory() - This member can be implemented by a pass if it wants to be able to release its memo...
void verifyAnalysis() const override
verifyAnalysis() - This member can be implemented by a analysis pass to check state of analysis infor...
static LLVM_ABI LoopGuards collect(const Loop *L, ScalarEvolution &SE)
Collect rewrite map for loop guards for loop L, together with flags indicating if NUW and NSW can be ...
LLVM_ABI const SCEV * rewrite(const SCEV *Expr) const
Try to apply the collected loop guards to Expr.
The main scalar evolution driver.
LLVM_ABI const SCEV * getUDivExpr(SCEVUse LHS, SCEVUse RHS)
Get a canonical unsigned division expression, or something simpler if possible.
const SCEV * getConstantMaxBackedgeTakenCount(const Loop *L)
When successful, this returns a SCEVConstant that is greater than or equal to (i.e.
const DataLayout & getDataLayout() const
Return the DataLayout associated with the module this SCEV instance is operating on.
LLVM_ABI bool isKnownNonNegative(const SCEV *S)
Test if the given expression is known to be non-negative.
LLVM_ABI bool isKnownOnEveryIteration(CmpPredicate Pred, const SCEVAddRecExpr *LHS, const SCEV *RHS)
Test if the condition described by Pred, LHS, RHS is known to be true on every iteration of the loop ...
static bool hasFlags(SCEVFlags Flags, SCEVFlags TestFlags)
LLVM_ABI std::optional< LoopInvariantPredicate > getLoopInvariantExitCondDuringFirstIterationsImpl(CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, const Loop *L, const Instruction *CtxI, const SCEV *MaxIter)
LLVM_ABI const SCEV * getZeroExtendExpr(SCEVUse Op, Type *Ty, unsigned Depth=0)
LLVM_ABI const SCEV * getUDivCeilSCEV(const SCEV *N, const SCEV *D)
Compute ceil(N / D).
LLVM_ABI std::optional< LoopInvariantPredicate > getLoopInvariantExitCondDuringFirstIterations(CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, const Loop *L, const Instruction *CtxI, const SCEV *MaxIter)
If the result of the predicate LHS Pred RHS is loop invariant with respect to L at given Context duri...
LLVM_ABI Type * getWiderType(Type *Ty1, Type *Ty2) const
LLVM_ABI const SCEV * getAbsExpr(const SCEV *Op, bool IsNSW)
LLVM_ABI bool isKnownNonPositive(const SCEV *S)
Test if the given expression is known to be non-positive.
LLVM_ABI const SCEV * getElementCount(Type *Ty, ElementCount EC, SCEVFlags Flags=SCEV::FlagNone)
LLVM_ABI bool isKnownNegative(const SCEV *S)
Test if the given expression is known to be negative.
LLVM_ABI const SCEV * getPredicatedConstantMaxBackedgeTakenCount(const Loop *L, SmallVectorImpl< const SCEVPredicate * > &Predicates)
Similar to getConstantMaxBackedgeTakenCount, except it will add a set of SCEV predicates to Predicate...
LLVM_ABI const SCEV * removePointerBase(const SCEV *S)
Compute an expression equivalent to S - getPointerBase(S).
LLVM_ABI bool isLoopEntryGuardedByCond(const Loop *L, CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS)
Test whether entry to the loop is protected by a conditional between LHS and RHS.
LLVM_ABI bool isKnownNonZero(const SCEV *S)
Test if the given expression is known to be non-zero.
LLVM_ABI const SCEV * getURemExpr(SCEVUse LHS, SCEVUse RHS)
Represents an unsigned remainder expression based on unsigned division.
LLVM_ABI const SCEV * getBackedgeTakenCount(const Loop *L, ExitCountKind Kind=Exact)
If the specified loop has a predictable backedge-taken count, return it, otherwise return a SCEVCould...
LLVM_ABI const SCEV * getSMinExpr(SCEVUse LHS, SCEVUse RHS)
LLVM_ABI const SCEV * getUMaxFromMismatchedTypes(const SCEV *LHS, const SCEV *RHS)
Promote the operands to the wider of the types using zero-extension, and then perform a umax operatio...
const SCEV * getZero(Type *Ty)
Return a SCEV for the constant 0 of a specific type.
LLVM_ABI bool willNotOverflow(Instruction::BinaryOps BinOp, bool Signed, const SCEV *LHS, const SCEV *RHS, const Instruction *CtxI=nullptr)
Is operation BinOp between LHS and RHS provably does not have a signed/unsigned overflow (Signed)?
LLVM_ABI const SCEV * getMinusSCEV(SCEVUse LHS, SCEVUse RHS, SCEVFlags Flags=SCEV::FlagNone, unsigned Depth=0)
Return LHS-RHS.
LLVM_ABI ExitLimit computeExitLimitFromCond(const Loop *L, Value *ExitCond, bool ExitIfTrue, bool ControlsOnlyExit, bool AllowPredicates=false)
Compute the number of times the backedge of the specified loop will execute if its exit condition wer...
LLVM_ABI const SCEV * getMinMaxExpr(SCEVTypes Kind, SmallVectorImpl< SCEVUse > &Operands)
LLVM_ABI const SCEVPredicate * getEqualPredicate(const SCEV *LHS, const SCEV *RHS)
LLVM_ABI unsigned getSmallConstantTripMultiple(const Loop *L, const SCEV *ExitCount)
Returns the largest constant divisor of the trip count as a normal unsigned value,...
LLVM_ABI SCEVUse getSCEVAtScope(const SCEV *S, const Loop *L)
Return a SCEV expression for the specified value at the specified scope in the program.
LLVM_ABI uint64_t getTypeSizeInBits(Type *Ty) const
Return the size in bits of the specified type, for which isSCEVable must return true.
LLVM_ABI void registerUser(const SCEV *User, ArrayRef< SCEVUse > Ops)
Notify this ScalarEvolution that User directly uses SCEVs in Ops.
LLVM_ABI const SCEV * getConstant(ConstantInt *V)
LLVM_ABI const SCEV * getPredicatedBackedgeTakenCount(const Loop *L, SmallVectorImpl< const SCEVPredicate * > &Predicates)
Similar to getBackedgeTakenCount, except it will add a set of SCEV predicates to Predicates that are ...
LLVM_ABI const SCEV * getSCEV(Value *V)
Return a SCEV expression for the full generality of the specified expression.
ConstantRange getSignedRange(const SCEV *S)
Determine the signed range for a particular SCEV.
LLVM_ABI const SCEV * getNoopOrSignExtend(const SCEV *V, Type *Ty)
Return a SCEV corresponding to a conversion of the input value to the specified type.
static SCEVFlags setFlags(SCEVFlags Flags, SCEVFlags OnFlags)
static LLVM_ABI bool isGuaranteedNotToBePoison(const SCEV *Op)
Returns true if Op is guaranteed to not be poison.
bool loopHasNoAbnormalExits(const Loop *L)
Return true if the loop has no abnormal exits.
LLVM_ABI const SCEV * getTripCountFromExitCount(const SCEV *ExitCount)
A version of getTripCountFromExitCount below which always picks an evaluation type which can not resu...
LLVM_ABI ScalarEvolution(Function &F, TargetLibraryInfo &TLI, AssumptionCache &AC, DominatorTree &DT, LoopInfo &LI)
const SCEV * getOne(Type *Ty)
Return a SCEV for the constant 1 of a specific type.
LLVM_ABI const SCEV * getTruncateOrNoop(const SCEV *V, Type *Ty)
Return a SCEV corresponding to a conversion of the input value to the specified type.
LLVM_ABI void forgetValues(ArrayRef< Value * > Values)
Batched forgetValue: invalidates all Values in one shared def-use walk, avoiding the redundant re-tra...
LLVM_ABI const SCEV * getSequentialMinMaxExpr(SCEVTypes Kind, SmallVectorImpl< SCEVUse > &Operands)
LLVM_ABI const SCEV * getCastExpr(SCEVTypes Kind, SCEVUse Op, Type *Ty)
LLVM_ABI std::optional< bool > evaluatePredicateAt(CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, const Instruction *CtxI)
Check whether the condition described by Pred, LHS, and RHS is true or false in the given Context.
LLVM_ABI SCEVUse getAddRecExpr(SCEVUse Start, SCEVUse Step, const Loop *L, SCEVFlagsPair Flags)
Get an add recurrence expression for the specified loop.
LLVM_ABI unsigned getSmallConstantMaxTripCount(const Loop *L, SmallVectorImpl< const SCEVPredicate * > *Predicates=nullptr)
Returns the upper bound of the loop trip count as a normal unsigned value.
LLVM_ABI bool isKnownMultipleOf(const SCEV *S, uint64_t M, SmallVectorImpl< const SCEVPredicate * > *Predicates=nullptr)
Check that S is a multiple of M.
LLVM_ABI bool isBackedgeTakenCountMaxOrZero(const Loop *L)
Return true if the backedge taken count is either the value returned by getConstantMaxBackedgeTakenCo...
LLVM_ABI void forgetLoop(const Loop *L)
This method should be called by the client when it has changed a loop in a way that may effect Scalar...
LLVM_ABI bool isLoopInvariant(const SCEV *S, const Loop *L)
Return true if the value of the given SCEV is unchanging in the specified loop.
LLVM_ABI bool isKnownPositive(const SCEV *S)
Test if the given expression is known to be positive.
LLVM_ABI bool SimplifyICmpOperands(CmpPredicate &Pred, SCEVUse &LHS, SCEVUse &RHS, unsigned Depth=0)
Simplify LHS and RHS in a comparison with predicate Pred.
APInt getUnsignedRangeMin(const SCEV *S)
Determine the min of the unsigned range for a particular SCEV.
static SCEVFlags clearFlags(SCEVFlags Flags, SCEVFlags OffFlags)
LLVM_ABI const SCEV * getOffsetOfExpr(Type *IntTy, StructType *STy, unsigned FieldNo)
Return an expression for offsetof on the given field with type IntTy.
LLVM_ABI LoopDisposition getLoopDisposition(const SCEV *S, const Loop *L)
Return the "disposition" of the given SCEV with respect to the given loop.
static SCEVFlags maskFlags(SCEVFlags Flags, SCEVFlags Mask)
Convenient SCEVFlags manipulation.
LLVM_ABI bool containsAddRecurrence(const SCEV *S)
Return true if the SCEV is a scAddRecExpr or it contains scAddRecExpr.
LLVM_ABI const SCEV * getTruncateExpr(SCEVUse Op, Type *Ty, unsigned Depth=0)
LLVM_ABI SCEVUse getAddExpr(SmallVectorImpl< SCEVUse > &Ops, SCEVFlagsPair Flags={}, unsigned Depth=0)
Get a canonical add expression, or something simpler if possible.
LLVM_ABI const SCEV * getZeroExtendExprImpl(SCEVUse Op, Type *Ty, unsigned Depth=0)
LLVM_ABI bool isSCEVable(Type *Ty) const
Test if values of the given type are analyzable within the SCEV framework.
LLVM_ABI Type * getEffectiveSCEVType(Type *Ty) const
Return a type with the same bitwidth as the given type and which represents how SCEV will treat the g...
LLVM_ABI const SCEVPredicate * getComparePredicate(ICmpInst::Predicate Pred, const SCEV *LHS, const SCEV *RHS)
LLVM_ABI bool haveSameSign(const SCEV *S1, const SCEV *S2)
Return true if we know that S1 and S2 must have the same sign.
LLVM_ABI const SCEV * getNotSCEV(const SCEV *V)
Return the SCEV object corresponding to ~V.
LLVM_ABI bool instructionCouldExistWithOperands(const SCEV *A, const SCEV *B)
Return true if there exists a point in the program at which both A and B could be operands to the sam...
LLVM_ABI std::optional< SCEVFlags > getStrengthenedNoWrapFlagsFromBinOp(const OverflowingBinaryOperator *OBO)
Parse NSW/NUW flags from add/sub/mul IR binary operation Op into SCEV no-wrap flags,...
ConstantRange getUnsignedRange(const SCEV *S)
Determine the unsigned range for a particular SCEV.
LLVM_ABI void print(raw_ostream &OS) const
LLVM_ABI const SCEV * getAnyExtendExpr(SCEVUse Op, Type *Ty)
getAnyExtendExpr - Return a SCEV for the given operand extended with unspecified bits out to the give...
LLVM_ABI const SCEV * getPredicatedExitCount(const Loop *L, const BasicBlock *ExitingBlock, SmallVectorImpl< const SCEVPredicate * > *Predicates, ExitCountKind Kind=Exact)
Same as above except this uses the predicated backedge taken info and may require predicates.
LLVM_ABI void forgetTopmostLoop(const Loop *L)
LLVM_ABI void forgetValue(Value *V)
This method should be called by the client when it has changed a value in a way that may effect its v...
APInt getSignedRangeMin(const SCEV *S)
Determine the min of the signed range for a particular SCEV.
LLVM_ABI bool isLoopUniform(const SCEV *S, const Loop *L)
Returns true if the given SCEV is loop-uniform with respect to the specified loop L.
LLVM_ABI const SCEV * getNoopOrAnyExtend(const SCEV *V, Type *Ty)
Return a SCEV corresponding to a conversion of the input value to the specified type.
LLVM_ABI void forgetBlockAndLoopDispositions(Value *V=nullptr)
Called when the client has changed the disposition of values in a loop or block.
LLVM_ABI const SCEV * getSignExtendExpr(SCEVUse Op, Type *Ty, unsigned Depth=0)
LLVM_ABI const SCEV * getUMaxExpr(SCEVUse LHS, SCEVUse RHS)
LLVM_ABI std::optional< LoopInvariantPredicate > getLoopInvariantPredicate(CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, const Loop *L, const Instruction *CtxI=nullptr)
If the result of the predicate LHS Pred RHS is loop invariant with respect to L, return a LoopInvaria...
LLVM_ABI const SCEV * getStoreSizeOfExpr(Type *IntTy, Type *StoreTy)
Return an expression for the store size of StoreTy that is type IntTy.
LLVM_ABI const SCEVPredicate * getWrapPredicate(const SCEVAddRecExpr *AR, SCEVWrapPredicate::IncrementWrapFlags AddedFlags)
LLVM_ABI bool isLoopBackedgeGuardedByCond(const Loop *L, CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS)
Test whether the backedge of the loop is protected by a conditional between LHS and RHS.
LLVM_ABI APInt getNonZeroConstantMultiple(const SCEV *S)
const SCEV * getMinusOne(Type *Ty)
Return a SCEV for the constant -1 of a specific type.
LLVM_ABI bool hasLoopInvariantBackedgeTakenCount(const Loop *L)
Return true if the specified loop has an analyzable loop-invariant backedge-taken count.
LLVM_ABI BlockDisposition getBlockDisposition(const SCEV *S, const BasicBlock *BB)
Return the "disposition" of the given SCEV with respect to the given block.
LLVM_ABI const SCEV * getNoopOrZeroExtend(const SCEV *V, Type *Ty)
Return a SCEV corresponding to a conversion of the input value to the specified type.
LLVM_ABI bool invalidate(Function &F, const PreservedAnalyses &PA, FunctionAnalysisManager::Invalidator &Inv)
LLVM_ABI const SCEV * getUMinFromMismatchedTypes(const SCEV *LHS, const SCEV *RHS, bool Sequential=false)
Promote the operands to the wider of the types using zero-extension, and then perform a umin operatio...
LLVM_ABI bool loopIsFiniteByAssumption(const Loop *L)
Return true if this loop is finite by assumption.
LLVM_ABI SCEVUse getSCEVAtExit(const SCEV *S, const Loop *L, const BasicBlock *ExitingBlock)
Return the SCEV expression at the specified loop exit.
LLVM_ABI const SCEV * getExistingSCEV(Value *V)
Return an existing SCEV for V if there is one, otherwise return nullptr.
LLVM_ABI APInt getConstantMultiple(const SCEV *S, const Instruction *CtxI=nullptr)
Returns the max constant multiple of S.
LoopDisposition
An enum describing the relationship between a SCEV and a loop.
@ LoopComputable
The SCEV varies predictably with the loop.
@ LoopVariant
The SCEV is loop-variant (unknown).
@ LoopInvariant
The SCEV is loop-invariant.
@ LoopUniform
The SCEV is loop-uniform.
LLVM_ABI bool isKnownToBeAPowerOfTwo(const SCEV *S, bool OrZero=false, bool OrNegative=false)
Test if the given expression is known to be a power of 2.
LLVM_ABI void forgetLcssaPhiWithNewPredecessor(Loop *L, PHINode *V)
Forget LCSSA phi node V of loop L to which a new predecessor was added, such that it may no longer be...
LLVM_ABI bool containsUndefs(const SCEV *S) const
Return true if the SCEV expression contains an undef value.
LLVM_ABI std::optional< MonotonicPredicateType > getMonotonicPredicateType(const SCEVAddRecExpr *LHS, ICmpInst::Predicate Pred)
If, for all loop invariant X, the predicate "LHS `Pred` X" is monotonically increasing or decreasing,...
LLVM_ABI const SCEV * getCouldNotCompute()
LLVM_ABI bool isAvailableAtLoopEntry(const SCEV *S, const Loop *L)
Determine if the SCEV can be evaluated at loop's entry.
LLVM_ABI uint32_t getMinTrailingZeros(const SCEV *S, const Instruction *CtxI=nullptr)
Determine the minimum number of zero bits that S is guaranteed to end in (at every loop iteration).
BlockDisposition
An enum describing the relationship between a SCEV and a basic block.
@ DominatesBlock
The SCEV dominates the block.
@ ProperlyDominatesBlock
The SCEV properly dominates the block.
@ DoesNotDominateBlock
The SCEV does not dominate the block.
LLVM_ABI const SCEV * getExitCount(const Loop *L, const BasicBlock *ExitingBlock, ExitCountKind Kind=Exact)
Return the number of times the backedge executes before the given exit would be taken; if not exactly...
LLVM_ABI void getPoisonGeneratingValues(SmallPtrSetImpl< const Value * > &Result, const SCEV *S)
Return the set of Values that, if poison, will definitively result in S being poison as well.
LLVM_ABI void setNoWrapFlags(SCEVAddRecExpr *AddRec, SCEVFlags Flags)
Update no-wrap flags of an AddRec.
LLVM_ABI void forgetLoopDispositions()
Called when the client has changed the disposition of values in this loop.
LLVM_ABI const SCEV * getVScale(Type *Ty)
LLVM_ABI SCEVUse getMulExpr(SmallVectorImpl< SCEVUse > &Ops, SCEVFlagsPair Flags={}, unsigned Depth=0)
Get a canonical multiply expression, or something simpler if possible.
LLVM_ABI unsigned getSmallConstantTripCount(const Loop *L)
Returns the exact trip count of the loop if we can compute it, and the result is a small constant.
LLVM_ABI bool hasComputableLoopEvolution(const SCEV *S, const Loop *L)
Return true if the given SCEV changes value in a known way in the specified loop.
LLVM_ABI const SCEV * getPointerBase(const SCEV *V)
Transitively follow the chain of pointer-type operands until reaching a SCEV that does not have a sin...
LLVM_ABI void forgetAllLoops()
LLVM_ABI const SCEV * getSignExtendExprImpl(SCEVUse Op, Type *Ty, unsigned Depth=0)
LLVM_ABI bool dominates(const SCEV *S, const BasicBlock *BB)
Return true if elements that makes up the given SCEV dominate the specified basic block.
APInt getUnsignedRangeMax(const SCEV *S)
Determine the max of the unsigned range for a particular SCEV.
ExitCountKind
The terms "backedge taken count" and "exit count" are used interchangeably to refer to the number of ...
@ SymbolicMaximum
An expression which provides an upper bound on the exact trip count.
@ ConstantMaximum
A constant which provides an upper bound on the exact trip count.
@ Exact
An expression exactly describing the number of times the backedge has executed when a loop is exited.
LLVM_ABI bool isKnownPredicate(CmpPredicate Pred, SCEVUse LHS, SCEVUse RHS)
Test if the given expression is known to satisfy the condition described by Pred, LHS,...
LLVM_ABI const SCEV * applyLoopGuards(const SCEV *Expr, const Loop *L)
Try to apply information from loop guards for L to Expr.
LLVM_ABI const SCEV * getPtrToAddrExpr(const SCEV *Op)
LLVM_ABI const SCEVAddRecExpr * convertSCEVToAddRecWithPredicates(const SCEV *S, const Loop *L, SmallVectorImpl< const SCEVPredicate * > &Preds)
Tries to convert the S expression to an AddRec expression, adding additional predicates to Preds as r...
LLVM_ABI const SCEV * getSMaxExpr(SCEVUse LHS, SCEVUse RHS)
LLVM_ABI const SCEV * getElementSize(Instruction *Inst)
Return the size of an element read or written by Inst.
LLVM_ABI const SCEV * getSizeOfExpr(Type *IntTy, TypeSize Size)
Return an expression for a TypeSize.
LLVM_ABI std::optional< bool > evaluatePredicate(CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS)
Check whether the condition described by Pred, LHS, and RHS is true or false.
LLVM_ABI const SCEV * getUnknown(Value *V)
LLVM_ABI std::optional< std::pair< const SCEV *, SmallVector< const SCEVPredicate *, 3 > > > createAddRecFromPHIWithCasts(const SCEVUnknown *SymbolicPHI)
Checks if SymbolicPHI can be rewritten as an AddRecExpr under some Predicates.
LLVM_ABI const SCEV * getTruncateOrZeroExtend(const SCEV *V, Type *Ty, unsigned Depth=0)
Return a SCEV corresponding to a conversion of the input value to the specified type.
LLVM_ABI bool isKnownViaInduction(CmpPredicate Pred, SCEVUse LHS, SCEVUse RHS)
We'd like to check the predicate on every iteration of the most dominated loop between loops used in ...
LLVM_ABI std::optional< APInt > computeConstantDifference(const SCEV *LHS, const SCEV *RHS)
Compute LHS - RHS and returns the result as an APInt if it is a constant, and std::nullopt if it isn'...
LLVM_ABI bool properlyDominates(const SCEV *S, const BasicBlock *BB)
Return true if elements that makes up the given SCEV properly dominate the specified basic block.
LLVM_ABI const SCEV * getNegativeSCEV(const SCEV *V, SCEVFlags Flags=SCEV::FlagNone)
Return the SCEV object corresponding to -V.
LLVM_ABI const SCEV * getUDivExactExpr(SCEVUse LHS, SCEVUse RHS)
Get a canonical unsigned division expression, or something simpler if possible.
LLVM_ABI const SCEV * rewriteUsingPredicate(const SCEV *S, const Loop *L, const SCEVPredicate &A)
Re-writes the SCEV according to the Predicates in A.
LLVM_ABI std::pair< const SCEV *, const SCEV * > SplitIntoInitAndPostInc(const Loop *L, const SCEV *S)
Splits SCEV expression S into two SCEVs.
LLVM_ABI bool canReuseInstruction(const SCEV *S, Instruction *I, SmallVectorImpl< Instruction * > &DropPoisonGeneratingInsts)
Check whether it is poison-safe to represent the expression S using the instruction I.
LLVM_ABI bool isKnownPredicateAt(CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, const Instruction *CtxI)
Test if the given expression is known to satisfy the condition described by Pred, LHS,...
LLVM_ABI const SCEV * getPredicatedSymbolicMaxBackedgeTakenCount(const Loop *L, SmallVectorImpl< const SCEVPredicate * > &Predicates)
Similar to getSymbolicMaxBackedgeTakenCount, except it will add a set of SCEV predicates to Predicate...
LLVM_ABI const SCEV * getGEPExpr(GEPOperator *GEP, ArrayRef< SCEVUse > IndexExprs)
Returns an expression for a GEP.
LLVM_ABI const SCEV * getUMinExpr(SCEVUse LHS, SCEVUse RHS, bool Sequential=false)
LLVM_ABI bool isBasicBlockEntryGuardedByCond(const BasicBlock *BB, CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS)
Test whether entry to the basic block is protected by a conditional between LHS and RHS.
LLVM_ABI const SCEV * getTruncateOrSignExtend(const SCEV *V, Type *Ty, unsigned Depth=0)
Return a SCEV corresponding to a conversion of the input value to the specified type.
LLVM_ABI bool containsErasedValue(const SCEV *S) const
Return true if the SCEV expression contains a Value that has been optimised out and is now a nullptr.
const SCEV * getSymbolicMaxBackedgeTakenCount(const Loop *L)
When successful, this returns a SCEV that is greater than or equal to (i.e.
APInt getSignedRangeMax(const SCEV *S)
Determine the max of the signed range for a particular SCEV.
LLVM_ABI void verify() const
LLVMContext & getContext() const
This class represents the LLVM 'select' instruction.
Implements a dense probed hash-table based set with some number of buckets stored inline.
Definition DenseSet.h:293
size_type size() const
A templated base class for SmallPtrSet which provides the typesafe interface that is common across al...
std::pair< iterator, bool > insert(PtrType Ptr)
Inserts Ptr if and only if there is no element in the container equal to Ptr.
bool contains(ConstPtrType Ptr) const
SmallPtrSet - This class implements a set which is optimized for holding SmallSize or less elements.
This class consists of common code factored out of the SmallVector class to reduce code duplication b...
reference emplace_back(ArgTypes &&... Args)
void reserve(size_type N)
iterator erase(const_iterator CI)
void append(ItTy in_start, ItTy in_end)
Add the specified range to the end of the SmallVector.
iterator insert(iterator I, T &&Elt)
void push_back(const T &Elt)
This is a 'vector' (really, a variable-sized array), optimized for the case when the array is small.
Used to lazily calculate structure layout information for a target machine, based on the DataLayout s...
Definition DataLayout.h:743
TypeSize getElementOffset(unsigned Idx) const
Definition DataLayout.h:774
TypeSize getSizeInBits() const
Definition DataLayout.h:754
Class to represent struct types.
Analysis pass providing the TargetLibraryInfo.
Provides information about what library functions are available for the current target.
The instances of the Type class are immutable: once they are created, they are never changed.
Definition Type.h:46
static LLVM_ABI IntegerType * getInt32Ty(LLVMContext &C)
Definition Type.cpp:299
bool isPointerTy() const
True if this is an instance of PointerType.
Definition Type.h:277
LLVM_ABI TypeSize getPrimitiveSizeInBits() const LLVM_READONLY
Return the basic size of this type if it is a primitive type.
Definition Type.cpp:187
static LLVM_ABI IntegerType * getInt1Ty(LLVMContext &C)
Definition Type.cpp:296
bool isIntegerTy() const
True if this is an instance of IntegerType.
Definition Type.h:252
static LLVM_ABI IntegerType * getIntNTy(LLVMContext &C, unsigned N)
Definition Type.cpp:303
A Use represents the edge between a Value definition and its users.
Definition Use.h:35
op_range operands()
Definition User.h:267
Use & Op()
Definition User.h:171
Value * getOperand(unsigned i) const
Definition User.h:207
LLVM Value Representation.
Definition Value.h:75
Type * getType() const
All values are typed, get the type of this value.
Definition Value.h:257
LLVMContext & getContext() const
All values hold a context through their type.
Definition Value.h:260
iterator_range< user_iterator > users()
Definition Value.h:428
unsigned getValueID() const
Return an ID for the concrete type of this object.
Definition Value.h:545
LLVM_ABI void printAsOperand(raw_ostream &O, bool PrintType=true, const Module *M=nullptr) const
Print the name of this Value out to the specified raw_ostream.
LLVM_ABI StringRef getName() const
Return a constant reference to the value's name.
Definition Value.cpp:319
constexpr bool isScalable() const
Returns whether the quantity is scaled by a runtime quantity (vscale).
Definition TypeSize.h:168
An efficient, type-erasing, non-owning reference to a callable.
const ParentTy * getParent() const
Definition ilist_node.h:34
This class implements an extremely fast bulk output stream that can only output to a stream.
Definition raw_ostream.h:53
raw_ostream & indent(unsigned NumSpaces)
indent - Insert 'NumSpaces' spaces.
Changed
#define llvm_unreachable(msg)
Marks that the current location is not supposed to be reachable.
constexpr char Align[]
Key for Kernel::Arg::Metadata::mAlign.
const APInt & smin(const APInt &A, const APInt &B)
Determine the smaller of two APInts considered to be signed.
Definition APInt.h:2274
const APInt & smax(const APInt &A, const APInt &B)
Determine the larger of two APInts considered to be signed.
Definition APInt.h:2279
const APInt & umin(const APInt &A, const APInt &B)
Determine the smaller of two APInts considered to be unsigned.
Definition APInt.h:2284
LLVM_ABI std::optional< APInt > SolveQuadraticEquationWrap(APInt A, APInt B, APInt C, unsigned RangeWidth)
Let q(n) = An^2 + Bn + C, and BW = bit width of the value range (e.g.
Definition APInt.cpp:2850
LLVM_ABI APInt GreatestCommonDivisor(APInt A, APInt B, bool IsSigned=false)
Compute GCD of two APInt values.
Definition APInt.cpp:826
const APInt & umax(const APInt &A, const APInt &B)
Determine the larger of two APInts considered to be unsigned.
Definition APInt.h:2289
constexpr bool any(E Val)
@ Entry
Definition COFF.h:862
@ BasicBlock
Various leaf nodes.
Definition ISDOpcodes.h:83
LLVM_ABI Function * getDeclarationIfExists(const Module *M, ID id)
Look up the Function declaration of the intrinsic id in the Module M and return it if it exists.
Predicate
Predicate - These are "(BI << 5) | BO" for various predicates.
match_combine_or< Ty... > m_CombineOr(const Ty &...Ps)
Combine pattern matchers matching any of Ps patterns.
BinaryOp_match< LHS, RHS, Instruction::AShr > m_AShr(const LHS &L, const RHS &R)
ap_match< APInt > m_APInt(const APInt *&Res)
Match a ConstantInt or splatted ConstantVector, binding the specified pointer to the contained APInt.
bool match(Val *V, const Pattern &P)
ThreeOps_match< Cond, LHS, RHS, Instruction::Select > m_Select(const Cond &C, const LHS &L, const RHS &R)
Matches SelectInst.
auto m_BasicBlock()
Match an arbitrary basic block value and ignore it.
ExtractValue_match< Ind, Val_t > m_ExtractValue(const Val_t &V)
Match a single index ExtractValue instruction.
auto m_Value()
Match an arbitrary value and ignore it.
auto m_LogicalOr()
Matches L || R where L and R are arbitrary values.
match_bind< WithOverflowInst > m_WithOverflowInst(WithOverflowInst *&I)
Match a with overflow intrinsic, capturing it if we match.
auto m_Intrinsic(const Ts &...Ops)
Match intrinsic calls like this: m_Intrinsic<Intrinsic::fabs>(m_Value(X))
BinaryOp_match< LHS, RHS, Instruction::SDiv > m_SDiv(const LHS &L, const RHS &R)
BinaryOp_match< LHS, RHS, Instruction::LShr > m_LShr(const LHS &L, const RHS &R)
BinaryOp_match< LHS, RHS, Instruction::Shl > m_Shl(const LHS &L, const RHS &R)
auto m_LogicalAnd()
Matches L && R where L and R are arbitrary values.
brc_match< Cond_t, match_bind< BasicBlock >, match_bind< BasicBlock > > m_Br(const Cond_t &C, BasicBlock *&T, BasicBlock *&F)
CastOperator_match< OpTy, Instruction::PtrToInt > m_PtrToInt(const OpTy &Op)
Matches PtrToInt.
auto m_ConstantInt()
Match an arbitrary ConstantInt and ignore it.
bind_cst_ty m_scev_APInt(const APInt *&C)
Match an SCEV constant and bind it to an APInt.
cst_pred_ty< is_all_ones > m_scev_AllOnes()
Match an integer with all bits set.
SCEVUnaryExpr_match< SCEVZeroExtendExpr, Op0_t > m_scev_ZExt(const Op0_t &Op0)
is_undef_or_poison m_scev_UndefOrPoison()
Match an SCEVUnknown wrapping undef or poison.
cst_pred_ty< is_one > m_scev_One()
Match an integer 1.
specificloop_ty m_SpecificLoop(const Loop *L)
SCEVUnaryExpr_match< SCEVSignExtendExpr, Op0_t > m_scev_SExt(const Op0_t &Op0)
match_bind< const SCEVMulExpr > m_scev_Mul(const SCEVMulExpr *&V)
cst_pred_ty< is_zero > m_scev_Zero()
Match an integer 0.
SCEVUnaryExpr_match< SCEVTruncateExpr, Op0_t > m_scev_Trunc(const Op0_t &Op0)
bool match(const SCEV *S, const Pattern &P)
SCEVBinaryExpr_match< SCEVUDivExpr, Op0_t, Op1_t > m_scev_UDiv(const Op0_t &Op0, const Op1_t &Op1)
specificscev_ty m_scev_Specific(const SCEV *S)
Match if we have a specific specified SCEV.
SCEVAffineAddRec_match< Op0_t, Op1_t, match_isa< const Loop > > m_scev_AffineAddRec(const Op0_t &Op0, const Op1_t &Op1)
match_bind< const SCEVUnknown > m_SCEVUnknown(const SCEVUnknown *&V)
SCEVBinaryExpr_match< SCEVMulExpr, Op0_t, Op1_t, SCEV::FlagNUW, true > m_scev_c_NUWMul(const Op0_t &Op0, const Op1_t &Op1)
SCEVBinaryExpr_match< SCEVSMaxExpr, Op0_t, Op1_t, SCEV::FlagNone, true > m_scev_SMax(const Op0_t &Op0, const Op1_t &Op1)
SCEVBinaryExpr_match< SCEVMulExpr, Op0_t, Op1_t, SCEV::FlagNone, true > m_scev_c_Mul(const Op0_t &Op0, const Op1_t &Op1)
match_bind< const SCEVAddExpr > m_scev_Add(const SCEVAddExpr *&V)
SCEVURem_match< Op0_t, Op1_t > m_scev_URem(Op0_t LHS, Op1_t RHS, ScalarEvolution &SE)
Match the mathematical pattern A - (A / B) * B, where A and B can be arbitrary expressions.
@ Valid
The data is already valid.
initializer< Ty > init(const Ty &Val)
LocationClass< Ty > location(Ty &L)
@ Switch
The "resume-switch" lowering, where there are separate resume and destroy functions that are shared b...
Definition CoroShape.h:32
constexpr double e
NodeAddr< PhiNode * > Phi
Definition RDFGraph.h:390
friend class Instruction
Iterator for Instructions in a `BasicBlock.
Definition BasicBlock.h:73
unsigned getOpcode(const VPValue *V)
Return the instruction opcode for the recipe defining V or 0 for unsupported recipes and VPValues not...
This is an optimization pass for GlobalISel generic memory operations.
void visitAll(const SCEV *Root, SV &Visitor)
Use SCEVTraversal to visit all nodes in the given expression tree.
auto drop_begin(T &&RangeOrContainer, size_t N=1)
Return a range covering RangeOrContainer with the first N elements excluded.
Definition STLExtras.h:316
@ Offset
Definition DWP.cpp:577
void stable_sort(R &&Range)
Definition STLExtras.h:2132
bool all_of(R &&range, UnaryPredicate P)
Provide wrappers to std::all_of which take ranges instead of having to pass begin/end explicitly.
Definition STLExtras.h:1755
SaveAndRestore(T &) -> SaveAndRestore< T >
Printable print(const GCNRegPressure &RP, const GCNSubtarget *ST=nullptr, unsigned DynamicVGPRBlockSize=0)
LLVM_ABI bool canCreatePoison(const Operator *Op, bool ConsiderFlagsAndMetadata=true)
LLVM_ABI bool mustTriggerUB(const Instruction *I, const SmallPtrSetImpl< const Value * > &KnownPoison)
Return true if the given instruction must trigger undefined behavior when I is executed with any oper...
RelativeUniformCounterPtr Values
Definition InstrProf.h:91
@ Known
Known to have no common set bits.
@ Dead
Unused definition.
InterleavedRange< Range > interleaved(const Range &R, StringRef Separator=", ", StringRef Prefix="", StringRef Suffix="")
Output range R as a sequence of interleaved elements.
decltype(auto) dyn_cast(const From &Val)
dyn_cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:643
LLVM_ABI bool verifyFunction(const Function &F, raw_ostream *OS=nullptr)
Check a function for errors, useful for use when debugging a pass.
auto successors(const MachineBasicBlock *BB)
scope_exit(Callable) -> scope_exit< Callable >
const Value * getLoadStorePointerOperand(const Value *V)
A helper function that returns the pointer operand of a load or store instruction.
@ BinaryOp
One of the operands is a binary op.
constexpr from_range_t from_range
auto dyn_cast_if_present(const Y &Val)
dyn_cast_if_present<X> - Functionally identical to dyn_cast, except that a null (or none in the case ...
Definition Casting.h:732
bool set_is_subset(const S1Ty &S1, const S2Ty &S2)
set_is_subset(A, B) - Return true iff A in B
void append_range(Container &C, Range &&R)
Wrapper function to append range R to container C.
Definition STLExtras.h:2224
constexpr bool isUIntN(unsigned N, uint64_t x)
Checks if an unsigned integer fits into the given (dynamic) bit width.
Definition MathExtras.h:244
LLVM_ABI void computeKnownBits(const Value *V, KnownBits &Known, const DataLayout &DL, AssumptionCache *AC=nullptr, const Instruction *CtxI=nullptr, const DominatorTree *DT=nullptr, bool UseInstrInfo=true, unsigned Depth=0)
Determine which bits of V are known to be either zero or one and return them in the KnownZero/KnownOn...
void * PointerTy
LLVM_ABI bool VerifySCEV
auto uninitialized_copy(R &&Src, IterTy Dst)
Definition STLExtras.h:2127
bool isa_and_nonnull(const Y &Val)
Definition Casting.h:676
LLVM_ABI unsigned ComputeNumSignBits(const Value *Op, const DataLayout &DL, AssumptionCache *AC=nullptr, const Instruction *CtxI=nullptr, const DominatorTree *DT=nullptr, bool UseInstrInfo=true, unsigned Depth=0)
Return the number of times the sign bit of the register is replicated into the other bits.
LLVM_ABI ConstantRange getConstantRangeFromMetadata(const MDNode &RangeMD)
Parse out a conservative ConstantRange from !range metadata.
LLVM_ABI bool canConstantFoldCallTo(const CallBase *Call, const Function *F, const TargetLibraryInfo *TLI=nullptr)
canConstantFoldCallTo - Return true if its even possible to fold a call to the specified function.
RelativeUniformCounterPtr ValuesPtrExpr VTableAddr Value
Definition InstrProf.h:143
int countr_zero(T Val)
Count number of 0's from the least significant bit to the most stopping at the first 1.
Definition bit.h:204
LLVM_ABI Value * simplifyInstruction(Instruction *I, const SimplifyQuery &Q)
See if we can compute a simplified version of this instruction.
LLVM_ABI bool isOverflowIntrinsicNoWrap(const WithOverflowInst *WO, const DominatorTree &DT)
Returns true if the arithmetic part of the WO 's result is used only along the paths control dependen...
DomTreeNodeBase< BasicBlock > DomTreeNode
Definition Dominators.h:65
LLVM_ABI bool matchSimpleRecurrence(const PHINode *P, BinaryOperator *&BO, Value *&Start, Value *&Step)
Attempt to match a simple first order recurrence cycle of the form: iv = phi Ty [Start,...
LLVM_ABI Constant * ConstantFoldCompareInstOperands(unsigned Predicate, Constant *LHS, Constant *RHS, const DataLayout &DL, const TargetLibraryInfo *TLI=nullptr, const Function *CtxF=nullptr)
Attempt to constant fold a compare instruction (icmp/fcmp) with the specified operands.
auto dyn_cast_or_null(const Y &Val)
Definition Casting.h:753
void erase(Container &C, ValueType V)
Wrapper function to remove a value from a container:
Definition STLExtras.h:2216
bool any_of(R &&range, UnaryPredicate P)
Provide wrappers to std::any_of which take ranges instead of having to pass begin/end explicitly.
Definition STLExtras.h:1762
auto reverse(ContainerTy &&C)
Definition STLExtras.h:408
LLVM_ABI bool isMustProgress(const Loop *L)
Return true if this loop can be assumed to make progress.
LLVM_ABI bool impliesPoison(const Value *ValAssumedPoison, const Value *V)
Return true if V is poison given that ValAssumedPoison is already poison.
LLVM_ABI bool isFinite(const Loop *L)
Return true if this loop can be assumed to run for a finite number of iterations.
unsigned short computeExpressionSize(ArrayRef< SCEVUse > Args)
LLVM_ABI bool programUndefinedIfPoison(const Instruction *Inst)
LLVM_ABI raw_ostream & dbgs()
dbgs() - This returns a reference to a raw_ostream for debugging messages.
Definition Debug.cpp:209
bool isPointerTy(const Type *T)
Definition SPIRVUtils.h:383
LLVM_ABI ConstantRange getVScaleRange(const Function *F, unsigned BitWidth)
Determine the possible constant range of vscale with the given bit width, based on the vscale_range f...
class LLVM_GSL_OWNER SmallVector
Forward declaration of SmallVector so that calculateSmallVectorDefaultInlinedElements can reference s...
bool isa(const From &Val)
isa<X> - Return true if the parameter to the template is an instance of one of the template type argu...
Definition Casting.h:547
LLVM_ATTRIBUTE_VISIBILITY_DEFAULT AnalysisKey InnerAnalysisManagerProxy< AnalysisManagerT, IRUnitT, ExtraArgTs... >::Key
LLVM_ABI bool isKnownNonZero(const Value *V, const SimplifyQuery &Q, unsigned Depth=0)
Return true if the given value is known to be non-zero when defined.
constexpr T divideCeil(U Numerator, V Denominator)
Returns the integer ceil(Numerator / Denominator).
Definition MathExtras.h:389
LLVM_ABI bool propagatesPoison(const Use &PoisonOp)
Return true if PoisonOp's user yields poison or raises UB if its operand PoisonOp is poison.
@ UMin
Unsigned integer min implemented in terms of select(cmp()).
@ Mul
Product of integers.
@ SMax
Signed integer max implemented in terms of select(cmp()).
@ SMin
Signed integer min implemented in terms of select(cmp()).
@ Add
Sum of integers.
@ UMax
Unsigned integer max implemented in terms of select(cmp()).
RelativeUniformCounterPtr ValuesPtrExpr VTableAddr Count
Definition InstrProf.h:145
auto count(R &&Range, const E &Element)
Wrapper function around std::count to count the number of times an element Element occurs in the give...
Definition STLExtras.h:2028
DWARFExpression::Operation Op
auto max_element(R &&Range)
Provide wrappers to std::max_element which take ranges instead of having to pass begin/end explicitly...
Definition STLExtras.h:2104
raw_ostream & operator<<(raw_ostream &OS, const APFixedPoint &FX)
SCEVFlags
SCEVFlags are bitfield indices into SCEV's SubclassData.
ArrayRef(const T &OneElt) -> ArrayRef< T >
constexpr unsigned BitWidth
OutputIt move(R &&Range, OutputIt Out)
Provide wrappers to std::move which take ranges instead of having to pass begin/end explicitly.
Definition STLExtras.h:1933
LLVM_ABI bool isGuaranteedToTransferExecutionToSuccessor(const Instruction *I)
Return true if this function can prove that the instruction I will always transfer execution to one o...
auto count_if(R &&Range, UnaryPredicate P)
Wrapper function around std::count_if to count the number of times an element satisfying a given pred...
Definition STLExtras.h:2035
decltype(auto) cast(const From &Val)
cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:559
constexpr auto seq(T Begin, T End)
Iterate over an integral type from Begin up to - but not including - End.
Definition Sequence.h:341
void erase_if(Container &C, UnaryPredicate P)
Provide a container algorithm similar to C++ Library Fundamentals v2's erase_if which is equivalent t...
Definition STLExtras.h:2208
constexpr bool isIntN(unsigned N, int64_t x)
Checks if an signed integer fits into the given (dynamic) bit width.
Definition MathExtras.h:249
auto predecessors(const MachineBasicBlock *BB)
bool is_contained(R &&Range, const E &Element)
Returns true if Element is found in Range.
Definition STLExtras.h:1963
Type * getLoadStoreType(const Value *I)
A helper function that returns the type of a load or store instruction.
iterator_range< df_iterator< T > > depth_first(const T &G)
AnalysisManager< Function > FunctionAnalysisManager
Convenience typedef for the Function analysis manager.
bool equal(L &&LRange, R &&RRange)
Wrapper function around std::equal to detect if pair-wise elements between two ranges are the same.
Definition STLExtras.h:2162
LLVM_ABI bool isGuaranteedNotToBePoison(const Value *V, AssumptionCache *AC=nullptr, const Instruction *CtxI=nullptr, const DominatorTree *DT=nullptr, unsigned Depth=0)
Returns true if V cannot be poison, but may be undef.
LLVM_ABI Constant * ConstantFoldInstOperands(const Instruction *I, ArrayRef< Constant * > Ops, const DataLayout &DL, const TargetLibraryInfo *TLI=nullptr, bool AllowNonDeterministic=true)
ConstantFoldInstOperands - Attempt to constant fold an instruction with the specified operands.
constexpr detail::IsaCheckPredicate< Types... > IsaPred
Function object wrapper for the llvm::isa type check.
Definition Casting.h:866
SCEVUseT< const SCEV * > SCEVUse
bool SCEVExprContains(const SCEV *Root, PredTy Pred)
Return true if any node in Root satisfies the predicate Pred.
Implement std::hash so that hash_code can be used in STL containers.
Definition BitVector.h:878
void swap(llvm::BitVector &LHS, llvm::BitVector &RHS)
Implement std::swap in terms of BitVector swap.
Definition BitVector.h:880
#define N
#define NC
Definition regutils.h:42
A special type used by analysis passes to provide an address that identifies that particular analysis...
Definition Analysis.h:29
static KnownBits makeConstant(const APInt &C)
Create known bits from a known constant.
Definition KnownBits.h:315
static LLVM_ABI KnownBits ashr(const KnownBits &LHS, const KnownBits &RHS, bool ShAmtNonZero=false, bool Exact=false)
Compute known bits for ashr(LHS, RHS).
static LLVM_ABI KnownBits lshr(const KnownBits &LHS, const KnownBits &RHS, bool ShAmtNonZero=false, bool Exact=false)
Compute known bits for lshr(LHS, RHS).
static LLVM_ABI KnownBits shl(const KnownBits &LHS, const KnownBits &RHS, bool NUW=false, bool NSW=false, bool ShAmtNonZero=false)
Compute known bits for shl(LHS, RHS).
An object of this class is returned by queries that could not be answered.
static LLVM_ABI bool classof(const SCEV *S)
Methods for support type inquiry through isa, cast, and dyn_cast:
The no-wrap flags to apply when creating a SCEV expression, to the expression and use respectively.
SCEVFlags ExprFlags
Flags applied directly to a SCEV expression, must be valid wherever the expression is valid.
SCEVFlags UseFlags
Flags only applied to a SCEVUse.
SCEVPtrT getPointer() const
This class defines a simple visitor class that may be used for various SCEV analysis purposes.
A utility class that uses RAII to save and restore the value of a variable.
Information about the number of loop iterations for which a loop exit's branch condition evaluates to...
LLVM_ABI ExitLimit(const SCEV *E)
Construct either an exact exit limit from a constant, or an unknown one from a SCEVCouldNotCompute.
SmallVector< const SCEVPredicate *, 4 > Predicates
A vector of predicate guards for this ExitLimit.