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 if (all_of(operands(), [](SCEVUse Op) { return Op.isCanonical(); })) {
271 CanonicalSCEV = this;
272 return;
273 }
274
276 map_range(operands(), [](SCEVUse Op) { return Op.getCanonical(); }));
277 // Rebuild the expression from the canonical operands, stripping use flags.
278 CanonicalSCEV = SE.getWithOperands(this, CanonOps);
279}
280
281//===----------------------------------------------------------------------===//
282// Implementation of the SCEV class.
283//
284
285#if !defined(NDEBUG) || defined(LLVM_ENABLE_DUMP)
287 print(dbgs());
288 dbgs() << '\n';
289}
290#endif
291
292void SCEV::print(raw_ostream &OS) const {
293 switch (getSCEVType()) {
294 case scConstant:
295 cast<SCEVConstant>(this)->getValue()->printAsOperand(OS, false);
296 return;
297 case scVScale:
298 OS << "vscale";
299 return;
300 case scPtrToAddr: {
301 const SCEVCastExpr *PtrCast = cast<SCEVCastExpr>(this);
302 SCEVUse Op = PtrCast->getOperand();
303 OS << "(ptrtoaddr " << *Op->getType() << " " << Op << " to "
304 << *PtrCast->getType() << ")";
305 return;
306 }
307 case scTruncate: {
308 const SCEVTruncateExpr *Trunc = cast<SCEVTruncateExpr>(this);
309 SCEVUse Op = Trunc->getOperand();
310 OS << "(trunc " << *Op->getType() << " " << Op << " to "
311 << *Trunc->getType() << ")";
312 return;
313 }
314 case scZeroExtend: {
316 SCEVUse Op = ZExt->getOperand();
317 OS << "(zext " << *Op->getType() << " " << Op << " to " << *ZExt->getType()
318 << ")";
319 return;
320 }
321 case scSignExtend: {
323 SCEVUse Op = SExt->getOperand();
324 OS << "(sext " << *Op->getType() << " " << Op << " to " << *SExt->getType()
325 << ")";
326 return;
327 }
328 case scAddRecExpr: {
329 const SCEVAddRecExpr *AR = cast<SCEVAddRecExpr>(this);
330 OS << "{" << AR->getOperand(0);
331 for (unsigned i = 1, e = AR->getNumOperands(); i != e; ++i)
332 OS << ",+," << AR->getOperand(i);
333 OS << "}<";
334 if (AR->hasNoUnsignedWrap())
335 OS << "nuw><";
336 if (AR->hasNoSignedWrap())
337 OS << "nsw><";
338 if (AR->hasNoSelfWrap() && !AR->hasNoUnsignedWrap() &&
339 !AR->hasNoSignedWrap())
340 OS << "nw><";
341 AR->getLoop()->getHeader()->printAsOperand(OS, /*PrintType=*/false);
342 OS << ">";
343 return;
344 }
345 case scAddExpr:
346 case scMulExpr:
347 case scUMaxExpr:
348 case scSMaxExpr:
349 case scUMinExpr:
350 case scSMinExpr:
352 const SCEVNAryExpr *NAry = cast<SCEVNAryExpr>(this);
353 const char *OpStr = nullptr;
354 switch (NAry->getSCEVType()) {
355 case scAddExpr: OpStr = " + "; break;
356 case scMulExpr: OpStr = " * "; break;
357 case scUMaxExpr: OpStr = " umax "; break;
358 case scSMaxExpr: OpStr = " smax "; break;
359 case scUMinExpr:
360 OpStr = " umin ";
361 break;
362 case scSMinExpr:
363 OpStr = " smin ";
364 break;
366 OpStr = " umin_seq ";
367 break;
368 default:
369 llvm_unreachable("There are no other nary expression types.");
370 }
371 OS << "(" << llvm::interleaved(NAry->operands(), OpStr) << ")";
372 switch (NAry->getSCEVType()) {
373 case scAddExpr:
374 case scMulExpr:
375 if (NAry->hasNoUnsignedWrap())
376 OS << "<nuw>";
377 if (NAry->hasNoSignedWrap())
378 OS << "<nsw>";
379 break;
380 default:
381 // Nothing to print for other nary expressions.
382 break;
383 }
384 return;
385 }
386 case scUDivExpr: {
387 const SCEVUDivExpr *UDiv = cast<SCEVUDivExpr>(this);
388 OS << "(" << UDiv->getLHS() << " /u " << UDiv->getRHS() << ")";
389 return;
390 }
391 case scUnknown:
392 cast<SCEVUnknown>(this)->getValue()->printAsOperand(OS, false);
393 return;
395 OS << "***COULDNOTCOMPUTE***";
396 return;
397 }
398 llvm_unreachable("Unknown SCEV kind!");
399}
400
402 switch (getSCEVType()) {
403 case scConstant:
404 case scVScale:
405 case scUnknown:
406 return {};
407 case scPtrToAddr:
408 case scTruncate:
409 case scZeroExtend:
410 case scSignExtend:
411 return cast<SCEVCastExpr>(this)->operands();
412 case scAddRecExpr:
413 case scAddExpr:
414 case scMulExpr:
415 case scUMaxExpr:
416 case scSMaxExpr:
417 case scUMinExpr:
418 case scSMinExpr:
420 return cast<SCEVNAryExpr>(this)->operands();
421 case scUDivExpr:
422 return cast<SCEVUDivExpr>(this)->operands();
424 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
425 }
426 llvm_unreachable("Unknown SCEV kind!");
427}
428
429bool SCEV::isZero() const { return match(this, m_scev_Zero()); }
430
431bool SCEV::isOne() const { return match(this, m_scev_One()); }
432
433bool SCEV::isAllOnesValue() const { return match(this, m_scev_AllOnes()); }
434
437 if (!Mul) return false;
438
439 // If there is a constant factor, it will be first.
440 const SCEVConstant *SC = dyn_cast<SCEVConstant>(Mul->getOperand(0));
441 if (!SC) return false;
442
443 // Return true if the value is negative, this matches things like (-42 * V).
444 return SC->getAPInt().isNegative();
445}
446
449
451 return S->getSCEVType() == scCouldNotCompute;
452}
453
455 auto &Entry = ConstantSCEVs[V];
456 if (Entry)
457 return Entry;
458
461 ID.AddPointer(V);
463 if (SCEVConstant *S =
464 static_cast<SCEVConstant *>(UniqueSCEVs.lookup(ID, Token)))
465 return Entry = S;
466 SCEVConstant *S =
467 new (SCEVAllocator) SCEVConstant(ID.Intern(SCEVAllocator), V);
468 UniqueSCEVs.insert(S, Token);
469 S->computeAndSetCanonical(*this);
470 return Entry = S;
471}
472
474 return getConstant(ConstantInt::get(getContext(), Val));
475}
476
477const SCEV *
480 // TODO: Avoid implicit trunc?
481 // See https://github.com/llvm/llvm-project/issues/112510.
482 return getConstant(
483 ConstantInt::get(ITy, V, isSigned, /*ImplicitTrunc=*/true));
484}
485
489 ID.AddPointer(Ty);
491 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
492 return S;
493 SCEV *S = new (SCEVAllocator) SCEVVScale(ID.Intern(SCEVAllocator), Ty);
494 UniqueSCEVs.insert(S, Token);
495 S->computeAndSetCanonical(*this);
496 return S;
497}
498
500 SCEVFlags Flags) {
501 const SCEV *Res = getConstant(Ty, EC.getKnownMinValue());
502 if (EC.isScalable())
503 Res = getMulExpr(Res, getVScale(Ty), Flags);
504 return Res;
505}
506
508 SCEVUse op, Type *ty)
509 : SCEV(ID, SCEVTy, computeExpressionSize(op), ty), Op(op) {}
510
511SCEVPtrToAddrExpr::SCEVPtrToAddrExpr(const FoldingSetNodeIDRef ID,
512 const SCEV *Op, Type *ITy)
513 : SCEVCastExpr(ID, scPtrToAddr, Op, ITy) {
514 assert(getOperand()->getType()->isPointerTy() && getType()->isIntegerTy() &&
515 "Must be a non-bit-width-changing pointer-to-integer cast!");
516}
517
522
523SCEVTruncateExpr::SCEVTruncateExpr(const FoldingSetNodeIDRef ID, SCEVUse op,
524 Type *ty)
526 assert(getOperand()->getType()->isIntOrPtrTy() && getType()->isIntOrPtrTy() &&
527 "Cannot truncate non-integer value!");
528}
529
530SCEVZeroExtendExpr::SCEVZeroExtendExpr(const FoldingSetNodeIDRef ID, SCEVUse op,
531 Type *ty)
533 assert(getOperand()->getType()->isIntOrPtrTy() && getType()->isIntOrPtrTy() &&
534 "Cannot zero extend non-integer value!");
535}
536
537SCEVSignExtendExpr::SCEVSignExtendExpr(const FoldingSetNodeIDRef ID, SCEVUse op,
538 Type *ty)
540 assert(getOperand()->getType()->isIntOrPtrTy() && getType()->isIntOrPtrTy() &&
541 "Cannot sign extend non-integer value!");
542}
543
545 // Clear this SCEVUnknown from various maps.
546 SE->forgetMemoizedResults({this});
547
548 // Remove this SCEVUnknown from the uniquing map.
549 SE->UniqueSCEVs.erase(this);
550
551 // Release the value.
552 setValPtr(nullptr);
553}
554
555void SCEVUnknown::allUsesReplacedWith(Value *New) {
556 // Clear this SCEVUnknown from various maps.
557 SE->forgetMemoizedResults({this});
558
559 // Remove this SCEVUnknown from the uniquing map.
560 SE->UniqueSCEVs.erase(this);
561
562 // Replace the value pointer in case someone is still using this SCEVUnknown.
563 setValPtr(New);
564}
565
566//===----------------------------------------------------------------------===//
567// SCEV Utilities
568//===----------------------------------------------------------------------===//
569
570/// Compare the two values \p LV and \p RV in terms of their "complexity" where
571/// "complexity" is a partial (and somewhat ad-hoc) relation used to order
572/// operands in SCEV expressions.
573static int CompareValueComplexity(const LoopInfo *const LI, Value *LV,
574 Value *RV, unsigned Depth) {
576 return 0;
577
578 // Order pointer values after integer values. This helps SCEVExpander form
579 // GEPs.
580 bool LIsPointer = LV->getType()->isPointerTy(),
581 RIsPointer = RV->getType()->isPointerTy();
582 if (LIsPointer != RIsPointer)
583 return (int)LIsPointer - (int)RIsPointer;
584
585 // Compare getValueID values.
586 unsigned LID = LV->getValueID(), RID = RV->getValueID();
587 if (LID != RID)
588 return (int)LID - (int)RID;
589
590 // Sort arguments by their position.
591 if (const auto *LA = dyn_cast<Argument>(LV)) {
592 const auto *RA = cast<Argument>(RV);
593 unsigned LArgNo = LA->getArgNo(), RArgNo = RA->getArgNo();
594 return (int)LArgNo - (int)RArgNo;
595 }
596
597 if (const auto *LGV = dyn_cast<GlobalValue>(LV)) {
598 const auto *RGV = cast<GlobalValue>(RV);
599
600 if (auto L = LGV->getLinkage() - RGV->getLinkage())
601 return L;
602
603 const auto IsGVNameSemantic = [&](const GlobalValue *GV) {
604 auto LT = GV->getLinkage();
605 return !(GlobalValue::isPrivateLinkage(LT) ||
607 };
608
609 // Use the names to distinguish the two values, but only if the
610 // names are semantically important.
611 if (IsGVNameSemantic(LGV) && IsGVNameSemantic(RGV))
612 return LGV->getName().compare(RGV->getName());
613 }
614
615 // For instructions, compare their loop depth, and their operand count. This
616 // is pretty loose.
617 if (const auto *LInst = dyn_cast<Instruction>(LV)) {
618 const auto *RInst = cast<Instruction>(RV);
619
620 // Compare loop depths.
621 const BasicBlock *LParent = LInst->getParent(),
622 *RParent = RInst->getParent();
623 if (LParent != RParent) {
624 unsigned LDepth = LI->getLoopDepth(LParent),
625 RDepth = LI->getLoopDepth(RParent);
626 if (LDepth != RDepth)
627 return (int)LDepth - (int)RDepth;
628 }
629
630 // Compare the number of operands.
631 unsigned LNumOps = LInst->getNumOperands(),
632 RNumOps = RInst->getNumOperands();
633 if (LNumOps != RNumOps)
634 return (int)LNumOps - (int)RNumOps;
635
636 for (unsigned Idx : seq(LNumOps)) {
637 int Result = CompareValueComplexity(LI, LInst->getOperand(Idx),
638 RInst->getOperand(Idx), Depth + 1);
639 if (Result != 0)
640 return Result;
641 }
642 }
643
644 return 0;
645}
646
647// Return negative, zero, or positive, if LHS is less than, equal to, or greater
648// than RHS, respectively. A three-way result allows recursive comparisons to be
649// more efficient.
650// If the max analysis depth was reached, return std::nullopt, assuming we do
651// not know if they are equivalent for sure.
652static std::optional<int>
653CompareSCEVComplexity(const LoopInfo *const LI, const SCEV *LHS,
654 const SCEV *RHS, DominatorTree &DT, unsigned Depth = 0) {
655 // Fast-path: SCEVs are uniqued so we can do a quick equality check.
656 if (LHS == RHS)
657 return 0;
658
659 // Primarily, sort the SCEVs by their getSCEVType().
660 SCEVTypes LType = LHS->getSCEVType(), RType = RHS->getSCEVType();
661 if (LType != RType)
662 return (int)LType - (int)RType;
663
665 return std::nullopt;
666
667 // Aside from the getSCEVType() ordering, the particular ordering
668 // isn't very important except that it's beneficial to be consistent,
669 // so that (a + b) and (b + a) don't end up as different expressions.
670 switch (LType) {
671 case scUnknown: {
672 const SCEVUnknown *LU = cast<SCEVUnknown>(LHS);
673 const SCEVUnknown *RU = cast<SCEVUnknown>(RHS);
674
675 int X =
676 CompareValueComplexity(LI, LU->getValue(), RU->getValue(), Depth + 1);
677 return X;
678 }
679
680 case scConstant: {
683
684 // Compare constant values.
685 const APInt &LA = LC->getAPInt();
686 const APInt &RA = RC->getAPInt();
687 unsigned LBitWidth = LA.getBitWidth(), RBitWidth = RA.getBitWidth();
688 if (LBitWidth != RBitWidth)
689 return (int)LBitWidth - (int)RBitWidth;
690 return LA.ult(RA) ? -1 : 1;
691 }
692
693 case scVScale: {
694 const auto *LTy = cast<IntegerType>(cast<SCEVVScale>(LHS)->getType());
695 const auto *RTy = cast<IntegerType>(cast<SCEVVScale>(RHS)->getType());
696 return LTy->getBitWidth() - RTy->getBitWidth();
697 }
698
699 case scAddRecExpr: {
702
703 // There is always a dominance between two recs that are used by one SCEV,
704 // so we can safely sort recs by loop header dominance. We require such
705 // order in getAddExpr.
706 const Loop *LLoop = LA->getLoop(), *RLoop = RA->getLoop();
707 if (LLoop != RLoop) {
708 const BasicBlock *LHead = LLoop->getHeader(), *RHead = RLoop->getHeader();
709 assert(LHead != RHead && "Two loops share the same header?");
710 if (DT.dominates(LHead, RHead))
711 return 1;
712 assert(DT.dominates(RHead, LHead) &&
713 "No dominance between recurrences used by one SCEV?");
714 return -1;
715 }
716
717 [[fallthrough]];
718 }
719
720 case scTruncate:
721 case scZeroExtend:
722 case scSignExtend:
723 case scPtrToAddr:
724 case scAddExpr:
725 case scMulExpr:
726 case scUDivExpr:
727 case scSMaxExpr:
728 case scUMaxExpr:
729 case scSMinExpr:
730 case scUMinExpr:
732 ArrayRef<SCEVUse> LOps = LHS->operands();
733 ArrayRef<SCEVUse> ROps = RHS->operands();
734
735 // Lexicographically compare n-ary-like expressions.
736 unsigned LNumOps = LOps.size(), RNumOps = ROps.size();
737 if (LNumOps != RNumOps)
738 return (int)LNumOps - (int)RNumOps;
739
740 for (unsigned i = 0; i != LNumOps; ++i) {
741 auto X = CompareSCEVComplexity(LI, LOps[i].getPointer(),
742 ROps[i].getPointer(), DT, Depth + 1);
743 if (X != 0)
744 return X;
745 }
746 return 0;
747 }
748
750 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
751 }
752 llvm_unreachable("Unknown SCEV kind!");
753}
754
755/// Given a list of SCEV objects, order them by their complexity, and group
756/// objects of the same complexity together by value. When this routine is
757/// finished, we know that any duplicates in the vector are consecutive and that
758/// complexity is monotonically increasing.
759///
760/// Note that we go take special precautions to ensure that we get deterministic
761/// results from this routine. In other words, we don't want the results of
762/// this to depend on where the addresses of various SCEV objects happened to
763/// land in memory.
765 DominatorTree &DT) {
766 if (Ops.size() < 2) return; // Noop
767
768 // Whether LHS has provably less complexity than RHS.
769 auto IsLessComplex = [&](SCEVUse LHS, SCEVUse RHS) {
770 auto Complexity = CompareSCEVComplexity(LI, LHS, RHS, DT);
771 return Complexity && *Complexity < 0;
772 };
773 if (Ops.size() == 2) {
774 // This is the common case, which also happens to be trivially simple.
775 // Special case it.
776 SCEVUse &LHS = Ops[0], &RHS = Ops[1];
777 if (IsLessComplex(RHS, LHS))
778 std::swap(LHS, RHS);
779 return;
780 }
781
782 // Do the rough sort by complexity.
784 Ops, [&](SCEVUse LHS, SCEVUse RHS) { return IsLessComplex(LHS, RHS); });
785
786 // Now that we are sorted by complexity, group elements of the same
787 // complexity. Note that this is, at worst, N^2, but the vector is likely to
788 // be extremely short in practice. Note that we take this approach because we
789 // do not want to depend on the addresses of the objects we are grouping.
790 for (unsigned i = 0, e = Ops.size(); i != e-2; ++i) {
791 const SCEV *S = Ops[i];
792 unsigned Complexity = S->getSCEVType();
793
794 // If there are any objects of the same complexity and same value as this
795 // one, group them.
796 for (unsigned j = i+1; j != e && Ops[j]->getSCEVType() == Complexity; ++j) {
797 if (Ops[j] == S) { // Found a duplicate.
798 // Move it to immediately after i'th element.
799 std::swap(Ops[i+1], Ops[j]);
800 ++i; // no need to rescan it.
801 if (i == e-2) return; // Done!
802 }
803 }
804 }
805}
806
807/// Returns true if \p Ops contains a huge SCEV (the subtree of S contains at
808/// least HugeExprThreshold nodes).
810 return any_of(Ops, [](const SCEV *S) {
812 });
813}
814
815/// Performs a number of common optimizations on the passed \p Ops. If the
816/// whole expression reduces down to a single operand, it will be returned.
817///
818/// The following optimizations are performed:
819/// * Fold constants using the \p Fold function.
820/// * Remove identity constants satisfying \p IsIdentity.
821/// * If a constant satisfies \p IsAbsorber, return it.
822/// * Sort operands by complexity.
823template <typename FoldT, typename IsIdentityT, typename IsAbsorberT>
824static const SCEV *
826 SmallVectorImpl<SCEVUse> &Ops, FoldT Fold,
827 IsIdentityT IsIdentity, IsAbsorberT IsAbsorber) {
828 const SCEVConstant *Folded = nullptr;
829 for (unsigned Idx = 0; Idx < Ops.size();) {
830 const SCEV *Op = Ops[Idx];
831 if (const auto *C = dyn_cast<SCEVConstant>(Op)) {
832 if (!Folded)
833 Folded = C;
834 else
835 Folded = cast<SCEVConstant>(
836 SE.getConstant(Fold(Folded->getAPInt(), C->getAPInt())));
837 Ops.erase(Ops.begin() + Idx);
838 continue;
839 }
840 ++Idx;
841 }
842
843 if (Ops.empty()) {
844 assert(Folded && "Must have folded value");
845 return Folded;
846 }
847
848 if (Folded && IsAbsorber(Folded->getAPInt()))
849 return Folded;
850
851 GroupByComplexity(Ops, &LI, DT);
852 if (Folded && !IsIdentity(Folded->getAPInt()))
853 Ops.insert(Ops.begin(), Folded);
854
855 return Ops.size() == 1 ? Ops[0] : nullptr;
856}
857
858//===----------------------------------------------------------------------===//
859// Simple SCEV method implementations
860//===----------------------------------------------------------------------===//
861
862/// Compute BC(It, K). The result has width W. Assume, K > 0.
863static const SCEV *BinomialCoefficient(const SCEV *It, unsigned K,
864 ScalarEvolution &SE,
865 Type *ResultTy) {
866 // Handle the simplest case efficiently.
867 if (K == 1)
868 return SE.getTruncateOrZeroExtend(It, ResultTy);
869
870 // We are using the following formula for BC(It, K):
871 //
872 // BC(It, K) = (It * (It - 1) * ... * (It - K + 1)) / K!
873 //
874 // Suppose, W is the bitwidth of the return value. We must be prepared for
875 // overflow. Hence, we must assure that the result of our computation is
876 // equal to the accurate one modulo 2^W. Unfortunately, division isn't
877 // safe in modular arithmetic.
878 //
879 // However, this code doesn't use exactly that formula; the formula it uses
880 // is something like the following, where T is the number of factors of 2 in
881 // K! (i.e. trailing zeros in the binary representation of K!), and ^ is
882 // exponentiation:
883 //
884 // BC(It, K) = (It * (It - 1) * ... * (It - K + 1)) / 2^T / (K! / 2^T)
885 //
886 // This formula is trivially equivalent to the previous formula. However,
887 // this formula can be implemented much more efficiently. The trick is that
888 // K! / 2^T is odd, and exact division by an odd number *is* safe in modular
889 // arithmetic. To do exact division in modular arithmetic, all we have
890 // to do is multiply by the inverse. Therefore, this step can be done at
891 // width W.
892 //
893 // The next issue is how to safely do the division by 2^T. The way this
894 // is done is by doing the multiplication step at a width of at least W + T
895 // bits. This way, the bottom W+T bits of the product are accurate. Then,
896 // when we perform the division by 2^T (which is equivalent to a right shift
897 // by T), the bottom W bits are accurate. Extra bits are okay; they'll get
898 // truncated out after the division by 2^T.
899 //
900 // In comparison to just directly using the first formula, this technique
901 // is much more efficient; using the first formula requires W * K bits,
902 // but this formula less than W + K bits. Also, the first formula requires
903 // a division step, whereas this formula only requires multiplies and shifts.
904 //
905 // It doesn't matter whether the subtraction step is done in the calculation
906 // width or the input iteration count's width; if the subtraction overflows,
907 // the result must be zero anyway. We prefer here to do it in the width of
908 // the induction variable because it helps a lot for certain cases; CodeGen
909 // isn't smart enough to ignore the overflow, which leads to much less
910 // efficient code if the width of the subtraction is wider than the native
911 // register width.
912 //
913 // (It's possible to not widen at all by pulling out factors of 2 before
914 // the multiplication; for example, K=2 can be calculated as
915 // It/2*(It+(It*INT_MIN/INT_MIN)+-1). However, it requires
916 // extra arithmetic, so it's not an obvious win, and it gets
917 // much more complicated for K > 3.)
918
919 // Protection from insane SCEVs; this bound is conservative,
920 // but it probably doesn't matter.
921 if (K > 1000)
922 return SE.getCouldNotCompute();
923
924 unsigned W = SE.getTypeSizeInBits(ResultTy);
925
926 // Calculate K! / 2^T and T; we divide out the factors of two before
927 // multiplying for calculating K! / 2^T to avoid overflow.
928 // Other overflow doesn't matter because we only care about the bottom
929 // W bits of the result.
930 APInt OddFactorial(W, 1);
931 unsigned T = 1;
932 for (unsigned i = 3; i <= K; ++i) {
933 unsigned TwoFactors = countr_zero(i);
934 T += TwoFactors;
935 OddFactorial *= (i >> TwoFactors);
936 }
937
938 // We need at least W + T bits for the multiplication step
939 unsigned CalculationBits = W + T;
940
941 // Calculate 2^T, at width T+W.
942 APInt DivFactor = APInt::getOneBitSet(CalculationBits, T);
943
944 // Calculate the multiplicative inverse of K! / 2^T;
945 // this multiplication factor will perform the exact division by
946 // K! / 2^T.
947 APInt MultiplyFactor = OddFactorial.multiplicativeInverse();
948
949 // Calculate the product, at width T+W
950 IntegerType *CalculationTy = IntegerType::get(SE.getContext(),
951 CalculationBits);
952 const SCEV *Dividend = SE.getTruncateOrZeroExtend(It, CalculationTy);
953 for (unsigned i = 1; i != K; ++i) {
954 const SCEV *S = SE.getMinusSCEV(It, SE.getConstant(It->getType(), i));
955 Dividend = SE.getMulExpr(Dividend,
956 SE.getTruncateOrZeroExtend(S, CalculationTy));
957 }
958
959 // Divide by 2^T
960 const SCEV *DivResult = SE.getUDivExpr(Dividend, SE.getConstant(DivFactor));
961
962 // Truncate the result, and divide by K! / 2^T.
963
964 return SE.getMulExpr(SE.getConstant(MultiplyFactor),
965 SE.getTruncateOrZeroExtend(DivResult, ResultTy));
966}
967
968/// Return the value of this chain of recurrences at the specified iteration
969/// number. We can evaluate this recurrence by multiplying each element in the
970/// chain by the binomial coefficient corresponding to it. In other words, we
971/// can evaluate {A,+,B,+,C,+,D} as:
972///
973/// A*BC(It, 0) + B*BC(It, 1) + C*BC(It, 2) + D*BC(It, 3)
974///
975/// where BC(It, k) stands for binomial coefficient.
977 ScalarEvolution &SE) const {
978 return evaluateAtIteration(operands(), It, SE);
979}
980
982 const SCEV *It, ScalarEvolution &SE,
983 SCEVFlags UseFlags) {
984 assert(Operands.size() > 0);
985 assert((Operands.size() == 2 || UseFlags == SCEV::FlagNone) &&
986 "use-specific flags only supported for affine AddRecs");
987 SCEVUse Result = Operands[0].getPointer();
988 for (unsigned i = 1, e = Operands.size(); i != e; ++i) {
989 // The computation is correct in the face of overflow provided that the
990 // multiplication is performed _after_ the evaluation of the binomial
991 // coefficient.
992 const SCEV *Coeff = BinomialCoefficient(It, i, SE, Result->getType());
993 if (isa<SCEVCouldNotCompute>(Coeff))
994 return Coeff;
995
996 SCEVUse Mul = SE.getMulExpr(Operands[i].getPointer(), Coeff,
997 {SCEV::FlagNone, UseFlags});
998 Result = SE.getAddExpr(Result, Mul, {SCEV::FlagNone, UseFlags});
999 }
1000 return Result;
1001}
1002
1004 const SCEV *BTC = SE.getBackedgeTakenCount(getLoop());
1005 if (isa<SCEVCouldNotCompute>(BTC))
1006 return BTC;
1007 // The loop reaches iteration BTC, so the value this recurrence computes there
1008 // is the value it had, and that did not wrap.
1009 return evaluateAtIteration(operands(), BTC, SE,
1011 : SCEV::FlagNone);
1012}
1013
1014//===----------------------------------------------------------------------===//
1015// SCEV Expression folder implementations
1016//===----------------------------------------------------------------------===//
1017
1018/// The SCEVCastSinkingRewriter takes a scalar evolution expression,
1019/// which computes a pointer-typed value, and rewrites the whole expression
1020/// tree so that *all* the computations are done on integers, and the only
1021/// pointer-typed operands in the expression are SCEVUnknown.
1022/// The CreatePtrCast callback is invoked to create the actual conversion
1023/// (ptrtoint or ptrtoaddr) at the SCEVUnknown leaves.
1025 : public SCEVRewriteVisitor<SCEVCastSinkingRewriter> {
1027 using ConversionFn = function_ref<const SCEV *(const SCEVUnknown *)>;
1028 Type *TargetTy;
1029 ConversionFn CreatePtrCast;
1030
1031public:
1033 ConversionFn CreatePtrCast)
1034 : Base(SE), TargetTy(TargetTy), CreatePtrCast(std::move(CreatePtrCast)) {}
1035
1036 static const SCEV *rewrite(const SCEV *Scev, ScalarEvolution &SE,
1037 Type *TargetTy, ConversionFn CreatePtrCast) {
1038 SCEVCastSinkingRewriter Rewriter(SE, TargetTy, std::move(CreatePtrCast));
1039 return Rewriter.visit(Scev);
1040 }
1041
1042 const SCEV *visit(const SCEV *S) {
1043 Type *STy = S->getType();
1044 // If the expression is not pointer-typed, just keep it as-is.
1045 if (!STy->isPointerTy())
1046 return S;
1047 // Else, recursively sink the cast down into it.
1048 return Base::visit(S);
1049 }
1050
1051 const SCEV *visitAddExpr(const SCEVAddExpr *Expr) {
1052 // Preserve wrap flags on rewritten SCEVAddExpr, which the default
1053 // implementation drops.
1055 bool Changed = false;
1056 for (SCEVUse Op : Expr->operands()) {
1057 Operands.push_back(visit(Op.getPointer()));
1058 Changed |= Op.getPointer() != Operands.back();
1059 }
1060 return !Changed ? Expr : SE.getAddExpr(Operands, Expr->getNoWrapFlags());
1061 }
1062
1063 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
1064 assert(Expr->getType()->isPointerTy() &&
1065 "Should only reach pointer-typed SCEVUnknown's.");
1066 // Perform some basic constant folding. If the operand of the cast is a
1067 // null pointer, don't create a cast SCEV expression (that will be left
1068 // as-is), but produce a zero constant.
1070 return SE.getZero(TargetTy);
1071 return CreatePtrCast(Expr);
1072 }
1073};
1074
1076 assert(Op->getType()->isPointerTy() && "Op must be a pointer");
1077
1078 // Treat pointers with unstable representation conservatively, since the
1079 // address bits may change.
1080 if (DL.hasUnstableRepresentation(Op->getType()))
1081 return getCouldNotCompute();
1082
1083 Type *Ty = DL.getAddressType(Op->getType());
1084
1085 // Use the rewriter to sink the cast down to SCEVUnknown leaves.
1086 // The rewriter handles null pointer constant folding.
1088 Op, *this, Ty, [this, Ty](const SCEVUnknown *U) {
1091 ID.AddPointer(U);
1092 ID.AddPointer(Ty);
1094 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
1095 return S;
1096 SCEV *S = new (SCEVAllocator)
1097 SCEVPtrToAddrExpr(ID.Intern(SCEVAllocator), U, Ty);
1098 UniqueSCEVs.insert(S, Token);
1099 S->computeAndSetCanonical(*this);
1100 registerUser(S, {U});
1101 return static_cast<const SCEV *>(S);
1102 });
1103 assert(IntOp->getType()->isIntegerTy() &&
1104 "We must have succeeded in sinking the cast, "
1105 "and ending up with an integer-typed expression!");
1106 return IntOp;
1107}
1108
1110 unsigned Depth) {
1111 assert(getTypeSizeInBits(Op->getType()) > getTypeSizeInBits(Ty) &&
1112 "This is not a truncating conversion!");
1113 assert(isSCEVable(Ty) &&
1114 "This is not a conversion to a SCEVable type!");
1115 assert(!Op->getType()->isPointerTy() && "Can't truncate pointer!");
1116 Ty = getEffectiveSCEVType(Ty);
1117
1120 ID.AddPointer(Op.getOpaqueValue());
1121 ID.AddPointer(Ty);
1123 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
1124 return S;
1125
1126 // Fold if the operand is constant.
1127 if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(Op))
1128 return getConstant(
1129 cast<ConstantInt>(ConstantExpr::getTrunc(SC->getValue(), Ty)));
1130
1131 // trunc(trunc(x)) --> trunc(x)
1133 return getTruncateExpr(ST->getOperand(), Ty, Depth + 1);
1134
1135 // trunc(sext(x)) --> sext(x) if widening or trunc(x) if narrowing
1137 return getTruncateOrSignExtend(SS->getOperand(), Ty, Depth + 1);
1138
1139 // trunc(zext(x)) --> zext(x) if widening or trunc(x) if narrowing
1141 return getTruncateOrZeroExtend(SZ->getOperand(), Ty, Depth + 1);
1142
1143 if (Depth > MaxCastDepth) {
1144 SCEV *S =
1145 new (SCEVAllocator) SCEVTruncateExpr(ID.Intern(SCEVAllocator), Op, Ty);
1146 UniqueSCEVs.insert(S, Token);
1147 S->computeAndSetCanonical(*this);
1148 registerUser(S, Op);
1149 return S;
1150 }
1151
1152 // trunc(x1 + ... + xN) --> trunc(x1) + ... + trunc(xN) and
1153 // trunc(x1 * ... * xN) --> trunc(x1) * ... * trunc(xN),
1154 // if after transforming we have at most one truncate, not counting truncates
1155 // that replace other casts.
1157 auto *CommOp = cast<SCEVCommutativeExpr>(Op);
1159 unsigned numTruncs = 0;
1160 for (unsigned i = 0, e = CommOp->getNumOperands(); i != e && numTruncs < 2;
1161 ++i) {
1162 const SCEV *S = getTruncateExpr(CommOp->getOperand(i), Ty, Depth + 1);
1163 if (!isa<SCEVIntegralCastExpr>(CommOp->getOperand(i)) &&
1165 numTruncs++;
1166 Operands.push_back(S);
1167 }
1168 if (numTruncs < 2) {
1169 if (isa<SCEVAddExpr>(Op))
1170 return getAddExpr(Operands);
1171 if (isa<SCEVMulExpr>(Op))
1172 return getMulExpr(Operands);
1173 llvm_unreachable("Unexpected SCEV type for Op.");
1174 }
1175 // Although we checked in the beginning that ID is not in the cache, it is
1176 // possible that during recursion and different modification ID was inserted
1177 // into the cache. So if we find it, just return it.
1178 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
1179 return S;
1180 }
1181
1182 // If the input value is a chrec scev, truncate the chrec's operands.
1183 if (const SCEVAddRecExpr *AddRec = dyn_cast<SCEVAddRecExpr>(Op)) {
1185 for (const SCEV *Op : AddRec->operands())
1186 Operands.push_back(getTruncateExpr(Op, Ty, Depth + 1));
1187 return getAddRecExpr(Operands, AddRec->getLoop(), SCEV::FlagNone);
1188 }
1189
1190 // Return zero if truncating to known zeros.
1191 uint32_t MinTrailingZeros = getMinTrailingZeros(Op);
1192 if (MinTrailingZeros >= getTypeSizeInBits(Ty))
1193 return getZero(Ty);
1194
1195 // The cast wasn't folded; create an explicit cast node. We can reuse
1196 // the existing insert position since if we get here, we won't have
1197 // made any changes which would invalidate it.
1198 SCEV *S = new (SCEVAllocator) SCEVTruncateExpr(ID.Intern(SCEVAllocator),
1199 Op, Ty);
1200 UniqueSCEVs.insert(S, Token);
1201 S->computeAndSetCanonical(*this);
1202 registerUser(S, Op);
1203 return S;
1204}
1205
1206// Get the limit of a recurrence such that incrementing by Step cannot cause
1207// signed overflow as long as the value of the recurrence within the
1208// loop does not exceed this limit before incrementing.
1209static const SCEV *getSignedOverflowLimitForStep(const SCEV *Step,
1210 ICmpInst::Predicate *Pred,
1211 ScalarEvolution *SE) {
1212 unsigned BitWidth = SE->getTypeSizeInBits(Step->getType());
1213 if (SE->isKnownPositive(Step)) {
1214 *Pred = ICmpInst::ICMP_SLT;
1216 SE->getSignedRangeMax(Step));
1217 }
1218 if (SE->isKnownNegative(Step)) {
1219 *Pred = ICmpInst::ICMP_SGT;
1221 SE->getSignedRangeMin(Step));
1222 }
1223 return nullptr;
1224}
1225
1226// Get the limit of a recurrence such that incrementing by Step cannot cause
1227// unsigned overflow as long as the value of the recurrence within the loop does
1228// not exceed this limit before incrementing.
1230 ICmpInst::Predicate *Pred,
1231 ScalarEvolution *SE) {
1232 unsigned BitWidth = SE->getTypeSizeInBits(Step->getType());
1233 *Pred = ICmpInst::ICMP_ULT;
1234
1236 SE->getUnsignedRangeMax(Step));
1237}
1238
1239namespace {
1240
1241struct ExtendOpTraitsBase {
1242 typedef const SCEV *(ScalarEvolution::*GetExtendExprTy)(SCEVUse, Type *,
1243 unsigned);
1244};
1245
1246// Used to make code generic over signed and unsigned overflow.
1247template <typename ExtendOp> struct ExtendOpTraits {
1248 // Members present:
1249 //
1250 // static const SCEVFlags WrapType;
1251 //
1252 // static const ExtendOpTraitsBase::GetExtendExprTy GetExtendExpr;
1253 //
1254 // static const SCEV *getOverflowLimitForStep(const SCEV *Step,
1255 // ICmpInst::Predicate *Pred,
1256 // ScalarEvolution *SE);
1257};
1258
1259template <>
1260struct ExtendOpTraits<SCEVSignExtendExpr> : public ExtendOpTraitsBase {
1261 static const SCEVFlags WrapType = SCEV::FlagNSW;
1262
1263 static const GetExtendExprTy GetExtendExpr;
1264
1265 static const SCEV *getOverflowLimitForStep(const SCEV *Step,
1266 ICmpInst::Predicate *Pred,
1267 ScalarEvolution *SE) {
1268 return getSignedOverflowLimitForStep(Step, Pred, SE);
1269 }
1270};
1271
1272const ExtendOpTraitsBase::GetExtendExprTy ExtendOpTraits<
1274
1275template <>
1276struct ExtendOpTraits<SCEVZeroExtendExpr> : public ExtendOpTraitsBase {
1277 static const SCEVFlags WrapType = SCEV::FlagNUW;
1278
1279 static const GetExtendExprTy GetExtendExpr;
1280
1281 static const SCEV *getOverflowLimitForStep(const SCEV *Step,
1282 ICmpInst::Predicate *Pred,
1283 ScalarEvolution *SE) {
1284 return getUnsignedOverflowLimitForStep(Step, Pred, SE);
1285 }
1286};
1287
1288const ExtendOpTraitsBase::GetExtendExprTy ExtendOpTraits<
1290
1291} // end anonymous namespace
1292
1293// The recurrence AR has been shown to have no signed/unsigned wrap or something
1294// close to it. Typically, if we can prove NSW/NUW for AR, then we can just as
1295// easily prove NSW/NUW for its preincrement or postincrement sibling. This
1296// allows normalizing a sign/zero extended AddRec as such: {sext/zext(Step +
1297// Start),+,Step} => {(Step + sext/zext(Start),+,Step} As a result, the
1298// expression "Step + sext/zext(PreIncAR)" is congruent with
1299// "sext/zext(PostIncAR)"
1300template <typename ExtendOpTy>
1302 ScalarEvolution *SE, unsigned Depth) {
1303 auto WrapType = ExtendOpTraits<ExtendOpTy>::WrapType;
1304 auto GetExtendExpr = ExtendOpTraits<ExtendOpTy>::GetExtendExpr;
1305
1306 const Loop *L = AR->getLoop();
1307 const SCEV *Start = AR->getStart();
1308 const SCEV *Step = AR->getStepRecurrence(*SE);
1309
1310 // Check for a simple looking step prior to loop entry.
1311 const SCEVAddExpr *SA = dyn_cast<SCEVAddExpr>(Start);
1312 if (!SA)
1313 return nullptr;
1314
1315 // Create an AddExpr for "PreStart" after subtracting Step. Full SCEV
1316 // subtraction is expensive. For this purpose, perform a quick and dirty
1317 // difference, by checking for Step in the operand list. Note, that
1318 // SA might have repeated ops, like %a + %a + ..., so only remove one.
1319 SmallVector<SCEVUse, 4> DiffOps(SA->operands());
1320 for (auto It = DiffOps.begin(); It != DiffOps.end(); ++It)
1321 if (*It == Step) {
1322 DiffOps.erase(It);
1323 break;
1324 }
1325
1326 if (DiffOps.size() == SA->getNumOperands())
1327 return nullptr;
1328
1329 // Try to prove `WrapType` (SCEV::FlagNSW or SCEV::FlagNUW) on `PreStart` +
1330 // `Step`:
1331
1332 // 1. NSW/NUW flags on the step increment.
1333 auto PreStartFlags =
1335 const SCEV *PreStart = SE->getAddExpr(DiffOps, PreStartFlags);
1337 SE->getAddRecExpr(PreStart, Step, L, SCEV::FlagNone));
1338
1339 // "{S,+,X} is <nsw>/<nuw>" and "the backedge is taken at least once" implies
1340 // "S+X does not sign/unsign-overflow".
1341 //
1342
1343 const SCEV *BECount = SE->getBackedgeTakenCount(L);
1344 if (PreAR && any(PreAR->getNoWrapFlags(WrapType)) &&
1345 !isa<SCEVCouldNotCompute>(BECount) && SE->isKnownPositive(BECount))
1346 return PreStart;
1347
1348 // 2. Direct overflow check on the step operation's expression.
1349 unsigned BitWidth = SE->getTypeSizeInBits(AR->getType());
1350 Type *WideTy = IntegerType::get(SE->getContext(), BitWidth * 2);
1351 const SCEV *OperandExtendedStart =
1352 SE->getAddExpr((SE->*GetExtendExpr)(PreStart, WideTy, Depth),
1353 (SE->*GetExtendExpr)(Step, WideTy, Depth));
1354 if ((SE->*GetExtendExpr)(Start, WideTy, Depth) == OperandExtendedStart) {
1355 if (PreAR && any(AR->getNoWrapFlags(WrapType))) {
1356 // If we know `AR` == {`PreStart`+`Step`,+,`Step`} is `WrapType` (FlagNSW
1357 // or FlagNUW) and that `PreStart` + `Step` is `WrapType` too, then
1358 // `PreAR` == {`PreStart`,+,`Step`} is also `WrapType`. Cache this fact.
1359 SE->setNoWrapFlags(const_cast<SCEVAddRecExpr *>(PreAR), WrapType);
1360 }
1361 return PreStart;
1362 }
1363
1364 // 3. Loop precondition.
1366 const SCEV *OverflowLimit =
1367 ExtendOpTraits<ExtendOpTy>::getOverflowLimitForStep(Step, &Pred, SE);
1368
1369 if (OverflowLimit &&
1370 SE->isLoopEntryGuardedByCond(L, Pred, PreStart, OverflowLimit))
1371 return PreStart;
1372
1373 return nullptr;
1374}
1375
1376// Get the normalized zero or sign extended expression for this AddRec's Start.
1377template <typename ExtendOpTy>
1378static const SCEV *getExtendAddRecStart(const SCEVAddRecExpr *AR, Type *Ty,
1379 ScalarEvolution *SE,
1380 unsigned Depth) {
1381 auto GetExtendExpr = ExtendOpTraits<ExtendOpTy>::GetExtendExpr;
1382
1383 const SCEV *PreStart = getPreStartForExtend<ExtendOpTy>(AR, SE, Depth);
1384 if (!PreStart)
1385 return (SE->*GetExtendExpr)(AR->getStart(), Ty, Depth);
1386
1387 return SE->getAddExpr((SE->*GetExtendExpr)(AR->getStepRecurrence(*SE), Ty,
1388 Depth),
1389 (SE->*GetExtendExpr)(PreStart, Ty, Depth));
1390}
1391
1392// Try to prove away overflow by looking at "nearby" add recurrences. A
1393// motivating example for this rule: if we know `{0,+,4}` is `ult` `-1` and it
1394// does not itself wrap then we can conclude that `{1,+,4}` is `nuw`.
1395//
1396// Formally:
1397//
1398// {S,+,X} == {S-T,+,X} + T
1399// => Ext({S,+,X}) == Ext({S-T,+,X} + T)
1400//
1401// If ({S-T,+,X} + T) does not overflow ... (1)
1402//
1403// RHS == Ext({S-T,+,X} + T) == Ext({S-T,+,X}) + Ext(T)
1404//
1405// If {S-T,+,X} does not overflow ... (2)
1406//
1407// RHS == Ext({S-T,+,X}) + Ext(T) == {Ext(S-T),+,Ext(X)} + Ext(T)
1408// == {Ext(S-T)+Ext(T),+,Ext(X)}
1409//
1410// If (S-T)+T does not overflow ... (3)
1411//
1412// RHS == {Ext(S-T)+Ext(T),+,Ext(X)} == {Ext(S-T+T),+,Ext(X)}
1413// == {Ext(S),+,Ext(X)} == LHS
1414//
1415// Thus, if (1), (2) and (3) are true for some T, then
1416// Ext({S,+,X}) == {Ext(S),+,Ext(X)}
1417//
1418// (3) is implied by (1) -- "(S-T)+T does not overflow" is simply "({S-T,+,X}+T)
1419// does not overflow" restricted to the 0th iteration. Therefore we only need
1420// to check for (1) and (2).
1421//
1422// In the current context, S is `Start`, X is `Step`, Ext is `ExtendOpTy` and T
1423// is `Delta` (defined below).
1424template <typename ExtendOpTy>
1425bool ScalarEvolution::proveNoWrapByVaryingStart(const SCEV *Start,
1426 const SCEV *Step,
1427 const Loop *L) {
1428 auto WrapType = ExtendOpTraits<ExtendOpTy>::WrapType;
1429
1430 // We restrict `Start` to a constant to prevent SCEV from spending too much
1431 // time here. It is correct (but more expensive) to continue with a
1432 // non-constant `Start` and do a general SCEV subtraction to compute
1433 // `PreStart` below.
1434 const SCEVConstant *StartC = dyn_cast<SCEVConstant>(Start);
1435 if (!StartC)
1436 return false;
1437
1438 APInt StartAI = StartC->getAPInt();
1439
1440 for (unsigned Delta : {-2, -1, 1, 2}) {
1441 const SCEV *PreStart = getConstant(StartAI - Delta);
1442 const auto *PreAR = static_cast<SCEVAddRecExpr *>(
1443 findExistingSCEVInCache(scAddRecExpr, {PreStart, Step}, L));
1444
1445 // Give up if we don't already have the add recurrence we need because
1446 // actually constructing an add recurrence is relatively expensive.
1447 if (PreAR && any(PreAR->getNoWrapFlags(WrapType))) { // proves (2)
1448 const SCEV *DeltaS = getConstant(StartC->getType(), Delta);
1450 const SCEV *Limit = ExtendOpTraits<ExtendOpTy>::getOverflowLimitForStep(
1451 DeltaS, &Pred, this);
1452 if (Limit && isKnownPredicate(Pred, PreAR, Limit)) // proves (1)
1453 return true;
1454 }
1455 }
1456
1457 return false;
1458}
1459
1460// Finds an integer D for an expression (C + x + y + ...) such that the top
1461// level addition in (D + (C - D + x + y + ...)) would not wrap (signed or
1462// unsigned) and the number of trailing zeros of (C - D + x + y + ...) is
1463// maximized, where C is the \p ConstantTerm, x, y, ... are arbitrary SCEVs, and
1464// the (C + x + y + ...) expression is \p WholeAddExpr.
1466 const SCEVConstant *ConstantTerm,
1467 const SCEVAddExpr *WholeAddExpr) {
1468 const APInt &C = ConstantTerm->getAPInt();
1469 const unsigned BitWidth = C.getBitWidth();
1470 // Find number of trailing zeros of (x + y + ...) w/o the C first:
1471 uint32_t TZ = BitWidth;
1472 for (unsigned I = 1, E = WholeAddExpr->getNumOperands(); I < E && TZ; ++I)
1473 TZ = std::min(TZ, SE.getMinTrailingZeros(WholeAddExpr->getOperand(I)));
1474 if (TZ) {
1475 // Set D to be as many least significant bits of C as possible while still
1476 // guaranteeing that adding D to (C - D + x + y + ...) won't cause a wrap:
1477 return TZ < BitWidth ? C.trunc(TZ).zext(BitWidth) : C;
1478 }
1479 return APInt(BitWidth, 0);
1480}
1481
1482// Finds an integer D for an affine AddRec expression {C,+,x} such that the top
1483// level addition in (D + {C-D,+,x}) would not wrap (signed or unsigned) and the
1484// number of trailing zeros of (C - D + x * n) is maximized, where C is the \p
1485// ConstantStart, x is an arbitrary \p Step, and n is the loop trip count.
1487 const APInt &ConstantStart,
1488 const SCEV *Step) {
1489 const unsigned BitWidth = ConstantStart.getBitWidth();
1490 const uint32_t TZ = SE.getMinTrailingZeros(Step);
1491 if (TZ)
1492 return TZ < BitWidth ? ConstantStart.trunc(TZ).zext(BitWidth)
1493 : ConstantStart;
1494 return APInt(BitWidth, 0);
1495}
1496
1498 const ScalarEvolution::FoldID &ID, const SCEV *S,
1501 &FoldCacheUser) {
1502 auto I = FoldCache.insert({ID, S});
1503 if (!I.second) {
1504 // Remove FoldCacheUser entry for ID when replacing an existing FoldCache
1505 // entry.
1506 auto &UserIDs = FoldCacheUser[I.first->second];
1507 assert(count(UserIDs, ID) == 1 && "unexpected duplicates in UserIDs");
1508 for (unsigned I = 0; I != UserIDs.size(); ++I)
1509 if (UserIDs[I] == ID) {
1510 std::swap(UserIDs[I], UserIDs.back());
1511 break;
1512 }
1513 UserIDs.pop_back();
1514 I.first->second = S;
1515 }
1516 FoldCacheUser[S].push_back(ID);
1517}
1518
1520 unsigned Depth) {
1521 assert(getTypeSizeInBits(Op->getType()) < getTypeSizeInBits(Ty) &&
1522 "This is not an extending conversion!");
1523 assert(isSCEVable(Ty) &&
1524 "This is not a conversion to a SCEVable type!");
1525 assert(!Op->getType()->isPointerTy() && "Can't extend pointer!");
1526 Ty = getEffectiveSCEVType(Ty);
1527
1528 FoldID ID(scZeroExtend, Op, Ty);
1529 if (const SCEV *S = FoldCache.lookup(ID))
1530 return S;
1531
1532 const SCEV *S = getZeroExtendExprImpl(Op, Ty, Depth);
1534 insertFoldCacheEntry(ID, S, FoldCache, FoldCacheUser);
1535 return S;
1536}
1537
1539 unsigned Depth) {
1540 assert(getTypeSizeInBits(Op->getType()) < getTypeSizeInBits(Ty) &&
1541 "This is not an extending conversion!");
1542 assert(isSCEVable(Ty) && "This is not a conversion to a SCEVable type!");
1543 assert(!Op->getType()->isPointerTy() && "Can't extend pointer!");
1544
1545 // Fold if the operand is constant.
1546 if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(Op))
1547 return getConstant(SC->getAPInt().zext(getTypeSizeInBits(Ty)));
1548
1549 // zext(zext(x)) --> zext(x)
1551 return getZeroExtendExpr(SZ->getOperand(), Ty, Depth + 1);
1552
1553 // If the operand is an affine AddRec with the no-unsigned-wrap flag, the
1554 // zero-extension distributes over the recurrence.
1555 const SCEV *Start, *Step;
1556 const Loop *L;
1557 if (Depth <= MaxCastDepth &&
1558 match(Op, m_scev_AffineAddRec(m_SCEV(Start), m_SCEV(Step), m_Loop(L)))) {
1559 const auto *AR = cast<SCEVAddRecExpr>(Op);
1560 if (AR->hasNoUnsignedWrap()) {
1561 Start = getExtendAddRecStart<SCEVZeroExtendExpr>(AR, Ty, this, Depth + 1);
1562 Step = getZeroExtendExpr(Step, Ty, Depth + 1);
1563 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1564 }
1565 }
1566
1567 // Before doing any expensive analysis, check to see if we've already
1568 // computed a SCEV for this Op and Ty.
1571 ID.AddPointer(Op.getOpaqueValue());
1572 ID.AddPointer(Ty);
1574 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
1575 return S;
1576 if (Depth > MaxCastDepth) {
1577 SCEV *S = new (SCEVAllocator) SCEVZeroExtendExpr(ID.Intern(SCEVAllocator),
1578 Op, Ty);
1579 UniqueSCEVs.insert(S, Token);
1580 S->computeAndSetCanonical(*this);
1581 registerUser(S, Op);
1582 return S;
1583 }
1584
1585 // zext(trunc(x)) --> zext(x) or x or trunc(x)
1587 // It's possible the bits taken off by the truncate were all zero bits. If
1588 // so, we should be able to simplify this further.
1589 const SCEV *X = ST->getOperand();
1591 unsigned TruncBits = getTypeSizeInBits(ST->getType());
1592 unsigned NewBits = getTypeSizeInBits(Ty);
1593 if (CR.truncate(TruncBits).zeroExtend(NewBits).contains(
1594 CR.zextOrTrunc(NewBits)))
1595 return getTruncateOrZeroExtend(X, Ty, Depth);
1596 }
1597
1598 // If the input value is a chrec scev, and we can prove that the value
1599 // did not overflow the old, smaller, value, we can zero extend all of the
1600 // operands (often constants). This allows analysis of something like
1601 // this: for (unsigned char X = 0; X < 100; ++X) { int Y = X; }
1602 if (match(Op, m_scev_AffineAddRec(m_SCEV(Start), m_SCEV(Step), m_Loop(L)))) {
1603 const auto *AR = cast<SCEVAddRecExpr>(Op);
1604 unsigned BitWidth = getTypeSizeInBits(AR->getType());
1605
1606 // The no-unsigned-wrap case is handled before the uniquing lookup above.
1607
1608 // Check whether the backedge-taken count is SCEVCouldNotCompute.
1609 // Note that this serves two purposes: It filters out loops that are
1610 // simply not analyzable, and it covers the case where this code is
1611 // being called from within backedge-taken count analysis, such that
1612 // attempting to ask for the backedge-taken count would likely result
1613 // in infinite recursion. In the later case, the analysis code will
1614 // cope with a conservative value, and it will take care to purge
1615 // that value once it has finished.
1616 const SCEV *MaxBECount = getConstantMaxBackedgeTakenCount(L);
1617 if (!isa<SCEVCouldNotCompute>(MaxBECount)) {
1618 // Manually compute the final value for AR, checking for overflow.
1619
1620 // Check whether the backedge-taken count can be losslessly casted to
1621 // the addrec's type. The count is always unsigned.
1622 const SCEV *CastedMaxBECount =
1623 getTruncateOrZeroExtend(MaxBECount, Start->getType(), Depth);
1624 const SCEV *RecastedMaxBECount = getTruncateOrZeroExtend(
1625 CastedMaxBECount, MaxBECount->getType(), Depth);
1626 if (MaxBECount == RecastedMaxBECount) {
1627 Type *WideTy = IntegerType::get(getContext(), BitWidth * 2);
1628 // Check whether Start+Step*MaxBECount has no unsigned overflow.
1629 const SCEV *ZMul =
1630 getMulExpr(CastedMaxBECount, Step, SCEV::FlagNone, Depth + 1);
1631 const SCEV *ZAdd = getZeroExtendExpr(
1632 getAddExpr(Start, ZMul, SCEV::FlagNone, Depth + 1), WideTy,
1633 Depth + 1);
1634 const SCEV *WideStart = getZeroExtendExpr(Start, WideTy, Depth + 1);
1635 const SCEV *WideMaxBECount =
1636 getZeroExtendExpr(CastedMaxBECount, WideTy, Depth + 1);
1637 const SCEV *OperandExtendedAdd =
1638 getAddExpr(WideStart,
1639 getMulExpr(WideMaxBECount,
1640 getZeroExtendExpr(Step, WideTy, Depth + 1),
1641 SCEV::FlagNone, Depth + 1),
1642 SCEV::FlagNone, Depth + 1);
1643 if (ZAdd == OperandExtendedAdd) {
1644 // Cache knowledge of AR NUW, which is propagated to this AddRec.
1645 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNUW);
1646 // Return the expression with the addrec on the outside.
1647 Start =
1649 Step = getZeroExtendExpr(Step, Ty, Depth + 1);
1650 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1651 }
1652 // Similar to above, only this time treat the step value as signed.
1653 // This covers loops that count down.
1654 OperandExtendedAdd =
1655 getAddExpr(WideStart,
1656 getMulExpr(WideMaxBECount,
1657 getSignExtendExpr(Step, WideTy, Depth + 1),
1658 SCEV::FlagNone, Depth + 1),
1659 SCEV::FlagNone, Depth + 1);
1660 if (ZAdd == OperandExtendedAdd) {
1661 // Cache knowledge of AR NW, which is propagated to this AddRec.
1662 // Negative step causes unsigned wrap, but it still can't self-wrap.
1663 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNW);
1664 // Return the expression with the addrec on the outside.
1665 Start =
1667 Step = getSignExtendExpr(Step, Ty, Depth + 1);
1668 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1669 }
1670 }
1671 }
1672
1673 // Normally, in the cases we can prove no-overflow via a
1674 // backedge guarding condition, we can also compute a backedge
1675 // taken count for the loop. The exceptions are assumptions and
1676 // guards present in the loop -- SCEV is not great at exploiting
1677 // these to compute max backedge taken counts, but can still use
1678 // these to prove lack of overflow. Use this fact to avoid
1679 // doing extra work that may not pay off.
1680 if (!isa<SCEVCouldNotCompute>(MaxBECount) || HasGuards ||
1681 !AC.assumptions().empty()) {
1682
1683 auto NewFlags = proveNoUnsignedWrapViaInduction(AR);
1684 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), NewFlags);
1685 if (AR->hasNoUnsignedWrap()) {
1686 // Same as nuw case above - duplicated here to avoid a compile time
1687 // issue. It's not clear that the order of checks does matter, but
1688 // it's one of two issue possible causes for a change which was
1689 // reverted. Be conservative for the moment.
1690 Start =
1692 Step = getZeroExtendExpr(Step, Ty, Depth + 1);
1693 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1694 }
1695
1696 // For a negative step, we can extend the operands iff doing so only
1697 // traverses values in the range zext([0,UINT_MAX]).
1698 if (isKnownNegative(Step)) {
1699 const SCEV *N =
1703 // Cache knowledge of AR NW, which is propagated to this
1704 // AddRec. Negative step causes unsigned wrap, but it
1705 // still can't self-wrap.
1706 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNW);
1707 // Return the expression with the addrec on the outside.
1708 Start =
1710 Step = getSignExtendExpr(Step, Ty, Depth + 1);
1711 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1712 }
1713 }
1714 }
1715
1716 // zext({C,+,Step}) --> (zext(D) + zext({C-D,+,Step}))<nuw><nsw>
1717 // if D + (C - D + Step * n) could be proven to not unsigned wrap
1718 // where D maximizes the number of trailing zeros of (C - D + Step * n)
1719 if (const auto *SC = dyn_cast<SCEVConstant>(Start)) {
1720 const APInt &C = SC->getAPInt();
1721 const APInt &D = extractConstantWithoutWrapping(*this, C, Step);
1722 if (D != 0) {
1723 const SCEV *SZExtD = getZeroExtendExpr(getConstant(D), Ty, Depth);
1724 const SCEV *SResidual =
1725 getAddRecExpr(getConstant(C - D), Step, L, AR->getNoWrapFlags());
1726 const SCEV *SZExtR = getZeroExtendExpr(SResidual, Ty, Depth + 1);
1727 return getAddExpr(SZExtD, SZExtR, SCEV::FlagNSW | SCEV::FlagNUW,
1728 Depth + 1);
1729 }
1730 }
1731
1732 if (proveNoWrapByVaryingStart<SCEVZeroExtendExpr>(Start, Step, L)) {
1733 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNUW);
1734 Start = getExtendAddRecStart<SCEVZeroExtendExpr>(AR, Ty, this, Depth + 1);
1735 Step = getZeroExtendExpr(Step, Ty, Depth + 1);
1736 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1737 }
1738 }
1739
1740 // zext(A % B) --> zext(A) % zext(B)
1741 {
1742 const SCEV *LHS;
1743 const SCEV *RHS;
1744 if (match(Op, m_scev_URem(m_SCEV(LHS), m_SCEV(RHS), *this)))
1745 return getURemExpr(getZeroExtendExpr(LHS, Ty, Depth + 1),
1746 getZeroExtendExpr(RHS, Ty, Depth + 1));
1747 }
1748
1749 // zext(A / B) --> zext(A) / zext(B).
1750 if (auto *Div = dyn_cast<SCEVUDivExpr>(Op))
1751 return getUDivExpr(getZeroExtendExpr(Div->getLHS(), Ty, Depth + 1),
1752 getZeroExtendExpr(Div->getRHS(), Ty, Depth + 1));
1753
1754 if (auto *SA = dyn_cast<SCEVAddExpr>(Op)) {
1755 // zext((A + B + ...)<nuw>) --> (zext(A) + zext(B) + ...)<nuw>
1756 if (SA->hasNoUnsignedWrap()) {
1757 // If the addition does not unsign overflow then we can, by definition,
1758 // commute the zero extension with the addition operation.
1760 for (SCEVUse Op : SA->operands())
1761 Ops.push_back(getZeroExtendExpr(Op, Ty, Depth + 1));
1762 return getAddExpr(Ops, SCEV::FlagNUW, Depth + 1);
1763 }
1764
1765 const APInt *C, *C2;
1766 // zext (C + A)<nsw> -> (sext(C) + sext(A))<nsw> if zext (C + A)<nsw> >=s 0.
1767 // Currently the non-negative check is done manually, as isKnownNonNegative
1768 // is too expensive.
1769 if (SA->hasNoSignedWrap() &&
1771 m_scev_SMax(m_scev_APInt(C2), m_SCEV()))) &&
1772 C->isNegative() && !C->isMinSignedValue() && C2->sge(C->abs())) {
1773 assert(isKnownNonNegative(SA) && "incorrectly determined non-negative");
1774 return getAddExpr(getSignExtendExpr(SA->getOperand(0), Ty, Depth + 1),
1775 getSignExtendExpr(SA->getOperand(1), Ty, Depth + 1),
1776 SCEV::FlagNSW, Depth + 1);
1777 }
1778
1779 // zext(C + x + y + ...) --> (zext(D) + zext((C - D) + x + y + ...))
1780 // if D + (C - D + x + y + ...) could be proven to not unsigned wrap
1781 // where D maximizes the number of trailing zeros of (C - D + x + y + ...)
1782 //
1783 // Often address arithmetics contain expressions like
1784 // (zext (add (shl X, C1), C2)), for instance, (zext (5 + (4 * X))).
1785 // This transformation is useful while proving that such expressions are
1786 // equal or differ by a small constant amount, see LoadStoreVectorizer pass.
1787 if (const auto *SC = dyn_cast<SCEVConstant>(SA->getOperand(0))) {
1788 const APInt &D = extractConstantWithoutWrapping(*this, SC, SA);
1789 if (D != 0) {
1790 const SCEV *SZExtD = getZeroExtendExpr(getConstant(D), Ty, Depth);
1791 const SCEV *SResidual =
1793 const SCEV *SZExtR = getZeroExtendExpr(SResidual, Ty, Depth + 1);
1794 return getAddExpr(SZExtD, SZExtR, (SCEV::FlagNSW | SCEV::FlagNUW),
1795 Depth + 1);
1796 }
1797 }
1798 }
1799
1800 if (auto *SM = dyn_cast<SCEVMulExpr>(Op)) {
1801 // zext((A * B * ...)<nuw>) --> (zext(A) * zext(B) * ...)<nuw>
1802 if (SM->hasNoUnsignedWrap()) {
1803 // If the multiply does not unsign overflow then we can, by definition,
1804 // commute the zero extension with the multiply operation.
1806 for (SCEVUse Op : SM->operands())
1807 Ops.push_back(getZeroExtendExpr(Op, Ty, Depth + 1));
1808 return getMulExpr(Ops, SCEV::FlagNUW, Depth + 1);
1809 }
1810
1811 // zext(2^K * (trunc X to iN)) to iM ->
1812 // 2^K * (zext(trunc X to i{N-K}) to iM)<nuw>
1813 //
1814 // Proof:
1815 //
1816 // zext(2^K * (trunc X to iN)) to iM
1817 // = zext((trunc X to iN) << K) to iM
1818 // = zext((trunc X to i{N-K}) << K)<nuw> to iM
1819 // (because shl removes the top K bits)
1820 // = zext((2^K * (trunc X to i{N-K}))<nuw>) to iM
1821 // = (2^K * (zext(trunc X to i{N-K}) to iM))<nuw>.
1822 //
1823 const APInt *C;
1824 const SCEV *TruncRHS;
1825 if (match(SM,
1826 m_scev_Mul(m_scev_APInt(C), m_scev_Trunc(m_SCEV(TruncRHS)))) &&
1827 C->isPowerOf2()) {
1828 int NewTruncBits =
1829 getTypeSizeInBits(SM->getOperand(1)->getType()) - C->logBase2();
1830 Type *NewTruncTy = IntegerType::get(getContext(), NewTruncBits);
1831 return getMulExpr(
1832 getZeroExtendExpr(SM->getOperand(0), Ty),
1833 getZeroExtendExpr(getTruncateExpr(TruncRHS, NewTruncTy), Ty),
1834 SCEV::FlagNUW, Depth + 1);
1835 }
1836 }
1837
1838 // zext(umin(x, y)) -> umin(zext(x), zext(y))
1839 // zext(umax(x, y)) -> umax(zext(x), zext(y))
1843 for (SCEVUse Operand : MinMax->operands())
1844 Operands.push_back(getZeroExtendExpr(Operand, Ty));
1846 return getUMinExpr(Operands);
1847 return getUMaxExpr(Operands);
1848 }
1849
1850 // zext(umin_seq(x, y)) -> umin_seq(zext(x), zext(y))
1852 assert(isa<SCEVSequentialUMinExpr>(MinMax) && "Not supported!");
1854 for (SCEVUse Operand : MinMax->operands())
1855 Operands.push_back(getZeroExtendExpr(Operand, Ty));
1856 return getUMinExpr(Operands, /*Sequential*/ true);
1857 }
1858
1859 // The cast wasn't folded; create an explicit cast node.
1860 // Recompute the insert position, as it may have been invalidated.
1861 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
1862 return S;
1863 SCEV *S = new (SCEVAllocator) SCEVZeroExtendExpr(ID.Intern(SCEVAllocator),
1864 Op, Ty);
1865 UniqueSCEVs.insert(S, Token);
1866 S->computeAndSetCanonical(*this);
1867 registerUser(S, Op);
1868 return S;
1869}
1870
1872 unsigned Depth) {
1873 assert(getTypeSizeInBits(Op->getType()) < getTypeSizeInBits(Ty) &&
1874 "This is not an extending conversion!");
1875 assert(isSCEVable(Ty) &&
1876 "This is not a conversion to a SCEVable type!");
1877 assert(!Op->getType()->isPointerTy() && "Can't extend pointer!");
1878 Ty = getEffectiveSCEVType(Ty);
1879
1880 FoldID ID(scSignExtend, Op, Ty);
1881 if (const SCEV *S = FoldCache.lookup(ID))
1882 return S;
1883
1884 const SCEV *S = getSignExtendExprImpl(Op, Ty, Depth);
1886 insertFoldCacheEntry(ID, S, FoldCache, FoldCacheUser);
1887 return S;
1888}
1889
1891 unsigned Depth) {
1892 assert(getTypeSizeInBits(Op->getType()) < getTypeSizeInBits(Ty) &&
1893 "This is not an extending conversion!");
1894 assert(isSCEVable(Ty) && "This is not a conversion to a SCEVable type!");
1895 assert(!Op->getType()->isPointerTy() && "Can't extend pointer!");
1896 Ty = getEffectiveSCEVType(Ty);
1897
1898 // Fold if the operand is constant.
1899 if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(Op))
1900 return getConstant(SC->getAPInt().sext(getTypeSizeInBits(Ty)));
1901
1902 // sext(sext(x)) --> sext(x)
1904 return getSignExtendExpr(SS->getOperand(), Ty, Depth + 1);
1905
1906 // sext(zext(x)) --> zext(x)
1908 return getZeroExtendExpr(SZ->getOperand(), Ty, Depth + 1);
1909
1910 // If the operand is an affine AddRec with the no-signed-wrap flag, the
1911 // sign-extension distributes over the recurrence.
1912 const SCEV *Start, *Step;
1913 const Loop *L;
1914 if (Depth <= MaxCastDepth &&
1915 match(Op, m_scev_AffineAddRec(m_SCEV(Start), m_SCEV(Step), m_Loop(L)))) {
1916 const auto *AR = cast<SCEVAddRecExpr>(Op);
1917 if (AR->hasNoSignedWrap()) {
1918 Start = getExtendAddRecStart<SCEVSignExtendExpr>(AR, Ty, this, Depth + 1);
1919 Step = getSignExtendExpr(Step, Ty, Depth + 1);
1920 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
1921 }
1922 }
1923
1924 // Before doing any expensive analysis, check to see if we've already
1925 // computed a SCEV for this Op and Ty.
1928 ID.AddPointer(Op.getOpaqueValue());
1929 ID.AddPointer(Ty);
1931 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
1932 return S;
1933 // Limit recursion depth.
1934 if (Depth > MaxCastDepth) {
1935 SCEV *S = new (SCEVAllocator) SCEVSignExtendExpr(ID.Intern(SCEVAllocator),
1936 Op, Ty);
1937 UniqueSCEVs.insert(S, Token);
1938 S->computeAndSetCanonical(*this);
1939 registerUser(S, Op);
1940 return S;
1941 }
1942
1943 // sext(trunc(x)) --> sext(x) or x or trunc(x)
1945 // It's possible the bits taken off by the truncate were all sign bits. If
1946 // so, we should be able to simplify this further.
1947 const SCEV *X = ST->getOperand();
1949 unsigned TruncBits = getTypeSizeInBits(ST->getType());
1950 unsigned NewBits = getTypeSizeInBits(Ty);
1951 if (CR.truncate(TruncBits).signExtend(NewBits).contains(
1952 CR.sextOrTrunc(NewBits)))
1953 return getTruncateOrSignExtend(X, Ty, Depth);
1954 }
1955
1956 if (auto *SA = dyn_cast<SCEVAddExpr>(Op)) {
1957 // sext((A + B + ...)<nsw>) --> (sext(A) + sext(B) + ...)<nsw>
1958 if (SA->hasNoSignedWrap()) {
1959 // If the addition does not sign overflow then we can, by definition,
1960 // commute the sign extension with the addition operation.
1962 for (SCEVUse Op : SA->operands())
1963 Ops.push_back(getSignExtendExpr(Op, Ty, Depth + 1));
1964 return getAddExpr(Ops, SCEV::FlagNSW, Depth + 1);
1965 }
1966
1967 // sext(C + x + y + ...) --> (sext(D) + sext((C - D) + x + y + ...))
1968 // if D + (C - D + x + y + ...) could be proven to not signed wrap
1969 // where D maximizes the number of trailing zeros of (C - D + x + y + ...)
1970 //
1971 // For instance, this will bring two seemingly different expressions:
1972 // 1 + sext(5 + 20 * %x + 24 * %y) and
1973 // sext(6 + 20 * %x + 24 * %y)
1974 // to the same form:
1975 // 2 + sext(4 + 20 * %x + 24 * %y)
1976 if (const auto *SC = dyn_cast<SCEVConstant>(SA->getOperand(0))) {
1977 const APInt &D = extractConstantWithoutWrapping(*this, SC, SA);
1978 if (D != 0) {
1979 const SCEV *SSExtD = getSignExtendExpr(getConstant(D), Ty, Depth);
1980 const SCEV *SResidual =
1982 const SCEV *SSExtR = getSignExtendExpr(SResidual, Ty, Depth + 1);
1983 return getAddExpr(SSExtD, SSExtR, (SCEV::FlagNSW | SCEV::FlagNUW),
1984 Depth + 1);
1985 }
1986 }
1987 }
1988 // If the input value is a chrec scev, and we can prove that the value
1989 // did not overflow the old, smaller, value, we can sign extend all of the
1990 // operands (often constants). This allows analysis of something like
1991 // this: for (signed char X = 0; X < 100; ++X) { int Y = X; }
1992 if (match(Op, m_scev_AffineAddRec(m_SCEV(Start), m_SCEV(Step), m_Loop(L)))) {
1993 const auto *AR = cast<SCEVAddRecExpr>(Op);
1994 unsigned BitWidth = getTypeSizeInBits(AR->getType());
1995
1996 // The no-signed-wrap case is handled before the uniquing lookup above.
1997
1998 // Check whether the backedge-taken count is SCEVCouldNotCompute.
1999 // Note that this serves two purposes: It filters out loops that are
2000 // simply not analyzable, and it covers the case where this code is
2001 // being called from within backedge-taken count analysis, such that
2002 // attempting to ask for the backedge-taken count would likely result
2003 // in infinite recursion. In the later case, the analysis code will
2004 // cope with a conservative value, and it will take care to purge
2005 // that value once it has finished.
2006 const SCEV *MaxBECount = getConstantMaxBackedgeTakenCount(L);
2007 if (!isa<SCEVCouldNotCompute>(MaxBECount)) {
2008 // Manually compute the final value for AR, checking for
2009 // overflow.
2010
2011 // Check whether the backedge-taken count can be losslessly casted to
2012 // the addrec's type. The count is always unsigned.
2013 const SCEV *CastedMaxBECount =
2014 getTruncateOrZeroExtend(MaxBECount, Start->getType(), Depth);
2015 const SCEV *RecastedMaxBECount = getTruncateOrZeroExtend(
2016 CastedMaxBECount, MaxBECount->getType(), Depth);
2017 if (MaxBECount == RecastedMaxBECount) {
2018 Type *WideTy = IntegerType::get(getContext(), BitWidth * 2);
2019 // Check whether Start+Step*MaxBECount has no signed overflow.
2020 const SCEV *SMul =
2021 getMulExpr(CastedMaxBECount, Step, SCEV::FlagNone, Depth + 1);
2022 const SCEV *SAdd = getSignExtendExpr(
2023 getAddExpr(Start, SMul, SCEV::FlagNone, Depth + 1), WideTy,
2024 Depth + 1);
2025 const SCEV *WideStart = getSignExtendExpr(Start, WideTy, Depth + 1);
2026 const SCEV *WideMaxBECount =
2027 getZeroExtendExpr(CastedMaxBECount, WideTy, Depth + 1);
2028 const SCEV *OperandExtendedAdd =
2029 getAddExpr(WideStart,
2030 getMulExpr(WideMaxBECount,
2031 getSignExtendExpr(Step, WideTy, Depth + 1),
2032 SCEV::FlagNone, Depth + 1),
2033 SCEV::FlagNone, Depth + 1);
2034 if (SAdd == OperandExtendedAdd) {
2035 // Cache knowledge of AR NSW, which is propagated to this AddRec.
2036 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNSW);
2037 // Return the expression with the addrec on the outside.
2038 Start =
2040 Step = getSignExtendExpr(Step, Ty, Depth + 1);
2041 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
2042 }
2043 // Similar to above, only this time treat the step value as unsigned.
2044 // This covers loops that count up with an unsigned step.
2045 OperandExtendedAdd =
2046 getAddExpr(WideStart,
2047 getMulExpr(WideMaxBECount,
2048 getZeroExtendExpr(Step, WideTy, Depth + 1),
2049 SCEV::FlagNone, Depth + 1),
2050 SCEV::FlagNone, Depth + 1);
2051 if (SAdd == OperandExtendedAdd) {
2052 // If AR wraps around then
2053 //
2054 // abs(Step) * MaxBECount > unsigned-max(AR->getType())
2055 // => SAdd != OperandExtendedAdd
2056 //
2057 // Thus (AR is not NW => SAdd != OperandExtendedAdd) <=>
2058 // (SAdd == OperandExtendedAdd => AR is NW)
2059
2060 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNW);
2061
2062 // Return the expression with the addrec on the outside.
2063 Start =
2065 Step = getZeroExtendExpr(Step, Ty, Depth + 1);
2066 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
2067 }
2068 }
2069 }
2070
2071 auto NewFlags = proveNoSignedWrapViaInduction(AR);
2072 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), NewFlags);
2073 if (AR->hasNoSignedWrap()) {
2074 // Same as nsw case above - duplicated here to avoid a compile time
2075 // issue. It's not clear that the order of checks does matter, but
2076 // it's one of two issue possible causes for a change which was
2077 // reverted. Be conservative for the moment.
2078 Start = getExtendAddRecStart<SCEVSignExtendExpr>(AR, Ty, this, Depth + 1);
2079 Step = getSignExtendExpr(Step, Ty, Depth + 1);
2080 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
2081 }
2082
2083 // sext({C,+,Step}) --> (sext(D) + sext({C-D,+,Step}))<nuw><nsw>
2084 // if D + (C - D + Step * n) could be proven to not signed wrap
2085 // where D maximizes the number of trailing zeros of (C - D + Step * n)
2086 if (const auto *SC = dyn_cast<SCEVConstant>(Start)) {
2087 const APInt &C = SC->getAPInt();
2088 const APInt &D = extractConstantWithoutWrapping(*this, C, Step);
2089 if (D != 0) {
2090 const SCEV *SSExtD = getSignExtendExpr(getConstant(D), Ty, Depth);
2091 const SCEV *SResidual =
2092 getAddRecExpr(getConstant(C - D), Step, L, AR->getNoWrapFlags());
2093 const SCEV *SSExtR = getSignExtendExpr(SResidual, Ty, Depth + 1);
2094 return getAddExpr(SSExtD, SSExtR, (SCEV::FlagNSW | SCEV::FlagNUW),
2095 Depth + 1);
2096 }
2097 }
2098
2099 if (proveNoWrapByVaryingStart<SCEVSignExtendExpr>(Start, Step, L)) {
2100 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), SCEV::FlagNSW);
2101 Start = getExtendAddRecStart<SCEVSignExtendExpr>(AR, Ty, this, Depth + 1);
2102 Step = getSignExtendExpr(Step, Ty, Depth + 1);
2103 return getAddRecExpr(Start, Step, L, AR->getNoWrapFlags());
2104 }
2105 }
2106
2107 // If the input value is provably positive and we could not simplify
2108 // away the sext build a zext instead.
2110 return getZeroExtendExpr(Op, Ty, Depth + 1);
2111
2112 // sext(smin(x, y)) -> smin(sext(x), sext(y))
2113 // sext(smax(x, y)) -> smax(sext(x), sext(y))
2117 for (SCEVUse Operand : MinMax->operands())
2118 Operands.push_back(getSignExtendExpr(Operand, Ty));
2120 return getSMinExpr(Operands);
2121 return getSMaxExpr(Operands);
2122 }
2123
2124 // The cast wasn't folded; create an explicit cast node.
2125 // Recompute the insert position, as it may have been invalidated.
2126 if (const SCEV *S = UniqueSCEVs.lookup(ID, Token))
2127 return S;
2128 SCEV *S = new (SCEVAllocator) SCEVSignExtendExpr(ID.Intern(SCEVAllocator),
2129 Op, Ty);
2130 UniqueSCEVs.insert(S, Token);
2131 S->computeAndSetCanonical(*this);
2132 registerUser(S, Op);
2133 return S;
2134}
2135
2137 switch (Kind) {
2138 case scTruncate:
2139 return getTruncateExpr(Op, Ty);
2140 case scZeroExtend:
2141 return getZeroExtendExpr(Op, Ty);
2142 case scSignExtend:
2143 return getSignExtendExpr(Op, Ty);
2144 case scPtrToAddr: {
2145 const SCEV *Expr = getPtrToAddrExpr(Op);
2146 assert(Expr->getType() == Ty && "requested type must match");
2147 return Expr;
2148 }
2149 default:
2150 llvm_unreachable("Not a SCEV cast expression!");
2151 }
2152}
2153
2154/// getAnyExtendExpr - Return a SCEV for the given operand extended with
2155/// unspecified bits out to the given type.
2157 assert(getTypeSizeInBits(Op->getType()) < getTypeSizeInBits(Ty) &&
2158 "This is not an extending conversion!");
2159 assert(isSCEVable(Ty) &&
2160 "This is not a conversion to a SCEVable type!");
2161 Ty = getEffectiveSCEVType(Ty);
2162
2163 // Sign-extend negative constants.
2164 if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(Op))
2165 if (SC->getAPInt().isNegative())
2166 return getSignExtendExpr(Op, Ty);
2167
2168 // Peel off a truncate cast.
2170 const SCEV *NewOp = T->getOperand();
2171 if (getTypeSizeInBits(NewOp->getType()) < getTypeSizeInBits(Ty))
2172 return getAnyExtendExpr(NewOp, Ty);
2173 return getTruncateOrNoop(NewOp, Ty);
2174 }
2175
2176 // Next try a zext cast. If the cast is folded, use it.
2177 const SCEV *ZExt = getZeroExtendExpr(Op, Ty);
2178 if (!isa<SCEVZeroExtendExpr>(ZExt))
2179 return ZExt;
2180
2181 // Next try a sext cast. If the cast is folded, use it.
2182 const SCEV *SExt = getSignExtendExpr(Op, Ty);
2183 if (!isa<SCEVSignExtendExpr>(SExt))
2184 return SExt;
2185
2186 // Force the cast to be folded into the operands of an addrec.
2187 if (const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(Op)) {
2189 for (const SCEV *Op : AR->operands())
2190 Ops.push_back(getAnyExtendExpr(Op, Ty));
2191 return getAddRecExpr(Ops, AR->getLoop(), SCEV::FlagNW);
2192 }
2193
2194 // If the expression is obviously signed, use the sext cast value.
2195 if (isa<SCEVSMaxExpr>(Op))
2196 return SExt;
2197
2198 // Absent any other information, use the zext cast value.
2199 return ZExt;
2200}
2201
2202/// Process the given Ops list, which is a list of operands to be added under
2203/// the given scale, update the given map. This is a helper function for
2204/// getAddRecExpr. As an example of what it does, given a sequence of operands
2205/// that would form an add expression like this:
2206///
2207/// m + n + 13 + (A * (o + p + (B * (q + m + 29)))) + r + (-1 * r)
2208///
2209/// where A and B are constants, update the map with these values:
2210///
2211/// (m, 1+A*B), (n, 1), (o, A), (p, A), (q, A*B), (r, 0)
2212///
2213/// and add 13 + A*B*29 to AccumulatedConstant.
2214/// This will allow getAddRecExpr to produce this:
2215///
2216/// 13+A*B*29 + n + (m * (1+A*B)) + ((o + p) * A) + (q * A*B)
2217///
2218/// This form often exposes folding opportunities that are hidden in
2219/// the original operand list.
2220///
2221/// Return true iff it appears that any interesting folding opportunities
2222/// may be exposed. This helps getAddRecExpr short-circuit extra work in
2223/// the common case where no interesting opportunities are present, and
2224/// is also used as a check to avoid infinite recursion.
2227 APInt &AccumulatedConstant,
2229 const APInt &Scale,
2230 ScalarEvolution &SE) {
2231 bool Interesting = false;
2232
2233 // Iterate over the add operands. They are sorted, with constants first.
2234 unsigned i = 0;
2235 while (const SCEVConstant *C = dyn_cast<SCEVConstant>(Ops[i])) {
2236 ++i;
2237 // Pull a buried constant out to the outside.
2238 if (Scale != 1 || AccumulatedConstant != 0 || C->getValue()->isZero())
2239 Interesting = true;
2240 AccumulatedConstant += Scale * C->getAPInt();
2241 }
2242
2243 // Next comes everything else. We're especially interested in multiplies
2244 // here, but they're in the middle, so just visit the rest with one loop.
2245 for (; i != Ops.size(); ++i) {
2247 if (Mul && isa<SCEVConstant>(Mul->getOperand(0))) {
2248 APInt NewScale =
2249 Scale * cast<SCEVConstant>(Mul->getOperand(0))->getAPInt();
2250 if (Mul->getNumOperands() == 2 && isa<SCEVAddExpr>(Mul->getOperand(1))) {
2251 // A multiplication of a constant with another add; recurse.
2252 const SCEVAddExpr *Add = cast<SCEVAddExpr>(Mul->getOperand(1));
2253 Interesting |= CollectAddOperandsWithScales(
2254 M, NewOps, AccumulatedConstant, Add->operands(), NewScale, SE);
2255 } else {
2256 // A multiplication of a constant with some other value. Update
2257 // the map.
2258 SmallVector<SCEVUse, 4> MulOps(drop_begin(Mul->operands()));
2259 const SCEV *Key = SE.getMulExpr(MulOps);
2260 auto Pair = M.insert({Key, NewScale});
2261 if (Pair.second) {
2262 NewOps.push_back(Pair.first->first);
2263 } else {
2264 Pair.first->second += NewScale;
2265 // The map already had an entry for this value, which may indicate
2266 // a folding opportunity.
2267 Interesting = true;
2268 }
2269 }
2270 } else {
2271 // An ordinary operand. Update the map.
2272 auto Pair = M.insert({Ops[i], Scale});
2273 if (Pair.second) {
2274 NewOps.push_back(Pair.first->first);
2275 } else {
2276 Pair.first->second += Scale;
2277 // The map already had an entry for this value, which may indicate
2278 // a folding opportunity.
2279 Interesting = true;
2280 }
2281 }
2282 }
2283
2284 return Interesting;
2285}
2286
2288 const SCEV *LHS, const SCEV *RHS,
2289 const Instruction *CtxI) {
2290 auto Operation = [this, BinOp](SCEVUse L, SCEVUse R) -> const SCEV * {
2291 switch (BinOp) {
2292 default:
2293 llvm_unreachable("Unsupported binary op");
2294 case Instruction::Add:
2295 return getAddExpr(L, R);
2296 case Instruction::Sub:
2297 return getMinusSCEV(L, R);
2298 case Instruction::Mul:
2299 return getMulExpr(L, R);
2300 }
2301 };
2302
2303 const SCEV *(ScalarEvolution::*Extension)(SCEVUse, Type *, unsigned) =
2306
2307 // Check ext(LHS op RHS) == ext(LHS) op ext(RHS)
2308 auto *NarrowTy = cast<IntegerType>(LHS->getType());
2309 auto *WideTy =
2310 IntegerType::get(NarrowTy->getContext(), NarrowTy->getBitWidth() * 2);
2311
2312 const SCEV *A = (this->*Extension)(Operation(LHS, RHS), WideTy, 0);
2313 const SCEV *LHSB = (this->*Extension)(LHS, WideTy, 0);
2314 const SCEV *RHSB = (this->*Extension)(RHS, WideTy, 0);
2315 const SCEV *B = Operation(LHSB, RHSB);
2316 if (A == B)
2317 return true;
2318 // Can we use context to prove the fact we need?
2319 if (!CtxI)
2320 return false;
2321 // TODO: Support mul.
2322 if (BinOp == Instruction::Mul)
2323 return false;
2324 auto *RHSC = dyn_cast<SCEVConstant>(RHS);
2325 // TODO: Lift this limitation.
2326 if (!RHSC)
2327 return false;
2328 APInt C = RHSC->getAPInt();
2329 unsigned NumBits = C.getBitWidth();
2330 bool IsSub = (BinOp == Instruction::Sub);
2331 bool IsNegativeConst = (Signed && C.isNegative());
2332 // Compute the direction and magnitude by which we need to check overflow.
2333 bool OverflowDown = IsSub ^ IsNegativeConst;
2334 APInt Magnitude = C;
2335 if (IsNegativeConst) {
2336 if (C == APInt::getSignedMinValue(NumBits))
2337 // TODO: SINT_MIN on inversion gives the same negative value, we don't
2338 // want to deal with that.
2339 return false;
2340 Magnitude = -C;
2341 }
2342
2344 if (OverflowDown) {
2345 // To avoid overflow down, we need to make sure that MIN + Magnitude <= LHS.
2346 APInt Min = Signed ? APInt::getSignedMinValue(NumBits)
2347 : APInt::getMinValue(NumBits);
2348 APInt Limit = Min + Magnitude;
2349 return isKnownPredicateAt(Pred, getConstant(Limit), LHS, CtxI);
2350 } else {
2351 // To avoid overflow up, we need to make sure that LHS <= MAX - Magnitude.
2352 APInt Max = Signed ? APInt::getSignedMaxValue(NumBits)
2353 : APInt::getMaxValue(NumBits);
2354 APInt Limit = Max - Magnitude;
2355 return isKnownPredicateAt(Pred, LHS, getConstant(Limit), CtxI);
2356 }
2357}
2358
2360 const OverflowingBinaryOperator *OBO) {
2361 // It cannot be done any better.
2362 if (OBO->hasNoUnsignedWrap() && OBO->hasNoSignedWrap())
2363 return std::nullopt;
2364
2366
2367 if (OBO->hasNoUnsignedWrap())
2369 if (OBO->hasNoSignedWrap())
2371
2372 bool Deduced = false;
2373
2375 const SCEV *LHS = getSCEV(OBO->getOperand(0));
2376 const SCEV *RHS = getSCEV(OBO->getOperand(1));
2377
2378 bool CanUseNSW = true;
2379 const APInt *ShiftAmt;
2380 // Treat `shl %a, C` as `mul %a, 1 << C`.
2381 if (match(OBO, m_Shl(m_Value(), m_APInt(ShiftAmt)))) {
2382 unsigned BitWidth = ShiftAmt->getBitWidth();
2383 if (ShiftAmt->uge(BitWidth))
2384 return std::nullopt;
2385 // NSW only transfers if the shift amount is < BitWidth - 1, as INT_MIN * -1
2386 // overflows.
2387 CanUseNSW = ShiftAmt->ult(BitWidth - 1);
2388 Opcode = Instruction::Mul;
2390 } else if (Opcode != Instruction::Add && Opcode != Instruction::Sub &&
2391 Opcode != Instruction::Mul) {
2392 return std::nullopt;
2393 }
2394
2395 const Instruction *CtxI =
2397 if (!OBO->hasNoUnsignedWrap() &&
2398 willNotOverflow(Opcode, /* Signed */ false, LHS, RHS, CtxI)) {
2400 Deduced = true;
2401 }
2402
2403 if (CanUseNSW && !OBO->hasNoSignedWrap() &&
2404 willNotOverflow(Opcode, /* Signed */ true, LHS, RHS, CtxI)) {
2406 Deduced = true;
2407 }
2408
2409 if (Deduced)
2410 return Flags;
2411 return std::nullopt;
2412}
2413
2414// We're trying to construct a SCEV of type `Type' with `Ops' as operands and
2415// `OldFlags' as can't-wrap behavior. Infer a more aggressive set of
2416// can't-overflow flags for the operation if possible.
2419 using namespace std::placeholders;
2420
2421 using OBO = OverflowingBinaryOperator;
2422
2423 bool CanAnalyze =
2425 (void)CanAnalyze;
2426 assert(CanAnalyze && "don't call from other places!");
2427
2428 SCEVFlags SignOrUnsignMask = SCEV::FlagNUW | SCEV::FlagNSW;
2429 SCEVFlags SignOrUnsignWrap =
2430 ScalarEvolution::maskFlags(Flags, SignOrUnsignMask);
2431
2432 // If FlagNSW is true and all the operands are non-negative, infer FlagNUW.
2433 auto IsKnownNonNegative = [&](SCEVUse U) {
2434 return SE->isKnownNonNegative(U);
2435 };
2436
2437 if (SignOrUnsignWrap == SCEV::FlagNSW && all_of(Ops, IsKnownNonNegative))
2438 Flags = ScalarEvolution::setFlags(Flags, SignOrUnsignMask);
2439
2440 SignOrUnsignWrap = ScalarEvolution::maskFlags(Flags, SignOrUnsignMask);
2441
2442 if (SignOrUnsignWrap != SignOrUnsignMask &&
2443 (Type == scAddExpr || Type == scMulExpr) && Ops.size() == 2 &&
2444 isa<SCEVConstant>(Ops[0])) {
2445
2446 auto Opcode = [&] {
2447 switch (Type) {
2448 case scAddExpr:
2449 return Instruction::Add;
2450 case scMulExpr:
2451 return Instruction::Mul;
2452 default:
2453 llvm_unreachable("Unexpected SCEV op.");
2454 }
2455 }();
2456
2457 const APInt &C = cast<SCEVConstant>(Ops[0])->getAPInt();
2458
2459 // (A <opcode> C) --> (A <opcode> C)<nsw> if the op doesn't sign overflow.
2460 if (!(SignOrUnsignWrap & SCEV::FlagNSW)) {
2461 auto NSWRegion =
2462 ConstantRange::makeExactNoWrapRegion(Opcode, C, OBO::NoSignedWrap);
2463 if (NSWRegion.contains(SE->getSignedRange(Ops[1])))
2465 }
2466
2467 // (A <opcode> C) --> (A <opcode> C)<nuw> if the op doesn't unsign overflow.
2468 if (!(SignOrUnsignWrap & SCEV::FlagNUW)) {
2469 auto NUWRegion =
2470 ConstantRange::makeExactNoWrapRegion(Opcode, C, OBO::NoUnsignedWrap);
2471 if (NUWRegion.contains(SE->getUnsignedRange(Ops[1])))
2473 }
2474 }
2475
2476 // <0,+,nonnegative><nw> is also nuw
2477 // TODO: Add corresponding nsw case
2479 !ScalarEvolution::hasFlags(Flags, SCEV::FlagNUW) && Ops.size() == 2 &&
2480 Ops[0]->isZero() && IsKnownNonNegative(Ops[1]))
2482
2483 // both (udiv X, Y) * Y and Y * (udiv X, Y) are always NUW
2485 Ops.size() == 2) {
2486 if (auto *UDiv = dyn_cast<SCEVUDivExpr>(Ops[0]))
2487 if (UDiv->getOperand(1) == Ops[1])
2489 if (auto *UDiv = dyn_cast<SCEVUDivExpr>(Ops[1]))
2490 if (UDiv->getOperand(1) == Ops[0])
2492 }
2493
2494 return Flags;
2495}
2496
2498 return isLoopInvariant(S, L) && properlyDominates(S, L->getHeader());
2499}
2500
2501/// Get a canonical add expression, or something simpler if possible.
2503 SCEVFlagsPair Flags, unsigned Depth) {
2504 SCEVFlags ExprFlags = Flags.ExprFlags;
2505 SCEVFlags UseFlags = Flags.UseFlags;
2506 assert(!(ExprFlags & ~(SCEV::FlagNUW | SCEV::FlagNSW)) &&
2507 "only nuw or nsw allowed");
2508 assert(!(UseFlags & ~(SCEV::FlagNUW | SCEV::FlagNSW)) &&
2509 "only nuw or nsw allowed");
2510 assert(!Ops.empty() && "Cannot get empty add!");
2511 if (Ops.size() == 1) return Ops[0];
2512#ifndef NDEBUG
2513 Type *ETy = getEffectiveSCEVType(Ops[0]->getType());
2514 for (unsigned i = 1, e = Ops.size(); i != e; ++i)
2515 assert(getEffectiveSCEVType(Ops[i]->getType()) == ETy &&
2516 "SCEVAddExpr operand types don't match!");
2517 unsigned NumPtrs = count_if(
2518 Ops, [](const SCEV *Op) { return Op->getType()->isPointerTy(); });
2519 assert(NumPtrs <= 1 && "add has at most one pointer operand");
2520#endif
2521
2522 const SCEV *Folded = constantFoldAndGroupOps(
2523 *this, LI, DT, Ops,
2524 [](const APInt &C1, const APInt &C2) { return C1 + C2; },
2525 [](const APInt &C) { return C.isZero(); }, // identity
2526 [](const APInt &C) { return false; }); // absorber
2527 if (Folded)
2528 return Folded;
2529
2530#ifndef NDEBUG
2531 // Keep track of operands after constant folding, for verification when adding
2532 // use-specific flags.
2533 const SmallVector<SCEVUse, 8> OrigOps(Ops.begin(), Ops.end());
2534#endif
2535
2536 unsigned Idx = isa<SCEVConstant>(Ops[0]) ? 1 : 0;
2537
2538 // Delay expensive flag strengthening until necessary.
2539 auto ComputeFlags = [this, ExprFlags](ArrayRef<SCEVUse> Ops) {
2540 return StrengthenNoWrapFlags(this, scAddExpr, Ops, ExprFlags);
2541 };
2542
2543 // Limit recursion calls depth.
2545 return {getOrCreateAddExpr(Ops, ComputeFlags(Ops)), UseFlags};
2546
2547 if (SCEV *S = findExistingSCEVInCache(scAddExpr, Ops)) {
2548 // Don't strengthen flags if we have no new information.
2549 SCEVAddExpr *Add = static_cast<SCEVAddExpr *>(S);
2550 if (Add->getNoWrapFlags(ExprFlags) != ExprFlags)
2551 Add->setNoWrapFlags(ComputeFlags(Ops));
2552 return {S, UseFlags};
2553 }
2554
2555 // Okay, check to see if the same value occurs in the operand list more than
2556 // once. If so, merge them together into an multiply expression. Since we
2557 // sorted the list, these values are required to be adjacent.
2558 Type *Ty = Ops[0]->getType();
2559 bool FoundMatch = false;
2560 for (unsigned i = 0, e = Ops.size(); i != e-1; ++i)
2561 if (Ops[i] == Ops[i+1]) { // X + Y + Y --> X + Y*2
2562 // Scan ahead to count how many equal operands there are.
2563 unsigned Count = 2;
2564 while (i+Count != e && Ops[i+Count] == Ops[i])
2565 ++Count;
2566 // Merge the values into a multiply.
2567 SCEVUse Scale = getConstant(Ty, Count);
2568 const SCEV *Mul = getMulExpr(Scale, Ops[i], SCEV::FlagNone, Depth + 1);
2569 if (Ops.size() == Count)
2570 return Mul;
2571 Ops[i] = Mul;
2572 Ops.erase(Ops.begin()+i+1, Ops.begin()+i+Count);
2573 --i; e -= Count - 1;
2574 FoundMatch = true;
2575 }
2576 if (FoundMatch)
2577 return getAddExpr(Ops, ExprFlags, Depth + 1);
2578
2579 // Check for truncates. If all the operands are truncated from the same
2580 // type, see if factoring out the truncate would permit the result to be
2581 // folded. eg., n*trunc(x) + m*trunc(y) --> trunc(trunc(m)*x + trunc(n)*y)
2582 // if the contents of the resulting outer trunc fold to something simple.
2583 auto FindTruncSrcType = [&]() -> Type * {
2584 // We're ultimately looking to fold an addrec of truncs and muls of only
2585 // constants and truncs, so if we find any other types of SCEV
2586 // as operands of the addrec then we bail and return nullptr here.
2587 // Otherwise, we return the type of the operand of a trunc that we find.
2588 if (auto *T = dyn_cast<SCEVTruncateExpr>(Ops[Idx]))
2589 return T->getOperand()->getType();
2590 if (const auto *Mul = dyn_cast<SCEVMulExpr>(Ops[Idx])) {
2591 SCEVUse LastOp = Mul->getOperand(Mul->getNumOperands() - 1);
2592 if (const auto *T = dyn_cast<SCEVTruncateExpr>(LastOp))
2593 return T->getOperand()->getType();
2594 }
2595 return nullptr;
2596 };
2597 if (auto *SrcType = FindTruncSrcType()) {
2598 SmallVector<SCEVUse, 8> LargeOps;
2599 bool Ok = true;
2600 // Check all the operands to see if they can be represented in the
2601 // source type of the truncate.
2602 for (const SCEV *Op : Ops) {
2604 if (T->getOperand()->getType() != SrcType) {
2605 Ok = false;
2606 break;
2607 }
2608 LargeOps.push_back(T->getOperand());
2609 } else if (const SCEVConstant *C = dyn_cast<SCEVConstant>(Op)) {
2610 LargeOps.push_back(getAnyExtendExpr(C, SrcType));
2611 } else if (const SCEVMulExpr *M = dyn_cast<SCEVMulExpr>(Op)) {
2612 SmallVector<SCEVUse, 8> LargeMulOps;
2613 for (unsigned j = 0, f = M->getNumOperands(); j != f && Ok; ++j) {
2614 if (const SCEVTruncateExpr *T =
2615 dyn_cast<SCEVTruncateExpr>(M->getOperand(j))) {
2616 if (T->getOperand()->getType() != SrcType) {
2617 Ok = false;
2618 break;
2619 }
2620 LargeMulOps.push_back(T->getOperand());
2621 } else if (const auto *C = dyn_cast<SCEVConstant>(M->getOperand(j))) {
2622 LargeMulOps.push_back(getAnyExtendExpr(C, SrcType));
2623 } else {
2624 Ok = false;
2625 break;
2626 }
2627 }
2628 if (Ok)
2629 LargeOps.push_back(
2630 getMulExpr(LargeMulOps, SCEV::FlagNone, Depth + 1));
2631 } else {
2632 Ok = false;
2633 break;
2634 }
2635 }
2636 if (Ok) {
2637 // Evaluate the expression in the larger type.
2638 const SCEV *Fold = getAddExpr(LargeOps, SCEV::FlagNone, Depth + 1);
2639 // If it folds to something simple, use it. Otherwise, don't.
2640 if (isa<SCEVConstant>(Fold) || isa<SCEVUnknown>(Fold))
2641 return getTruncateExpr(Fold, Ty);
2642 }
2643 }
2644
2645 if (Ops.size() == 2) {
2646 // Check if we have an expression of the form ((X + C1) - C2), where C1 and
2647 // C2 can be folded in a way that allows retaining wrapping flags of (X +
2648 // C1).
2649 const SCEV *A = Ops[0];
2650 const SCEV *B = Ops[1];
2651 auto *AddExpr = dyn_cast<SCEVAddExpr>(B);
2652 auto *C = dyn_cast<SCEVConstant>(A);
2653 if (AddExpr && C && isa<SCEVConstant>(AddExpr->getOperand(0))) {
2654 auto C1 = cast<SCEVConstant>(AddExpr->getOperand(0))->getAPInt();
2655 auto C2 = C->getAPInt();
2656 SCEVFlags PreservedFlags = SCEV::FlagNone;
2657
2658 APInt ConstAdd = C1 + C2;
2659 auto AddFlags = AddExpr->getNoWrapFlags();
2660 // Adding a smaller constant is NUW if the original AddExpr was NUW.
2662 ConstAdd.ule(C1)) {
2663 PreservedFlags =
2665 }
2666
2667 // Adding a constant with the same sign and small magnitude is NSW, if the
2668 // original AddExpr was NSW.
2670 C1.isSignBitSet() == ConstAdd.isSignBitSet() &&
2671 ConstAdd.abs().ule(C1.abs())) {
2672 PreservedFlags =
2674 }
2675
2676 if (PreservedFlags != SCEV::FlagNone) {
2677 SmallVector<SCEVUse, 4> NewOps(AddExpr->operands());
2678 NewOps[0] = getConstant(ConstAdd);
2679 return getAddExpr(NewOps, PreservedFlags);
2680 }
2681 }
2682
2683 // Try to push the constant operand into a ZExt: A + zext (-A + B) -> zext
2684 // (B), if trunc (A) + -A + B does not unsigned-wrap.
2685 const SCEVAddExpr *InnerAdd;
2686 if (match(B, m_scev_ZExt(m_scev_Add(InnerAdd)))) {
2687 const SCEV *NarrowA = getTruncateExpr(A, InnerAdd->getType());
2688 if (NarrowA == getNegativeSCEV(InnerAdd->getOperand(0)) &&
2689 getZeroExtendExpr(NarrowA, B->getType()) == A &&
2690 hasFlags(StrengthenNoWrapFlags(this, scAddExpr, {NarrowA, InnerAdd},
2692 SCEV::FlagNUW)) {
2693 return getZeroExtendExpr(getAddExpr(NarrowA, InnerAdd), B->getType());
2694 }
2695 }
2696 }
2697
2698 // Canonicalize (-1 * urem X, Y) + X --> (Y * X/Y)
2699 const SCEV *Y;
2700 if (Ops.size() == 2 &&
2701 match(Ops[0],
2703 m_scev_URem(m_scev_Specific(Ops[1]), m_SCEV(Y), *this))))
2704 return getMulExpr(Y, getUDivExpr(Ops[1], Y));
2705
2706 // Skip past any other cast SCEVs.
2707 while (Idx < Ops.size() && Ops[Idx]->getSCEVType() < scAddExpr)
2708 ++Idx;
2709
2710 // If there are add operands they would be next.
2711 if (Idx < Ops.size()) {
2712 bool DeletedAdd = false;
2713 // If the original flags and all inlined SCEVAddExprs are NUW, use the
2714 // common NUW flag for expression after inlining. Other flags cannot be
2715 // preserved, because they may depend on the original order of operations.
2716 SCEVFlags CommonFlags = maskFlags(ExprFlags, SCEV::FlagNUW);
2717 while (const SCEVAddExpr *Add = dyn_cast<SCEVAddExpr>(Ops[Idx])) {
2718 if (Ops.size() > AddOpsInlineThreshold ||
2719 Add->getNumOperands() > AddOpsInlineThreshold)
2720 break;
2721 // If we have an add, expand the add operands onto the end of the operands
2722 // list.
2723 Ops.erase(Ops.begin()+Idx);
2724 append_range(Ops, Add->operands());
2725 DeletedAdd = true;
2726 CommonFlags = maskFlags(CommonFlags, Add->getNoWrapFlags());
2727 }
2728
2729 // If we deleted at least one add, we added operands to the end of the list,
2730 // and they are not necessarily sorted. Recurse to resort and resimplify
2731 // any operands we just acquired.
2732 if (DeletedAdd)
2733 return getAddExpr(Ops, CommonFlags, Depth + 1);
2734 }
2735
2736 // Skip over the add expression until we get to a multiply.
2737 while (Idx < Ops.size() && Ops[Idx]->getSCEVType() < scMulExpr)
2738 ++Idx;
2739
2740 // Check to see if there are any folding opportunities present with
2741 // operands multiplied by constant values.
2742 if (Idx < Ops.size() && isa<SCEVMulExpr>(Ops[Idx])) {
2743 uint64_t BitWidth = getTypeSizeInBits(Ty);
2746 APInt AccumulatedConstant(BitWidth, 0);
2747 if (CollectAddOperandsWithScales(M, NewOps, AccumulatedConstant,
2748 Ops, APInt(BitWidth, 1), *this)) {
2749 struct APIntCompare {
2750 bool operator()(const APInt &LHS, const APInt &RHS) const {
2751 return LHS.ult(RHS);
2752 }
2753 };
2754
2755 // Some interesting folding opportunity is present, so its worthwhile to
2756 // re-generate the operands list. Group the operands by constant scale,
2757 // to avoid multiplying by the same constant scale multiple times.
2758 std::map<APInt, SmallVector<SCEVUse, 4>, APIntCompare> MulOpLists;
2759 for (SCEVUse NewOp : NewOps)
2760 MulOpLists[M.find(NewOp)->second].push_back(NewOp);
2761 // Re-generate the operands list.
2762 Ops.clear();
2763 if (AccumulatedConstant != 0)
2764 Ops.push_back(getConstant(AccumulatedConstant));
2765 for (auto &MulOp : MulOpLists) {
2766 if (MulOp.first == 1) {
2767 Ops.push_back(getAddExpr(MulOp.second, SCEV::FlagNone, Depth + 1));
2768 } else if (MulOp.first != 0) {
2769 Ops.push_back(
2770 getMulExpr(getConstant(MulOp.first),
2771 getAddExpr(MulOp.second, SCEV::FlagNone, Depth + 1),
2772 SCEV::FlagNone, Depth + 1));
2773 }
2774 }
2775 if (Ops.empty())
2776 return getZero(Ty);
2777 if (Ops.size() == 1)
2778 return Ops[0];
2779 return getAddExpr(Ops, SCEV::FlagNone, Depth + 1);
2780 }
2781 }
2782
2783 // Given a SCEVMulExpr and an operand index, return the product of all
2784 // operands except the one at OpIdx.
2785 auto StripFactor = [&](const SCEVMulExpr *M, unsigned OpIdx) -> SCEVUse {
2786 if (M->getNumOperands() == 2)
2787 return M->getOperand(OpIdx == 0);
2788 SmallVector<SCEVUse, 4> Remaining(M->operands().take_front(OpIdx));
2789 append_range(Remaining, M->operands().drop_front(OpIdx + 1));
2790 return getMulExpr(Remaining, SCEV::FlagNone, Depth + 1);
2791 };
2792
2793 // If we are adding something to a multiply expression, make sure the
2794 // something is not already an operand of the multiply. If so, merge it into
2795 // the multiply.
2796 for (; Idx < Ops.size() && isa<SCEVMulExpr>(Ops[Idx]); ++Idx) {
2797 const SCEVMulExpr *Mul = cast<SCEVMulExpr>(Ops[Idx]);
2798 for (unsigned MulOp = 0, e = Mul->getNumOperands(); MulOp != e; ++MulOp) {
2799 // Scan all terms to find every occurrence of common factor MulOpSCEV
2800 // and fold them in one shot:
2801 // A1*X + A2*X + ... + An*X --> X * (A1 + A2 + ... + An)
2802 const SCEV *MulOpSCEV = Mul->getOperand(MulOp);
2803 if (isa<SCEVConstant>(MulOpSCEV))
2804 continue;
2805
2806 // Cofactors: 1 for bare addends matching MulOpSCEV, or the
2807 // remaining product for multiply terms containing MulOpSCEV.
2808 SmallVector<SCEVUse, 4> Cofactors;
2809 SmallVector<unsigned, 4> DeadIndices;
2810 for (unsigned AddOp = 0, e = Ops.size(); AddOp != e; ++AddOp) {
2811 if (MulOpSCEV == Ops[AddOp]) {
2812 // W + X + (X * Y * Z) --> W + (X * ((Y*Z)+1))
2813 Cofactors.push_back(getOne(Ty));
2814 DeadIndices.push_back(AddOp);
2815 continue;
2816 }
2817
2818 if (AddOp <= Idx || !isa<SCEVMulExpr>(Ops[AddOp]))
2819 continue;
2820
2821 const SCEVMulExpr *OtherMul = cast<SCEVMulExpr>(Ops[AddOp]);
2822 for (unsigned OMulOp = 0, OE = OtherMul->getNumOperands(); OMulOp != OE;
2823 ++OMulOp) {
2824 if (OtherMul->getOperand(OMulOp) == MulOpSCEV) {
2825 // (A*B*C) + (A*D*E) --> A * (B*C + D*E)
2826 Cofactors.push_back(StripFactor(OtherMul, OMulOp));
2827 DeadIndices.push_back(AddOp);
2828 break;
2829 }
2830 }
2831 }
2832
2833 // Fold all collected cofactors with the anchor multiply's cofactor:
2834 // MulOpSCEV * (Cofactor_1 + ... + Cofactor_n + AnchorCofactor)
2835 if (!Cofactors.empty()) {
2836 Cofactors.push_back(StripFactor(Mul, MulOp));
2837
2838 SCEVUse InnerSum = getAddExpr(Cofactors, SCEV::FlagNone, Depth + 1);
2839 SCEVUse OuterMul =
2840 getMulExpr(MulOpSCEV, InnerSum, SCEV::FlagNone, Depth + 1);
2841
2842 // DeadIndices does not include Idx (the anchor), hence +1.
2843 if (Ops.size() == DeadIndices.size() + 1)
2844 return OuterMul;
2845
2846 // Erase Ops[Idx] first, then erase DeadIndices in reverse order.
2847 // The -1 adjustment accounts for the shift from removing Idx;
2848 // reverse order means each erasure only shifts later positions,
2849 // which have already been processed.
2850 Ops.erase(Ops.begin() + Idx);
2851 for (unsigned Dead : reverse(DeadIndices))
2852 Ops.erase(Ops.begin() + (Dead > Idx ? Dead - 1 : Dead));
2853
2854 Ops.push_back(OuterMul);
2855 return getAddExpr(Ops, SCEV::FlagNone, Depth + 1);
2856 }
2857 }
2858 }
2859
2860 // If there are any add recurrences in the operands list, see if any other
2861 // added values are loop invariant. If so, we can fold them into the
2862 // recurrence.
2863 while (Idx < Ops.size() && Ops[Idx]->getSCEVType() < scAddRecExpr)
2864 ++Idx;
2865
2866 // Scan over all recurrences, trying to fold loop invariants into them.
2867 for (; Idx < Ops.size() && isa<SCEVAddRecExpr>(Ops[Idx]); ++Idx) {
2868 // Scan all of the other operands to this add and add them to the vector if
2869 // they are loop invariant w.r.t. the recurrence.
2871 const SCEVAddRecExpr *AddRec = cast<SCEVAddRecExpr>(Ops[Idx]);
2872 const Loop *AddRecLoop = AddRec->getLoop();
2873 for (unsigned i = 0, e = Ops.size(); i != e; ++i)
2874 if (isAvailableAtLoopEntry(Ops[i], AddRecLoop)) {
2875 LIOps.push_back(Ops[i]);
2876 Ops.erase(Ops.begin()+i);
2877 --i; --e;
2878 }
2879
2880 // If we found some loop invariants, fold them into the recurrence.
2881 if (!LIOps.empty()) {
2882 // Compute nowrap flags for the addition of the loop-invariant ops and
2883 // the addrec. Temporarily push it as an operand for that purpose. These
2884 // flags are valid in the scope of the addrec only.
2885 LIOps.push_back(AddRec);
2886 SCEVFlags Flags = ComputeFlags(LIOps);
2887 LIOps.pop_back();
2888
2889 // NLI + LI + {Start,+,Step} --> NLI + {LI+Start,+,Step}
2890 LIOps.push_back(AddRec->getStart());
2891
2892 SmallVector<SCEVUse, 4> AddRecOps(AddRec->operands());
2893
2894 // It is not in general safe to propagate flags valid on an add within
2895 // the addrec scope to one outside it. We must prove that the inner
2896 // scope is guaranteed to execute if the outer one does to be able to
2897 // safely propagate. We know the program is undefined if poison is
2898 // produced on the inner scoped addrec. We also know that *for this use*
2899 // the outer scoped add can't overflow (because of the flags we just
2900 // computed for the inner scoped add) without the program being undefined.
2901 // Proving that entry to the outer scope neccesitates entry to the inner
2902 // scope, thus proves the program undefined if the flags would be violated
2903 // in the outer scope.
2904 SCEVFlags AddFlags = Flags;
2905 if (AddFlags != SCEV::FlagNone) {
2906 auto *DefI = getDefiningScopeBound(LIOps);
2907 auto *ReachI = &*AddRecLoop->getHeader()->begin();
2908 if (!isGuaranteedToTransferExecutionTo(DefI, ReachI))
2909 AddFlags = SCEV::FlagNone;
2910 }
2911 AddRecOps[0] = getAddExpr(LIOps, AddFlags, Depth + 1);
2912
2913 // Build the new addrec. Propagate the NUW and NSW flags if both the
2914 // outer add and the inner addrec are guaranteed to have no overflow.
2915 // Always propagate NW.
2916 Flags = AddRec->getNoWrapFlags(setFlags(Flags, SCEV::FlagNW));
2917 const SCEV *NewRec = getAddRecExpr(AddRecOps, AddRecLoop, Flags);
2918
2919 // If all of the other operands were loop invariant, we are done.
2920 if (Ops.size() == 1) return NewRec;
2921
2922 // Otherwise, add the folded AddRec by the non-invariant parts.
2923 for (unsigned i = 0;; ++i)
2924 if (Ops[i] == AddRec) {
2925 Ops[i] = NewRec;
2926 break;
2927 }
2928 return getAddExpr(Ops, SCEV::FlagNone, Depth + 1);
2929 }
2930
2931 // Okay, if there weren't any loop invariants to be folded, check to see if
2932 // there are multiple AddRec's with the same loop induction variable being
2933 // added together. If so, we can fold them.
2934 for (unsigned OtherIdx = Idx+1;
2935 OtherIdx < Ops.size() && isa<SCEVAddRecExpr>(Ops[OtherIdx]);
2936 ++OtherIdx) {
2937 // We expect the AddRecExpr's to be sorted in reverse dominance order,
2938 // so that the 1st found AddRecExpr is dominated by all others.
2939 assert(DT.dominates(
2940 cast<SCEVAddRecExpr>(Ops[OtherIdx])->getLoop()->getHeader(),
2941 AddRec->getLoop()->getHeader()) &&
2942 "AddRecExprs are not sorted in reverse dominance order?");
2943 if (AddRecLoop == cast<SCEVAddRecExpr>(Ops[OtherIdx])->getLoop()) {
2944 // Other + {A,+,B}<L> + {C,+,D}<L> --> Other + {A+C,+,B+D}<L>
2945 SmallVector<SCEVUse, 4> AddRecOps(AddRec->operands());
2946 for (; OtherIdx != Ops.size() && isa<SCEVAddRecExpr>(Ops[OtherIdx]);
2947 ++OtherIdx) {
2948 const auto *OtherAddRec = cast<SCEVAddRecExpr>(Ops[OtherIdx]);
2949 if (OtherAddRec->getLoop() == AddRecLoop) {
2950 for (unsigned i = 0, e = OtherAddRec->getNumOperands();
2951 i != e; ++i) {
2952 if (i >= AddRecOps.size()) {
2953 append_range(AddRecOps, OtherAddRec->operands().drop_front(i));
2954 break;
2955 }
2956 AddRecOps[i] =
2957 getAddExpr(AddRecOps[i], OtherAddRec->getOperand(i),
2958 SCEV::FlagNone, Depth + 1);
2959 }
2960 Ops.erase(Ops.begin() + OtherIdx); --OtherIdx;
2961 }
2962 }
2963 // Step size has changed, so we cannot guarantee no self-wraparound.
2964 Ops[Idx] = getAddRecExpr(AddRecOps, AddRecLoop, SCEV::FlagNone);
2965 return getAddExpr(Ops, SCEV::FlagNone, Depth + 1);
2966 }
2967 }
2968
2969 // Otherwise couldn't fold anything into this recurrence. Move onto the
2970 // next one.
2971 }
2972
2973 // Okay, it looks like we really DO need an add expr. Check to see if we
2974 // already have one, otherwise create a new one.
2975 assert((UseFlags == SCEV::FlagNone || equal(OrigOps, Ops)) &&
2976 "Tried to add SCEVUse flags after operands changed");
2977 return {getOrCreateAddExpr(Ops, ComputeFlags(Ops)), UseFlags};
2978}
2979
2980const SCEV *ScalarEvolution::getOrCreateAddExpr(ArrayRef<SCEVUse> Ops,
2981 SCEVFlags Flags) {
2984 for (SCEVUse Op : Ops)
2985 ID.AddPointer(Op.getOpaqueValue());
2987 SCEVAddExpr *S = static_cast<SCEVAddExpr *>(UniqueSCEVs.lookup(ID, Token));
2988 if (!S) {
2989 SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
2991 S = new (SCEVAllocator)
2992 SCEVAddExpr(ID.Intern(SCEVAllocator), O, Ops.size());
2993 UniqueSCEVs.insert(S, Token);
2994 S->computeAndSetCanonical(*this);
2995 registerUser(S, Ops);
2996 }
2997 S->setNoWrapFlags(Flags);
2998 return S;
2999}
3000
3001const SCEV *ScalarEvolution::getOrCreateAddRecExpr(ArrayRef<SCEVUse> Ops,
3002 const Loop *L,
3003 SCEVFlags Flags) {
3004 FoldingSetNodeID ID;
3005 ID.AddInteger(scAddRecExpr);
3006 for (SCEVUse Op : Ops)
3007 ID.AddPointer(Op.getOpaqueValue());
3008 ID.AddPointer(L);
3009 FoldingSetInsertToken Token;
3010 SCEVAddRecExpr *S =
3011 static_cast<SCEVAddRecExpr *>(UniqueSCEVs.lookup(ID, Token));
3012 if (!S) {
3013 SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
3015 S = new (SCEVAllocator)
3016 SCEVAddRecExpr(ID.Intern(SCEVAllocator), O, Ops.size(), L);
3017 UniqueSCEVs.insert(S, Token);
3018 S->computeAndSetCanonical(*this);
3019 LoopUsers[L].push_back(S);
3020 registerUser(S, Ops);
3021 }
3022 setNoWrapFlags(S, Flags);
3023 return S;
3024}
3025
3026const SCEV *ScalarEvolution::getOrCreateMulExpr(ArrayRef<SCEVUse> Ops,
3027 SCEVFlags Flags) {
3028 FoldingSetNodeID ID;
3029 ID.AddInteger(scMulExpr);
3030 for (SCEVUse Op : Ops)
3031 ID.AddPointer(Op.getOpaqueValue());
3032 FoldingSetInsertToken Token;
3033 SCEVMulExpr *S = static_cast<SCEVMulExpr *>(UniqueSCEVs.lookup(ID, Token));
3034 if (!S) {
3035 SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
3037 S = new (SCEVAllocator) SCEVMulExpr(ID.Intern(SCEVAllocator),
3038 O, Ops.size());
3039 UniqueSCEVs.insert(S, Token);
3040 S->computeAndSetCanonical(*this);
3041 registerUser(S, Ops);
3042 }
3043 S->setNoWrapFlags(Flags);
3044 return S;
3045}
3046
3047const SCEV *ScalarEvolution::getOrCreateUDivExpr(SCEVUse LHS, SCEVUse RHS) {
3048 FoldingSetNodeID ID;
3049 ID.AddInteger(scUDivExpr);
3050 ID.AddPointer(LHS.getOpaqueValue());
3051 ID.AddPointer(RHS.getOpaqueValue());
3052 FoldingSetInsertToken Token;
3053 SCEV *S = UniqueSCEVs.lookup(ID, Token);
3054 if (!S) {
3055 S = new (SCEVAllocator) SCEVUDivExpr(ID.Intern(SCEVAllocator), LHS, RHS);
3056 UniqueSCEVs.insert(S, Token);
3057 S->computeAndSetCanonical(*this);
3058 registerUser(S, {LHS, RHS});
3059 }
3060 return S;
3061}
3062
3063static uint64_t umul_ov(uint64_t i, uint64_t j, bool &Overflow) {
3064 uint64_t k = i*j;
3065 if (j > 1 && k / j != i) Overflow = true;
3066 return k;
3067}
3068
3069/// Compute the result of "n choose k", the binomial coefficient. If an
3070/// intermediate computation overflows, Overflow will be set and the return will
3071/// be garbage. Overflow is not cleared on absence of overflow.
3072static uint64_t Choose(uint64_t n, uint64_t k, bool &Overflow) {
3073 // We use the multiplicative formula:
3074 // n(n-1)(n-2)...(n-(k-1)) / k(k-1)(k-2)...1 .
3075 // At each iteration, we take the n-th term of the numeral and divide by the
3076 // (k-n)th term of the denominator. This division will always produce an
3077 // integral result, and helps reduce the chance of overflow in the
3078 // intermediate computations. However, we can still overflow even when the
3079 // final result would fit.
3080
3081 if (n == 0 || n == k) return 1;
3082 if (k > n) return 0;
3083
3084 if (k > n/2)
3085 k = n-k;
3086
3087 uint64_t r = 1;
3088 for (uint64_t i = 1; i <= k; ++i) {
3089 r = umul_ov(r, n-(i-1), Overflow);
3090 r /= i;
3091 }
3092 return r;
3093}
3094
3095/// Determine if any of the operands in this SCEV are a constant or if
3096/// any of the add or multiply expressions in this SCEV contain a constant.
3097static bool containsConstantInAddMulChain(const SCEV *StartExpr) {
3098 struct FindConstantInAddMulChain {
3099 bool FoundConstant = false;
3100
3101 bool follow(const SCEV *S) {
3102 FoundConstant |= isa<SCEVConstant>(S);
3103 return isa<SCEVAddExpr>(S) || isa<SCEVMulExpr>(S);
3104 }
3105
3106 bool isDone() const {
3107 return FoundConstant;
3108 }
3109 };
3110
3111 FindConstantInAddMulChain F;
3113 ST.visitAll(StartExpr);
3114 return F.FoundConstant;
3115}
3116
3117/// Get a canonical multiply expression, or something simpler if possible.
3119 SCEVFlagsPair Flags, unsigned Depth) {
3120 SCEVFlags ExprFlags = Flags.ExprFlags;
3121 SCEVFlags UseFlags = Flags.UseFlags;
3122 assert(ExprFlags == maskFlags(ExprFlags, SCEV::FlagNUW | SCEV::FlagNSW) &&
3123 "only nuw or nsw allowed");
3124 assert(UseFlags == maskFlags(UseFlags, SCEV::FlagNUW | SCEV::FlagNSW) &&
3125 "only nuw or nsw allowed");
3126 assert(!Ops.empty() && "Cannot get empty mul!");
3127 if (Ops.size() == 1) return Ops[0];
3128#ifndef NDEBUG
3129 Type *ETy = Ops[0]->getType();
3130 assert(!ETy->isPointerTy());
3131 for (unsigned i = 1, e = Ops.size(); i != e; ++i)
3132 assert(Ops[i]->getType() == ETy &&
3133 "SCEVMulExpr operand types don't match!");
3134#endif
3135
3136 const SCEV *Folded = constantFoldAndGroupOps(
3137 *this, LI, DT, Ops,
3138 [](const APInt &C1, const APInt &C2) { return C1 * C2; },
3139 [](const APInt &C) { return C.isOne(); }, // identity
3140 [](const APInt &C) { return C.isZero(); }); // absorber
3141 if (Folded)
3142 return Folded;
3143
3144#ifndef NDEBUG
3145 // Keep track of operands after constant folding, for verification when adding
3146 // use-specific flags.
3147 const SmallVector<SCEVUse, 8> OrigOps(Ops.begin(), Ops.end());
3148#endif
3149
3150 // Delay expensive flag strengthening until necessary.
3151 auto ComputeFlags = [this, ExprFlags](const ArrayRef<SCEVUse> Ops) {
3152 return StrengthenNoWrapFlags(this, scMulExpr, Ops, ExprFlags);
3153 };
3154
3155 // Limit recursion calls depth.
3157 return {getOrCreateMulExpr(Ops, ComputeFlags(Ops)), UseFlags};
3158
3159 if (SCEV *S = findExistingSCEVInCache(scMulExpr, Ops)) {
3160 // Don't strengthen flags if we have no new information.
3161 SCEVMulExpr *Mul = static_cast<SCEVMulExpr *>(S);
3162 if (Mul->getNoWrapFlags(ExprFlags) != ExprFlags)
3163 Mul->setNoWrapFlags(ComputeFlags(Ops));
3164 return {S, UseFlags};
3165 }
3166
3167 if (const SCEVConstant *LHSC = dyn_cast<SCEVConstant>(Ops[0])) {
3168 if (Ops.size() == 2) {
3169 // C1*(C2+V) -> C1*C2 + C1*V
3170 // If any of Add's ops are Adds or Muls with a constant, apply this
3171 // transformation as well.
3172 //
3173 // TODO: There are some cases where this transformation is not
3174 // profitable; for example, Add = (C0 + X) * Y + Z. Maybe the scope of
3175 // this transformation should be narrowed down.
3176 const SCEV *Op0, *Op1;
3177 if (match(Ops[1], m_scev_Add(m_SCEV(Op0), m_SCEV(Op1))) &&
3179 const SCEV *LHS = getMulExpr(LHSC, Op0, SCEV::FlagNone, Depth + 1);
3180 const SCEV *RHS = getMulExpr(LHSC, Op1, SCEV::FlagNone, Depth + 1);
3181 return getAddExpr(LHS, RHS, SCEV::FlagNone, Depth + 1);
3182 }
3183
3184 if (Ops[0]->isAllOnesValue()) {
3185 // If we have a mul by -1 of an add, try distributing the -1 among the
3186 // add operands.
3187 if (const SCEVAddExpr *Add = dyn_cast<SCEVAddExpr>(Ops[1])) {
3189 bool AnyFolded = false;
3190 for (const SCEV *AddOp : Add->operands()) {
3191 const SCEV *Mul =
3192 getMulExpr(Ops[0], SCEVUse(AddOp), SCEV::FlagNone, Depth + 1);
3193 if (!isa<SCEVMulExpr>(Mul)) AnyFolded = true;
3194 NewOps.push_back(Mul);
3195 }
3196 if (AnyFolded)
3197 return getAddExpr(NewOps, SCEV::FlagNone, Depth + 1);
3198 } else if (const auto *AddRec = dyn_cast<SCEVAddRecExpr>(Ops[1])) {
3199 // Negation preserves a recurrence's no self-wrap property.
3201 for (const SCEV *AddRecOp : AddRec->operands())
3202 Operands.push_back(getMulExpr(Ops[0], SCEVUse(AddRecOp),
3203 SCEV::FlagNone, Depth + 1));
3204 // Let M be the minimum representable signed value. AddRec with nsw
3205 // multiplied by -1 can have signed overflow if and only if it takes a
3206 // value of M: M * (-1) would stay M and (M + 1) * (-1) would be the
3207 // maximum signed value. In all other cases signed overflow is
3208 // impossible.
3209 auto FlagsMask = SCEV::FlagNW;
3210 if (AddRec->hasNoSignedWrap()) {
3211 auto MinInt =
3212 APInt::getSignedMinValue(getTypeSizeInBits(AddRec->getType()));
3213 if (getSignedRangeMin(AddRec) != MinInt)
3215 }
3216 return getAddRecExpr(Operands, AddRec->getLoop(),
3217 AddRec->getNoWrapFlags(FlagsMask));
3218 }
3219 }
3220
3221 // Try to push the constant operand into a ZExt: C * zext (A + B) ->
3222 // zext (C*A + C*B) if trunc (C) * (A + B) does not unsigned-wrap.
3223 const SCEVAddExpr *InnerAdd;
3224 if (match(Ops[1], m_scev_ZExt(m_scev_Add(InnerAdd)))) {
3225 const SCEV *NarrowC = getTruncateExpr(LHSC, InnerAdd->getType());
3226 if (isa<SCEVConstant>(InnerAdd->getOperand(0)) &&
3227 getZeroExtendExpr(NarrowC, Ops[1]->getType()) == LHSC &&
3228 hasFlags(StrengthenNoWrapFlags(this, scMulExpr, {NarrowC, InnerAdd},
3230 SCEV::FlagNUW)) {
3231 const SCEV *Res =
3232 getMulExpr(NarrowC, InnerAdd, SCEV::FlagNUW, Depth + 1);
3233 return getZeroExtendExpr(Res, Ops[1]->getType(), Depth + 1);
3234 };
3235 }
3236
3237 // Try to fold (C1 * D /u C2) -> C1/C2 * D, if C1 and C2 are powers-of-2,
3238 // D is a multiple of C2, and C1 is a multiple of C2. If C2 is a multiple
3239 // of C1, fold to (D /u (C2 /u C1)).
3240 const SCEV *D;
3241 APInt C1V = LHSC->getAPInt();
3242 // (C1 * D /u C2) == -1 * -C1 * D /u C2 when C1 != INT_MIN. Don't treat -1
3243 // as -1 * 1, as it won't enable additional folds.
3244 if (C1V.isNegative() && !C1V.isMinSignedValue() && !C1V.isAllOnes())
3245 C1V = C1V.abs();
3246 const SCEVConstant *C2;
3247 if (C1V.isPowerOf2() &&
3249 C2->getAPInt().isPowerOf2() &&
3250 C1V.logBase2() <= getMinTrailingZeros(D)) {
3251 const SCEV *NewMul = nullptr;
3252 if (C1V.uge(C2->getAPInt())) {
3253 NewMul = getMulExpr(getUDivExpr(getConstant(C1V), C2), D);
3254 } else if (C2->getAPInt().logBase2() <= getMinTrailingZeros(D)) {
3255 assert(C1V.ugt(1) && "C1 <= 1 should have been folded earlier");
3256 NewMul = getUDivExpr(D, getUDivExpr(C2, getConstant(C1V)));
3257 }
3258 if (NewMul)
3259 return C1V == LHSC->getAPInt() ? NewMul : getNegativeSCEV(NewMul);
3260 }
3261 }
3262 }
3263
3264 // Skip over the add expression until we get to a multiply.
3265 unsigned Idx = 0;
3266 while (Idx < Ops.size() && Ops[Idx]->getSCEVType() < scMulExpr)
3267 ++Idx;
3268
3269 // If there are mul operands inline them all into this expression.
3270 if (Idx < Ops.size()) {
3271 bool DeletedMul = false;
3272 while (const SCEVMulExpr *Mul = dyn_cast<SCEVMulExpr>(Ops[Idx])) {
3273 if (Ops.size() > MulOpsInlineThreshold)
3274 break;
3275 // If we have an mul, expand the mul operands onto the end of the
3276 // operands list.
3277 Ops.erase(Ops.begin()+Idx);
3278 append_range(Ops, Mul->operands());
3279 DeletedMul = true;
3280 }
3281
3282 // If we deleted at least one mul, we added operands to the end of the
3283 // list, and they are not necessarily sorted. Recurse to resort and
3284 // resimplify any operands we just acquired.
3285 if (DeletedMul)
3286 return getMulExpr(Ops, SCEV::FlagNone, Depth + 1);
3287 }
3288
3289 // If there are any add recurrences in the operands list, see if any other
3290 // added values are loop invariant. If so, we can fold them into the
3291 // recurrence.
3292 while (Idx < Ops.size() && Ops[Idx]->getSCEVType() < scAddRecExpr)
3293 ++Idx;
3294
3295 // Scan over all recurrences, trying to fold loop invariants into them.
3296 for (; Idx < Ops.size() && isa<SCEVAddRecExpr>(Ops[Idx]); ++Idx) {
3297 // Scan all of the other operands to this mul and add them to the vector
3298 // if they are loop invariant w.r.t. the recurrence.
3300 const SCEVAddRecExpr *AddRec = cast<SCEVAddRecExpr>(Ops[Idx]);
3301 for (unsigned i = 0, e = Ops.size(); i != e; ++i)
3302 if (isAvailableAtLoopEntry(Ops[i], AddRec->getLoop())) {
3303 LIOps.push_back(Ops[i]);
3304 Ops.erase(Ops.begin()+i);
3305 --i; --e;
3306 }
3307
3308 // If we found some loop invariants, fold them into the recurrence.
3309 if (!LIOps.empty()) {
3310 // NLI * LI * {Start,+,Step} --> NLI * {LI*Start,+,LI*Step}
3312 NewOps.reserve(AddRec->getNumOperands());
3313 const SCEV *Scale = getMulExpr(LIOps, SCEV::FlagNone, Depth + 1);
3314
3315 // If both the mul and addrec are nuw, we can preserve nuw.
3316 // If both the mul and addrec are nsw, we can only preserve nsw if either
3317 // a) they are also nuw, or
3318 // b) all multiplications of addrec operands with scale are nsw.
3319 SCEVFlags Flags = AddRec->getNoWrapFlags(ComputeFlags({Scale, AddRec}));
3320
3321 for (unsigned i = 0, e = AddRec->getNumOperands(); i != e; ++i) {
3322 NewOps.push_back(getMulExpr(Scale, AddRec->getOperand(i),
3323 SCEV::FlagNone, Depth + 1));
3324
3325 if (hasFlags(Flags, SCEV::FlagNSW) && !hasFlags(Flags, SCEV::FlagNUW)) {
3327 Instruction::Mul, getSignedRange(Scale),
3329 if (!NSWRegion.contains(getSignedRange(AddRec->getOperand(i))))
3330 Flags = clearFlags(Flags, SCEV::FlagNSW);
3331 }
3332 }
3333
3334 const SCEV *NewRec = getAddRecExpr(NewOps, AddRec->getLoop(), Flags);
3335
3336 // If all of the other operands were loop invariant, we are done.
3337 if (Ops.size() == 1) return NewRec;
3338
3339 // Otherwise, multiply the folded AddRec by the non-invariant parts.
3340 for (unsigned i = 0;; ++i)
3341 if (Ops[i] == AddRec) {
3342 Ops[i] = NewRec;
3343 break;
3344 }
3345 return getMulExpr(Ops, SCEV::FlagNone, Depth + 1);
3346 }
3347
3348 // Okay, if there weren't any loop invariants to be folded, check to see
3349 // if there are multiple AddRec's with the same loop induction variable
3350 // being multiplied together. If so, we can fold them.
3351
3352 // {A1,+,A2,+,...,+,An}<L> * {B1,+,B2,+,...,+,Bn}<L>
3353 // = {x=1 in [ sum y=x..2x [ sum z=max(y-x, y-n)..min(x,n) [
3354 // choose(x, 2x)*choose(2x-y, x-z)*A_{y-z}*B_z
3355 // ]]],+,...up to x=2n}.
3356 // Note that the arguments to choose() are always integers with values
3357 // known at compile time, never SCEV objects.
3358 //
3359 // The implementation avoids pointless extra computations when the two
3360 // addrec's are of different length (mathematically, it's equivalent to
3361 // an infinite stream of zeros on the right).
3362 bool OpsModified = false;
3363 for (unsigned OtherIdx = Idx+1;
3364 OtherIdx != Ops.size() && isa<SCEVAddRecExpr>(Ops[OtherIdx]);
3365 ++OtherIdx) {
3366 const SCEVAddRecExpr *OtherAddRec =
3367 dyn_cast<SCEVAddRecExpr>(Ops[OtherIdx]);
3368 if (!OtherAddRec || OtherAddRec->getLoop() != AddRec->getLoop())
3369 continue;
3370
3371 // Limit max number of arguments to avoid creation of unreasonably big
3372 // SCEVAddRecs with very complex operands.
3373 if (AddRec->getNumOperands() + OtherAddRec->getNumOperands() - 1 >
3374 MaxAddRecSize || hasHugeExpression({AddRec, OtherAddRec}))
3375 continue;
3376
3377 bool Overflow = false;
3378 Type *Ty = AddRec->getType();
3379 bool LargerThan64Bits = getTypeSizeInBits(Ty) > 64;
3380 SmallVector<SCEVUse, 7> AddRecOps;
3381 for (int x = 0, xe = AddRec->getNumOperands() +
3382 OtherAddRec->getNumOperands() - 1; x != xe && !Overflow; ++x) {
3384 for (int y = x, ye = 2*x+1; y != ye && !Overflow; ++y) {
3385 uint64_t Coeff1 = Choose(x, 2*x - y, Overflow);
3386 for (int z = std::max(y-x, y-(int)AddRec->getNumOperands()+1),
3387 ze = std::min(x+1, (int)OtherAddRec->getNumOperands());
3388 z < ze && !Overflow; ++z) {
3389 uint64_t Coeff2 = Choose(2*x - y, x-z, Overflow);
3390 uint64_t Coeff;
3391 if (LargerThan64Bits)
3392 Coeff = umul_ov(Coeff1, Coeff2, Overflow);
3393 else
3394 Coeff = Coeff1*Coeff2;
3395 const SCEV *CoeffTerm = getConstant(Ty, Coeff);
3396 const SCEV *Term1 = AddRec->getOperand(y-z);
3397 const SCEV *Term2 = OtherAddRec->getOperand(z);
3398 SumOps.push_back(
3399 getMulExpr(CoeffTerm, Term1, Term2, SCEV::FlagNone, Depth + 1));
3400 }
3401 }
3402 if (SumOps.empty())
3403 SumOps.push_back(getZero(Ty));
3404 AddRecOps.push_back(getAddExpr(SumOps, SCEV::FlagNone, Depth + 1));
3405 }
3406 if (!Overflow) {
3407 const SCEV *NewAddRec =
3408 getAddRecExpr(AddRecOps, AddRec->getLoop(), SCEV::FlagNone);
3409 if (Ops.size() == 2) return NewAddRec;
3410 Ops[Idx] = NewAddRec;
3411 Ops.erase(Ops.begin() + OtherIdx); --OtherIdx;
3412 OpsModified = true;
3413 AddRec = dyn_cast<SCEVAddRecExpr>(NewAddRec);
3414 if (!AddRec)
3415 break;
3416 }
3417 }
3418 if (OpsModified)
3419 return getMulExpr(Ops, SCEV::FlagNone, Depth + 1);
3420
3421 // Otherwise couldn't fold anything into this recurrence. Move onto the
3422 // next one.
3423 }
3424
3425 // Okay, it looks like we really DO need an mul expr. Check to see if we
3426 // already have one, otherwise create a new one.
3427 assert((UseFlags == SCEV::FlagNone || equal(OrigOps, Ops)) &&
3428 "Tried to add SCEVUse flags after operands changed");
3429 return {getOrCreateMulExpr(Ops, ComputeFlags(Ops)), UseFlags};
3430}
3431
3432/// Represents an unsigned remainder expression based on unsigned division.
3434 assert(getEffectiveSCEVType(LHS->getType()) ==
3435 getEffectiveSCEVType(RHS->getType()) &&
3436 "SCEVURemExpr operand types don't match!");
3437
3438 // Short-circuit easy cases
3439 if (const SCEVConstant *RHSC = dyn_cast<SCEVConstant>(RHS)) {
3440 // If constant is one, the result is trivial
3441 if (RHSC->getValue()->isOne())
3442 return getZero(LHS->getType()); // X urem 1 --> 0
3443
3444 // If constant is a power of two, fold into a zext(trunc(LHS)).
3445 if (RHSC->getAPInt().isPowerOf2()) {
3446 Type *FullTy = LHS->getType();
3447 Type *TruncTy =
3448 IntegerType::get(getContext(), RHSC->getAPInt().logBase2());
3449 return getZeroExtendExpr(getTruncateExpr(LHS, TruncTy), FullTy);
3450 }
3451 }
3452
3453 // Fallback to %a == %x urem %y == %x -<nuw> ((%x udiv %y) *<nuw> %y)
3454 const SCEV *UDiv = getUDivExpr(LHS, RHS);
3455 const SCEV *Mult = getMulExpr(UDiv, RHS, SCEV::FlagNUW);
3456 return getMinusSCEV(LHS, Mult, SCEV::FlagNUW);
3457}
3458
3459/// Get a canonical unsigned division expression, or something simpler if
3460/// possible.
3462 assert(!LHS->getType()->isPointerTy() &&
3463 "SCEVUDivExpr operand can't be pointer!");
3464 assert(LHS->getType() == RHS->getType() &&
3465 "SCEVUDivExpr operand types don't match!");
3466
3467 if (SCEV *S = findExistingSCEVInCache(scUDivExpr, {LHS, RHS}))
3468 return S;
3469
3470 // 0 udiv Y == 0
3471 if (match(LHS, m_scev_Zero()))
3472 return LHS;
3473
3474 if (const SCEVConstant *RHSC = dyn_cast<SCEVConstant>(RHS)) {
3475 if (RHSC->getValue()->isOne())
3476 return LHS; // X udiv 1 --> x
3477 // If the denominator is zero, the result of the udiv is undefined. Don't
3478 // try to analyze it, because the resolution chosen here may differ from
3479 // the resolution chosen in other parts of the compiler.
3480 if (!RHSC->getValue()->isZero()) {
3481 // Determine if the division can be folded into the operands of
3482 // its operands.
3483 // TODO: Generalize this to non-constants by using known-bits information.
3484 Type *Ty = LHS->getType();
3485 unsigned LZ = RHSC->getAPInt().countl_zero();
3486 unsigned MaxShiftAmt = getTypeSizeInBits(Ty) - LZ - 1;
3487 // For non-power-of-two values, effectively round the value up to the
3488 // nearest power of two.
3489 if (!RHSC->getAPInt().isPowerOf2())
3490 ++MaxShiftAmt;
3491 IntegerType *ExtTy =
3492 IntegerType::get(getContext(), getTypeSizeInBits(Ty) + MaxShiftAmt);
3493 if (const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(LHS))
3494 if (const SCEVConstant *Step =
3495 dyn_cast<SCEVConstant>(AR->getStepRecurrence(*this))) {
3496 // {X,+,N}/C --> {X/C,+,N/C} if safe and N/C can be folded.
3497 const APInt &StepInt = Step->getAPInt();
3498 const APInt &DivInt = RHSC->getAPInt();
3499 if (!StepInt.urem(DivInt) &&
3500 getZeroExtendExpr(AR, ExtTy) ==
3501 getAddRecExpr(getZeroExtendExpr(AR->getStart(), ExtTy),
3502 getZeroExtendExpr(Step, ExtTy), AR->getLoop(),
3503 SCEV::FlagNone)) {
3505 for (const SCEV *Op : AR->operands())
3506 Operands.push_back(getUDivExpr(Op, RHS));
3507 return getAddRecExpr(Operands, AR->getLoop(), SCEV::FlagNW);
3508 }
3509 /// Get a canonical UDivExpr for a recurrence.
3510 /// {X,+,N}/C => {Y,+,N}/C where Y=X-(X%N). Safe when C%N=0.
3511 const APInt *StartRem;
3512 if (!DivInt.urem(StepInt) && match(getURemExpr(AR->getStart(), Step),
3513 m_scev_APInt(StartRem))) {
3514 bool NoWrap =
3515 getZeroExtendExpr(AR, ExtTy) ==
3516 getAddRecExpr(getZeroExtendExpr(AR->getStart(), ExtTy),
3517 getZeroExtendExpr(Step, ExtTy), AR->getLoop(),
3519
3520 // With N <= C and both N, C as powers-of-2, the transformation
3521 // {X,+,N}/C => {(X - X%N),+,N}/C preserves division results even
3522 // if wrapping occurs, as the division results remain equivalent for
3523 // all offsets in [[(X - X%N), X).
3524 bool CanFoldWithWrap = StepInt.ule(DivInt) && // N <= C
3525 StepInt.isPowerOf2() && DivInt.isPowerOf2();
3526 // Only fold if the subtraction can be folded in the start
3527 // expression.
3528 const SCEV *NewStart =
3529 getMinusSCEV(AR->getStart(), getConstant(*StartRem));
3530 if (*StartRem != 0 && (NoWrap || CanFoldWithWrap) &&
3531 !isa<SCEVAddExpr>(NewStart)) {
3532 const SCEV *NewLHS =
3533 getAddRecExpr(NewStart, Step, AR->getLoop(),
3534 NoWrap ? SCEV::FlagNW : SCEV::FlagNone);
3535 if (LHS != NewLHS)
3536 return getUDivExpr(NewLHS, RHS);
3537 }
3538 }
3539 }
3540 // (A*B)/C --> A*(B/C) if safe and B/C can be folded.
3541 if (const SCEVMulExpr *M = dyn_cast<SCEVMulExpr>(LHS)) {
3542 if (M->hasNoUnsignedWrap()) {
3543 // Find an operand that's safely divisible.
3544 for (unsigned i = 0, e = M->getNumOperands(); i != e; ++i) {
3545 const SCEV *Op = M->getOperand(i);
3546 const SCEV *Div = getUDivExpr(Op, RHSC);
3547 if (!isa<SCEVUDivExpr>(Div) && getMulExpr(Div, RHSC) == Op) {
3548 SmallVector<SCEVUse, 4> Operands(M->operands());
3549 Operands[i] = Div;
3550 return getMulExpr(Operands);
3551 }
3552 }
3553
3554 // Even if it's not divisible, try to remove a common factor.
3555 if (const auto *LHSC = dyn_cast<SCEVConstant>(M->getOperand(0))) {
3556 APInt Factor = APIntOps::GreatestCommonDivisor(LHSC->getAPInt(),
3557 RHSC->getAPInt());
3558 if (!Factor.isIntN(1)) {
3559 SmallVector<SCEVUse, 2> NewOperands;
3560 NewOperands.push_back(getConstant(LHSC->getAPInt().udiv(Factor)));
3561 append_range(NewOperands, M->operands().drop_front());
3562 const SCEV *NewMul = getMulExpr(NewOperands);
3563 return getUDivExpr(NewMul,
3564 getConstant(RHSC->getAPInt().udiv(Factor)));
3565 }
3566 }
3567 }
3568 }
3569
3570 // (A/B)/C --> A/(B*C) if safe and B*C can be folded.
3571 if (const SCEVUDivExpr *OtherDiv = dyn_cast<SCEVUDivExpr>(LHS)) {
3572 if (auto *DivisorConstant =
3573 dyn_cast<SCEVConstant>(OtherDiv->getRHS())) {
3574 bool Overflow = false;
3575 APInt NewRHS =
3576 DivisorConstant->getAPInt().umul_ov(RHSC->getAPInt(), Overflow);
3577 if (Overflow) {
3578 return getConstant(RHSC->getType(), 0, false);
3579 }
3580 return getUDivExpr(OtherDiv->getLHS(), getConstant(NewRHS));
3581 }
3582 }
3583
3584 // (A+B)/C --> (A/C + B/C) if the add does not unsigned wrap and A/C and
3585 // B/C can be folded.
3586 if (const SCEVAddExpr *A = dyn_cast<SCEVAddExpr>(LHS)) {
3587 if (A->hasNoUnsignedWrap()) {
3589 for (unsigned i = 0, e = A->getNumOperands(); i != e; ++i) {
3590 const SCEV *Op = getUDivExpr(A->getOperand(i), RHS);
3591 if (isa<SCEVUDivExpr>(Op) ||
3592 getMulExpr(Op, RHS) != A->getOperand(i))
3593 break;
3594 Operands.push_back(Op);
3595 }
3596 if (Operands.size() == A->getNumOperands())
3597 return getAddExpr(Operands);
3598 }
3599 }
3600
3601 // ((N - M) + (M * A)) / N --> ((N - 1) + (M * A)) / N
3602 // This is an idiom for rounding A up to the next multiple of N, where A
3603 // is aready known to be a multiple of M. In this case, instcombine can
3604 // see that some low bits of the added constant are unused, so can clear
3605 // them, but we want to canonicalise to set the low bits. This makes the
3606 // pattern easier to match, without needing to check for known bits in
3607 // A*M.
3608 const APInt &N = RHSC->getAPInt();
3609 const APInt *NMinusM, *M;
3610 const SCEV *A;
3611 if (match(LHS, m_scev_Add(m_scev_APInt(NMinusM),
3612 m_scev_Mul(m_scev_APInt(M), m_SCEV(A))))) {
3613 if (N.isPowerOf2() && M->isPowerOf2() && M->ult(N) &&
3614 *NMinusM == N - *M) {
3615 return getUDivExpr(
3617 RHS);
3618 }
3619 }
3620
3621 // Fold if both operands are constant.
3622 if (const SCEVConstant *LHSC = dyn_cast<SCEVConstant>(LHS))
3623 return getConstant(LHSC->getAPInt().udiv(RHSC->getAPInt()));
3624 }
3625 }
3626
3627 // ((-C + (C smax %x)) /u %x) evaluates to zero, for any positive constant C.
3628 const APInt *NegC, *C;
3629 if (match(LHS,
3632 NegC->isNegative() && !NegC->isMinSignedValue() && *C == -*NegC)
3633 return getZero(LHS->getType());
3634
3635 // (%a * %b)<nuw> / %b -> %a
3636 const auto *Mul = dyn_cast<SCEVMulExpr>(LHS);
3637 if (Mul && Mul->hasNoUnsignedWrap()) {
3638 for (int i = 0, e = Mul->getNumOperands(); i != e; ++i) {
3639 if (Mul->getOperand(i) == RHS) {
3641 append_range(Operands, Mul->operands().take_front(i));
3642 append_range(Operands, Mul->operands().drop_front(i + 1));
3643 return getMulExpr(Operands);
3644 }
3645 }
3646 }
3647
3648 // TODO: Generalize to handle any common factors.
3649 // udiv (mul nuw a, vscale), (mul nuw b, vscale) --> udiv a, b
3650 const SCEV *NewLHS, *NewRHS;
3651 if (match(LHS, m_scev_c_NUWMul(m_SCEV(NewLHS), m_SCEVVScale())) &&
3652 match(RHS, m_scev_c_NUWMul(m_SCEV(NewRHS), m_SCEVVScale())))
3653 return getUDivExpr(NewLHS, NewRHS);
3654
3655 return getOrCreateUDivExpr(LHS, RHS);
3656}
3657
3658/// Get a canonical unsigned division expression, or something simpler if
3659/// possible. There is no representation for an exact udiv in SCEV IR, but we
3660/// can attempt to optimize it prior to construction.
3662 // Currently there is no exact specific logic.
3663
3664 return getUDivExpr(LHS, RHS);
3665}
3666
3667/// Get an add recurrence expression for the specified loop. Simplify the
3668/// expression as much as possible.
3670 const Loop *L, SCEVFlagsPair Flags) {
3672 Operands.push_back(Start);
3673 if (const SCEVAddRecExpr *StepChrec = dyn_cast<SCEVAddRecExpr>(Step))
3674 if (StepChrec->getLoop() == L) {
3675 append_range(Operands, StepChrec->operands());
3676 // The use flags describe the two-operand recurrence, not the flattened
3677 // one built here, so drop them just like the expression's NUW/NSW.
3678 return getAddRecExpr(Operands, L,
3679 maskFlags(Flags.ExprFlags, SCEV::FlagNW));
3680 }
3681
3682 Operands.push_back(Step);
3683 return getAddRecExpr(Operands, L, Flags);
3684}
3685
3686/// Get an add recurrence expression for the specified loop. Simplify the
3687/// expression as much as possible.
3689 const Loop *L, SCEVFlagsPair NWFlags) {
3690 SCEVFlags ExprFlags = NWFlags.ExprFlags;
3691 SCEVFlags UseFlags = NWFlags.UseFlags;
3692 assert(!(UseFlags & ~(SCEV::FlagNUW | SCEV::FlagNSW)) &&
3693 "only nuw or nsw allowed");
3694 if (Operands.size() == 1) return Operands[0];
3695#ifndef NDEBUG
3697 for (const SCEV *Op : llvm::drop_begin(Operands)) {
3698 assert(getEffectiveSCEVType(Op->getType()) == ETy &&
3699 "SCEVAddRecExpr operand types don't match!");
3700 assert(!Op->getType()->isPointerTy() && "Step must be integer");
3701 }
3702 for (const SCEV *Op : Operands)
3704 "SCEVAddRecExpr operand is not available at loop entry!");
3705
3706 // Keep track of the original operands, for verification when adding
3707 // use-specific flags.
3708 const SmallVector<SCEVUse, 4> OrigOperands(Operands.begin(), Operands.end());
3709#endif
3710
3711 if (Operands.back()->isZero()) {
3712 Operands.pop_back();
3713 return getAddRecExpr(Operands, L, SCEV::FlagNone); // {X,+,0} --> X
3714 }
3715
3716 // It's tempting to want to call getConstantMaxBackedgeTakenCount count here and
3717 // use that information to infer NUW and NSW flags. However, computing a
3718 // BE count requires calling getAddRecExpr, so we may not yet have a
3719 // meaningful BE count at this point (and if we don't, we'd be stuck
3720 // with a SCEVCouldNotCompute as the cached BE count).
3721
3722 ExprFlags = StrengthenNoWrapFlags(this, scAddRecExpr, Operands, ExprFlags);
3723
3724 // Canonicalize nested AddRecs in by nesting them in order of loop depth.
3725 if (const SCEVAddRecExpr *NestedAR = dyn_cast<SCEVAddRecExpr>(Operands[0])) {
3726 const Loop *NestedLoop = NestedAR->getLoop();
3727 if (L->contains(NestedLoop)
3728 ? (L->getLoopDepth() < NestedLoop->getLoopDepth())
3729 : (!NestedLoop->contains(L) &&
3730 DT.dominates(L->getHeader(), NestedLoop->getHeader()))) {
3731 SmallVector<SCEVUse, 4> NestedOperands(NestedAR->operands());
3732 Operands[0] = NestedAR->getStart();
3733 // AddRecs require their operands be loop-invariant with respect to their
3734 // loops. Don't perform this transformation if it would break this
3735 // requirement.
3736 bool AllInvariant = all_of(
3737 Operands, [&](const SCEV *Op) { return isLoopInvariant(Op, L); });
3738
3739 if (AllInvariant) {
3740 // Create a recurrence for the outer loop with the same step size.
3741 //
3742 // The outer recurrence keeps its NW flag but only keeps NUW/NSW if the
3743 // inner recurrence has the same property.
3744 SCEVFlags OuterFlags =
3745 maskFlags(ExprFlags, SCEV::FlagNW | NestedAR->getNoWrapFlags());
3746
3747 NestedOperands[0] = getAddRecExpr(Operands, L, OuterFlags);
3748 AllInvariant = all_of(NestedOperands, [&](const SCEV *Op) {
3749 return isLoopInvariant(Op, NestedLoop);
3750 });
3751
3752 if (AllInvariant) {
3753 // Ok, both add recurrences are valid after the transformation.
3754 //
3755 // The inner recurrence keeps its NW flag but only keeps NUW/NSW if
3756 // the outer recurrence has the same property.
3757 SCEVFlags InnerFlags =
3758 maskFlags(NestedAR->getNoWrapFlags(), SCEV::FlagNW | ExprFlags);
3759 return getAddRecExpr(NestedOperands, NestedLoop, InnerFlags);
3760 }
3761 }
3762 // Reset Operands to its original state.
3763 Operands[0] = NestedAR;
3764 }
3765 }
3766
3767 // Okay, it looks like we really DO need an addrec expr. Check to see if we
3768 // already have one, otherwise create a new one.
3769 assert((UseFlags == SCEV::FlagNone || equal(OrigOperands, Operands)) &&
3770 "Tried to add SCEVUse flags after operands changed");
3771 return {getOrCreateAddRecExpr(Operands, L, ExprFlags), UseFlags};
3772}
3773
3775 ArrayRef<SCEVUse> IndexExprs) {
3776 const SCEV *BaseExpr = getSCEV(GEP->getPointerOperand());
3777 // getSCEV(Base)->getType() has the same address space as Base->getType()
3778 // because SCEV::getType() preserves the address space.
3779 GEPNoWrapFlags NW = GEP->getNoWrapFlags();
3780 if (NW != GEPNoWrapFlags::none()) {
3781 // We'd like to propagate flags from the IR to the corresponding SCEV nodes,
3782 // but to do that, we have to ensure that said flag is valid in the entire
3783 // defined scope of the SCEV.
3784 // TODO: non-instructions have global scope. We might be able to prove
3785 // some global scope cases
3786 auto *GEPI = dyn_cast<Instruction>(GEP);
3787 if (!GEPI || !isSCEVExprNeverPoison(GEPI))
3788 NW = GEPNoWrapFlags::none();
3789 }
3790
3791 return getGEPExpr(BaseExpr, IndexExprs, GEP->getSourceElementType(), NW);
3792}
3793
3795 ArrayRef<SCEVUse> IndexExprs,
3796 Type *SrcElementTy, GEPNoWrapFlags NW) {
3797 SCEVFlags OffsetWrap = SCEV::FlagNone;
3798 if (NW.hasNoUnsignedSignedWrap())
3799 OffsetWrap = setFlags(OffsetWrap, SCEV::FlagNSW);
3800 if (NW.hasNoUnsignedWrap())
3801 OffsetWrap = setFlags(OffsetWrap, SCEV::FlagNUW);
3802
3803 Type *CurTy = BaseExpr->getType();
3804 Type *IntIdxTy = getEffectiveSCEVType(BaseExpr->getType());
3805 bool FirstIter = true;
3807 for (SCEVUse IndexExpr : IndexExprs) {
3808 // Compute the (potentially symbolic) offset in bytes for this index.
3809 if (StructType *STy = dyn_cast<StructType>(CurTy)) {
3810 // For a struct, add the member offset.
3811 ConstantInt *Index = cast<SCEVConstant>(IndexExpr)->getValue();
3812 unsigned FieldNo = Index->getZExtValue();
3813 const SCEV *FieldOffset = getOffsetOfExpr(IntIdxTy, STy, FieldNo);
3814 Offsets.push_back(FieldOffset);
3815
3816 // Update CurTy to the type of the field at Index.
3817 CurTy = STy->getTypeAtIndex(Index);
3818 } else {
3819 // Update CurTy to its element type.
3820 if (FirstIter) {
3821 assert(isa<PointerType>(CurTy) &&
3822 "The first index of a GEP indexes a pointer");
3823 CurTy = SrcElementTy;
3824 FirstIter = false;
3825 } else {
3826 CurTy = GetElementPtrInst::getTypeAtIndex(CurTy, (uint64_t)0);
3827 }
3828 // For an array, add the element offset, explicitly scaled.
3829 const SCEV *ElementSize = getSizeOfExpr(IntIdxTy, CurTy);
3830 // Getelementptr indices are signed.
3831 IndexExpr = getTruncateOrSignExtend(IndexExpr, IntIdxTy);
3832
3833 // Multiply the index by the element size to compute the element offset.
3834 const SCEV *LocalOffset = getMulExpr(IndexExpr, ElementSize, OffsetWrap);
3835 Offsets.push_back(LocalOffset);
3836 }
3837 }
3838
3839 // Handle degenerate case of GEP without offsets.
3840 if (Offsets.empty())
3841 return BaseExpr;
3842
3843 // Add the offsets together, assuming nsw if inbounds.
3844 const SCEV *Offset = getAddExpr(Offsets, OffsetWrap);
3845 // Add the base address and the offset. We cannot use the nsw flag, as the
3846 // base address is unsigned. However, if we know that the offset is
3847 // non-negative, we can use nuw.
3848 bool NUW = NW.hasNoUnsignedWrap() ||
3850 SCEVFlags BaseWrap = NUW ? SCEV::FlagNUW : SCEV::FlagNone;
3851 const SCEV *GEPExpr = getAddExpr(BaseExpr, Offset, BaseWrap);
3852 assert(BaseExpr->getType() == GEPExpr->getType() &&
3853 "GEP should not change type mid-flight.");
3854 return GEPExpr;
3855}
3856
3857SCEV *ScalarEvolution::findExistingSCEVInCache(SCEVTypes SCEVType,
3859 const Loop *L) {
3860 assert((SCEVType != scAddRecExpr || L) &&
3861 "L must be passed to find existing AddRecs");
3863 ID.AddInteger(SCEVType);
3864 for (SCEVUse Op : Ops)
3865 ID.AddPointer(Op.getOpaqueValue());
3866 if (L)
3867 ID.AddPointer(L);
3869 return UniqueSCEVs.lookup(ID, Token);
3870}
3871
3872const SCEV *ScalarEvolution::getAbsExpr(const SCEV *Op, bool IsNSW) {
3873 SCEVFlags Flags = IsNSW ? SCEV::FlagNSW : SCEV::FlagNone;
3874 return getSMaxExpr(Op, getNegativeSCEV(Op, Flags));
3875}
3876
3879 assert(SCEVMinMaxExpr::isMinMaxType(Kind) && "Not a SCEVMinMaxExpr!");
3880 assert(!Ops.empty() && "Cannot get empty (u|s)(min|max)!");
3881 if (Ops.size() == 1) return Ops[0];
3882#ifndef NDEBUG
3883 Type *ETy = getEffectiveSCEVType(Ops[0]->getType());
3884 for (unsigned i = 1, e = Ops.size(); i != e; ++i) {
3885 assert(getEffectiveSCEVType(Ops[i]->getType()) == ETy &&
3886 "Operand types don't match!");
3887 assert(Ops[0]->getType()->isPointerTy() ==
3888 Ops[i]->getType()->isPointerTy() &&
3889 "min/max should be consistently pointerish");
3890 }
3891#endif
3892
3893 bool IsSigned = Kind == scSMaxExpr || Kind == scSMinExpr;
3894 bool IsMax = Kind == scSMaxExpr || Kind == scUMaxExpr;
3895
3896 const SCEV *Folded = constantFoldAndGroupOps(
3897 *this, LI, DT, Ops,
3898 [&](const APInt &C1, const APInt &C2) {
3899 switch (Kind) {
3900 case scSMaxExpr:
3901 return APIntOps::smax(C1, C2);
3902 case scSMinExpr:
3903 return APIntOps::smin(C1, C2);
3904 case scUMaxExpr:
3905 return APIntOps::umax(C1, C2);
3906 case scUMinExpr:
3907 return APIntOps::umin(C1, C2);
3908 default:
3909 llvm_unreachable("Unknown SCEV min/max opcode");
3910 }
3911 },
3912 [&](const APInt &C) {
3913 // identity
3914 if (IsMax)
3915 return IsSigned ? C.isMinSignedValue() : C.isMinValue();
3916 else
3917 return IsSigned ? C.isMaxSignedValue() : C.isMaxValue();
3918 },
3919 [&](const APInt &C) {
3920 // absorber
3921 if (IsMax)
3922 return IsSigned ? C.isMaxSignedValue() : C.isMaxValue();
3923 else
3924 return IsSigned ? C.isMinSignedValue() : C.isMinValue();
3925 });
3926 if (Folded)
3927 return Folded;
3928
3929 // Check if we have created the same expression before.
3930 if (const SCEV *S = findExistingSCEVInCache(Kind, Ops)) {
3931 return S;
3932 }
3933
3934 // Find the first operation of the same kind
3935 unsigned Idx = 0;
3936 while (Idx < Ops.size() && Ops[Idx]->getSCEVType() < Kind)
3937 ++Idx;
3938
3939 // Check to see if one of the operands is of the same kind. If so, expand its
3940 // operands onto our operand list, and recurse to simplify.
3941 if (Idx < Ops.size()) {
3942 bool DeletedAny = false;
3943 while (Ops[Idx]->getSCEVType() == Kind) {
3944 const SCEVMinMaxExpr *SMME = cast<SCEVMinMaxExpr>(Ops[Idx]);
3945 Ops.erase(Ops.begin()+Idx);
3946 append_range(Ops, SMME->operands());
3947 DeletedAny = true;
3948 }
3949
3950 if (DeletedAny)
3951 return getMinMaxExpr(Kind, Ops);
3952 }
3953
3954 // Okay, check to see if the same value occurs in the operand list twice. If
3955 // so, delete one. Since we sorted the list, these values are required to
3956 // be adjacent.
3961 llvm::CmpInst::Predicate FirstPred = IsMax ? GEPred : LEPred;
3962 llvm::CmpInst::Predicate SecondPred = IsMax ? LEPred : GEPred;
3963 for (unsigned i = 0, e = Ops.size() - 1; i != e; ++i) {
3964 if (Ops[i] == Ops[i + 1] ||
3965 isKnownViaNonRecursiveReasoning(FirstPred, Ops[i], Ops[i + 1])) {
3966 // X op Y op Y --> X op Y
3967 // X op Y --> X, if we know X, Y are ordered appropriately
3968 Ops.erase(Ops.begin() + i + 1, Ops.begin() + i + 2);
3969 --i;
3970 --e;
3971 } else if (isKnownViaNonRecursiveReasoning(SecondPred, Ops[i],
3972 Ops[i + 1])) {
3973 // X op Y --> Y, if we know X, Y are ordered appropriately
3974 Ops.erase(Ops.begin() + i, Ops.begin() + i + 1);
3975 --i;
3976 --e;
3977 }
3978 }
3979
3980 if (Ops.size() == 1) return Ops[0];
3981
3982 assert(!Ops.empty() && "Reduced smax down to nothing!");
3983
3984 // Okay, it looks like we really DO need an expr. Check to see if we
3985 // already have one, otherwise create a new one.
3987 ID.AddInteger(Kind);
3988 for (SCEVUse Op : Ops)
3989 ID.AddPointer(Op.getOpaqueValue());
3991 const SCEV *ExistingSCEV = UniqueSCEVs.lookup(ID, Token);
3992 if (ExistingSCEV)
3993 return ExistingSCEV;
3994 SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
3996 SCEV *S = new (SCEVAllocator)
3997 SCEVMinMaxExpr(ID.Intern(SCEVAllocator), Kind, O, Ops.size());
3998
3999 UniqueSCEVs.insert(S, Token);
4000 S->computeAndSetCanonical(*this);
4001 registerUser(S, Ops);
4002 return S;
4003}
4004
4005namespace {
4006
4007class SCEVSequentialMinMaxDeduplicatingVisitor final
4008 : public SCEVVisitor<SCEVSequentialMinMaxDeduplicatingVisitor,
4009 std::optional<const SCEV *>> {
4010 using RetVal = std::optional<const SCEV *>;
4011
4012 ScalarEvolution &SE;
4013 const SCEVTypes RootKind; // Must be a sequential min/max expression.
4014 const SCEVTypes NonSequentialRootKind; // Non-sequential variant of RootKind.
4016
4017 bool canRecurseInto(SCEVTypes Kind) const {
4018 // We can only recurse into the SCEV expression of the same effective type
4019 // as the type of our root SCEV expression.
4020 return RootKind == Kind || NonSequentialRootKind == Kind;
4021 };
4022
4023 RetVal visit(const SCEV *S) {
4024 // Has the whole operand been seen already?
4025 if (!SeenOps.insert(S).second)
4026 return std::nullopt;
4028 SCEVTypes Kind = S->getSCEVType();
4029
4030 if (!canRecurseInto(Kind))
4031 return S;
4032
4033 auto *NAry = cast<SCEVNAryExpr>(S);
4034 SmallVector<SCEVUse> NewOps;
4035 bool Changed = visit(Kind, NAry->operands(), NewOps);
4036
4037 if (!Changed)
4038 return S;
4039 if (NewOps.empty())
4040 return std::nullopt;
4041
4043 ? SE.getSequentialMinMaxExpr(Kind, NewOps)
4044 : SE.getMinMaxExpr(Kind, NewOps);
4045 }
4046 return S;
4047 }
4048
4049public:
4050 SCEVSequentialMinMaxDeduplicatingVisitor(ScalarEvolution &SE,
4051 SCEVTypes RootKind)
4052 : SE(SE), RootKind(RootKind),
4053 NonSequentialRootKind(
4054 SCEVSequentialMinMaxExpr::getEquivalentNonSequentialSCEVType(
4055 RootKind)) {}
4056
4057 bool /*Changed*/ visit(SCEVTypes Kind, ArrayRef<SCEVUse> OrigOps,
4058 SmallVectorImpl<SCEVUse> &NewOps) {
4059 bool Changed = false;
4061 Ops.reserve(OrigOps.size());
4062
4063 for (const SCEV *Op : OrigOps) {
4064 RetVal NewOp = visit(Op);
4065 if (NewOp != Op)
4066 Changed = true;
4067 if (NewOp)
4068 Ops.emplace_back(*NewOp);
4069 }
4070
4071 if (Changed)
4072 NewOps = std::move(Ops);
4073 return Changed;
4074 }
4075};
4076
4077} // namespace
4078
4080 switch (Kind) {
4081 case scConstant:
4082 case scVScale:
4083 case scTruncate:
4084 case scZeroExtend:
4085 case scSignExtend:
4086 case scPtrToAddr:
4087 case scAddExpr:
4088 case scMulExpr:
4089 case scUDivExpr:
4090 case scAddRecExpr:
4091 case scUMaxExpr:
4092 case scSMaxExpr:
4093 case scUMinExpr:
4094 case scSMinExpr:
4095 case scUnknown:
4096 // If any operand is poison, the whole expression is poison.
4097 return true;
4099 // FIXME: if the *first* operand is poison, the whole expression is poison.
4100 return false; // Pessimistically, say that it does not propagate poison.
4101 case scCouldNotCompute:
4102 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
4103 }
4104 llvm_unreachable("Unknown SCEV kind!");
4105}
4106
4107namespace {
4108// The only way poison may be introduced in a SCEV expression is from a
4109// poison SCEVUnknown (ConstantExprs are also represented as SCEVUnknown,
4110// not SCEVConstant). Notably, SCEVFlags on SCEV nodes can *not*
4111// introduce poison -- they encode guaranteed, non-speculated knowledge.
4112//
4113// Additionally, all SCEV nodes propagate poison from inputs to outputs,
4114// with the notable exception of umin_seq, where only poison from the first
4115// operand is (unconditionally) propagated.
4116struct SCEVPoisonCollector {
4117 bool LookThroughMaybePoisonBlocking;
4118 SmallPtrSet<const SCEVUnknown *, 4> MaybePoison;
4119 SCEVPoisonCollector(bool LookThroughMaybePoisonBlocking)
4120 : LookThroughMaybePoisonBlocking(LookThroughMaybePoisonBlocking) {}
4121
4122 bool follow(const SCEV *S) {
4123 if (!LookThroughMaybePoisonBlocking &&
4125 return false;
4126
4127 if (auto *SU = dyn_cast<SCEVUnknown>(S)) {
4128 if (!isGuaranteedNotToBePoison(SU->getValue()))
4129 MaybePoison.insert(SU);
4130 }
4131 return true;
4132 }
4133 bool isDone() const { return false; }
4134};
4135} // namespace
4136
4137/// Return true if V is poison given that AssumedPoison is already poison.
4138static bool impliesPoison(const SCEV *AssumedPoison, const SCEV *S) {
4139 // First collect all SCEVs that might result in AssumedPoison to be poison.
4140 // We need to look through potentially poison-blocking operations here,
4141 // because we want to find all SCEVs that *might* result in poison, not only
4142 // those that are *required* to.
4143 SCEVPoisonCollector PC1(/* LookThroughMaybePoisonBlocking */ true);
4144 visitAll(AssumedPoison, PC1);
4145
4146 // AssumedPoison is never poison. As the assumption is false, the implication
4147 // is true. Don't bother walking the other SCEV in this case.
4148 if (PC1.MaybePoison.empty())
4149 return true;
4150
4151 // Collect all SCEVs in S that, if poison, *will* result in S being poison
4152 // as well. We cannot look through potentially poison-blocking operations
4153 // here, as their arguments only *may* make the result poison.
4154 SCEVPoisonCollector PC2(/* LookThroughMaybePoisonBlocking */ false);
4155 visitAll(S, PC2);
4156
4157 // Make sure that no matter which SCEV in PC1.MaybePoison is actually poison,
4158 // it will also make S poison by being part of PC2.MaybePoison.
4159 return llvm::set_is_subset(PC1.MaybePoison, PC2.MaybePoison);
4160}
4161
4163 SmallPtrSetImpl<const Value *> &Result, const SCEV *S) {
4164 SCEVPoisonCollector PC(/* LookThroughMaybePoisonBlocking */ false);
4165 visitAll(S, PC);
4166 for (const SCEVUnknown *SU : PC.MaybePoison)
4167 Result.insert(SU->getValue());
4168}
4169
4171 const SCEV *S, Instruction *I,
4172 SmallVectorImpl<Instruction *> &DropPoisonGeneratingInsts) {
4173 // If the instruction cannot be poison, it's always safe to reuse.
4175 return true;
4176
4177 // Otherwise, it is possible that I is more poisonous that S. Collect the
4178 // poison-contributors of S, and then check whether I has any additional
4179 // poison-contributors. Poison that is contributed through poison-generating
4180 // flags is handled by dropping those flags instead.
4182 getPoisonGeneratingValues(PoisonVals, S);
4183
4184 SmallVector<Value *> Worklist;
4186 Worklist.push_back(I);
4187 while (!Worklist.empty()) {
4188 Value *V = Worklist.pop_back_val();
4189 if (!Visited.insert(V).second)
4190 continue;
4191
4192 // Avoid walking large instruction graphs.
4193 if (Visited.size() > 16)
4194 return false;
4195
4196 // Either the value can't be poison, or the S would also be poison if it
4197 // is.
4198 if (PoisonVals.contains(V) || ::isGuaranteedNotToBePoison(V))
4199 continue;
4200
4201 auto *I = dyn_cast<Instruction>(V);
4202 if (!I)
4203 return false;
4204
4205 // Disjoint or instructions are interpreted as adds by SCEV. However, we
4206 // can't replace an arbitrary add with disjoint or, even if we drop the
4207 // flag. We would need to convert the or into an add.
4208 if (auto *PDI = dyn_cast<PossiblyDisjointInst>(I))
4209 if (PDI->isDisjoint())
4210 return false;
4211
4212 // FIXME: Ignore vscale, even though it technically could be poison. Do this
4213 // because SCEV currently assumes it can't be poison. Remove this special
4214 // case once we proper model when vscale can be poison.
4215 if (auto *II = dyn_cast<IntrinsicInst>(I);
4216 II && II->getIntrinsicID() == Intrinsic::vscale)
4217 continue;
4218
4219 if (canCreatePoison(cast<Operator>(I), /*ConsiderFlagsAndMetadata*/ false))
4220 return false;
4221
4222 // If the instruction can't create poison, we can recurse to its operands.
4223 if (I->hasPoisonGeneratingAnnotations())
4224 DropPoisonGeneratingInsts.push_back(I);
4225
4226 llvm::append_range(Worklist, I->operands());
4227 }
4228 return true;
4229}
4230
4231const SCEV *
4234 assert(SCEVSequentialMinMaxExpr::isSequentialMinMaxType(Kind) &&
4235 "Not a SCEVSequentialMinMaxExpr!");
4236 assert(!Ops.empty() && "Cannot get empty (u|s)(min|max)!");
4237 if (Ops.size() == 1)
4238 return Ops[0];
4239#ifndef NDEBUG
4240 Type *ETy = getEffectiveSCEVType(Ops[0]->getType());
4241 for (unsigned i = 1, e = Ops.size(); i != e; ++i) {
4242 assert(getEffectiveSCEVType(Ops[i]->getType()) == ETy &&
4243 "Operand types don't match!");
4244 assert(Ops[0]->getType()->isPointerTy() ==
4245 Ops[i]->getType()->isPointerTy() &&
4246 "min/max should be consistently pointerish");
4247 }
4248#endif
4249
4250 // Note that SCEVSequentialMinMaxExpr is *NOT* commutative,
4251 // so we can *NOT* do any kind of sorting of the expressions!
4252
4253 // Check if we have created the same expression before.
4254 if (const SCEV *S = findExistingSCEVInCache(Kind, Ops))
4255 return S;
4256
4257 // FIXME: there are *some* simplifications that we can do here.
4258
4259 // Keep only the first instance of an operand.
4260 {
4261 SCEVSequentialMinMaxDeduplicatingVisitor Deduplicator(*this, Kind);
4262 bool Changed = Deduplicator.visit(Kind, Ops, Ops);
4263 if (Changed)
4264 return getSequentialMinMaxExpr(Kind, Ops);
4265 }
4266
4267 // Check to see if one of the operands is of the same kind. If so, expand its
4268 // operands onto our operand list, and recurse to simplify.
4269 {
4270 unsigned Idx = 0;
4271 bool DeletedAny = false;
4272 while (Idx < Ops.size()) {
4273 if (Ops[Idx]->getSCEVType() != Kind) {
4274 ++Idx;
4275 continue;
4276 }
4277 const auto *SMME = cast<SCEVSequentialMinMaxExpr>(Ops[Idx]);
4278 Ops.erase(Ops.begin() + Idx);
4279 Ops.insert(Ops.begin() + Idx, SMME->operands().begin(),
4280 SMME->operands().end());
4281 DeletedAny = true;
4282 }
4283
4284 if (DeletedAny)
4285 return getSequentialMinMaxExpr(Kind, Ops);
4286 }
4287
4288 const SCEV *SaturationPoint;
4290 switch (Kind) {
4292 SaturationPoint = getZero(Ops[0]->getType());
4293 Pred = ICmpInst::ICMP_ULE;
4294 break;
4295 default:
4296 llvm_unreachable("Not a sequential min/max type.");
4297 }
4298
4299 for (unsigned i = 1, e = Ops.size(); i != e; ++i) {
4300 if (!isGuaranteedNotToCauseUB(Ops[i]))
4301 continue;
4302 // We can replace %x umin_seq %y with %x umin %y if either:
4303 // * %y being poison implies %x is also poison.
4304 // * %x cannot be the saturating value (e.g. zero for umin).
4305 if (::impliesPoison(Ops[i], Ops[i - 1]) ||
4306 isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_NE, Ops[i - 1],
4307 SaturationPoint)) {
4308 SmallVector<SCEVUse, 2> SeqOps = {Ops[i - 1], Ops[i]};
4309 Ops[i - 1] = getMinMaxExpr(
4311 SeqOps);
4312 Ops.erase(Ops.begin() + i);
4313 return getSequentialMinMaxExpr(Kind, Ops);
4314 }
4315 // Fold %x umin_seq %y to %x if %x ule %y.
4316 // TODO: We might be able to prove the predicate for a later operand.
4317 if (isKnownViaNonRecursiveReasoning(Pred, Ops[i - 1], Ops[i])) {
4318 Ops.erase(Ops.begin() + i);
4319 return getSequentialMinMaxExpr(Kind, Ops);
4320 }
4321 }
4322
4323 // Okay, it looks like we really DO need an expr. Check to see if we
4324 // already have one, otherwise create a new one.
4326 ID.AddInteger(Kind);
4327 for (SCEVUse Op : Ops)
4328 ID.AddPointer(Op.getOpaqueValue());
4330 const SCEV *ExistingSCEV = UniqueSCEVs.lookup(ID, Token);
4331 if (ExistingSCEV)
4332 return ExistingSCEV;
4333
4334 SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
4336 SCEV *S = new (SCEVAllocator)
4337 SCEVSequentialMinMaxExpr(ID.Intern(SCEVAllocator), Kind, O, Ops.size());
4338
4339 UniqueSCEVs.insert(S, Token);
4340 S->computeAndSetCanonical(*this);
4341 registerUser(S, Ops);
4342 return S;
4343}
4344
4349
4353
4358
4362
4367
4371
4373 bool Sequential) {
4374 SmallVector<SCEVUse, 2> Ops = {LHS, RHS};
4375 return getUMinExpr(Ops, Sequential);
4376}
4377
4383
4384const SCEV *
4386 const SCEV *Res = getConstant(IntTy, Size.getKnownMinValue());
4387 if (Size.isScalable())
4388 Res = getMulExpr(Res, getVScale(IntTy));
4389 return Res;
4390}
4391
4393 return getSizeOfExpr(IntTy, getDataLayout().getTypeAllocSize(AllocTy));
4394}
4395
4397 return getSizeOfExpr(IntTy, getDataLayout().getTypeStoreSize(StoreTy));
4398}
4399
4401 StructType *STy,
4402 unsigned FieldNo) {
4403 // We can bypass creating a target-independent constant expression and then
4404 // folding it back into a ConstantInt. This is just a compile-time
4405 // optimization.
4406 const StructLayout *SL = getDataLayout().getStructLayout(STy);
4407 assert(!SL->getSizeInBits().isScalable() &&
4408 "Cannot get offset for structure containing scalable vector types");
4409 return getConstant(IntTy, SL->getElementOffset(FieldNo));
4410}
4411
4413 // Don't attempt to do anything other than create a SCEVUnknown object
4414 // here. createSCEV only calls getUnknown after checking for all other
4415 // interesting possibilities, and any other code that calls getUnknown
4416 // is doing so in order to hide a value from SCEV canonicalization.
4417
4420 ID.AddPointer(V);
4422 if (SCEV *S = UniqueSCEVs.lookup(ID, Token)) {
4423 assert(cast<SCEVUnknown>(S)->getValue() == V &&
4424 "Stale SCEVUnknown in uniquing map!");
4425 return S;
4426 }
4427 SCEV *S = new (SCEVAllocator) SCEVUnknown(ID.Intern(SCEVAllocator), V, this,
4428 FirstUnknown);
4429 FirstUnknown = cast<SCEVUnknown>(S);
4430 UniqueSCEVs.insert(S, Token);
4431 S->computeAndSetCanonical(*this);
4432 return S;
4433}
4434
4435//===----------------------------------------------------------------------===//
4436// Basic SCEV Analysis and PHI Idiom Recognition Code
4437//
4438
4439/// Test if values of the given type are analyzable within the SCEV
4440/// framework. This primarily includes integer types, and it can optionally
4441/// include pointer types if the ScalarEvolution class has access to
4442/// target-specific information.
4444 // Integers and pointers are always SCEVable.
4445 return Ty->isIntOrPtrTy();
4446}
4447
4448/// Return the size in bits of the specified type, for which isSCEVable must
4449/// return true.
4451 assert(isSCEVable(Ty) && "Type is not SCEVable!");
4452 if (Ty->isPointerTy())
4454 return getDataLayout().getTypeSizeInBits(Ty);
4455}
4456
4457/// Return a type with the same bitwidth as the given type and which represents
4458/// how SCEV will treat the given type, for which isSCEVable must return
4459/// true. For pointer types, this is the pointer index sized integer type.
4461 assert(isSCEVable(Ty) && "Type is not SCEVable!");
4462
4463 if (Ty->isIntegerTy())
4464 return Ty;
4465
4466 // The only other support type is pointer.
4467 assert(Ty->isPointerTy() && "Unexpected non-pointer non-integer type!");
4468 return getDataLayout().getIndexType(Ty);
4469}
4470
4472 return getTypeSizeInBits(T1) >= getTypeSizeInBits(T2) ? T1 : T2;
4473}
4474
4476 const SCEV *B) {
4477 /// For a valid use point to exist, the defining scope of one operand
4478 /// must dominate the other.
4479 bool PreciseA, PreciseB;
4480 auto *ScopeA = getDefiningScopeBound({A}, PreciseA);
4481 auto *ScopeB = getDefiningScopeBound({B}, PreciseB);
4482 if (!PreciseA || !PreciseB)
4483 // Can't tell.
4484 return false;
4485 return (ScopeA == ScopeB) || DT.dominates(ScopeA, ScopeB) ||
4486 DT.dominates(ScopeB, ScopeA);
4487}
4488
4490 return CouldNotCompute.get();
4491}
4492
4493bool ScalarEvolution::checkValidity(const SCEV *S) const {
4494 bool ContainsNulls = SCEVExprContains(S, [](const SCEV *S) {
4495 auto *SU = dyn_cast<SCEVUnknown>(S);
4496 return SU && SU->getValue() == nullptr;
4497 });
4498
4499 return !ContainsNulls;
4500}
4501
4503 HasRecMapType::iterator I = HasRecMap.find(S);
4504 if (I != HasRecMap.end())
4505 return I->second;
4506
4507 bool FoundAddRec =
4508 SCEVExprContains(S, [](const SCEV *S) { return isa<SCEVAddRecExpr>(S); });
4509 HasRecMap.insert({S, FoundAddRec});
4510 return FoundAddRec;
4511}
4512
4513/// Return the ValueOffsetPair set for \p S. \p S can be represented
4514/// by the value and offset from any ValueOffsetPair in the set.
4515ArrayRef<Value *> ScalarEvolution::getSCEVValues(const SCEV *S) {
4516 ExprValueMapType::iterator SI = ExprValueMap.find_as(S);
4517 if (SI == ExprValueMap.end())
4518 return {};
4519 return SI->second.getArrayRef();
4520}
4521
4522/// Erase Value from ValueExprMap and ExprValueMap. ValueExprMap.erase(V)
4523/// cannot be used separately. eraseValueFromMap should be used to remove
4524/// V from ValueExprMap and ExprValueMap at the same time.
4525void ScalarEvolution::eraseValueFromMap(Value *V) {
4526 ValueExprMapType::iterator I = ValueExprMap.find_as(V);
4527 if (I != ValueExprMap.end()) {
4528 auto EVIt = ExprValueMap.find(I->second);
4529 bool Removed = EVIt->second.remove(V);
4530 (void) Removed;
4531 assert(Removed && "Value not in ExprValueMap?");
4532 ValueExprMap.erase(I);
4533 }
4534}
4535
4536void ScalarEvolution::insertValueToMap(Value *V, const SCEV *S) {
4537 // A recursive query may have already computed the SCEV. It should be
4538 // equivalent, but may not necessarily be exactly the same, e.g. due to lazily
4539 // inferred nowrap flags.
4540 auto It = ValueExprMap.find_as(V);
4541 if (It == ValueExprMap.end()) {
4542 ValueExprMap.insert({SCEVCallbackVH(V, this), S});
4543 ExprValueMap[S].insert(V);
4544 }
4545}
4546
4547/// Return an existing SCEV if it exists, otherwise analyze the expression and
4548/// create a new one.
4550 assert(isSCEVable(V->getType()) && "Value is not SCEVable!");
4551
4552 if (const SCEV *S = getExistingSCEV(V))
4553 return S;
4554 return createSCEVIter(V);
4555}
4556
4558 assert(isSCEVable(V->getType()) && "Value is not SCEVable!");
4559
4560 ValueExprMapType::iterator I = ValueExprMap.find_as(V);
4561 if (I != ValueExprMap.end()) {
4562 const SCEV *S = I->second;
4563 assert(checkValidity(S) &&
4564 "existing SCEV has not been properly invalidated");
4565 return S;
4566 }
4567 return nullptr;
4568}
4569
4570/// Return a SCEV corresponding to -V = -1*V
4572 if (const SCEVConstant *VC = dyn_cast<SCEVConstant>(V))
4573 return getConstant(
4574 cast<ConstantInt>(ConstantExpr::getNeg(VC->getValue())));
4575
4576 Type *Ty = V->getType();
4577 Ty = getEffectiveSCEVType(Ty);
4578 return getMulExpr(V, getMinusOne(Ty), Flags);
4579}
4580
4581/// If Expr computes ~A, return A else return nullptr
4582static const SCEV *MatchNotExpr(const SCEV *Expr) {
4583 const SCEV *MulOp;
4584 if (match(Expr, m_scev_Add(m_scev_AllOnes(),
4585 m_scev_Mul(m_scev_AllOnes(), m_SCEV(MulOp)))))
4586 return MulOp;
4587 return nullptr;
4588}
4589
4590/// Return a SCEV corresponding to ~V = -1-V
4592 assert(!V->getType()->isPointerTy() && "Can't negate pointer");
4593
4594 if (const SCEVConstant *VC = dyn_cast<SCEVConstant>(V))
4595 return getConstant(
4596 cast<ConstantInt>(ConstantExpr::getNot(VC->getValue())));
4597
4598 // Fold ~(u|s)(min|max)(~x, ~y) to (u|s)(max|min)(x, y)
4599 if (const SCEVMinMaxExpr *MME = dyn_cast<SCEVMinMaxExpr>(V)) {
4600 auto MatchMinMaxNegation = [&](const SCEVMinMaxExpr *MME) {
4601 SmallVector<SCEVUse, 2> MatchedOperands;
4602 for (const SCEV *Operand : MME->operands()) {
4603 const SCEV *Matched = MatchNotExpr(Operand);
4604 if (!Matched)
4605 return (const SCEV *)nullptr;
4606 MatchedOperands.push_back(Matched);
4607 }
4608 return getMinMaxExpr(SCEVMinMaxExpr::negate(MME->getSCEVType()),
4609 MatchedOperands);
4610 };
4611 if (const SCEV *Replaced = MatchMinMaxNegation(MME))
4612 return Replaced;
4613 }
4614
4615 Type *Ty = V->getType();
4616 Ty = getEffectiveSCEVType(Ty);
4617 return getMinusSCEV(getMinusOne(Ty), V);
4618}
4619
4621 assert(P->getType()->isPointerTy());
4622
4623 if (auto *AddRec = dyn_cast<SCEVAddRecExpr>(P)) {
4624 // The base of an AddRec is the first operand.
4625 SmallVector<SCEVUse> Ops{AddRec->operands()};
4626 Ops[0] = removePointerBase(Ops[0]);
4627 // Don't try to transfer nowrap flags for now. We could in some cases
4628 // (for example, if pointer operand of the AddRec is a SCEVUnknown).
4629 return getAddRecExpr(Ops, AddRec->getLoop(), SCEV::FlagNone);
4630 }
4631 if (auto *Add = dyn_cast<SCEVAddExpr>(P)) {
4632 // The base of an Add is the pointer operand.
4633 SmallVector<SCEVUse> Ops{Add->operands()};
4634 SCEVUse *PtrOp = nullptr;
4635 for (SCEVUse &AddOp : Ops) {
4636 if (AddOp->getType()->isPointerTy()) {
4637 assert(!PtrOp && "Cannot have multiple pointer ops");
4638 PtrOp = &AddOp;
4639 }
4640 }
4641 *PtrOp = removePointerBase(*PtrOp);
4642 // Don't try to transfer nowrap flags for now. We could in some cases
4643 // (for example, if the pointer operand of the Add is a SCEVUnknown).
4644 return getAddExpr(Ops);
4645 }
4646 // Any other expression must be a pointer base.
4647 return getZero(P->getType());
4648}
4649
4651 SCEVFlags Flags, unsigned Depth) {
4652 // Fast path: X - X --> 0.
4653 if (LHS == RHS)
4654 return getZero(LHS->getType());
4655
4656 // If we subtract two pointers with different pointer bases, bail.
4657 // Eventually, we're going to add an assertion to getMulExpr that we
4658 // can't multiply by a pointer.
4659 if (RHS->getType()->isPointerTy()) {
4660 if (!LHS->getType()->isPointerTy() ||
4661 getPointerBase(LHS) != getPointerBase(RHS))
4662 return getCouldNotCompute();
4663 LHS = removePointerBase(LHS);
4664 RHS = removePointerBase(RHS);
4665 }
4666
4667 // We represent LHS - RHS as LHS + (-1)*RHS. This transformation
4668 // makes it so that we cannot make much use of NUW.
4669 auto AddFlags = SCEV::FlagNone;
4670 const bool RHSIsNotMinSigned =
4672 if (hasFlags(Flags, SCEV::FlagNSW)) {
4673 // Let M be the minimum representable signed value. Then (-1)*RHS
4674 // signed-wraps if and only if RHS is M. That can happen even for
4675 // a NSW subtraction because e.g. (-1)*M signed-wraps even though
4676 // -1 - M does not. So to transfer NSW from LHS - RHS to LHS +
4677 // (-1)*RHS, we need to prove that RHS != M.
4678 //
4679 // If LHS is non-negative and we know that LHS - RHS does not
4680 // signed-wrap, then RHS cannot be M. So we can rule out signed-wrap
4681 // either by proving that RHS > M or that LHS >= 0.
4682 if (RHSIsNotMinSigned || isKnownNonNegative(LHS)) {
4683 AddFlags = SCEV::FlagNSW;
4684 }
4685 }
4686
4687 // FIXME: Find a correct way to transfer NSW to (-1)*M when LHS -
4688 // RHS is NSW and LHS >= 0.
4689 //
4690 // The difficulty here is that the NSW flag may have been proven
4691 // relative to a loop that is to be found in a recurrence in LHS and
4692 // not in RHS. Applying NSW to (-1)*M may then let the NSW have a
4693 // larger scope than intended.
4694 auto NegFlags = RHSIsNotMinSigned ? SCEV::FlagNSW : SCEV::FlagNone;
4695
4696 return getAddExpr(LHS, getNegativeSCEV(RHS, NegFlags), AddFlags, Depth);
4697}
4698
4700 unsigned Depth) {
4701 Type *SrcTy = V->getType();
4702 assert(SrcTy->isIntOrPtrTy() && Ty->isIntOrPtrTy() &&
4703 "Cannot truncate or zero extend with non-integer arguments!");
4704 if (getTypeSizeInBits(SrcTy) == getTypeSizeInBits(Ty))
4705 return V; // No conversion
4706 if (getTypeSizeInBits(SrcTy) > getTypeSizeInBits(Ty))
4707 return getTruncateExpr(V, Ty, Depth);
4708 return getZeroExtendExpr(V, Ty, Depth);
4709}
4710
4712 unsigned Depth) {
4713 Type *SrcTy = V->getType();
4714 assert(SrcTy->isIntOrPtrTy() && Ty->isIntOrPtrTy() &&
4715 "Cannot truncate or zero extend with non-integer arguments!");
4716 if (getTypeSizeInBits(SrcTy) == getTypeSizeInBits(Ty))
4717 return V; // No conversion
4718 if (getTypeSizeInBits(SrcTy) > getTypeSizeInBits(Ty))
4719 return getTruncateExpr(V, Ty, Depth);
4720 return getSignExtendExpr(V, Ty, Depth);
4721}
4722
4724 Type *SrcTy = V->getType();
4725 assert(SrcTy->isIntOrPtrTy() && Ty->isIntOrPtrTy() &&
4726 "Cannot noop or zero extend with non-integer arguments!");
4728 "getNoopOrZeroExtend cannot truncate!");
4729 if (getTypeSizeInBits(SrcTy) == getTypeSizeInBits(Ty))
4730 return V; // No conversion
4731 return getZeroExtendExpr(V, Ty);
4732}
4733
4735 Type *SrcTy = V->getType();
4736 assert(SrcTy->isIntOrPtrTy() && Ty->isIntOrPtrTy() &&
4737 "Cannot noop or sign extend with non-integer arguments!");
4739 "getNoopOrSignExtend cannot truncate!");
4740 if (getTypeSizeInBits(SrcTy) == getTypeSizeInBits(Ty))
4741 return V; // No conversion
4742 return getSignExtendExpr(V, Ty);
4743}
4744
4746 Type *SrcTy = V->getType();
4747 assert(SrcTy->isIntOrPtrTy() && Ty->isIntOrPtrTy() &&
4748 "Cannot noop or any extend with non-integer arguments!");
4750 "getNoopOrAnyExtend cannot truncate!");
4751 if (getTypeSizeInBits(SrcTy) == getTypeSizeInBits(Ty))
4752 return V; // No conversion
4753 return getAnyExtendExpr(V, Ty);
4754}
4755
4757 Type *SrcTy = V->getType();
4758 assert(SrcTy->isIntOrPtrTy() && Ty->isIntOrPtrTy() &&
4759 "Cannot truncate or noop with non-integer arguments!");
4761 "getTruncateOrNoop cannot extend!");
4762 if (getTypeSizeInBits(SrcTy) == getTypeSizeInBits(Ty))
4763 return V; // No conversion
4764 return getTruncateExpr(V, Ty);
4765}
4766
4768 const SCEV *RHS) {
4769 const SCEV *PromotedLHS = LHS;
4770 const SCEV *PromotedRHS = RHS;
4771
4772 if (getTypeSizeInBits(LHS->getType()) > getTypeSizeInBits(RHS->getType()))
4773 PromotedRHS = getZeroExtendExpr(RHS, LHS->getType());
4774 else
4775 PromotedLHS = getNoopOrZeroExtend(LHS, RHS->getType());
4776
4777 return getUMaxExpr(PromotedLHS, PromotedRHS);
4778}
4779
4781 const SCEV *RHS,
4782 bool Sequential) {
4783 SmallVector<SCEVUse, 2> Ops = {LHS, RHS};
4784 return getUMinFromMismatchedTypes(Ops, Sequential);
4785}
4786
4787const SCEV *
4789 bool Sequential) {
4790 assert(!Ops.empty() && "At least one operand must be!");
4791 // Trivial case.
4792 if (Ops.size() == 1)
4793 return Ops[0];
4794
4795 // Find the max type first.
4796 Type *MaxType = nullptr;
4797 for (SCEVUse S : Ops)
4798 if (MaxType)
4799 MaxType = getWiderType(MaxType, S->getType());
4800 else
4801 MaxType = S->getType();
4802 assert(MaxType && "Failed to find maximum type!");
4803
4804 // Extend all ops to max type.
4805 SmallVector<SCEVUse, 2> PromotedOps;
4806 for (SCEVUse S : Ops)
4807 PromotedOps.push_back(getNoopOrZeroExtend(S, MaxType));
4808
4809 // Generate umin.
4810 return getUMinExpr(PromotedOps, Sequential);
4811}
4812
4814 // A pointer operand may evaluate to a nonpointer expression, such as null.
4815 if (!V->getType()->isPointerTy())
4816 return V;
4817
4818 while (true) {
4819 if (auto *AddRec = dyn_cast<SCEVAddRecExpr>(V)) {
4820 V = AddRec->getStart();
4821 } else if (auto *Add = dyn_cast<SCEVAddExpr>(V)) {
4822 const SCEV *PtrOp = nullptr;
4823 for (const SCEV *AddOp : Add->operands()) {
4824 if (AddOp->getType()->isPointerTy()) {
4825 assert(!PtrOp && "Cannot have multiple pointer ops");
4826 PtrOp = AddOp;
4827 }
4828 }
4829 assert(PtrOp && "Must have pointer op");
4830 V = PtrOp;
4831 } else // Not something we can look further into.
4832 return V;
4833 }
4834}
4835
4836/// Push users of the given Instruction onto the given Worklist.
4840 // Push the def-use children onto the Worklist stack.
4841 for (User *U : I->users()) {
4842 auto *UserInsn = cast<Instruction>(U);
4843 if (Visited.insert(UserInsn).second)
4844 Worklist.push_back(UserInsn);
4845 }
4846}
4847
4848namespace {
4849
4850/// Takes SCEV S and Loop L. For each AddRec sub-expression, use its start
4851/// expression in case its Loop is L. If it is not L then
4852/// if IgnoreOtherLoops is true then use AddRec itself
4853/// otherwise rewrite cannot be done.
4854/// If SCEV contains non-invariant unknown SCEV rewrite cannot be done.
4855class SCEVInitRewriter : public SCEVRewriteVisitor<SCEVInitRewriter> {
4856public:
4857 static const SCEV *rewrite(const SCEV *S, const Loop *L, ScalarEvolution &SE,
4858 bool IgnoreOtherLoops = true) {
4859 SCEVInitRewriter Rewriter(L, SE);
4860 const SCEV *Result = Rewriter.visit(S);
4861 if (Rewriter.hasSeenLoopVariantSCEVUnknown())
4862 return SE.getCouldNotCompute();
4863 return Rewriter.hasSeenOtherLoops() && !IgnoreOtherLoops
4864 ? SE.getCouldNotCompute()
4865 : Result;
4866 }
4867
4868 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
4869 if (!SE.isLoopInvariant(Expr, L))
4870 SeenLoopVariantSCEVUnknown = true;
4871 return Expr;
4872 }
4873
4874 const SCEV *visitAddRecExpr(const SCEVAddRecExpr *Expr) {
4875 // Only re-write AddRecExprs for this loop.
4876 if (Expr->getLoop() == L)
4877 return Expr->getStart();
4878 SeenOtherLoops = true;
4879 return Expr;
4880 }
4881
4882 bool hasSeenLoopVariantSCEVUnknown() { return SeenLoopVariantSCEVUnknown; }
4883
4884 bool hasSeenOtherLoops() { return SeenOtherLoops; }
4885
4886private:
4887 explicit SCEVInitRewriter(const Loop *L, ScalarEvolution &SE)
4888 : SCEVRewriteVisitor(SE), L(L) {}
4889
4890 const Loop *L;
4891 bool SeenLoopVariantSCEVUnknown = false;
4892 bool SeenOtherLoops = false;
4893};
4894
4895/// Takes SCEV S and Loop L. For each AddRec sub-expression, use its post
4896/// increment expression in case its Loop is L. If it is not L then
4897/// use AddRec itself.
4898/// If SCEV contains non-invariant unknown SCEV rewrite cannot be done.
4899class SCEVPostIncRewriter : public SCEVRewriteVisitor<SCEVPostIncRewriter> {
4900public:
4901 static const SCEV *rewrite(const SCEV *S, const Loop *L, ScalarEvolution &SE) {
4902 SCEVPostIncRewriter Rewriter(L, SE);
4903 const SCEV *Result = Rewriter.visit(S);
4904 return Rewriter.hasSeenLoopVariantSCEVUnknown()
4905 ? SE.getCouldNotCompute()
4906 : Result;
4907 }
4908
4909 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
4910 if (!SE.isLoopInvariant(Expr, L))
4911 SeenLoopVariantSCEVUnknown = true;
4912 return Expr;
4913 }
4914
4915 const SCEV *visitAddRecExpr(const SCEVAddRecExpr *Expr) {
4916 // Only re-write AddRecExprs for this loop.
4917 if (Expr->getLoop() == L)
4918 return Expr->getPostIncExpr(SE);
4919 SeenOtherLoops = true;
4920 return Expr;
4921 }
4922
4923 bool hasSeenLoopVariantSCEVUnknown() { return SeenLoopVariantSCEVUnknown; }
4924
4925 bool hasSeenOtherLoops() { return SeenOtherLoops; }
4926
4927private:
4928 explicit SCEVPostIncRewriter(const Loop *L, ScalarEvolution &SE)
4929 : SCEVRewriteVisitor(SE), L(L) {}
4930
4931 const Loop *L;
4932 bool SeenLoopVariantSCEVUnknown = false;
4933 bool SeenOtherLoops = false;
4934};
4935
4936/// This class evaluates the compare condition by matching it against the
4937/// condition of loop latch. If there is a match we assume a true value
4938/// for the condition while building SCEV nodes.
4939class SCEVBackedgeConditionFolder
4940 : public SCEVRewriteVisitor<SCEVBackedgeConditionFolder> {
4941public:
4942 static const SCEV *rewrite(const SCEV *S, const Loop *L,
4943 ScalarEvolution &SE) {
4944 bool IsPosBECond = false;
4945 Value *BECond = nullptr;
4946 if (BasicBlock *Latch = L->getLoopLatch()) {
4947 if (CondBrInst *BI = dyn_cast<CondBrInst>(Latch->getTerminator())) {
4948 assert(BI->getSuccessor(0) != BI->getSuccessor(1) &&
4949 "Both outgoing branches should not target same header!");
4950 BECond = BI->getCondition();
4951 IsPosBECond = BI->getSuccessor(0) == L->getHeader();
4952 } else {
4953 return S;
4954 }
4955 }
4956 SCEVBackedgeConditionFolder Rewriter(L, BECond, IsPosBECond, SE);
4957 return Rewriter.visit(S);
4958 }
4959
4960 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
4961 const SCEV *Result = Expr;
4962 bool InvariantF = SE.isLoopInvariant(Expr, L);
4963
4964 if (!InvariantF) {
4966 switch (I->getOpcode()) {
4967 case Instruction::Select: {
4968 SelectInst *SI = cast<SelectInst>(I);
4969 std::optional<const SCEV *> Res =
4970 compareWithBackedgeCondition(SI->getCondition());
4971 if (Res) {
4972 bool IsOne = cast<SCEVConstant>(*Res)->getValue()->isOne();
4973 Result = SE.getSCEV(IsOne ? SI->getTrueValue() : SI->getFalseValue());
4974 }
4975 break;
4976 }
4977 default: {
4978 std::optional<const SCEV *> Res = compareWithBackedgeCondition(I);
4979 if (Res)
4980 Result = *Res;
4981 break;
4982 }
4983 }
4984 }
4985 return Result;
4986 }
4987
4988private:
4989 explicit SCEVBackedgeConditionFolder(const Loop *L, Value *BECond,
4990 bool IsPosBECond, ScalarEvolution &SE)
4991 : SCEVRewriteVisitor(SE), L(L), BackedgeCond(BECond),
4992 IsPositiveBECond(IsPosBECond) {}
4993
4994 std::optional<const SCEV *> compareWithBackedgeCondition(Value *IC);
4995
4996 const Loop *L;
4997 /// Loop back condition.
4998 Value *BackedgeCond = nullptr;
4999 /// Set to true if loop back is on positive branch condition.
5000 bool IsPositiveBECond;
5001};
5002
5003std::optional<const SCEV *>
5004SCEVBackedgeConditionFolder::compareWithBackedgeCondition(Value *IC) {
5005
5006 // If value matches the backedge condition for loop latch,
5007 // then return a constant evolution node based on loopback
5008 // branch taken.
5009 if (BackedgeCond == IC)
5010 return IsPositiveBECond ? SE.getOne(Type::getInt1Ty(SE.getContext()))
5012 return std::nullopt;
5013}
5014
5015class SCEVShiftRewriter : public SCEVRewriteVisitor<SCEVShiftRewriter> {
5016public:
5017 static const SCEV *rewrite(const SCEV *S, const Loop *L,
5018 ScalarEvolution &SE) {
5019 SCEVShiftRewriter Rewriter(L, SE);
5020 const SCEV *Result = Rewriter.visit(S);
5021 return Rewriter.isValid() ? Result : SE.getCouldNotCompute();
5022 }
5023
5024 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
5025 // Only allow AddRecExprs for this loop.
5026 if (!SE.isLoopInvariant(Expr, L))
5027 Valid = false;
5028 return Expr;
5029 }
5030
5031 const SCEV *visitAddRecExpr(const SCEVAddRecExpr *Expr) {
5032 if (Expr->getLoop() == L && Expr->isAffine())
5033 return SE.getMinusSCEV(Expr, Expr->getStepRecurrence(SE));
5034 Valid = false;
5035 return Expr;
5036 }
5037
5038 bool isValid() { return Valid; }
5039
5040private:
5041 explicit SCEVShiftRewriter(const Loop *L, ScalarEvolution &SE)
5042 : SCEVRewriteVisitor(SE), L(L) {}
5043
5044 const Loop *L;
5045 bool Valid = true;
5046};
5047
5048} // end anonymous namespace
5049
5050void ScalarEvolution::inferNoWrapViaConstantRanges(const SCEVAddRecExpr *AR) {
5051 if (!AR->isAffine())
5052 return;
5053
5054 // Force computation of ranges, which will also perform range-based flag
5055 // inference.
5056 if (!AR->hasNoSignedWrap())
5057 (void)getSignedRange(AR);
5058
5059 if (!AR->hasNoUnsignedWrap())
5060 (void)getUnsignedRange(AR);
5061
5062 if (!AR->hasNoSelfWrap()) {
5063 const SCEV *BECount = getConstantMaxBackedgeTakenCount(AR->getLoop());
5064 if (const SCEVConstant *BECountMax = dyn_cast<SCEVConstant>(BECount)) {
5065 ConstantRange StepCR = getSignedRange(AR->getStepRecurrence(*this));
5066 const APInt &BECountAP = BECountMax->getAPInt();
5067 unsigned NoOverflowBitWidth =
5068 BECountAP.getActiveBits() + StepCR.getMinSignedBits();
5069 if (NoOverflowBitWidth <= getTypeSizeInBits(AR->getType()))
5070 const_cast<SCEVAddRecExpr *>(AR)->setNoWrapFlags(SCEV::FlagNW);
5071 }
5072 }
5073}
5074
5076ScalarEvolution::proveNoSignedWrapViaInduction(const SCEVAddRecExpr *AR) {
5078
5079 if (AR->hasNoSignedWrap())
5080 return Result;
5081
5082 if (!AR->isAffine())
5083 return Result;
5084
5085 // This function can be expensive, only try to prove NSW once per AddRec.
5086 if (!SignedWrapViaInductionTried.insert(AR).second)
5087 return Result;
5088
5089 const SCEV *Step = AR->getStepRecurrence(*this);
5090 const Loop *L = AR->getLoop();
5091
5092 // Check whether the backedge-taken count is SCEVCouldNotCompute.
5093 // Note that this serves two purposes: It filters out loops that are
5094 // simply not analyzable, and it covers the case where this code is
5095 // being called from within backedge-taken count analysis, such that
5096 // attempting to ask for the backedge-taken count would likely result
5097 // in infinite recursion. In the later case, the analysis code will
5098 // cope with a conservative value, and it will take care to purge
5099 // that value once it has finished.
5100 const SCEV *MaxBECount = getConstantMaxBackedgeTakenCount(L);
5101
5102 // Normally, in the cases we can prove no-overflow via a
5103 // backedge guarding condition, we can also compute a backedge
5104 // taken count for the loop. The exceptions are assumptions and
5105 // guards present in the loop -- SCEV is not great at exploiting
5106 // these to compute max backedge taken counts, but can still use
5107 // these to prove lack of overflow. Use this fact to avoid
5108 // doing extra work that may not pay off.
5109
5110 if (isa<SCEVCouldNotCompute>(MaxBECount) && !HasGuards &&
5111 AC.assumptions().empty())
5112 return Result;
5113
5114 // If the backedge is guarded by a comparison with the pre-inc value the
5115 // addrec is safe. Also, if the entry is guarded by a comparison with the
5116 // start value and the backedge is guarded by a comparison with the post-inc
5117 // value, the addrec is safe.
5119 const SCEV *OverflowLimit =
5120 getSignedOverflowLimitForStep(Step, &Pred, this);
5121 if (OverflowLimit &&
5122 (isLoopBackedgeGuardedByCond(L, Pred, AR, OverflowLimit) ||
5123 isKnownOnEveryIteration(Pred, AR, OverflowLimit))) {
5124 Result = setFlags(Result, SCEV::FlagNSW);
5125 }
5126 return Result;
5127}
5129ScalarEvolution::proveNoUnsignedWrapViaInduction(const SCEVAddRecExpr *AR) {
5131
5132 if (AR->hasNoUnsignedWrap())
5133 return Result;
5134
5135 if (!AR->isAffine())
5136 return Result;
5137
5138 // This function can be expensive, only try to prove NUW once per AddRec.
5139 if (!UnsignedWrapViaInductionTried.insert(AR).second)
5140 return Result;
5141
5142 const SCEV *Step = AR->getStepRecurrence(*this);
5143 const Loop *L = AR->getLoop();
5144
5145 // Check whether the backedge-taken count is SCEVCouldNotCompute.
5146 // Note that this serves two purposes: It filters out loops that are
5147 // simply not analyzable, and it covers the case where this code is
5148 // being called from within backedge-taken count analysis, such that
5149 // attempting to ask for the backedge-taken count would likely result
5150 // in infinite recursion. In the later case, the analysis code will
5151 // cope with a conservative value, and it will take care to purge
5152 // that value once it has finished.
5153 const SCEV *MaxBECount = getConstantMaxBackedgeTakenCount(L);
5154
5155 // Normally, in the cases we can prove no-overflow via a
5156 // backedge guarding condition, we can also compute a backedge
5157 // taken count for the loop. The exceptions are assumptions and
5158 // guards present in the loop -- SCEV is not great at exploiting
5159 // these to compute max backedge taken counts, but can still use
5160 // these to prove lack of overflow. Use this fact to avoid
5161 // doing extra work that may not pay off.
5162
5163 if (isa<SCEVCouldNotCompute>(MaxBECount) && !HasGuards &&
5164 AC.assumptions().empty())
5165 return Result;
5166
5167 // If the backedge is guarded by a comparison with the pre-inc value the
5168 // addrec is safe. Also, if the entry is guarded by a comparison with the
5169 // start value and the backedge is guarded by a comparison with the post-inc
5170 // value, the addrec is safe.
5171 if (isKnownPositive(Step)) {
5173 const SCEV *OverflowLimit =
5174 getUnsignedOverflowLimitForStep(Step, &Pred, this);
5175 if (isLoopBackedgeGuardedByCond(L, Pred, AR, OverflowLimit) ||
5176 isKnownOnEveryIteration(Pred, AR, OverflowLimit))
5177 Result = setFlags(Result, SCEV::FlagNUW);
5178 }
5179 return Result;
5180}
5181
5182namespace {
5183
5184/// Represents an abstract binary operation. This may exist as a
5185/// normal instruction or constant expression, or may have been
5186/// derived from an expression tree.
5187struct BinaryOp {
5188 unsigned Opcode;
5189 Value *LHS;
5190 Value *RHS;
5191 bool IsNSW = false;
5192 bool IsNUW = false;
5193
5194 /// Op is set if this BinaryOp corresponds to a concrete LLVM instruction or
5195 /// constant expression.
5196 Operator *Op = nullptr;
5197
5198 explicit BinaryOp(Operator *Op)
5199 : Opcode(Op->getOpcode()), LHS(Op->getOperand(0)), RHS(Op->getOperand(1)),
5200 Op(Op) {
5201 if (auto *OBO = dyn_cast<OverflowingBinaryOperator>(Op)) {
5202 IsNSW = OBO->hasNoSignedWrap();
5203 IsNUW = OBO->hasNoUnsignedWrap();
5204 }
5205 }
5206
5207 explicit BinaryOp(unsigned Opcode, Value *LHS, Value *RHS, bool IsNSW = false,
5208 bool IsNUW = false)
5209 : Opcode(Opcode), LHS(LHS), RHS(RHS), IsNSW(IsNSW), IsNUW(IsNUW) {}
5210};
5211
5212} // end anonymous namespace
5213
5214/// Try to map \p V into a BinaryOp, and return \c std::nullopt on failure.
5215static std::optional<BinaryOp> MatchBinaryOp(Value *V, const DataLayout &DL,
5216 AssumptionCache &AC,
5217 const DominatorTree &DT,
5218 const Instruction *CtxI) {
5219 auto *Op = dyn_cast<Operator>(V);
5220 if (!Op)
5221 return std::nullopt;
5222
5223 // Implementation detail: all the cleverness here should happen without
5224 // creating new SCEV expressions -- our caller knowns tricks to avoid creating
5225 // SCEV expressions when possible, and we should not break that.
5226
5227 switch (Op->getOpcode()) {
5228 case Instruction::Add:
5229 case Instruction::Sub:
5230 case Instruction::Mul:
5231 case Instruction::UDiv:
5232 case Instruction::URem:
5233 case Instruction::And:
5234 case Instruction::AShr:
5235 case Instruction::Shl:
5236 return BinaryOp(Op);
5237
5238 case Instruction::Or: {
5239 // Convert or disjoint into add nuw nsw.
5240 if (cast<PossiblyDisjointInst>(Op)->isDisjoint()) {
5241 BinaryOp BinOp(Instruction::Add, Op->getOperand(0), Op->getOperand(1),
5242 /*IsNSW=*/true, /*IsNUW=*/true);
5243 // Keep the reference to the original instruction so that we can later
5244 // check whether it can produce poison value or not.
5245 BinOp.Op = Op;
5246 return BinOp;
5247 }
5248 return BinaryOp(Op);
5249 }
5250
5251 case Instruction::Xor:
5252 if (auto *RHSC = dyn_cast<ConstantInt>(Op->getOperand(1)))
5253 // If the RHS of the xor is a signmask, then this is just an add.
5254 // Instcombine turns add of signmask into xor as a strength reduction step.
5255 if (RHSC->getValue().isSignMask())
5256 return BinaryOp(Instruction::Add, Op->getOperand(0), Op->getOperand(1));
5257 // Binary `xor` is a bit-wise `add`.
5258 if (V->getType()->isIntegerTy(1))
5259 return BinaryOp(Instruction::Add, Op->getOperand(0), Op->getOperand(1));
5260 return BinaryOp(Op);
5261
5262 case Instruction::LShr:
5263 // Turn logical shift right of a constant into a unsigned divide.
5264 if (ConstantInt *SA = dyn_cast<ConstantInt>(Op->getOperand(1))) {
5265 uint32_t BitWidth = cast<IntegerType>(Op->getType())->getBitWidth();
5266
5267 // If the shift count is not less than the bitwidth, the result of
5268 // the shift is undefined. Don't try to analyze it, because the
5269 // resolution chosen here may differ from the resolution chosen in
5270 // other parts of the compiler.
5271 if (SA->getValue().ult(BitWidth)) {
5272 Constant *X =
5273 ConstantInt::get(SA->getContext(),
5274 APInt::getOneBitSet(BitWidth, SA->getZExtValue()));
5275 return BinaryOp(Instruction::UDiv, Op->getOperand(0), X);
5276 }
5277 }
5278 return BinaryOp(Op);
5279
5280 case Instruction::ExtractValue: {
5281 auto *EVI = cast<ExtractValueInst>(Op);
5282 if (EVI->getNumIndices() != 1 || EVI->getIndices()[0] != 0)
5283 break;
5284
5285 auto *WO = dyn_cast<WithOverflowInst>(EVI->getAggregateOperand());
5286 if (!WO)
5287 break;
5288
5289 Instruction::BinaryOps BinOp = WO->getBinaryOp();
5290 bool Signed = WO->isSigned();
5291 // TODO: Should add nuw/nsw flags for mul as well.
5292 if (BinOp == Instruction::Mul || !isOverflowIntrinsicNoWrap(WO, DT))
5293 return BinaryOp(BinOp, WO->getLHS(), WO->getRHS());
5294
5295 // Now that we know that all uses of the arithmetic-result component of
5296 // CI are guarded by the overflow check, we can go ahead and pretend
5297 // that the arithmetic is non-overflowing.
5298 return BinaryOp(BinOp, WO->getLHS(), WO->getRHS(),
5299 /* IsNSW = */ Signed, /* IsNUW = */ !Signed);
5300 }
5301
5302 default:
5303 break;
5304 }
5305
5306 // Recognise intrinsic loop.decrement.reg, and as this has exactly the same
5307 // semantics as a Sub, return a binary sub expression.
5308 if (auto *II = dyn_cast<IntrinsicInst>(V))
5309 if (II->getIntrinsicID() == Intrinsic::loop_decrement_reg)
5310 return BinaryOp(Instruction::Sub, II->getOperand(0), II->getOperand(1));
5311
5312 return std::nullopt;
5313}
5314
5315/// Helper function to createAddRecFromPHIWithCasts. We have a phi
5316/// node whose symbolic (unknown) SCEV is \p SymbolicPHI, which is updated via
5317/// the loop backedge by a SCEVAddExpr, possibly also with a few casts on the
5318/// way. This function checks if \p Op, an operand of this SCEVAddExpr,
5319/// follows one of the following patterns:
5320/// Op == (SExt ix (Trunc iy (%SymbolicPHI) to ix) to iy)
5321/// Op == (ZExt ix (Trunc iy (%SymbolicPHI) to ix) to iy)
5322/// If the SCEV expression of \p Op conforms with one of the expected patterns
5323/// we return the type of the truncation operation, and indicate whether the
5324/// truncated type should be treated as signed/unsigned by setting
5325/// \p Signed to true/false, respectively.
5326static Type *isSimpleCastedPHI(const SCEV *Op, const SCEVUnknown *SymbolicPHI,
5327 bool &Signed, ScalarEvolution &SE) {
5328 // The case where Op == SymbolicPHI (that is, with no type conversions on
5329 // the way) is handled by the regular add recurrence creating logic and
5330 // would have already been triggered in createAddRecForPHI. Reaching it here
5331 // means that createAddRecFromPHI had failed for this PHI before (e.g.,
5332 // because one of the other operands of the SCEVAddExpr updating this PHI is
5333 // not invariant).
5334 //
5335 // Here we look for the case where Op = (ext(trunc(SymbolicPHI))), and in
5336 // this case predicates that allow us to prove that Op == SymbolicPHI will
5337 // be added.
5338 if (Op == SymbolicPHI)
5339 return nullptr;
5340
5341 unsigned SourceBits = SE.getTypeSizeInBits(SymbolicPHI->getType());
5342 unsigned NewBits = SE.getTypeSizeInBits(Op->getType());
5343 if (SourceBits != NewBits)
5344 return nullptr;
5345
5346 if (match(Op, m_scev_SExt(m_scev_Trunc(m_scev_Specific(SymbolicPHI))))) {
5347 Signed = true;
5348 return cast<SCEVCastExpr>(Op)->getOperand()->getType();
5349 }
5350 if (match(Op, m_scev_ZExt(m_scev_Trunc(m_scev_Specific(SymbolicPHI))))) {
5351 Signed = false;
5352 return cast<SCEVCastExpr>(Op)->getOperand()->getType();
5353 }
5354 return nullptr;
5355}
5356
5357static const Loop *isIntegerLoopHeaderPHI(const PHINode *PN, LoopInfo &LI) {
5358 if (!PN->getType()->isIntegerTy())
5359 return nullptr;
5360 const Loop *L = LI.getLoopFor(PN->getParent());
5361 if (!L || L->getHeader() != PN->getParent())
5362 return nullptr;
5363 return L;
5364}
5365
5366// Analyze \p SymbolicPHI, a SCEV expression of a phi node, and check if the
5367// computation that updates the phi follows the following pattern:
5368// (SExt/ZExt ix (Trunc iy (%SymbolicPHI) to ix) to iy) + InvariantAccum
5369// which correspond to a phi->trunc->sext/zext->add->phi update chain.
5370// If so, try to see if it can be rewritten as an AddRecExpr under some
5371// Predicates. If successful, return them as a pair. Also cache the results
5372// of the analysis.
5373//
5374// Example usage scenario:
5375// Say the Rewriter is called for the following SCEV:
5376// 8 * ((sext i32 (trunc i64 %X to i32) to i64) + %Step)
5377// where:
5378// %X = phi i64 (%Start, %BEValue)
5379// It will visitMul->visitAdd->visitSExt->visitTrunc->visitUnknown(%X),
5380// and call this function with %SymbolicPHI = %X.
5381//
5382// The analysis will find that the value coming around the backedge has
5383// the following SCEV:
5384// BEValue = ((sext i32 (trunc i64 %X to i32) to i64) + %Step)
5385// Upon concluding that this matches the desired pattern, the function
5386// will return the pair {NewAddRec, SmallPredsVec} where:
5387// NewAddRec = {%Start,+,%Step}
5388// SmallPredsVec = {P1, P2, P3} as follows:
5389// P1(WrapPred): AR: {trunc(%Start),+,(trunc %Step)}<nsw> Flags: <nssw>
5390// P2(EqualPred): %Start == (sext i32 (trunc i64 %Start to i32) to i64)
5391// P3(EqualPred): %Step == (sext i32 (trunc i64 %Step to i32) to i64)
5392// The returned pair means that SymbolicPHI can be rewritten into NewAddRec
5393// under the predicates {P1,P2,P3}.
5394// This predicated rewrite will be cached in PredicatedSCEVRewrites:
5395// PredicatedSCEVRewrites[{%X,L}] = {NewAddRec, {P1,P2,P3)}
5396//
5397// TODO's:
5398//
5399// 1) Extend the Induction descriptor to also support inductions that involve
5400// casts: When needed (namely, when we are called in the context of the
5401// vectorizer induction analysis), a Set of cast instructions will be
5402// populated by this method, and provided back to isInductionPHI. This is
5403// needed to allow the vectorizer to properly record them to be ignored by
5404// the cost model and to avoid vectorizing them (otherwise these casts,
5405// which are redundant under the runtime overflow checks, will be
5406// vectorized, which can be costly).
5407//
5408// 2) Support additional induction/PHISCEV patterns: We also want to support
5409// inductions where the sext-trunc / zext-trunc operations (partly) occur
5410// after the induction update operation (the induction increment):
5411//
5412// (Trunc iy (SExt/ZExt ix (%SymbolicPHI + InvariantAccum) to iy) to ix)
5413// which correspond to a phi->add->trunc->sext/zext->phi update chain.
5414//
5415// (Trunc iy ((SExt/ZExt ix (%SymbolicPhi) to iy) + InvariantAccum) to ix)
5416// which correspond to a phi->trunc->add->sext/zext->phi update chain.
5417//
5418// 3) Outline common code with createAddRecFromPHI to avoid duplication.
5419std::optional<std::pair<const SCEV *, SmallVector<const SCEVPredicate *, 3>>>
5420ScalarEvolution::createAddRecFromPHIWithCastsImpl(const SCEVUnknown *SymbolicPHI) {
5422
5423 // *** Part1: Analyze if we have a phi-with-cast pattern for which we can
5424 // return an AddRec expression under some predicate.
5425
5426 auto *PN = cast<PHINode>(SymbolicPHI->getValue());
5427 const Loop *L = isIntegerLoopHeaderPHI(PN, LI);
5428 assert(L && "Expecting an integer loop header phi");
5429
5430 // The loop may have multiple entrances or multiple exits; we can analyze
5431 // this phi as an addrec if it has a unique entry value and a unique
5432 // backedge value.
5433 Value *BEValueV = nullptr, *StartValueV = nullptr;
5434 for (unsigned i = 0, e = PN->getNumIncomingValues(); i != e; ++i) {
5435 Value *V = PN->getIncomingValue(i);
5436 if (L->contains(PN->getIncomingBlock(i))) {
5437 if (!BEValueV) {
5438 BEValueV = V;
5439 } else if (BEValueV != V) {
5440 BEValueV = nullptr;
5441 break;
5442 }
5443 } else if (!StartValueV) {
5444 StartValueV = V;
5445 } else if (StartValueV != V) {
5446 StartValueV = nullptr;
5447 break;
5448 }
5449 }
5450 if (!BEValueV || !StartValueV)
5451 return std::nullopt;
5452
5453 const SCEV *BEValue = getSCEV(BEValueV);
5454
5455 // If the value coming around the backedge is an add with the symbolic
5456 // value we just inserted, possibly with casts that we can ignore under
5457 // an appropriate runtime guard, then we found a simple induction variable!
5458 const auto *Add = dyn_cast<SCEVAddExpr>(BEValue);
5459 if (!Add)
5460 return std::nullopt;
5461
5462 // If there is a single occurrence of the symbolic value, possibly
5463 // casted, replace it with a recurrence.
5464 unsigned FoundIndex = Add->getNumOperands();
5465 Type *TruncTy = nullptr;
5466 bool Signed;
5467 for (unsigned i = 0, e = Add->getNumOperands(); i != e; ++i)
5468 if ((TruncTy =
5469 isSimpleCastedPHI(Add->getOperand(i), SymbolicPHI, Signed, *this)))
5470 if (FoundIndex == e) {
5471 FoundIndex = i;
5472 break;
5473 }
5474
5475 if (FoundIndex == Add->getNumOperands())
5476 return std::nullopt;
5477
5478 // Create an add with everything but the specified operand.
5480 for (unsigned i = 0, e = Add->getNumOperands(); i != e; ++i)
5481 if (i != FoundIndex)
5482 Ops.push_back(Add->getOperand(i));
5483 const SCEV *Accum = getAddExpr(Ops);
5484
5485 // The runtime checks will not be valid if the step amount is
5486 // varying inside the loop.
5487 if (!isLoopInvariant(Accum, L))
5488 return std::nullopt;
5489
5490 // *** Part2: Create the predicates
5491
5492 // Analysis was successful: we have a phi-with-cast pattern for which we
5493 // can return an AddRec expression under the following predicates:
5494 //
5495 // P1: A Wrap predicate that guarantees that Trunc(Start) + i*Trunc(Accum)
5496 // fits within the truncated type (does not overflow) for i = 0 to n-1.
5497 // P2: An Equal predicate that guarantees that
5498 // Start = (Ext ix (Trunc iy (Start) to ix) to iy)
5499 // P3: An Equal predicate that guarantees that
5500 // Accum = (Ext ix (Trunc iy (Accum) to ix) to iy)
5501 //
5502 // As we next prove, the above predicates guarantee that:
5503 // Start + i*Accum = (Ext ix (Trunc iy ( Start + i*Accum ) to ix) to iy)
5504 //
5505 //
5506 // More formally, we want to prove that:
5507 // Expr(i+1) = Start + (i+1) * Accum
5508 // = (Ext ix (Trunc iy (Expr(i)) to ix) to iy) + Accum
5509 //
5510 // Given that:
5511 // 1) Expr(0) = Start
5512 // 2) Expr(1) = Start + Accum
5513 // = (Ext ix (Trunc iy (Start) to ix) to iy) + Accum :: from P2
5514 // 3) Induction hypothesis (step i):
5515 // Expr(i) = (Ext ix (Trunc iy (Expr(i-1)) to ix) to iy) + Accum
5516 //
5517 // Proof:
5518 // Expr(i+1) =
5519 // = Start + (i+1)*Accum
5520 // = (Start + i*Accum) + Accum
5521 // = Expr(i) + Accum
5522 // = (Ext ix (Trunc iy (Expr(i-1)) to ix) to iy) + Accum + Accum
5523 // :: from step i
5524 //
5525 // = (Ext ix (Trunc iy (Start + (i-1)*Accum) to ix) to iy) + Accum + Accum
5526 //
5527 // = (Ext ix (Trunc iy (Start + (i-1)*Accum) to ix) to iy)
5528 // + (Ext ix (Trunc iy (Accum) to ix) to iy)
5529 // + Accum :: from P3
5530 //
5531 // = (Ext ix (Trunc iy ((Start + (i-1)*Accum) + Accum) to ix) to iy)
5532 // + Accum :: from P1: Ext(x)+Ext(y)=>Ext(x+y)
5533 //
5534 // = (Ext ix (Trunc iy (Start + i*Accum) to ix) to iy) + Accum
5535 // = (Ext ix (Trunc iy (Expr(i)) to ix) to iy) + Accum
5536 //
5537 // By induction, the same applies to all iterations 1<=i<n:
5538 //
5539
5540 // Create a truncated addrec for which we will add a no overflow check (P1).
5541 const SCEV *StartVal = getSCEV(StartValueV);
5542 const SCEV *PHISCEV =
5543 getAddRecExpr(getTruncateExpr(StartVal, TruncTy),
5544 getTruncateExpr(Accum, TruncTy), L, SCEV::FlagNone);
5545
5546 // PHISCEV can be either a SCEVConstant or a SCEVAddRecExpr.
5547 // ex: If truncated Accum is 0 and StartVal is a constant, then PHISCEV
5548 // will be constant.
5549 //
5550 // If PHISCEV is a constant, then P1 degenerates into P2 or P3, so we don't
5551 // add P1.
5552 if (const auto *AR = dyn_cast<SCEVAddRecExpr>(PHISCEV)) {
5556 const SCEVPredicate *AddRecPred = getWrapPredicate(AR, AddedFlags);
5557 Predicates.push_back(AddRecPred);
5558 }
5559
5560 // Create the Equal Predicates P2,P3:
5561
5562 // It is possible that the predicates P2 and/or P3 are computable at
5563 // compile time due to StartVal and/or Accum being constants.
5564 // If either one is, then we can check that now and escape if either P2
5565 // or P3 is false.
5566
5567 // Construct the extended SCEV: (Ext ix (Trunc iy (Expr) to ix) to iy)
5568 // for each of StartVal and Accum
5569 auto getExtendedExpr = [&](const SCEV *Expr,
5570 bool CreateSignExtend) -> const SCEV * {
5571 assert(isLoopInvariant(Expr, L) && "Expr is expected to be invariant");
5572 const SCEV *TruncatedExpr = getTruncateExpr(Expr, TruncTy);
5573 const SCEV *ExtendedExpr =
5574 CreateSignExtend ? getSignExtendExpr(TruncatedExpr, Expr->getType())
5575 : getZeroExtendExpr(TruncatedExpr, Expr->getType());
5576 return ExtendedExpr;
5577 };
5578
5579 // Given:
5580 // ExtendedExpr = (Ext ix (Trunc iy (Expr) to ix) to iy
5581 // = getExtendedExpr(Expr)
5582 // Determine whether the predicate P: Expr == ExtendedExpr
5583 // is known to be false at compile time
5584 auto PredIsKnownFalse = [&](const SCEV *Expr,
5585 const SCEV *ExtendedExpr) -> bool {
5586 return Expr != ExtendedExpr &&
5587 isKnownPredicate(ICmpInst::ICMP_NE, Expr, ExtendedExpr);
5588 };
5589
5590 const SCEV *StartExtended = getExtendedExpr(StartVal, Signed);
5591 if (PredIsKnownFalse(StartVal, StartExtended)) {
5592 LLVM_DEBUG(dbgs() << "P2 is compile-time false\n";);
5593 return std::nullopt;
5594 }
5595
5596 // The Step is always Signed (because the overflow checks are either
5597 // NSSW or NUSW)
5598 const SCEV *AccumExtended = getExtendedExpr(Accum, /*CreateSignExtend=*/true);
5599 if (PredIsKnownFalse(Accum, AccumExtended)) {
5600 LLVM_DEBUG(dbgs() << "P3 is compile-time false\n";);
5601 return std::nullopt;
5602 }
5603
5604 auto AppendPredicate = [&](const SCEV *Expr,
5605 const SCEV *ExtendedExpr) -> void {
5606 if (Expr != ExtendedExpr &&
5607 !isKnownPredicate(ICmpInst::ICMP_EQ, Expr, ExtendedExpr)) {
5608 const SCEVPredicate *Pred = getEqualPredicate(Expr, ExtendedExpr);
5609 LLVM_DEBUG(dbgs() << "Added Predicate: " << *Pred);
5610 Predicates.push_back(Pred);
5611 }
5612 };
5613
5614 AppendPredicate(StartVal, StartExtended);
5615 AppendPredicate(Accum, AccumExtended);
5616
5617 // *** Part3: Predicates are ready. Now go ahead and create the new addrec in
5618 // which the casts had been folded away. The caller can rewrite SymbolicPHI
5619 // into NewAR if it will also add the runtime overflow checks specified in
5620 // Predicates.
5621 const SCEV *NewAR = getAddRecExpr(StartVal, Accum, L, SCEV::FlagNone);
5622
5623 std::pair<const SCEV *, SmallVector<const SCEVPredicate *, 3>> PredRewrite =
5624 std::make_pair(NewAR, Predicates);
5625 // Remember the result of the analysis for this SCEV at this locayyytion.
5626 PredicatedSCEVRewrites[{SymbolicPHI, L}] = PredRewrite;
5627 return PredRewrite;
5628}
5629
5630std::optional<std::pair<const SCEV *, SmallVector<const SCEVPredicate *, 3>>>
5632 auto *PN = cast<PHINode>(SymbolicPHI->getValue());
5633 const Loop *L = isIntegerLoopHeaderPHI(PN, LI);
5634 if (!L)
5635 return std::nullopt;
5636
5637 // Check to see if we already analyzed this PHI.
5638 auto I = PredicatedSCEVRewrites.find({SymbolicPHI, L});
5639 if (I != PredicatedSCEVRewrites.end()) {
5640 std::pair<const SCEV *, SmallVector<const SCEVPredicate *, 3>> Rewrite =
5641 I->second;
5642 // Analysis was done before and failed to create an AddRec:
5643 if (Rewrite.first == SymbolicPHI)
5644 return std::nullopt;
5645 // Analysis was done before and succeeded to create an AddRec under
5646 // a predicate:
5647 assert(isa<SCEVAddRecExpr>(Rewrite.first) && "Expected an AddRec");
5648 assert(!(Rewrite.second).empty() && "Expected to find Predicates");
5649 return Rewrite;
5650 }
5651
5652 std::optional<std::pair<const SCEV *, SmallVector<const SCEVPredicate *, 3>>>
5653 Rewrite = createAddRecFromPHIWithCastsImpl(SymbolicPHI);
5654
5655 // Record in the cache that the analysis failed
5656 if (!Rewrite) {
5658 PredicatedSCEVRewrites[{SymbolicPHI, L}] = {SymbolicPHI, Predicates};
5659 return std::nullopt;
5660 }
5661
5662 return Rewrite;
5663}
5664
5665// FIXME: This utility is currently required because the Rewriter currently
5666// does not rewrite this expression:
5667// {0, +, (sext ix (trunc iy to ix) to iy)}
5668// into {0, +, %step},
5669// even when the following Equal predicate exists:
5670// "%step == (sext ix (trunc iy to ix) to iy)".
5672 const SCEVAddRecExpr *AR1, const SCEVAddRecExpr *AR2,
5673 ArrayRef<const SCEVPredicate *> NoWrapPreds) const {
5674 if (AR1 == AR2)
5675 return true;
5676
5677 SCEVUnionPredicate NoWrapUnionPred(NoWrapPreds, SE);
5678 SCEVUnionPredicate AllPreds = Preds->getUnionWith(&NoWrapUnionPred, SE);
5679 auto areExprsEqual = [&](const SCEV *Expr1, const SCEV *Expr2) -> bool {
5680 if (Expr1 != Expr2 &&
5681 !AllPreds.implies(SE.getEqualPredicate(Expr1, Expr2), SE) &&
5682 !AllPreds.implies(SE.getEqualPredicate(Expr2, Expr1), SE))
5683 return false;
5684 return true;
5685 };
5686
5687 if (!areExprsEqual(AR1->getStart(), AR2->getStart()) ||
5688 !areExprsEqual(AR1->getStepRecurrence(SE), AR2->getStepRecurrence(SE)))
5689 return false;
5690 return true;
5691}
5692
5694 ScalarEvolution &SE) {
5695 SCEVFlags Flags = SCEV::FlagNone;
5696 GEPNoWrapFlags NW = GEP->getNoWrapFlags();
5697 // If the increment has any nowrap flags, then we know the address
5698 // space cannot be wrapped around.
5699 if (NW != GEPNoWrapFlags::none())
5701 // If the GEP is nuw or nusw with non-negative offset, we know that
5702 // no unsigned wrap occurs. We cannot set the nsw flag as only the
5703 // offset is treated as signed, while the base is unsigned.
5704 if (NW.hasNoUnsignedWrap() ||
5705 (NW.hasNoUnsignedSignedWrap() && SE.isKnownNonNegative(Accum)))
5707
5708 return Flags;
5709}
5710
5711/// A helper function for createAddRecFromPHI to handle simple cases.
5712///
5713/// This function tries to find an AddRec expression for the simplest (yet most
5714/// common) cases: PN = PHI(Start, OP(Self, LoopInvariant)).
5715/// If it fails, createAddRecFromPHI will use a more general, but slow,
5716/// technique for finding the AddRec expression.
5717const SCEV *ScalarEvolution::createSimpleAffineAddRec(PHINode *PN,
5718 Value *BEValueV,
5719 Value *StartValueV) {
5720 const Loop *L = LI.getLoopFor(PN->getParent());
5721 assert(L && L->getHeader() == PN->getParent());
5722 assert(BEValueV && StartValueV);
5723
5724 const SCEV *Accum = nullptr;
5726 if (auto BO = MatchBinaryOp(BEValueV, getDataLayout(), AC, DT, PN)) {
5727 if (BO->Opcode != Instruction::Add)
5728 return nullptr;
5729
5730 if (BO->LHS == PN && L->isLoopInvariant(BO->RHS))
5731 Accum = getSCEV(BO->RHS);
5732 else if (BO->RHS == PN && L->isLoopInvariant(BO->LHS))
5733 Accum = getSCEV(BO->LHS);
5734
5735 if (!Accum)
5736 return nullptr;
5737
5738 if (BO->IsNUW)
5739 Flags = setFlags(Flags, SCEV::FlagNUW);
5740 if (BO->IsNSW)
5741 Flags = setFlags(Flags, SCEV::FlagNSW);
5742 } else {
5743 // Handle pointer induction variable: PN = PHI(Start, gep PN,
5744 // LoopInvariant).
5745 auto *GEP = dyn_cast<GEPOperator>(BEValueV);
5746 if (!GEP || GEP->getPointerOperand() != PN || GEP->getNumIndices() != 1)
5747 return nullptr;
5748 Value *Idx = *GEP->idx_begin();
5749 if (!L->isLoopInvariant(Idx))
5750 return nullptr;
5751
5752 Type *IntIdxTy = getEffectiveSCEVType(GEP->getType());
5753 Accum = getMulExpr(getTruncateOrSignExtend(getSCEV(Idx), IntIdxTy),
5754 getSizeOfExpr(IntIdxTy, GEP->getSourceElementType()));
5755 Flags = getNoWrapFlagsForGEP(GEP, Accum, *this);
5756 }
5757
5758 const SCEV *StartVal = getSCEV(StartValueV);
5759 const SCEV *PHISCEV = getAddRecExpr(StartVal, Accum, L, Flags);
5760 insertValueToMap(PN, PHISCEV);
5761
5762 if (auto *AR = dyn_cast<SCEVAddRecExpr>(PHISCEV))
5763 inferNoWrapViaConstantRanges(AR);
5764
5765 // We can add Flags to the post-inc expression only if we
5766 // know that it is *undefined behavior* for BEValueV to
5767 // overflow.
5768 if (auto *BEInst = dyn_cast<Instruction>(BEValueV)) {
5769 assert(isLoopInvariant(Accum, L) &&
5770 "Accum is defined outside L, but is not invariant?");
5771 if (isAddRecNeverPoison(BEInst, L))
5772 (void)getAddRecExpr(getAddExpr(StartVal, Accum), Accum, L, Flags);
5773 }
5774
5775 return PHISCEV;
5776}
5777
5778const SCEV *ScalarEvolution::createAddRecFromPHI(PHINode *PN) {
5779 const Loop *L = LI.getLoopFor(PN->getParent());
5780 if (!L || L->getHeader() != PN->getParent())
5781 return nullptr;
5782
5783 // The loop may have multiple entrances or multiple exits; we can analyze
5784 // this phi as an addrec if it has a unique entry value and a unique
5785 // backedge value.
5786 Value *BEValueV = nullptr, *StartValueV = nullptr;
5787 for (unsigned i = 0, e = PN->getNumIncomingValues(); i != e; ++i) {
5788 Value *V = PN->getIncomingValue(i);
5789 if (L->contains(PN->getIncomingBlock(i))) {
5790 if (!BEValueV) {
5791 BEValueV = V;
5792 } else if (BEValueV != V) {
5793 BEValueV = nullptr;
5794 break;
5795 }
5796 } else if (!StartValueV) {
5797 StartValueV = V;
5798 } else if (StartValueV != V) {
5799 StartValueV = nullptr;
5800 break;
5801 }
5802 }
5803 if (!BEValueV || !StartValueV)
5804 return nullptr;
5805
5806 assert(ValueExprMap.find_as(PN) == ValueExprMap.end() &&
5807 "PHI node already processed?");
5808
5809 // First, try to find AddRec expression without creating a fictituos symbolic
5810 // value for PN.
5811 if (auto *S = createSimpleAffineAddRec(PN, BEValueV, StartValueV))
5812 return S;
5813
5814 // Handle PHI node value symbolically.
5815 const SCEV *SymbolicName = getUnknown(PN);
5816 insertValueToMap(PN, SymbolicName);
5817
5818 // Using this symbolic name for the PHI, analyze the value coming around
5819 // the back-edge.
5820 const SCEV *BEValue = getSCEV(BEValueV);
5821
5822 // NOTE: If BEValue is loop invariant, we know that the PHI node just
5823 // has a special value for the first iteration of the loop.
5824
5825 // If the value coming around the backedge is an add with the symbolic
5826 // value we just inserted, then we found a simple induction variable!
5827 if (const SCEVAddExpr *Add = dyn_cast<SCEVAddExpr>(BEValue)) {
5828 // If there is a single occurrence of the symbolic value, replace it
5829 // with a recurrence.
5830 unsigned FoundIndex = Add->getNumOperands();
5831 for (unsigned i = 0, e = Add->getNumOperands(); i != e; ++i)
5832 if (Add->getOperand(i) == SymbolicName)
5833 if (FoundIndex == e) {
5834 FoundIndex = i;
5835 break;
5836 }
5837
5838 if (FoundIndex != Add->getNumOperands()) {
5839 // Create an add with everything but the specified operand.
5841 for (unsigned i = 0, e = Add->getNumOperands(); i != e; ++i)
5842 if (i != FoundIndex)
5843 Ops.push_back(SCEVBackedgeConditionFolder::rewrite(Add->getOperand(i),
5844 L, *this));
5845 const SCEV *Accum = getAddExpr(Ops);
5846
5847 // This is not a valid addrec if the step amount is varying each
5848 // loop iteration, but is not itself an addrec in this loop.
5849 if (isLoopInvariant(Accum, L) ||
5850 (isa<SCEVAddRecExpr>(Accum) &&
5851 cast<SCEVAddRecExpr>(Accum)->getLoop() == L)) {
5853
5854 if (auto BO = MatchBinaryOp(BEValueV, getDataLayout(), AC, DT, PN)) {
5855 if (BO->Opcode == Instruction::Add && BO->LHS == PN) {
5856 if (BO->IsNUW)
5857 Flags = setFlags(Flags, SCEV::FlagNUW);
5858 if (BO->IsNSW)
5859 Flags = setFlags(Flags, SCEV::FlagNSW);
5860 }
5861 } else if (GEPOperator *GEP = dyn_cast<GEPOperator>(BEValueV)) {
5862 if (GEP->getOperand(0) == PN)
5863 Flags = getNoWrapFlagsForGEP(GEP, Accum, *this);
5864
5865 // We cannot transfer nuw and nsw flags from subtraction
5866 // operations -- sub nuw X, Y is not the same as add nuw X, -Y
5867 // for instance.
5868 }
5869
5870 const SCEV *StartVal = getSCEV(StartValueV);
5871 const SCEV *PHISCEV = getAddRecExpr(StartVal, Accum, L, Flags);
5872
5873 // Okay, for the entire analysis of this edge we assumed the PHI
5874 // to be symbolic. We now need to go back and purge all of the
5875 // entries for the scalars that use the symbolic expression.
5876 forgetMemoizedResults({SymbolicName});
5877 insertValueToMap(PN, PHISCEV);
5878
5879 if (auto *AR = dyn_cast<SCEVAddRecExpr>(PHISCEV))
5880 inferNoWrapViaConstantRanges(AR);
5881
5882 // We can add Flags to the post-inc expression only if we
5883 // know that it is *undefined behavior* for BEValueV to
5884 // overflow.
5885 if (auto *BEInst = dyn_cast<Instruction>(BEValueV))
5886 if (isLoopInvariant(Accum, L) && isAddRecNeverPoison(BEInst, L))
5887 (void)getAddRecExpr(getAddExpr(StartVal, Accum), Accum, L, Flags);
5888
5889 return PHISCEV;
5890 }
5891 }
5892 } else {
5893 // Otherwise, this could be a loop like this:
5894 // i = 0; for (j = 1; ..; ++j) { .... i = j; }
5895 // In this case, j = {1,+,1} and BEValue is j.
5896 // Because the other in-value of i (0) fits the evolution of BEValue
5897 // i really is an addrec evolution.
5898 //
5899 // We can generalize this saying that i is the shifted value of BEValue
5900 // by one iteration:
5901 // PHI(f(0), f({1,+,1})) --> f({0,+,1})
5902
5903 // Do not allow refinement in rewriting of BEValue.
5904 const SCEV *Shifted = SCEVShiftRewriter::rewrite(BEValue, L, *this);
5905 const SCEV *Start = SCEVInitRewriter::rewrite(Shifted, L, *this, false);
5906 if (Shifted != getCouldNotCompute() && Start != getCouldNotCompute() &&
5907 isGuaranteedNotToCauseUB(Shifted) && ::impliesPoison(Shifted, Start)) {
5908 const SCEV *StartVal = getSCEV(StartValueV);
5909 if (Start == StartVal) {
5910 // Okay, for the entire analysis of this edge we assumed the PHI
5911 // to be symbolic. We now need to go back and purge all of the
5912 // entries for the scalars that use the symbolic expression.
5913 forgetMemoizedResults({SymbolicName});
5914 insertValueToMap(PN, Shifted);
5915 return Shifted;
5916 }
5917 }
5918 }
5919
5920 // Remove the temporary PHI node SCEV that has been inserted while intending
5921 // to create an AddRecExpr for this PHI node. We can not keep this temporary
5922 // as it will prevent later (possibly simpler) SCEV expressions to be added
5923 // to the ValueExprMap.
5924 eraseValueFromMap(PN);
5925
5926 return nullptr;
5927}
5928
5929// Try to match a control flow sequence that branches out at BI and merges back
5930// at Merge into a "C ? LHS : RHS" select pattern. Return true on a successful
5931// match.
5933 Value *&C, Value *&LHS, Value *&RHS) {
5934 C = BI->getCondition();
5935
5936 BasicBlockEdge LeftEdge(BI->getParent(), BI->getSuccessor(0));
5937 BasicBlockEdge RightEdge(BI->getParent(), BI->getSuccessor(1));
5938
5939 Use &LeftUse = Merge->getOperandUse(0);
5940 Use &RightUse = Merge->getOperandUse(1);
5941
5942 if (DT.dominates(LeftEdge, LeftUse) && DT.dominates(RightEdge, RightUse)) {
5943 LHS = LeftUse;
5944 RHS = RightUse;
5945 return true;
5946 }
5947
5948 if (DT.dominates(LeftEdge, RightUse) && DT.dominates(RightEdge, LeftUse)) {
5949 LHS = RightUse;
5950 RHS = LeftUse;
5951 return true;
5952 }
5953
5954 return false;
5955}
5956
5958 Value *&Cond, Value *&LHS,
5959 Value *&RHS) {
5960 auto IsReachable =
5961 [&](BasicBlock *BB) { return DT.isReachableFromEntry(BB); };
5962 if (PN->getNumIncomingValues() == 2 && all_of(PN->blocks(), IsReachable)) {
5963 // Try to match
5964 //
5965 // br %cond, label %left, label %right
5966 // left:
5967 // br label %merge
5968 // right:
5969 // br label %merge
5970 // merge:
5971 // V = phi [ %x, %left ], [ %y, %right ]
5972 //
5973 // as "select %cond, %x, %y"
5974
5975 BasicBlock *IDom = DT[PN->getParent()]->getIDom()->getBlock();
5976 assert(IDom && "At least the entry block should dominate PN");
5977
5978 auto *BI = dyn_cast<CondBrInst>(IDom->getTerminator());
5979 return BI && BrPHIToSelect(DT, BI, PN, Cond, LHS, RHS);
5980 }
5981 return false;
5982}
5983
5984const SCEV *ScalarEvolution::createNodeFromSelectLikePHI(PHINode *PN) {
5985 Value *Cond = nullptr, *LHS = nullptr, *RHS = nullptr;
5986 if (getOperandsForSelectLikePHI(DT, PN, Cond, LHS, RHS) &&
5989 return createNodeForSelectOrPHI(PN, Cond, LHS, RHS);
5990
5991 return nullptr;
5992}
5993
5995 BinaryOperator *CommonInst = nullptr;
5996 // Check if instructions are identical.
5997 for (Value *Incoming : PN->incoming_values()) {
5998 auto *IncomingInst = dyn_cast<BinaryOperator>(Incoming);
5999 if (!IncomingInst)
6000 return nullptr;
6001 if (CommonInst) {
6002 if (!CommonInst->isIdenticalToWhenDefined(IncomingInst))
6003 return nullptr; // Not identical, give up
6004 } else {
6005 // Remember binary operator
6006 CommonInst = IncomingInst;
6007 }
6008 }
6009 return CommonInst;
6010}
6011
6012/// Returns SCEV for the first operand of a phi if all phi operands have
6013/// identical opcodes and operands
6014/// eg.
6015/// a: %add = %a + %b
6016/// br %c
6017/// b: %add1 = %a + %b
6018/// br %c
6019/// c: %phi = phi [%add, a], [%add1, b]
6020/// scev(%phi) => scev(%add)
6021const SCEV *
6022ScalarEvolution::createNodeForPHIWithIdenticalOperands(PHINode *PN) {
6023 BinaryOperator *CommonInst = getCommonInstForPHI(PN);
6024 if (!CommonInst)
6025 return nullptr;
6026
6027 // Check if SCEV exprs for instructions are identical.
6028 const SCEV *CommonSCEV = getSCEV(CommonInst);
6029 bool SCEVExprsIdentical =
6031 [this, CommonSCEV](Value *V) { return CommonSCEV == getSCEV(V); });
6032 return SCEVExprsIdentical ? CommonSCEV : nullptr;
6033}
6034
6035const SCEV *ScalarEvolution::createNodeForPHI(PHINode *PN) {
6036 if (const SCEV *S = createAddRecFromPHI(PN))
6037 return S;
6038
6039 // We do not allow simplifying phi (undef, X) to X here, to avoid reusing the
6040 // phi node for X.
6041 if (Value *V = simplifyInstruction(
6042 PN, {getDataLayout(), &TLI, &DT, &AC, /*CtxI=*/nullptr,
6043 /*UseInstrInfo=*/true, /*CanUseUndef=*/false}))
6044 return getSCEV(V);
6045
6046 if (const SCEV *S = createNodeForPHIWithIdenticalOperands(PN))
6047 return S;
6048
6049 if (const SCEV *S = createNodeFromSelectLikePHI(PN))
6050 return S;
6051
6052 // If it's not a loop phi, we can't handle it yet.
6053 return getUnknown(PN);
6054}
6055
6056bool SCEVMinMaxExprContains(const SCEV *Root, const SCEV *OperandToFind,
6057 SCEVTypes RootKind) {
6058 struct FindClosure {
6059 const SCEV *OperandToFind;
6060 const SCEVTypes RootKind; // Must be a sequential min/max expression.
6061 const SCEVTypes NonSequentialRootKind; // Non-seq variant of RootKind.
6062
6063 bool Found = false;
6064
6065 bool canRecurseInto(SCEVTypes Kind) const {
6066 // We can only recurse into the SCEV expression of the same effective type
6067 // as the type of our root SCEV expression, and into zero-extensions.
6068 return RootKind == Kind || NonSequentialRootKind == Kind ||
6069 scZeroExtend == Kind;
6070 };
6071
6072 FindClosure(const SCEV *OperandToFind, SCEVTypes RootKind)
6073 : OperandToFind(OperandToFind), RootKind(RootKind),
6074 NonSequentialRootKind(
6076 RootKind)) {}
6077
6078 bool follow(const SCEV *S) {
6079 Found = S == OperandToFind;
6080
6081 return !isDone() && canRecurseInto(S->getSCEVType());
6082 }
6083
6084 bool isDone() const { return Found; }
6085 };
6086
6087 FindClosure FC(OperandToFind, RootKind);
6088 visitAll(Root, FC);
6089 return FC.Found;
6090}
6091
6092std::optional<const SCEV *>
6093ScalarEvolution::createNodeForSelectOrPHIInstWithICmpInstCond(Type *Ty,
6094 ICmpInst *Cond,
6095 Value *TrueVal,
6096 Value *FalseVal) {
6097 // Try to match some simple smax or umax patterns.
6098 auto *ICI = Cond;
6099
6100 Value *LHS = ICI->getOperand(0);
6101 Value *RHS = ICI->getOperand(1);
6102
6103 switch (ICI->getPredicate()) {
6104 case ICmpInst::ICMP_SLT:
6105 case ICmpInst::ICMP_SLE:
6106 case ICmpInst::ICMP_ULT:
6107 case ICmpInst::ICMP_ULE:
6108 std::swap(LHS, RHS);
6109 [[fallthrough]];
6110 case ICmpInst::ICMP_SGT:
6111 case ICmpInst::ICMP_SGE:
6112 case ICmpInst::ICMP_UGT:
6113 case ICmpInst::ICMP_UGE:
6114 // a > b ? a+x : b+x -> max(a, b)+x
6115 // a > b ? b+x : a+x -> min(a, b)+x
6117 bool Signed = ICI->isSigned();
6118 const SCEV *LA = getSCEV(TrueVal);
6119 const SCEV *RA = getSCEV(FalseVal);
6120 const SCEV *LS = getSCEV(LHS);
6121 const SCEV *RS = getSCEV(RHS);
6122 if (LA->getType()->isPointerTy()) {
6123 // FIXME: Handle cases where LS/RS are pointers not equal to LA/RA.
6124 // Need to make sure we can't produce weird expressions involving
6125 // negated pointers.
6126 if (LA == LS && RA == RS)
6127 return Signed ? getSMaxExpr(LS, RS) : getUMaxExpr(LS, RS);
6128 if (LA == RS && RA == LS)
6129 return Signed ? getSMinExpr(LS, RS) : getUMinExpr(LS, RS);
6130 }
6131 auto CoerceOperand = [&](const SCEV *Op) -> const SCEV * {
6132 if (Op->getType()->isPointerTy()) {
6135 return Op;
6136 }
6137 if (Signed)
6138 Op = getNoopOrSignExtend(Op, Ty);
6139 else
6140 Op = getNoopOrZeroExtend(Op, Ty);
6141 return Op;
6142 };
6143 LS = CoerceOperand(LS);
6144 RS = CoerceOperand(RS);
6146 break;
6147 const SCEV *LDiff = getMinusSCEV(LA, LS);
6148 const SCEV *RDiff = getMinusSCEV(RA, RS);
6149 if (LDiff == RDiff)
6150 return getAddExpr(Signed ? getSMaxExpr(LS, RS) : getUMaxExpr(LS, RS),
6151 LDiff);
6152 LDiff = getMinusSCEV(LA, RS);
6153 RDiff = getMinusSCEV(RA, LS);
6154 if (LDiff == RDiff)
6155 return getAddExpr(Signed ? getSMinExpr(LS, RS) : getUMinExpr(LS, RS),
6156 LDiff);
6157 }
6158 break;
6159 case ICmpInst::ICMP_NE:
6160 // x != 0 ? x+y : C+y -> x == 0 ? C+y : x+y
6161 std::swap(TrueVal, FalseVal);
6162 [[fallthrough]];
6163 case ICmpInst::ICMP_EQ:
6164 // x == 0 ? C+y : x+y -> umax(x, C)+y iff C u<= 1
6167 const SCEV *X = getNoopOrZeroExtend(getSCEV(LHS), Ty);
6168 const SCEV *TrueValExpr = getSCEV(TrueVal); // C+y
6169 const SCEV *FalseValExpr = getSCEV(FalseVal); // x+y
6170 const SCEV *Y = getMinusSCEV(FalseValExpr, X); // y = (x+y)-x
6171 const SCEV *C = getMinusSCEV(TrueValExpr, Y); // C = (C+y)-y
6172 if (isa<SCEVConstant>(C) && cast<SCEVConstant>(C)->getAPInt().ule(1))
6173 return getAddExpr(getUMaxExpr(X, C), Y);
6174 }
6175 // x == 0 ? 0 : umin (..., x, ...) -> umin_seq(x, umin (...))
6176 // x == 0 ? 0 : umin_seq(..., x, ...) -> umin_seq(x, umin_seq(...))
6177 // x == 0 ? 0 : umin (..., umin_seq(..., x, ...), ...)
6178 // -> umin_seq(x, umin (..., umin_seq(...), ...))
6180 isa<ConstantInt>(TrueVal) && cast<ConstantInt>(TrueVal)->isZero()) {
6181 const SCEV *X = getSCEV(LHS);
6182 while (auto *ZExt = dyn_cast<SCEVZeroExtendExpr>(X))
6183 X = ZExt->getOperand();
6184 if (getTypeSizeInBits(X->getType()) <= getTypeSizeInBits(Ty)) {
6185 const SCEV *FalseValExpr = getSCEV(FalseVal);
6186 if (SCEVMinMaxExprContains(FalseValExpr, X, scSequentialUMinExpr))
6187 return getUMinExpr(getNoopOrZeroExtend(X, Ty), FalseValExpr,
6188 /*Sequential=*/true);
6189 }
6190 }
6191 break;
6192 default:
6193 break;
6194 }
6195
6196 return std::nullopt;
6197}
6198
6199static std::optional<const SCEV *>
6201 const SCEV *TrueExpr, const SCEV *FalseExpr) {
6202 assert(CondExpr->getType()->isIntegerTy(1) &&
6203 TrueExpr->getType() == FalseExpr->getType() &&
6204 TrueExpr->getType()->isIntegerTy(1) &&
6205 "Unexpected operands of a select.");
6206
6207 // i1 cond ? i1 x : i1 C --> C + (i1 cond ? (i1 x - i1 C) : i1 0)
6208 // --> C + (umin_seq cond, x - C)
6209 //
6210 // i1 cond ? i1 C : i1 x --> C + (i1 cond ? i1 0 : (i1 x - i1 C))
6211 // --> C + (i1 ~cond ? (i1 x - i1 C) : i1 0)
6212 // --> C + (umin_seq ~cond, x - C)
6213
6214 // FIXME: while we can't legally model the case where both of the hands
6215 // are fully variable, we only require that the *difference* is constant.
6216 if (!isa<SCEVConstant>(TrueExpr) && !isa<SCEVConstant>(FalseExpr))
6217 return std::nullopt;
6218
6219 const SCEV *X, *C;
6220 if (isa<SCEVConstant>(TrueExpr)) {
6221 CondExpr = SE->getNotSCEV(CondExpr);
6222 X = FalseExpr;
6223 C = TrueExpr;
6224 } else {
6225 X = TrueExpr;
6226 C = FalseExpr;
6227 }
6228 return SE->getAddExpr(C, SE->getUMinExpr(CondExpr, SE->getMinusSCEV(X, C),
6229 /*Sequential=*/true));
6230}
6231
6232static std::optional<const SCEV *>
6234 Value *FalseVal) {
6235 if (!isa<ConstantInt>(TrueVal) && !isa<ConstantInt>(FalseVal))
6236 return std::nullopt;
6237
6238 const auto *SECond = SE->getSCEV(Cond);
6239 const auto *SETrue = SE->getSCEV(TrueVal);
6240 const auto *SEFalse = SE->getSCEV(FalseVal);
6241 return createNodeForSelectViaUMinSeq(SE, SECond, SETrue, SEFalse);
6242}
6243
6244const SCEV *ScalarEvolution::createNodeForSelectOrPHIViaUMinSeq(
6245 Value *V, Value *Cond, Value *TrueVal, Value *FalseVal) {
6246 assert(Cond->getType()->isIntegerTy(1) && "Select condition is not an i1?");
6247 assert(TrueVal->getType() == FalseVal->getType() &&
6248 V->getType() == TrueVal->getType() &&
6249 "Types of select hands and of the result must match.");
6250
6251 // For now, only deal with i1-typed `select`s.
6252 if (!V->getType()->isIntegerTy(1))
6253 return getUnknown(V);
6254
6255 if (std::optional<const SCEV *> S =
6256 createNodeForSelectViaUMinSeq(this, Cond, TrueVal, FalseVal))
6257 return *S;
6258
6259 return getUnknown(V);
6260}
6261
6262const SCEV *ScalarEvolution::createNodeForSelectOrPHI(Value *V, Value *Cond,
6263 Value *TrueVal,
6264 Value *FalseVal) {
6265 // Handle "constant" branch or select. This can occur for instance when a
6266 // loop pass transforms an inner loop and moves on to process the outer loop.
6267 if (auto *CI = dyn_cast<ConstantInt>(Cond))
6268 return getSCEV(CI->isOne() ? TrueVal : FalseVal);
6269
6270 if (auto *I = dyn_cast<Instruction>(V)) {
6271 if (auto *ICI = dyn_cast<ICmpInst>(Cond)) {
6272 if (std::optional<const SCEV *> S =
6273 createNodeForSelectOrPHIInstWithICmpInstCond(I->getType(), ICI,
6274 TrueVal, FalseVal))
6275 return *S;
6276 }
6277 }
6278
6279 return createNodeForSelectOrPHIViaUMinSeq(V, Cond, TrueVal, FalseVal);
6280}
6281
6282/// Expand GEP instructions into add and multiply operations. This allows them
6283/// to be analyzed by regular SCEV code.
6284const SCEV *ScalarEvolution::createNodeForGEP(GEPOperator *GEP) {
6285 assert(GEP->getSourceElementType()->isSized() &&
6286 "GEP source element type must be sized");
6287
6288 SmallVector<SCEVUse, 4> IndexExprs;
6289 for (Value *Index : GEP->indices())
6290 IndexExprs.push_back(getSCEV(Index));
6291 return getGEPExpr(GEP, IndexExprs);
6292}
6293
6294APInt ScalarEvolution::getConstantMultipleImpl(const SCEV *S,
6295 const Instruction *CtxI) {
6297 auto GetShiftedByZeros = [BitWidth](uint32_t TrailingZeros) {
6298 return TrailingZeros >= BitWidth
6300 : APInt::getOneBitSet(BitWidth, TrailingZeros);
6301 };
6302 auto GetGCDMultiple = [this, CtxI](const SCEVNAryExpr *N) {
6303 // The result is GCD of all operands results.
6304 APInt Res = getConstantMultiple(N->getOperand(0), CtxI);
6305 for (unsigned I = 1, E = N->getNumOperands(); I < E && Res != 1; ++I)
6307 Res, getConstantMultiple(N->getOperand(I), CtxI));
6308 return Res;
6309 };
6310
6311 switch (S->getSCEVType()) {
6312 case scConstant:
6313 return cast<SCEVConstant>(S)->getAPInt();
6314 case scPtrToAddr:
6315 return getConstantMultiple(cast<SCEVCastExpr>(S)->getOperand());
6316 case scUDivExpr:
6317 case scVScale:
6318 return APInt(BitWidth, 1);
6319 case scTruncate: {
6320 // Only multiples that are a power of 2 will hold after truncation.
6321 const SCEVTruncateExpr *T = cast<SCEVTruncateExpr>(S);
6322 uint32_t TZ = getMinTrailingZeros(T->getOperand(), CtxI);
6323 return GetShiftedByZeros(TZ);
6324 }
6325 case scZeroExtend: {
6326 const SCEVZeroExtendExpr *Z = cast<SCEVZeroExtendExpr>(S);
6327 return getConstantMultiple(Z->getOperand(), CtxI).zext(BitWidth);
6328 }
6329 case scSignExtend: {
6330 // Only multiples that are a power of 2 will hold after sext.
6331 const SCEVSignExtendExpr *E = cast<SCEVSignExtendExpr>(S);
6332 uint32_t TZ = getMinTrailingZeros(E->getOperand(), CtxI);
6333 return GetShiftedByZeros(TZ);
6334 }
6335 case scMulExpr: {
6336 const SCEVMulExpr *M = cast<SCEVMulExpr>(S);
6337 if (M->hasNoUnsignedWrap()) {
6338 // The result is the product of all operand results.
6339 APInt Res = getConstantMultiple(M->getOperand(0), CtxI);
6340 for (const SCEV *Operand : M->operands().drop_front())
6341 Res = Res * getConstantMultiple(Operand, CtxI);
6342 return Res;
6343 }
6344
6345 // If there are no wrap guarentees, find the trailing zeros, which is the
6346 // sum of trailing zeros for all its operands.
6347 uint32_t TZ = 0;
6348 for (const SCEV *Operand : M->operands())
6349 TZ += getMinTrailingZeros(Operand, CtxI);
6350 return GetShiftedByZeros(TZ);
6351 }
6352 case scAddExpr:
6353 case scAddRecExpr: {
6354 const SCEVNAryExpr *N = cast<SCEVNAryExpr>(S);
6355 if (N->hasNoUnsignedWrap())
6356 return GetGCDMultiple(N);
6357 // Find the trailing bits, which is the minimum of its operands.
6358 uint32_t TZ = getMinTrailingZeros(N->getOperand(0), CtxI);
6359 for (const SCEV *Operand : N->operands().drop_front())
6360 TZ = std::min(TZ, getMinTrailingZeros(Operand, CtxI));
6361 return GetShiftedByZeros(TZ);
6362 }
6363 case scUMaxExpr:
6364 case scSMaxExpr:
6365 case scUMinExpr:
6366 case scSMinExpr:
6368 return GetGCDMultiple(cast<SCEVNAryExpr>(S));
6369 case scUnknown: {
6370 // Ask ValueTracking for known bits. SCEVUnknown only become available at
6371 // the point their underlying IR instruction has been defined. If CtxI was
6372 // not provided, use:
6373 // * the first instruction in the entry block if it is an argument
6374 // * the instruction itself otherwise.
6375 const SCEVUnknown *U = cast<SCEVUnknown>(S);
6376 if (!CtxI) {
6377 if (isa<Argument>(U->getValue()))
6378 CtxI = &*F.getEntryBlock().begin();
6379 else if (auto *I = dyn_cast<Instruction>(U->getValue()))
6380 CtxI = I;
6381 }
6382 unsigned Known =
6383 computeKnownBits(U->getValue(),
6384 SimplifyQuery(getDataLayout(), &DT, &AC, CtxI)
6385 .allowEphemerals(true))
6386 .countMinTrailingZeros();
6387 return GetShiftedByZeros(Known);
6388 }
6389 case scCouldNotCompute:
6390 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
6391 }
6392 llvm_unreachable("Unknown SCEV kind!");
6393}
6394
6396 const Instruction *CtxI) {
6397 // Skip looking up and updating the cache if there is a context instruction,
6398 // as the result will only be valid in the specified context.
6399 if (CtxI)
6400 return getConstantMultipleImpl(S, CtxI);
6401
6402 auto I = ConstantMultipleCache.find(S);
6403 if (I != ConstantMultipleCache.end())
6404 return I->second;
6405
6406 APInt Result = getConstantMultipleImpl(S, CtxI);
6407 auto InsertPair = ConstantMultipleCache.insert({S, Result});
6408 assert(InsertPair.second && "Should insert a new key");
6409 return InsertPair.first->second;
6410}
6411
6413 APInt Multiple = getConstantMultiple(S);
6414 return Multiple == 0 ? APInt(Multiple.getBitWidth(), 1) : Multiple;
6415}
6416
6418 const Instruction *CtxI) {
6419 return std::min(getConstantMultiple(S, CtxI).countTrailingZeros(),
6420 (unsigned)getTypeSizeInBits(S->getType()));
6421}
6422
6423/// Helper method to assign a range to V from metadata present in the IR.
6424static std::optional<ConstantRange> GetRangeFromMetadata(Value *V) {
6426 if (MDNode *MD = I->getMetadata(LLVMContext::MD_range))
6427 return getConstantRangeFromMetadata(*MD);
6428 if (const auto *CB = dyn_cast<CallBase>(V))
6429 if (std::optional<ConstantRange> Range = CB->getRange())
6430 return Range;
6431 }
6432 if (auto *A = dyn_cast<Argument>(V))
6433 if (std::optional<ConstantRange> Range = A->getRange())
6434 return Range;
6435
6436 return std::nullopt;
6437}
6438
6440 SCEVFlags NWFlags = Flags & SCEV::FlagsNoWrapMask;
6441 if (AddRec->getNoWrapFlags(NWFlags) != NWFlags) {
6442 AddRec->setNoWrapFlags(NWFlags);
6443 UnsignedRanges.erase(AddRec);
6444 SignedRanges.erase(AddRec);
6445 ConstantMultipleCache.erase(AddRec);
6446 }
6447}
6448
6449ConstantRange ScalarEvolution::
6450getRangeForUnknownRecurrence(const SCEVUnknown *U) {
6451 const DataLayout &DL = getDataLayout();
6452
6453 unsigned BitWidth = getTypeSizeInBits(U->getType());
6454 const ConstantRange FullSet(BitWidth, /*isFullSet=*/true);
6455
6456 // Match a simple recurrence of the form: <start, ShiftOp, Step>, and then
6457 // use information about the trip count to improve our available range. Note
6458 // that the trip count independent cases are already handled by known bits.
6459 // WARNING: The definition of recurrence used here is subtly different than
6460 // the one used by AddRec (and thus most of this file). Step is allowed to
6461 // be arbitrarily loop varying here, where AddRec allows only loop invariant
6462 // and other addrecs in the same loop (for non-affine addrecs). The code
6463 // below intentionally handles the case where step is not loop invariant.
6464 auto *P = dyn_cast<PHINode>(U->getValue());
6465 if (!P)
6466 return FullSet;
6467
6468 // Make sure that no Phi input comes from an unreachable block. Otherwise,
6469 // even the values that are not available in these blocks may come from them,
6470 // and this leads to false-positive recurrence test.
6471 for (auto *Pred : predecessors(P->getParent()))
6472 if (!DT.isReachableFromEntry(Pred))
6473 return FullSet;
6474
6475 BinaryOperator *BO;
6476 Value *Start, *Step;
6477 if (!matchSimpleRecurrence(P, BO, Start, Step))
6478 return FullSet;
6479
6480 // If we found a recurrence in reachable code, we must be in a loop. Note
6481 // that BO might be in some subloop of L, and that's completely okay.
6482 auto *L = LI.getLoopFor(P->getParent());
6483 assert(L && L->getHeader() == P->getParent());
6484 if (!L->contains(BO->getParent()))
6485 // NOTE: This bailout should be an assert instead. However, asserting
6486 // the condition here exposes a case where LoopFusion is querying SCEV
6487 // with malformed loop information during the midst of the transform.
6488 // There doesn't appear to be an obvious fix, so for the moment bailout
6489 // until the caller issue can be fixed. PR49566 tracks the bug.
6490 return FullSet;
6491
6492 // TODO: Extend to other opcodes such as mul, and div
6493 switch (BO->getOpcode()) {
6494 default:
6495 return FullSet;
6496 case Instruction::AShr:
6497 case Instruction::LShr:
6498 case Instruction::Shl:
6499 break;
6500 };
6501
6502 if (BO->getOperand(0) != P)
6503 // TODO: Handle the power function forms some day.
6504 return FullSet;
6505
6506 unsigned TC = getSmallConstantMaxTripCount(L);
6507 if (!TC || TC >= BitWidth)
6508 return FullSet;
6509
6510 auto KnownStart = computeKnownBits(Start, DL, &AC, nullptr, &DT);
6511 auto KnownStep = computeKnownBits(Step, DL, &AC, nullptr, &DT);
6512 assert(KnownStart.getBitWidth() == BitWidth &&
6513 KnownStep.getBitWidth() == BitWidth);
6514
6515 // Compute total shift amount, being careful of overflow and bitwidths.
6516 auto MaxShiftAmt = KnownStep.getMaxValue();
6517 APInt TCAP(BitWidth, TC-1);
6518 bool Overflow = false;
6519 auto TotalShift = MaxShiftAmt.umul_ov(TCAP, Overflow);
6520 if (Overflow)
6521 return FullSet;
6522
6523 switch (BO->getOpcode()) {
6524 default:
6525 llvm_unreachable("filtered out above");
6526 case Instruction::AShr: {
6527 // For each ashr, three cases:
6528 // shift = 0 => unchanged value
6529 // saturation => 0 or -1
6530 // other => a value closer to zero (of the same sign)
6531 // Thus, the end value is closer to zero than the start.
6532 auto KnownEnd = KnownBits::ashr(KnownStart,
6533 KnownBits::makeConstant(TotalShift));
6534 if (KnownStart.isNonNegative())
6535 // Analogous to lshr (simply not yet canonicalized)
6536 return ConstantRange::getNonEmpty(KnownEnd.getMinValue(),
6537 KnownStart.getMaxValue() + 1);
6538 if (KnownStart.isNegative())
6539 // End >=u Start && End <=s Start
6540 return ConstantRange::getNonEmpty(KnownStart.getMinValue(),
6541 KnownEnd.getMaxValue() + 1);
6542 break;
6543 }
6544 case Instruction::LShr: {
6545 // For each lshr, three cases:
6546 // shift = 0 => unchanged value
6547 // saturation => 0
6548 // other => a smaller positive number
6549 // Thus, the low end of the unsigned range is the last value produced.
6550 auto KnownEnd = KnownBits::lshr(KnownStart,
6551 KnownBits::makeConstant(TotalShift));
6552 return ConstantRange::getNonEmpty(KnownEnd.getMinValue(),
6553 KnownStart.getMaxValue() + 1);
6554 }
6555 case Instruction::Shl: {
6556 // Iff no bits are shifted out, value increases on every shift.
6557 auto KnownEnd = KnownBits::shl(KnownStart,
6558 KnownBits::makeConstant(TotalShift));
6559 if (TotalShift.ult(KnownStart.countMinLeadingZeros()))
6560 return ConstantRange(KnownStart.getMinValue(),
6561 KnownEnd.getMaxValue() + 1);
6562 break;
6563 }
6564 };
6565 return FullSet;
6566}
6567
6568// The goal of this function is to check if recursively visiting the operands
6569// of this PHI might lead to an infinite loop. If we do see such a loop,
6570// there's no good way to break it, so we avoid analyzing such cases.
6571//
6572// getRangeRef previously used a visited set to avoid infinite loops, but this
6573// caused other issues: the result was dependent on the order of getRangeRef
6574// calls, and the interaction with createSCEVIter could cause a stack overflow
6575// in some cases (see issue #148253).
6576//
6577// FIXME: The way this is implemented is overly conservative; this checks
6578// for a few obviously safe patterns, but anything that doesn't lead to
6579// recursion is fine.
6581 Value *Cond = nullptr, *LHS = nullptr, *RHS = nullptr;
6583 return true;
6584
6585 if (all_of(PHI->operands(),
6586 [&](Value *Operand) { return DT.dominates(Operand, PHI); }))
6587 return true;
6588
6589 return false;
6590}
6591
6592const ConstantRange &
6593ScalarEvolution::getRangeRefIter(const SCEV *S,
6594 ScalarEvolution::RangeSignHint SignHint) {
6595 DenseMap<const SCEV *, ConstantRange> &Cache =
6596 SignHint == ScalarEvolution::HINT_RANGE_UNSIGNED ? UnsignedRanges
6597 : SignedRanges;
6598 SmallVector<SCEVUse> WorkList;
6599 SmallPtrSet<const SCEV *, 8> Seen;
6600
6601 // Add Expr to the worklist, if Expr is either an N-ary expression or a
6602 // SCEVUnknown PHI node.
6603 auto AddToWorklist = [&WorkList, &Seen, &Cache](const SCEV *Expr) {
6604 if (!Seen.insert(Expr).second)
6605 return;
6606 if (Cache.contains(Expr))
6607 return;
6608 switch (Expr->getSCEVType()) {
6609 case scUnknown:
6611 break;
6612 [[fallthrough]];
6613 case scConstant:
6614 case scVScale:
6615 case scTruncate:
6616 case scZeroExtend:
6617 case scSignExtend:
6618 case scPtrToAddr:
6619 case scAddExpr:
6620 case scMulExpr:
6621 case scUDivExpr:
6622 case scAddRecExpr:
6623 case scUMaxExpr:
6624 case scSMaxExpr:
6625 case scUMinExpr:
6626 case scSMinExpr:
6628 WorkList.push_back(Expr);
6629 break;
6630 case scCouldNotCompute:
6631 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
6632 }
6633 };
6634 AddToWorklist(S);
6635
6636 // Build worklist by queuing operands of N-ary expressions and phi nodes.
6637 for (unsigned I = 0; I != WorkList.size(); ++I) {
6638 const SCEV *P = WorkList[I];
6639 auto *UnknownS = dyn_cast<SCEVUnknown>(P);
6640 // If it is not a `SCEVUnknown`, just recurse into operands.
6641 if (!UnknownS) {
6642 for (const SCEV *Op : P->operands())
6643 AddToWorklist(Op);
6644 continue;
6645 }
6646 // `SCEVUnknown`'s require special treatment.
6647 if (PHINode *P = dyn_cast<PHINode>(UnknownS->getValue())) {
6648 if (!RangeRefPHIAllowedOperands(DT, P))
6649 continue;
6650 for (auto &Op : reverse(P->operands()))
6651 AddToWorklist(getSCEV(Op));
6652 }
6653 }
6654
6655 if (!WorkList.empty()) {
6656 // Use getRangeRef to compute ranges for items in the worklist in reverse
6657 // order. This will force ranges for earlier operands to be computed before
6658 // their users in most cases.
6659 for (const SCEV *P : reverse(drop_begin(WorkList))) {
6660 getRangeRef(P, SignHint);
6661 }
6662 }
6663
6664 return getRangeRef(S, SignHint, 0);
6665}
6666
6667const APInt *ScalarEvolution::getConstantAPIntOrNull(const SCEV *S) {
6668 if (const auto *C = dyn_cast<SCEVConstant>(S))
6669 return &C->getAPInt();
6670 return nullptr;
6671}
6672
6673/// Determine the range for a particular SCEV. If SignHint is
6674/// HINT_RANGE_UNSIGNED (resp. HINT_RANGE_SIGNED) then getRange prefers ranges
6675/// with a "cleaner" unsigned (resp. signed) representation.
6676const ConstantRange &ScalarEvolution::getRangeRef(
6677 const SCEV *S, ScalarEvolution::RangeSignHint SignHint, unsigned Depth) {
6678 DenseMap<const SCEV *, ConstantRange> &Cache =
6679 SignHint == ScalarEvolution::HINT_RANGE_UNSIGNED ? UnsignedRanges
6680 : SignedRanges;
6682 SignHint == ScalarEvolution::HINT_RANGE_UNSIGNED ? ConstantRange::Unsigned
6684
6685 // See if we've computed this range already.
6686 auto I = Cache.find(S);
6687 if (I != Cache.end())
6688 return I->second;
6689
6690 if (const SCEVConstant *C = dyn_cast<SCEVConstant>(S))
6691 return setRange(C, SignHint, ConstantRange(C->getAPInt()));
6692
6693 // Switch to iteratively computing the range for S, if it is part of a deeply
6694 // nested expression.
6696 return getRangeRefIter(S, SignHint);
6697
6698 unsigned BitWidth = getTypeSizeInBits(S->getType());
6699 ConstantRange ConservativeResult(BitWidth, /*isFullSet=*/true);
6700 using OBO = OverflowingBinaryOperator;
6701
6702 // If the value has known zeros, the maximum value will have those known zeros
6703 // as well.
6704 if (SignHint == ScalarEvolution::HINT_RANGE_UNSIGNED) {
6705 APInt Multiple = getNonZeroConstantMultiple(S);
6706 APInt Remainder = APInt::getMaxValue(BitWidth).urem(Multiple);
6707 if (!Remainder.isZero())
6708 ConservativeResult =
6709 ConstantRange(APInt::getMinValue(BitWidth),
6710 APInt::getMaxValue(BitWidth) - Remainder + 1);
6711 }
6712 else {
6713 uint32_t TZ = getMinTrailingZeros(S);
6714 if (TZ != 0) {
6715 ConservativeResult = ConstantRange(
6717 APInt::getSignedMaxValue(BitWidth).ashr(TZ).shl(TZ) + 1);
6718 }
6719 }
6720
6721 switch (S->getSCEVType()) {
6722 case scConstant:
6723 llvm_unreachable("Already handled above.");
6724 case scVScale:
6725 return setRange(S, SignHint, getVScaleRange(&F, BitWidth));
6726 case scTruncate: {
6727 const SCEVTruncateExpr *Trunc = cast<SCEVTruncateExpr>(S);
6728 ConstantRange X = getRangeRef(Trunc->getOperand(), SignHint, Depth + 1);
6729 return setRange(
6730 Trunc, SignHint,
6731 ConservativeResult.intersectWith(X.truncate(BitWidth), RangeType));
6732 }
6733 case scZeroExtend: {
6734 const SCEVZeroExtendExpr *ZExt = cast<SCEVZeroExtendExpr>(S);
6735 ConstantRange X = getRangeRef(ZExt->getOperand(), SignHint, Depth + 1);
6736 return setRange(
6737 ZExt, SignHint,
6738 ConservativeResult.intersectWith(X.zeroExtend(BitWidth), RangeType));
6739 }
6740 case scSignExtend: {
6741 const SCEVSignExtendExpr *SExt = cast<SCEVSignExtendExpr>(S);
6742 ConstantRange X = getRangeRef(SExt->getOperand(), SignHint, Depth + 1);
6743 return setRange(
6744 SExt, SignHint,
6745 ConservativeResult.intersectWith(X.signExtend(BitWidth), RangeType));
6746 }
6747 case scPtrToAddr: {
6748 const SCEVCastExpr *Cast = cast<SCEVCastExpr>(S);
6749 ConstantRange X = getRangeRef(Cast->getOperand(), SignHint, Depth + 1);
6750 return setRange(Cast, SignHint, X);
6751 }
6752 case scAddExpr: {
6753 const SCEVAddExpr *Add = cast<SCEVAddExpr>(S);
6754 // Check if this is a URem pattern: A - (A / B) * B, which is always < B.
6755 const SCEV *URemLHS = nullptr, *URemRHS = nullptr;
6756 if (SignHint == ScalarEvolution::HINT_RANGE_UNSIGNED &&
6757 match(S, m_scev_URem(m_SCEV(URemLHS), m_SCEV(URemRHS), *this))) {
6758 ConstantRange LHSRange = getRangeRef(URemLHS, SignHint, Depth + 1);
6759 ConstantRange RHSRange = getRangeRef(URemRHS, SignHint, Depth + 1);
6760 ConservativeResult =
6761 ConservativeResult.intersectWith(LHSRange.urem(RHSRange), RangeType);
6762 }
6763 ConstantRange X = getRangeRef(Add->getOperand(0), SignHint, Depth + 1);
6764 unsigned WrapType = OBO::AnyWrap;
6765 if (Add->hasNoSignedWrap())
6766 WrapType |= OBO::NoSignedWrap;
6767 if (Add->hasNoUnsignedWrap())
6768 WrapType |= OBO::NoUnsignedWrap;
6769 for (const SCEV *Op : drop_begin(Add->operands()))
6770 X = X.addWithNoWrap(getRangeRef(Op, SignHint, Depth + 1), WrapType,
6771 RangeType);
6772 return setRange(Add, SignHint,
6773 ConservativeResult.intersectWith(X, RangeType));
6774 }
6775 case scMulExpr: {
6776 const SCEVMulExpr *Mul = cast<SCEVMulExpr>(S);
6777 ConstantRange X = getRangeRef(Mul->getOperand(0), SignHint, Depth + 1);
6778 for (const SCEV *Op : drop_begin(Mul->operands()))
6779 X = X.multiply(getRangeRef(Op, SignHint, Depth + 1));
6780 return setRange(Mul, SignHint,
6781 ConservativeResult.intersectWith(X, RangeType));
6782 }
6783 case scUDivExpr: {
6784 const SCEVUDivExpr *UDiv = cast<SCEVUDivExpr>(S);
6785 ConstantRange X = getRangeRef(UDiv->getLHS(), SignHint, Depth + 1);
6786 ConstantRange Y = getRangeRef(UDiv->getRHS(), SignHint, Depth + 1);
6787 return setRange(UDiv, SignHint,
6788 ConservativeResult.intersectWith(X.udiv(Y), RangeType));
6789 }
6790 case scAddRecExpr: {
6791 const SCEVAddRecExpr *AddRec = cast<SCEVAddRecExpr>(S);
6792 // If there's no unsigned wrap, the value will never be less than its
6793 // initial value.
6794 if (AddRec->hasNoUnsignedWrap()) {
6795 APInt UnsignedMinValue = getUnsignedRangeMin(AddRec->getStart());
6796 if (!UnsignedMinValue.isZero())
6797 ConservativeResult = ConservativeResult.intersectWith(
6798 ConstantRange(UnsignedMinValue, APInt(BitWidth, 0)), RangeType);
6799 }
6800
6801 // If there's no signed wrap, and all the operands except initial value have
6802 // the same sign or zero, the value won't ever be:
6803 // 1: smaller than initial value if operands are non negative,
6804 // 2: bigger than initial value if operands are non positive.
6805 // For both cases, value can not cross signed min/max boundary.
6806 if (AddRec->hasNoSignedWrap()) {
6807 bool AllNonNeg = true;
6808 bool AllNonPos = true;
6809 for (unsigned i = 1, e = AddRec->getNumOperands(); i != e; ++i) {
6810 if (!isKnownNonNegative(AddRec->getOperand(i)))
6811 AllNonNeg = false;
6812 if (!isKnownNonPositive(AddRec->getOperand(i)))
6813 AllNonPos = false;
6814 }
6815 if (AllNonNeg)
6816 ConservativeResult = ConservativeResult.intersectWith(
6819 RangeType);
6820 else if (AllNonPos)
6821 ConservativeResult = ConservativeResult.intersectWith(
6823 getSignedRangeMax(AddRec->getStart()) +
6824 1),
6825 RangeType);
6826 }
6827
6828 // TODO: non-affine addrec
6829 if (AddRec->isAffine()) {
6830 const SCEV *MaxBEScev =
6832 if (!isa<SCEVCouldNotCompute>(MaxBEScev)) {
6833 APInt MaxBECount = cast<SCEVConstant>(MaxBEScev)->getAPInt();
6834
6835 // Adjust MaxBECount to the same bitwidth as AddRec. We can truncate if
6836 // MaxBECount's active bits are all <= AddRec's bit width.
6837 if (MaxBECount.getBitWidth() > BitWidth &&
6838 MaxBECount.getActiveBits() <= BitWidth)
6839 MaxBECount = MaxBECount.trunc(BitWidth);
6840 else if (MaxBECount.getBitWidth() < BitWidth)
6841 MaxBECount = MaxBECount.zext(BitWidth);
6842
6843 if (MaxBECount.getBitWidth() == BitWidth) {
6844 auto [RangeFromAffine, Flags] = getRangeForAffineAR(
6845 AddRec->getStart(), AddRec->getStepRecurrence(*this), MaxBECount);
6846 ConservativeResult =
6847 ConservativeResult.intersectWith(RangeFromAffine, RangeType);
6848 const_cast<SCEVAddRecExpr *>(AddRec)->setNoWrapFlags(Flags);
6849
6850 auto RangeFromFactoring = getRangeViaFactoring(
6851 AddRec->getStart(), AddRec->getStepRecurrence(*this), MaxBECount);
6852 ConservativeResult =
6853 ConservativeResult.intersectWith(RangeFromFactoring, RangeType);
6854 }
6855 }
6856
6857 // Now try symbolic BE count and more powerful methods.
6859 const SCEV *SymbolicMaxBECount =
6861 if (!isa<SCEVCouldNotCompute>(SymbolicMaxBECount) &&
6862 getTypeSizeInBits(MaxBEScev->getType()) <= BitWidth &&
6863 AddRec->hasNoSelfWrap()) {
6864 auto RangeFromAffineNew = getRangeForAffineNoSelfWrappingAR(
6865 AddRec, SymbolicMaxBECount, BitWidth, SignHint);
6866 ConservativeResult =
6867 ConservativeResult.intersectWith(RangeFromAffineNew, RangeType);
6868 }
6869 }
6870 }
6871
6872 return setRange(AddRec, SignHint, std::move(ConservativeResult));
6873 }
6874 case scUMaxExpr:
6875 case scSMaxExpr:
6876 case scUMinExpr:
6877 case scSMinExpr:
6878 case scSequentialUMinExpr: {
6880 switch (S->getSCEVType()) {
6881 case scUMaxExpr:
6882 ID = Intrinsic::umax;
6883 break;
6884 case scSMaxExpr:
6885 ID = Intrinsic::smax;
6886 break;
6887 case scUMinExpr:
6889 ID = Intrinsic::umin;
6890 break;
6891 case scSMinExpr:
6892 ID = Intrinsic::smin;
6893 break;
6894 default:
6895 llvm_unreachable("Unknown SCEVMinMaxExpr/SCEVSequentialMinMaxExpr.");
6896 }
6897
6898 const auto *NAry = cast<SCEVNAryExpr>(S);
6899 ConstantRange X = getRangeRef(NAry->getOperand(0), SignHint, Depth + 1);
6900 for (unsigned i = 1, e = NAry->getNumOperands(); i != e; ++i)
6901 X = X.intrinsic(
6902 ID, {X, getRangeRef(NAry->getOperand(i), SignHint, Depth + 1)});
6903 return setRange(S, SignHint,
6904 ConservativeResult.intersectWith(X, RangeType));
6905 }
6906 case scUnknown: {
6907 const SCEVUnknown *U = cast<SCEVUnknown>(S);
6908 Value *V = U->getValue();
6909
6910 // Check if the IR explicitly contains !range metadata.
6911 std::optional<ConstantRange> MDRange = GetRangeFromMetadata(V);
6912 if (MDRange)
6913 ConservativeResult =
6914 ConservativeResult.intersectWith(*MDRange, RangeType);
6915
6916 // Use facts about recurrences in the underlying IR. Note that add
6917 // recurrences are AddRecExprs and thus don't hit this path. This
6918 // primarily handles shift recurrences.
6919 auto CR = getRangeForUnknownRecurrence(U);
6920 ConservativeResult = ConservativeResult.intersectWith(CR);
6921
6922 // See if ValueTracking can give us a useful range.
6923 const DataLayout &DL = getDataLayout();
6924 KnownBits Known = computeKnownBits(V, DL, &AC, nullptr, &DT);
6925 if (Known.getBitWidth() != BitWidth)
6926 Known = Known.zextOrTrunc(BitWidth);
6927
6928 // ValueTracking may be able to compute a tighter result for the number of
6929 // sign bits than for the value of those sign bits.
6930 unsigned NS = ComputeNumSignBits(V, DL, &AC, nullptr, &DT);
6931 if (U->getType()->isPointerTy()) {
6932 // NS counts the sign bits of the whole pointer; drop those above the
6933 // index bits.
6934 unsigned PtrIdxDiff =
6935 DL.getPointerTypeSizeInBits(U->getType()) - BitWidth;
6936 NS = NS > PtrIdxDiff ? NS - PtrIdxDiff : 1;
6937 }
6938
6939 if (NS > 1) {
6940 // If we know any of the sign bits, we know all of the sign bits.
6941 if (!Known.Zero.getHiBits(NS).isZero())
6942 Known.Zero.setHighBits(NS);
6943 if (!Known.One.getHiBits(NS).isZero())
6944 Known.One.setHighBits(NS);
6945 }
6946
6947 if (Known.getMinValue() != Known.getMaxValue() + 1)
6948 ConservativeResult = ConservativeResult.intersectWith(
6949 ConstantRange(Known.getMinValue(), Known.getMaxValue() + 1),
6950 RangeType);
6951 if (NS > 1)
6952 ConservativeResult = ConservativeResult.intersectWith(
6953 ConstantRange(APInt::getSignedMinValue(BitWidth).ashr(NS - 1),
6954 APInt::getSignedMaxValue(BitWidth).ashr(NS - 1) + 1),
6955 RangeType);
6956
6957 if (U->getType()->isPointerTy() && SignHint == HINT_RANGE_UNSIGNED) {
6958 // Strengthen the range if the underlying IR value is a
6959 // global/alloca/heap allocation using the size of the object.
6960 bool CanBeNull;
6961 uint64_t DerefBytes = V->getPointerDereferenceableBytes(
6962 DL, CanBeNull, /*CanBeFreed=*/nullptr);
6963 if (DerefBytes > 1 && isUIntN(BitWidth, DerefBytes)) {
6964 // The highest address the object can start is DerefBytes bytes before
6965 // the end (unsigned max value). If this value is not a multiple of the
6966 // alignment, the last possible start value is the next lowest multiple
6967 // of the alignment. Note: The computations below cannot overflow,
6968 // because if they would there's no possible start address for the
6969 // object.
6970 APInt MaxVal =
6971 APInt::getMaxValue(BitWidth) - APInt(BitWidth, DerefBytes);
6972 uint64_t Align = U->getValue()->getPointerAlignment(DL).value();
6973 uint64_t Rem = MaxVal.urem(Align);
6974 MaxVal -= APInt(BitWidth, Rem);
6975 APInt MinVal = APInt::getZero(BitWidth);
6976 if (llvm::isKnownNonZero(V, DL))
6977 MinVal = Align;
6978 ConservativeResult = ConservativeResult.intersectWith(
6979 ConstantRange::getNonEmpty(MinVal, MaxVal + 1), RangeType);
6980 }
6981 }
6982
6983 // A range of Phi is a subset of union of all ranges of its input.
6984 if (PHINode *Phi = dyn_cast<PHINode>(V)) {
6985 // SCEVExpander sometimes creates SCEVUnknowns that are secretly
6986 // AddRecs; return the range for the corresponding AddRec.
6987 if (auto *AR = dyn_cast<SCEVAddRecExpr>(getSCEV(V)))
6988 return getRangeRef(AR, SignHint, Depth + 1);
6989
6990 // Make sure that we do not run over cycled Phis.
6991 if (RangeRefPHIAllowedOperands(DT, Phi)) {
6992 ConstantRange RangeFromOps(BitWidth, /*isFullSet=*/false);
6993
6994 for (const auto &Op : Phi->operands()) {
6995 auto OpRange = getRangeRef(getSCEV(Op), SignHint, Depth + 1);
6996 RangeFromOps = RangeFromOps.unionWith(OpRange);
6997 // No point to continue if we already have a full set.
6998 if (RangeFromOps.isFullSet())
6999 break;
7000 }
7001 ConservativeResult =
7002 ConservativeResult.intersectWith(RangeFromOps, RangeType);
7003 }
7004 }
7005
7006 // vscale can't be equal to zero
7007 if (const auto *II = dyn_cast<IntrinsicInst>(V))
7008 if (II->getIntrinsicID() == Intrinsic::vscale) {
7009 ConstantRange Disallowed = APInt::getZero(BitWidth);
7010 ConservativeResult = ConservativeResult.difference(Disallowed);
7011 }
7012
7013 return setRange(U, SignHint, std::move(ConservativeResult));
7014 }
7015 case scCouldNotCompute:
7016 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
7017 }
7018
7019 return setRange(S, SignHint, std::move(ConservativeResult));
7020}
7021
7022// Given a StartRange, Step and MaxBECount for an expression compute a range of
7023// values that the expression can take. Initially, the expression has a value
7024// from StartRange and then is changed by Step up to MaxBECount times. Signed
7025// argument defines if we treat Step as signed or unsigned. The second return
7026// value indicates that no wrapping occurred.
7027static std::pair<ConstantRange, bool>
7029 const APInt &MaxBECount, bool Signed) {
7030 unsigned BitWidth = Step.getBitWidth();
7031 assert(BitWidth == StartRange.getBitWidth() &&
7032 BitWidth == MaxBECount.getBitWidth() && "mismatched bit widths");
7033 // If either Step or MaxBECount is 0, then the expression won't change, and we
7034 // just need to return the initial range.
7035 if (Step == 0 || MaxBECount == 0)
7036 return {StartRange, true};
7037
7038 // If we don't know anything about the initial value (i.e. StartRange is
7039 // FullRange), then we don't know anything about the final range either.
7040 // Return FullRange.
7041 if (StartRange.isFullSet())
7042 return {ConstantRange::getFull(BitWidth), false};
7043
7044 // If Step is signed and negative, then we use its absolute value, but we also
7045 // note that we're moving in the opposite direction.
7046 bool Descending = Signed && Step.isNegative();
7047
7048 if (Signed)
7049 // This is correct even for INT_SMIN. Let's look at i8 to illustrate this:
7050 // abs(INT_SMIN) = abs(-128) = abs(0x80) = -0x80 = 0x80 = 128.
7051 // This equations hold true due to the well-defined wrap-around behavior of
7052 // APInt.
7053 Step = Step.abs();
7054
7055 // Check if Offset is more than full span of BitWidth. If it is, the
7056 // expression is guaranteed to overflow.
7057 if (APInt::getMaxValue(StartRange.getBitWidth()).udiv(Step).ult(MaxBECount))
7058 return {ConstantRange::getFull(BitWidth), false};
7059
7060 // Offset is by how much the expression can change. Checks above guarantee no
7061 // overflow here.
7062 APInt Offset = Step * MaxBECount;
7063
7064 // Minimum value of the final range will match the minimal value of StartRange
7065 // if the expression is increasing and will be decreased by Offset otherwise.
7066 // Maximum value of the final range will match the maximal value of StartRange
7067 // if the expression is decreasing and will be increased by Offset otherwise.
7068 APInt StartLower = StartRange.getLower();
7069 APInt StartUpper = StartRange.getUpper() - 1;
7070 bool Overflow;
7071 APInt MovedBoundary;
7072 if (Signed) {
7073 // This does not use sadd_ov, as we want to check overflow for a signed
7074 // start with an unsigned offset.
7075 if (Descending) {
7076 MovedBoundary = StartLower - std::move(Offset);
7077 Overflow = MovedBoundary.sgt(StartLower) || StartRange.isSignWrappedSet();
7078 } else {
7079 MovedBoundary = StartUpper + std::move(Offset);
7080 Overflow = MovedBoundary.slt(StartUpper) || StartRange.isSignWrappedSet();
7081 }
7082 } else {
7083 MovedBoundary = StartUpper.uadd_ov(std::move(Offset), Overflow);
7084 Overflow |= StartRange.isWrappedSet();
7085 }
7086
7087 // It's possible that the new minimum/maximum value will fall into the initial
7088 // range (due to wrap around). This means that the expression can take any
7089 // value in this bitwidth, and we have to return full range.
7090 if (StartRange.contains(MovedBoundary))
7091 return {ConstantRange::getFull(BitWidth), false};
7092
7093 APInt NewLower =
7094 Descending ? std::move(MovedBoundary) : std::move(StartLower);
7095 APInt NewUpper =
7096 Descending ? std::move(StartUpper) : std::move(MovedBoundary);
7097 NewUpper += 1;
7098
7099 // No overflow detected, return [StartLower, StartUpper + Offset + 1) range.
7100 return {ConstantRange::getNonEmpty(std::move(NewLower), std::move(NewUpper)),
7101 !Overflow};
7102}
7103
7104std::pair<ConstantRange, SCEVFlags>
7105ScalarEvolution::getRangeForAffineAR(const SCEV *Start, const SCEV *Step,
7106 const APInt &MaxBECount) {
7107 assert(getTypeSizeInBits(Start->getType()) ==
7108 getTypeSizeInBits(Step->getType()) &&
7109 getTypeSizeInBits(Start->getType()) == MaxBECount.getBitWidth() &&
7110 "mismatched bit widths");
7111
7112 // First, consider step signed.
7113 ConstantRange StartSRange = getSignedRange(Start);
7114 ConstantRange StepSRange = getSignedRange(Step);
7115
7116 // If Step can be both positive and negative, we need to find ranges for the
7117 // maximum absolute step values in both directions and union them.
7118 auto [SR1, NSW1] = getRangeForAffineARHelper(
7119 StepSRange.getSignedMin(), StartSRange, MaxBECount, /*Signed=*/true);
7120 auto [SR2, NSW2] = getRangeForAffineARHelper(StepSRange.getSignedMax(),
7121 StartSRange, MaxBECount,
7122 /*Signed=*/true);
7123 ConstantRange SR = SR1.unionWith(SR2);
7124
7125 // Next, consider step unsigned.
7126 auto [UR, NUW] = getRangeForAffineARHelper(
7127 getUnsignedRangeMax(Step), getUnsignedRange(Start), MaxBECount,
7128 /*Signed=*/false);
7129
7131 if (NUW)
7133 if (NSW1 && NSW2)
7135
7136 // Finally, intersect signed and unsigned ranges.
7138}
7139
7140ConstantRange ScalarEvolution::getRangeForAffineNoSelfWrappingAR(
7141 const SCEVAddRecExpr *AddRec, const SCEV *MaxBECount, unsigned BitWidth,
7142 ScalarEvolution::RangeSignHint SignHint) {
7143 assert(AddRec->isAffine() && "Non-affine AddRecs are not suppored!\n");
7144 assert(AddRec->hasNoSelfWrap() &&
7145 "This only works for non-self-wrapping AddRecs!");
7146 const bool IsSigned = SignHint == HINT_RANGE_SIGNED;
7147 const SCEV *Step = AddRec->getStepRecurrence(*this);
7148 // Only deal with constant step to save compile time.
7149 if (!isa<SCEVConstant>(Step))
7150 return ConstantRange::getFull(BitWidth);
7151 // Let's make sure that we can prove that we do not self-wrap during
7152 // MaxBECount iterations. We need this because MaxBECount is a maximum
7153 // iteration count estimate, and we might infer nw from some exit for which we
7154 // do not know max exit count (or any other side reasoning).
7155 // TODO: Turn into assert at some point.
7156 if (getTypeSizeInBits(MaxBECount->getType()) >
7157 getTypeSizeInBits(AddRec->getType()))
7158 return ConstantRange::getFull(BitWidth);
7159 MaxBECount = getNoopOrZeroExtend(MaxBECount, AddRec->getType());
7160 const SCEV *RangeWidth = getMinusOne(AddRec->getType());
7161 const SCEV *StepAbs = getUMinExpr(Step, getNegativeSCEV(Step));
7162 const SCEV *MaxItersWithoutWrap = getUDivExpr(RangeWidth, StepAbs);
7163 if (!isKnownPredicateViaConstantRanges(ICmpInst::ICMP_ULE, MaxBECount,
7164 MaxItersWithoutWrap))
7165 return ConstantRange::getFull(BitWidth);
7166
7167 ICmpInst::Predicate LEPred =
7169 ICmpInst::Predicate GEPred =
7171 const SCEV *End = AddRec->evaluateAtIteration(MaxBECount, *this);
7172
7173 // We know that there is no self-wrap. Let's take Start and End values and
7174 // look at all intermediate values V1, V2, ..., Vn that IndVar takes during
7175 // the iteration. They either lie inside the range [Min(Start, End),
7176 // Max(Start, End)] or outside it:
7177 //
7178 // Case 1: RangeMin ... Start V1 ... VN End ... RangeMax;
7179 // Case 2: RangeMin Vk ... V1 Start ... End Vn ... Vk + 1 RangeMax;
7180 //
7181 // No self wrap flag guarantees that the intermediate values cannot be BOTH
7182 // outside and inside the range [Min(Start, End), Max(Start, End)]. Using that
7183 // knowledge, let's try to prove that we are dealing with Case 1. It is so if
7184 // Start <= End and step is positive, or Start >= End and step is negative.
7185 const SCEV *Start = applyLoopGuards(AddRec->getStart(), AddRec->getLoop());
7186 ConstantRange StartRange = getRangeRef(Start, SignHint);
7187 ConstantRange EndRange = getRangeRef(End, SignHint);
7188 ConstantRange RangeBetween = StartRange.unionWith(EndRange);
7189 // If they already cover full iteration space, we will know nothing useful
7190 // even if we prove what we want to prove.
7191 if (RangeBetween.isFullSet())
7192 return RangeBetween;
7193 // Only deal with ranges that do not wrap (i.e. RangeMin < RangeMax).
7194 bool IsWrappedSet = IsSigned ? RangeBetween.isSignWrappedSet()
7195 : RangeBetween.isWrappedSet();
7196 if (IsWrappedSet)
7197 return ConstantRange::getFull(BitWidth);
7198
7199 if (isKnownPositive(Step) &&
7200 isKnownPredicateViaConstantRanges(LEPred, Start, End))
7201 return RangeBetween;
7202 if (isKnownNegative(Step) &&
7203 isKnownPredicateViaConstantRanges(GEPred, Start, End))
7204 return RangeBetween;
7205 return ConstantRange::getFull(BitWidth);
7206}
7207
7208ConstantRange ScalarEvolution::getRangeViaFactoring(const SCEV *Start,
7209 const SCEV *Step,
7210 const APInt &MaxBECount) {
7211 // RangeOf({C?A:B,+,C?P:Q}) == RangeOf(C?{A,+,P}:{B,+,Q})
7212 // == RangeOf({A,+,P}) union RangeOf({B,+,Q})
7213
7214 unsigned BitWidth = MaxBECount.getBitWidth();
7215 assert(getTypeSizeInBits(Start->getType()) == BitWidth &&
7216 getTypeSizeInBits(Step->getType()) == BitWidth &&
7217 "mismatched bit widths");
7218
7219 struct SelectPattern {
7220 Value *Condition = nullptr;
7221 APInt TrueValue;
7222 APInt FalseValue;
7223
7224 explicit SelectPattern(ScalarEvolution &SE, unsigned BitWidth,
7225 const SCEV *S) {
7226 std::optional<unsigned> CastOp;
7227 APInt Offset(BitWidth, 0);
7228
7230 "Should be!");
7231
7232 // Peel off a constant offset. In the future we could consider being
7233 // smarter here and handle {Start+Step,+,Step} too.
7234 const APInt *Off;
7235 if (match(S, m_scev_Add(m_scev_APInt(Off), m_SCEV(S))))
7236 Offset = *Off;
7237
7238 // Peel off a cast operation
7239 if (auto *SCast = dyn_cast<SCEVIntegralCastExpr>(S)) {
7240 CastOp = SCast->getSCEVType();
7241 S = SCast->getOperand();
7242 }
7243
7244 using namespace llvm::PatternMatch;
7245
7246 auto *SU = dyn_cast<SCEVUnknown>(S);
7247 const APInt *TrueVal, *FalseVal;
7248 if (!SU ||
7249 !match(SU->getValue(), m_Select(m_Value(Condition), m_APInt(TrueVal),
7250 m_APInt(FalseVal)))) {
7251 Condition = nullptr;
7252 return;
7253 }
7254
7255 TrueValue = *TrueVal;
7256 FalseValue = *FalseVal;
7257
7258 // Re-apply the cast we peeled off earlier
7259 if (CastOp)
7260 switch (*CastOp) {
7261 default:
7262 llvm_unreachable("Unknown SCEV cast type!");
7263
7264 case scTruncate:
7265 TrueValue = TrueValue.trunc(BitWidth);
7266 FalseValue = FalseValue.trunc(BitWidth);
7267 break;
7268 case scZeroExtend:
7269 TrueValue = TrueValue.zext(BitWidth);
7270 FalseValue = FalseValue.zext(BitWidth);
7271 break;
7272 case scSignExtend:
7273 TrueValue = TrueValue.sext(BitWidth);
7274 FalseValue = FalseValue.sext(BitWidth);
7275 break;
7276 }
7277
7278 // Re-apply the constant offset we peeled off earlier
7279 TrueValue += Offset;
7280 FalseValue += Offset;
7281 }
7282
7283 bool isRecognized() { return Condition != nullptr; }
7284 };
7285
7286 SelectPattern StartPattern(*this, BitWidth, Start);
7287 if (!StartPattern.isRecognized())
7288 return ConstantRange::getFull(BitWidth);
7289
7290 SelectPattern StepPattern(*this, BitWidth, Step);
7291 if (!StepPattern.isRecognized())
7292 return ConstantRange::getFull(BitWidth);
7293
7294 if (StartPattern.Condition != StepPattern.Condition) {
7295 // We don't handle this case today; but we could, by considering four
7296 // possibilities below instead of two. I'm not sure if there are cases where
7297 // that will help over what getRange already does, though.
7298 return ConstantRange::getFull(BitWidth);
7299 }
7300
7301 // NB! Calling ScalarEvolution::getConstant is fine, but we should not try to
7302 // construct arbitrary general SCEV expressions here. This function is called
7303 // from deep in the call stack, and calling getSCEV (on a sext instruction,
7304 // say) can end up caching a suboptimal value.
7305
7306 // FIXME: without the explicit `this` receiver below, MSVC errors out with
7307 // C2352 and C2512 (otherwise it isn't needed).
7308
7309 const SCEV *TrueStart = this->getConstant(StartPattern.TrueValue);
7310 const SCEV *TrueStep = this->getConstant(StepPattern.TrueValue);
7311 const SCEV *FalseStart = this->getConstant(StartPattern.FalseValue);
7312 const SCEV *FalseStep = this->getConstant(StepPattern.FalseValue);
7313
7314 ConstantRange TrueRange =
7315 this->getRangeForAffineAR(TrueStart, TrueStep, MaxBECount).first;
7316 ConstantRange FalseRange =
7317 this->getRangeForAffineAR(FalseStart, FalseStep, MaxBECount).first;
7318
7319 return TrueRange.unionWith(FalseRange);
7320}
7321
7322SCEVFlags ScalarEvolution::getNoWrapFlagsFromUB(const Value *V) {
7323 if (isa<ConstantExpr>(V))
7324 return SCEV::FlagNone;
7325 const BinaryOperator *BinOp = cast<BinaryOperator>(V);
7326
7327 // Return early if there are no flags to propagate to the SCEV.
7329 if (auto *PDI = dyn_cast<PossiblyDisjointInst>(BinOp);
7330 PDI && PDI->isDisjoint()) {
7332 } else {
7333 if (BinOp->hasNoUnsignedWrap())
7335 if (BinOp->hasNoSignedWrap())
7337 }
7338 if (Flags == SCEV::FlagNone)
7339 return SCEV::FlagNone;
7340
7341 return isSCEVExprNeverPoison(BinOp) ? Flags : SCEV::FlagNone;
7342}
7343
7344const Instruction *
7345ScalarEvolution::getNonTrivialDefiningScopeBound(const SCEV *S) {
7346 if (auto *AddRec = dyn_cast<SCEVAddRecExpr>(S))
7347 return &*AddRec->getLoop()->getHeader()->begin();
7348 if (auto *U = dyn_cast<SCEVUnknown>(S))
7349 if (auto *I = dyn_cast<Instruction>(U->getValue()))
7350 return I;
7351 return nullptr;
7352}
7353
7354const Instruction *ScalarEvolution::getDefiningScopeBound(ArrayRef<SCEVUse> Ops,
7355 bool &Precise) {
7356 Precise = true;
7357 // Do a bounded search of the def relation of the requested SCEVs.
7358 SmallPtrSet<const SCEV *, 16> Visited;
7359 SmallVector<SCEVUse> Worklist;
7360 auto pushOp = [&](const SCEV *S) {
7361 if (!Visited.insert(S).second)
7362 return;
7363 // Threshold of 30 here is arbitrary.
7364 if (Visited.size() > 30) {
7365 Precise = false;
7366 return;
7367 }
7368 Worklist.push_back(S);
7369 };
7370
7371 for (SCEVUse S : Ops)
7372 pushOp(S);
7373
7374 const Instruction *Bound = nullptr;
7375 while (!Worklist.empty()) {
7376 SCEVUse S = Worklist.pop_back_val();
7377 if (auto *DefI = getNonTrivialDefiningScopeBound(S)) {
7378 if (!Bound || DT.dominates(Bound, DefI))
7379 Bound = DefI;
7380 } else {
7381 for (SCEVUse Op : S->operands())
7382 pushOp(Op);
7383 }
7384 }
7385 return Bound ? Bound : &*F.getEntryBlock().begin();
7386}
7387
7388const Instruction *
7389ScalarEvolution::getDefiningScopeBound(ArrayRef<SCEVUse> Ops) {
7390 bool Discard;
7391 return getDefiningScopeBound(Ops, Discard);
7392}
7393
7394bool ScalarEvolution::isGuaranteedToTransferExecutionTo(const Instruction *A,
7395 const Instruction *B) {
7396 if (A->getParent() == B->getParent() &&
7398 B->getIterator()))
7399 return true;
7400
7401 auto *BLoop = LI.getLoopFor(B->getParent());
7402 if (BLoop && BLoop->getHeader() == B->getParent() &&
7403 BLoop->getLoopPreheader() == A->getParent() &&
7405 A->getParent()->end()) &&
7406 isGuaranteedToTransferExecutionToSuccessor(B->getParent()->begin(),
7407 B->getIterator()))
7408 return true;
7409 return false;
7410}
7411
7413 SCEVPoisonCollector PC(/* LookThroughMaybePoisonBlocking */ true);
7414 visitAll(Op, PC);
7415 return PC.MaybePoison.empty();
7416}
7417
7418bool ScalarEvolution::isGuaranteedNotToCauseUB(const SCEV *Op) {
7419 return !SCEVExprContains(Op, [this](const SCEV *S) {
7420 const SCEV *Op1;
7421 bool M = match(S, m_scev_UDiv(m_SCEV(), m_SCEV(Op1)));
7422 // The UDiv may be UB if the divisor is poison or zero. Unless the divisor
7423 // is a non-zero constant, we have to assume the UDiv may be UB.
7424 return M && (!isKnownNonZero(Op1) || !isGuaranteedNotToBePoison(Op1));
7425 });
7426}
7427
7428bool ScalarEvolution::isSCEVExprNeverPoison(const Instruction *I) {
7429 // Only proceed if we can prove that I does not yield poison.
7431 return false;
7432
7433 // At this point we know that if I is executed, then it does not wrap
7434 // according to at least one of NSW or NUW. If I is not executed, then we do
7435 // not know if the calculation that I represents would wrap. Multiple
7436 // instructions can map to the same SCEV. If we apply NSW or NUW from I to
7437 // the SCEV, we must guarantee no wrapping for that SCEV also when it is
7438 // derived from other instructions that map to the same SCEV. We cannot make
7439 // that guarantee for cases where I is not executed. So we need to find a
7440 // upper bound on the defining scope for the SCEV, and prove that I is
7441 // executed every time we enter that scope. When the bounding scope is a
7442 // loop (the common case), this is equivalent to proving I executes on every
7443 // iteration of that loop.
7444 SmallVector<SCEVUse> SCEVOps;
7445 for (const Use &Op : I->operands()) {
7446 // I could be an extractvalue from a call to an overflow intrinsic.
7447 // TODO: We can do better here in some cases.
7448 if (isSCEVable(Op->getType()))
7449 SCEVOps.push_back(getSCEV(Op));
7450 }
7451 auto *DefI = getDefiningScopeBound(SCEVOps);
7452 return isGuaranteedToTransferExecutionTo(DefI, I);
7453}
7454
7455bool ScalarEvolution::isAddRecNeverPoison(const Instruction *I, const Loop *L) {
7456 // If we know that \c I can never be poison period, then that's enough.
7457 if (isSCEVExprNeverPoison(I))
7458 return true;
7459
7460 // If the loop only has one exit, then we know that, if the loop is entered,
7461 // any instruction dominating that exit will be executed. If any such
7462 // instruction would result in UB, the addrec cannot be poison.
7463 //
7464 // This is basically the same reasoning as in isSCEVExprNeverPoison(), but
7465 // also handles uses outside the loop header (they just need to dominate the
7466 // single exit).
7467
7468 auto *ExitingBB = L->getExitingBlock();
7469 if (!ExitingBB || !loopHasNoAbnormalExits(L))
7470 return false;
7471
7472 SmallPtrSet<const Value *, 16> KnownPoison;
7474
7475 // We start by assuming \c I, the post-inc add recurrence, is poison. Only
7476 // things that are known to be poison under that assumption go on the
7477 // Worklist.
7478 KnownPoison.insert(I);
7479 Worklist.push_back(I);
7480
7481 while (!Worklist.empty()) {
7482 const Instruction *Poison = Worklist.pop_back_val();
7483
7484 for (const Use &U : Poison->uses()) {
7485 const Instruction *PoisonUser = cast<Instruction>(U.getUser());
7486 if (mustTriggerUB(PoisonUser, KnownPoison) &&
7487 DT.dominates(PoisonUser->getParent(), ExitingBB))
7488 return true;
7489
7490 if (propagatesPoison(U) && L->contains(PoisonUser))
7491 if (KnownPoison.insert(PoisonUser).second)
7492 Worklist.push_back(PoisonUser);
7493 }
7494 }
7495
7496 return false;
7497}
7498
7499ScalarEvolution::LoopProperties
7500ScalarEvolution::getLoopProperties(const Loop *L) {
7501 using LoopProperties = ScalarEvolution::LoopProperties;
7502
7503 auto Itr = LoopPropertiesCache.find(L);
7504 if (Itr == LoopPropertiesCache.end()) {
7505 auto HasSideEffects = [](Instruction *I) {
7506 if (auto *SI = dyn_cast<StoreInst>(I))
7507 return !SI->isSimple();
7508
7509 if (I->mayThrow())
7510 return true;
7511
7512 // Non-volatile memset / memcpy do not count as side-effect for forward
7513 // progress.
7514 if (isa<MemIntrinsic>(I) && !I->isVolatile())
7515 return false;
7516
7517 return I->mayWriteToMemory();
7518 };
7519
7520 LoopProperties LP = {/* HasNoAbnormalExits */ true,
7521 /*HasNoSideEffects*/ true};
7522
7523 for (auto *BB : L->getBlocks())
7524 for (auto &I : *BB) {
7526 LP.HasNoAbnormalExits = false;
7527 if (HasSideEffects(&I))
7528 LP.HasNoSideEffects = false;
7529 if (!LP.HasNoAbnormalExits && !LP.HasNoSideEffects)
7530 break; // We're already as pessimistic as we can get.
7531 }
7532
7533 auto InsertPair = LoopPropertiesCache.insert({L, LP});
7534 assert(InsertPair.second && "We just checked!");
7535 Itr = InsertPair.first;
7536 }
7537
7538 return Itr->second;
7539}
7540
7542 // A mustprogress loop without side effects must be finite.
7543 // TODO: The check used here is very conservative. It's only *specific*
7544 // side effects which are well defined in infinite loops.
7545 return isFinite(L) || (isMustProgress(L) && loopHasNoSideEffects(L));
7546}
7547
7548const SCEV *ScalarEvolution::createSCEVIter(Value *V) {
7549 // Worklist item with a Value and a bool indicating whether all operands have
7550 // been visited already.
7553
7554 Stack.emplace_back(V, false);
7555 while (!Stack.empty()) {
7556 auto E = Stack.back();
7557 Value *CurV = E.getPointer();
7558
7559 if (getExistingSCEV(CurV)) {
7560 Stack.pop_back();
7561 continue;
7562 }
7563
7565 const SCEV *CreatedSCEV = nullptr;
7566 // If all operands have been visited already, create the SCEV.
7567 if (E.getInt()) {
7568 CreatedSCEV = createSCEV(CurV);
7569 } else {
7570 // Otherwise get the operands we need to create SCEV's for before creating
7571 // the SCEV for CurV. If the SCEV for CurV can be constructed trivially,
7572 // just use it.
7573 CreatedSCEV = getOperandsToCreate(CurV, Ops);
7574 }
7575
7576 if (CreatedSCEV) {
7577 insertValueToMap(CurV, CreatedSCEV);
7578 Stack.pop_back();
7579 } else {
7580 Stack.back().setInt(true);
7581 // Queue its operands which need to be constructed.
7582 for (Value *Op : Ops)
7583 Stack.emplace_back(Op, false);
7584 }
7585 }
7586
7587 return getExistingSCEV(V);
7588}
7589
7590const SCEV *
7591ScalarEvolution::getOperandsToCreate(Value *V, SmallVectorImpl<Value *> &Ops) {
7592 if (!isSCEVable(V->getType()))
7593 return getUnknown(V);
7594
7595 if (Instruction *I = dyn_cast<Instruction>(V)) {
7596 // Don't attempt to analyze instructions in blocks that aren't
7597 // reachable. Such instructions don't matter, and they aren't required
7598 // to obey basic rules for definitions dominating uses which this
7599 // analysis depends on.
7600 if (!DT.isReachableFromEntry(I->getParent()))
7601 return getUnknown(PoisonValue::get(V->getType()));
7602 } else if (ConstantInt *CI = dyn_cast<ConstantInt>(V))
7603 return getConstant(CI);
7604 else if (isa<GlobalAlias>(V))
7605 return getUnknown(V);
7606 else if (!isa<ConstantExpr>(V))
7607 return getUnknown(V);
7608
7610 if (auto BO =
7612 bool IsConstArg = isa<ConstantInt>(BO->RHS);
7613 switch (BO->Opcode) {
7614 case Instruction::Add:
7615 case Instruction::Mul: {
7616 // For additions and multiplications, traverse add/mul chains for which we
7617 // can potentially create a single SCEV, to reduce the number of
7618 // get{Add,Mul}Expr calls.
7619 do {
7620 if (BO->Op) {
7621 if (BO->Op != V && getExistingSCEV(BO->Op)) {
7622 Ops.push_back(BO->Op);
7623 break;
7624 }
7625 }
7626 Ops.push_back(BO->RHS);
7627 auto NewBO = MatchBinaryOp(BO->LHS, getDataLayout(), AC, DT,
7629 if (!NewBO ||
7630 (BO->Opcode == Instruction::Add &&
7631 (NewBO->Opcode != Instruction::Add &&
7632 NewBO->Opcode != Instruction::Sub)) ||
7633 (BO->Opcode == Instruction::Mul &&
7634 NewBO->Opcode != Instruction::Mul)) {
7635 Ops.push_back(BO->LHS);
7636 break;
7637 }
7638 // CreateSCEV calls getNoWrapFlagsFromUB, which under certain conditions
7639 // requires a SCEV for the LHS.
7640 if (BO->Op && (BO->IsNSW || BO->IsNUW)) {
7641 auto *I = dyn_cast<Instruction>(BO->Op);
7642 if (I && programUndefinedIfPoison(I)) {
7643 Ops.push_back(BO->LHS);
7644 break;
7645 }
7646 }
7647 BO = NewBO;
7648 } while (true);
7649 return nullptr;
7650 }
7651 case Instruction::Sub:
7652 case Instruction::UDiv:
7653 case Instruction::URem:
7654 break;
7655 case Instruction::AShr:
7656 case Instruction::Shl:
7657 case Instruction::Xor:
7658 if (!IsConstArg)
7659 return nullptr;
7660 break;
7661 case Instruction::And:
7662 case Instruction::Or:
7663 if (!IsConstArg && !BO->LHS->getType()->isIntegerTy(1))
7664 return nullptr;
7665 break;
7666 case Instruction::LShr:
7667 return getUnknown(V);
7668 default:
7669 llvm_unreachable("Unhandled binop");
7670 break;
7671 }
7672
7673 Ops.push_back(BO->LHS);
7674 Ops.push_back(BO->RHS);
7675 return nullptr;
7676 }
7677
7678 switch (U->getOpcode()) {
7679 case Instruction::Trunc:
7680 case Instruction::ZExt:
7681 case Instruction::SExt:
7682 case Instruction::PtrToAddr:
7683 case Instruction::PtrToInt:
7684 Ops.push_back(U->getOperand(0));
7685 return nullptr;
7686
7687 case Instruction::BitCast:
7688 if (isSCEVable(U->getType()) && isSCEVable(U->getOperand(0)->getType())) {
7689 Ops.push_back(U->getOperand(0));
7690 return nullptr;
7691 }
7692 return getUnknown(V);
7693
7694 case Instruction::SDiv:
7695 case Instruction::SRem:
7696 Ops.push_back(U->getOperand(0));
7697 Ops.push_back(U->getOperand(1));
7698 return nullptr;
7699
7700 case Instruction::GetElementPtr:
7701 assert(cast<GEPOperator>(U)->getSourceElementType()->isSized() &&
7702 "GEP source element type must be sized");
7703 llvm::append_range(Ops, U->operands());
7704 return nullptr;
7705
7706 case Instruction::IntToPtr:
7707 return getUnknown(V);
7708
7709 case Instruction::PHI:
7710 // getNodeForPHI has four ways to turn a PHI into a SCEV; retrieve the
7711 // relevant nodes for each of them.
7712 //
7713 // The first is just to call simplifyInstruction, and get something back
7714 // that isn't a PHI.
7715 if (Value *V = simplifyInstruction(
7716 cast<PHINode>(U),
7717 {getDataLayout(), &TLI, &DT, &AC, /*CtxI=*/nullptr,
7718 /*UseInstrInfo=*/true, /*CanUseUndef=*/false})) {
7719 assert(V);
7720 Ops.push_back(V);
7721 return nullptr;
7722 }
7723 // The second is createNodeForPHIWithIdenticalOperands: this looks for
7724 // operands which all perform the same operation, but haven't been
7725 // CSE'ed for whatever reason.
7726 if (BinaryOperator *BO = getCommonInstForPHI(cast<PHINode>(U))) {
7727 assert(BO);
7728 Ops.push_back(BO);
7729 return nullptr;
7730 }
7731 // The third is createNodeFromSelectLikePHI; this takes a PHI which
7732 // is equivalent to a select, and analyzes it like a select.
7733 {
7734 Value *Cond = nullptr, *LHS = nullptr, *RHS = nullptr;
7736 assert(Cond);
7737 assert(LHS);
7738 assert(RHS);
7739 if (auto *CondICmp = dyn_cast<ICmpInst>(Cond)) {
7740 Ops.push_back(CondICmp->getOperand(0));
7741 Ops.push_back(CondICmp->getOperand(1));
7742 }
7743 Ops.push_back(Cond);
7744 Ops.push_back(LHS);
7745 Ops.push_back(RHS);
7746 return nullptr;
7747 }
7748 }
7749 // The fourth way is createAddRecFromPHI. It's complicated to handle here,
7750 // so just construct it recursively.
7751 //
7752 // In addition to getNodeForPHI, also construct nodes which might be needed
7753 // by getRangeRef.
7755 for (Value *V : cast<PHINode>(U)->operands())
7756 Ops.push_back(V);
7757 return nullptr;
7758 }
7759 return nullptr;
7760
7761 case Instruction::Select: {
7762 // Check if U is a select that can be simplified to a SCEVUnknown.
7763 auto CanSimplifyToUnknown = [this, U]() {
7764 if (U->getType()->isIntegerTy(1) || isa<ConstantInt>(U->getOperand(0)))
7765 return false;
7766
7767 auto *ICI = dyn_cast<ICmpInst>(U->getOperand(0));
7768 if (!ICI)
7769 return false;
7770 Value *LHS = ICI->getOperand(0);
7771 Value *RHS = ICI->getOperand(1);
7772 if (ICI->getPredicate() == CmpInst::ICMP_EQ ||
7773 ICI->getPredicate() == CmpInst::ICMP_NE) {
7775 return true;
7776 } else if (getTypeSizeInBits(LHS->getType()) >
7777 getTypeSizeInBits(U->getType()))
7778 return true;
7779 return false;
7780 };
7781 if (CanSimplifyToUnknown())
7782 return getUnknown(U);
7783
7784 llvm::append_range(Ops, U->operands());
7785 return nullptr;
7786 break;
7787 }
7788 case Instruction::Call:
7789 case Instruction::Invoke:
7790 if (Value *RV = cast<CallBase>(U)->getReturnedArgOperand()) {
7791 Ops.push_back(RV);
7792 return nullptr;
7793 }
7794
7795 if (auto *II = dyn_cast<IntrinsicInst>(U)) {
7796 switch (II->getIntrinsicID()) {
7797 case Intrinsic::abs:
7798 Ops.push_back(II->getArgOperand(0));
7799 return nullptr;
7800 case Intrinsic::umax:
7801 case Intrinsic::umin:
7802 case Intrinsic::smax:
7803 case Intrinsic::smin:
7804 case Intrinsic::usub_sat:
7805 case Intrinsic::uadd_sat:
7806 Ops.push_back(II->getArgOperand(0));
7807 Ops.push_back(II->getArgOperand(1));
7808 return nullptr;
7809 case Intrinsic::start_loop_iterations:
7810 case Intrinsic::annotation:
7811 case Intrinsic::ptr_annotation:
7812 Ops.push_back(II->getArgOperand(0));
7813 return nullptr;
7814 default:
7815 break;
7816 }
7817 }
7818 break;
7819 }
7820
7821 return nullptr;
7822}
7823
7824const SCEV *ScalarEvolution::createSCEV(Value *V) {
7825 if (!isSCEVable(V->getType()))
7826 return getUnknown(V);
7827
7828 if (Instruction *I = dyn_cast<Instruction>(V)) {
7829 // Don't attempt to analyze instructions in blocks that aren't
7830 // reachable. Such instructions don't matter, and they aren't required
7831 // to obey basic rules for definitions dominating uses which this
7832 // analysis depends on.
7833 if (!DT.isReachableFromEntry(I->getParent()))
7834 return getUnknown(PoisonValue::get(V->getType()));
7835 } else if (ConstantInt *CI = dyn_cast<ConstantInt>(V))
7836 return getConstant(CI);
7837 else if (isa<GlobalAlias>(V))
7838 return getUnknown(V);
7839 else if (!isa<ConstantExpr>(V))
7840 return getUnknown(V);
7841
7842 const SCEV *LHS;
7843 const SCEV *RHS;
7844
7846 if (auto BO =
7848 switch (BO->Opcode) {
7849 case Instruction::Add: {
7850 // The simple thing to do would be to just call getSCEV on both operands
7851 // and call getAddExpr with the result. However if we're looking at a
7852 // bunch of things all added together, this can be quite inefficient,
7853 // because it leads to N-1 getAddExpr calls for N ultimate operands.
7854 // Instead, gather up all the operands and make a single getAddExpr call.
7855 // LLVM IR canonical form means we need only traverse the left operands.
7857 do {
7858 if (BO->Op) {
7859 if (auto *OpSCEV = getExistingSCEV(BO->Op)) {
7860 AddOps.push_back(OpSCEV);
7861 break;
7862 }
7863
7864 // If a NUW or NSW flag can be applied to the SCEV for this
7865 // addition, then compute the SCEV for this addition by itself
7866 // with a separate call to getAddExpr. We need to do that
7867 // instead of pushing the operands of the addition onto AddOps,
7868 // since the flags are only known to apply to this particular
7869 // addition - they may not apply to other additions that can be
7870 // formed with operands from AddOps.
7871 const SCEV *RHS = getSCEV(BO->RHS);
7872 SCEVFlags Flags = getNoWrapFlagsFromUB(BO->Op);
7873 if (Flags != SCEV::FlagNone) {
7874 const SCEV *LHS = getSCEV(BO->LHS);
7875 if (BO->Opcode == Instruction::Sub)
7876 AddOps.push_back(getMinusSCEV(LHS, RHS, Flags));
7877 else
7878 AddOps.push_back(getAddExpr(LHS, RHS, Flags));
7879 break;
7880 }
7881 }
7882
7883 if (BO->Opcode == Instruction::Sub)
7884 AddOps.push_back(getNegativeSCEV(getSCEV(BO->RHS)));
7885 else
7886 AddOps.push_back(getSCEV(BO->RHS));
7887
7888 auto NewBO = MatchBinaryOp(BO->LHS, getDataLayout(), AC, DT,
7890 if (!NewBO || (NewBO->Opcode != Instruction::Add &&
7891 NewBO->Opcode != Instruction::Sub)) {
7892 AddOps.push_back(getSCEV(BO->LHS));
7893 break;
7894 }
7895 BO = NewBO;
7896 } while (true);
7897
7898 return getAddExpr(AddOps);
7899 }
7900
7901 case Instruction::Mul: {
7903 do {
7904 if (BO->Op) {
7905 if (auto *OpSCEV = getExistingSCEV(BO->Op)) {
7906 MulOps.push_back(OpSCEV);
7907 break;
7908 }
7909
7910 SCEVFlags Flags = getNoWrapFlagsFromUB(BO->Op);
7911 if (Flags != SCEV::FlagNone) {
7912 LHS = getSCEV(BO->LHS);
7913 RHS = getSCEV(BO->RHS);
7914 MulOps.push_back(getMulExpr(LHS, RHS, Flags));
7915 break;
7916 }
7917 }
7918
7919 MulOps.push_back(getSCEV(BO->RHS));
7920 auto NewBO = MatchBinaryOp(BO->LHS, getDataLayout(), AC, DT,
7922 if (!NewBO || NewBO->Opcode != Instruction::Mul) {
7923 MulOps.push_back(getSCEV(BO->LHS));
7924 break;
7925 }
7926 BO = NewBO;
7927 } while (true);
7928
7929 return getMulExpr(MulOps);
7930 }
7931 case Instruction::UDiv:
7932 LHS = getSCEV(BO->LHS);
7933 RHS = getSCEV(BO->RHS);
7934 return getUDivExpr(LHS, RHS);
7935 case Instruction::URem:
7936 LHS = getSCEV(BO->LHS);
7937 RHS = getSCEV(BO->RHS);
7938 return getURemExpr(LHS, RHS);
7939 case Instruction::Sub: {
7941 if (BO->Op)
7942 Flags = getNoWrapFlagsFromUB(BO->Op);
7943
7944 // Try to use ptrtoaddr for subtracts with at least one ptrtoint
7945 // operand. While we don't model ptrtoint directly in SCEV, the
7946 // difference between two pointer addresses is well-defined.
7947 Value *PtrLHS = nullptr, *PtrRHS = nullptr;
7948 bool HasPtrLHS = match(BO->LHS, m_PtrToInt(m_Value(PtrLHS)));
7949 bool HasPtrRHS = match(BO->RHS, m_PtrToInt(m_Value(PtrRHS)));
7950 if (HasPtrLHS || HasPtrRHS) {
7951 // Convert a ptrtoint operand (OrigOp) to ptrtoaddr of its pointer
7952 // PtrOp. When only one side is ptrtoint (BothPtr is false), skip
7953 // SCEVUnknown pointers since wrapping them in ptrtoaddr adds no
7954 // useful structure.
7955 auto GetOp = [&](bool HasPtr, Value *PtrOp, Value *OrigOp,
7956 bool BothPtr) -> const SCEV * {
7957 if (!HasPtr)
7958 return getSCEV(OrigOp);
7959 const SCEV *PtrSCEV = getSCEV(PtrOp);
7960 if (BothPtr || !isa<SCEVUnknown>(PtrSCEV)) {
7961 const SCEV *Addr = getPtrToAddrExpr(PtrSCEV);
7962 if (!isa<SCEVCouldNotCompute>(Addr) &&
7963 getTypeSizeInBits(OrigOp->getType()) <=
7964 getTypeSizeInBits(Addr->getType()))
7965 return getTruncateOrNoop(Addr, OrigOp->getType());
7966 }
7967 return getSCEV(OrigOp);
7968 };
7969 const SCEV *L = GetOp(HasPtrLHS, PtrLHS, BO->LHS, HasPtrRHS);
7970 const SCEV *R = GetOp(HasPtrRHS, PtrRHS, BO->RHS, HasPtrLHS);
7971 return getMinusSCEV(L, R, Flags);
7972 }
7973
7974 LHS = getSCEV(BO->LHS);
7975 RHS = getSCEV(BO->RHS);
7976 return getMinusSCEV(LHS, RHS, Flags);
7977 }
7978 case Instruction::And:
7979 // For an expression like x&255 that merely masks off the high bits,
7980 // use zext(trunc(x)) as the SCEV expression.
7981 if (ConstantInt *CI = dyn_cast<ConstantInt>(BO->RHS)) {
7982 if (CI->isZero())
7983 return getSCEV(BO->RHS);
7984 if (CI->isMinusOne())
7985 return getSCEV(BO->LHS);
7986 const APInt &A = CI->getValue();
7987
7988 // Instcombine's ShrinkDemandedConstant may strip bits out of
7989 // constants, obscuring what would otherwise be a low-bits mask.
7990 // Use computeKnownBits to compute what ShrinkDemandedConstant
7991 // knew about to reconstruct a low-bits mask value.
7992 unsigned LZ = A.countl_zero();
7993 unsigned TZ = A.countr_zero();
7994 unsigned BitWidth = A.getBitWidth();
7995 KnownBits Known(BitWidth);
7996 computeKnownBits(BO->LHS, Known, getDataLayout(), &AC, nullptr, &DT);
7997
7998 APInt EffectiveMask =
7999 APInt::getLowBitsSet(BitWidth, BitWidth - LZ - TZ).shl(TZ);
8000 if ((LZ != 0 || TZ != 0) && !((~A & ~Known.Zero) & EffectiveMask)) {
8001 const SCEV *MulCount = getConstant(APInt::getOneBitSet(BitWidth, TZ));
8002 const SCEV *LHS = getSCEV(BO->LHS);
8003 const SCEV *ShiftedLHS = nullptr;
8004 if (auto *LHSMul = dyn_cast<SCEVMulExpr>(LHS)) {
8005 if (auto *OpC = dyn_cast<SCEVConstant>(LHSMul->getOperand(0))) {
8006 // For an expression like (x * 8) & 8, simplify the multiply.
8007 unsigned MulZeros = OpC->getAPInt().countr_zero();
8008 unsigned GCD = std::min(MulZeros, TZ);
8009 APInt DivAmt = APInt::getOneBitSet(BitWidth, TZ - GCD);
8011 MulOps.push_back(getConstant(OpC->getAPInt().ashr(GCD)));
8012 append_range(MulOps, LHSMul->operands().drop_front());
8013 const SCEV *NewMul = getMulExpr(MulOps, LHSMul->getNoWrapFlags());
8014 ShiftedLHS = getUDivExpr(NewMul, getConstant(DivAmt));
8015 }
8016 }
8017 if (!ShiftedLHS)
8018 ShiftedLHS = getUDivExpr(LHS, MulCount);
8019 return getMulExpr(
8021 getTruncateExpr(ShiftedLHS,
8022 IntegerType::get(getContext(), BitWidth - LZ - TZ)),
8023 BO->LHS->getType()),
8024 MulCount);
8025 }
8026 }
8027 // Binary `and` is a bit-wise `umin`.
8028 if (BO->LHS->getType()->isIntegerTy(1)) {
8029 LHS = getSCEV(BO->LHS);
8030 RHS = getSCEV(BO->RHS);
8031 return getUMinExpr(LHS, RHS);
8032 }
8033 break;
8034
8035 case Instruction::Or:
8036 // Binary `or` is a bit-wise `umax`.
8037 if (BO->LHS->getType()->isIntegerTy(1)) {
8038 LHS = getSCEV(BO->LHS);
8039 RHS = getSCEV(BO->RHS);
8040 return getUMaxExpr(LHS, RHS);
8041 }
8042 break;
8043
8044 case Instruction::Xor:
8045 if (ConstantInt *CI = dyn_cast<ConstantInt>(BO->RHS)) {
8046 // If the RHS of xor is -1, then this is a not operation.
8047 if (CI->isMinusOne())
8048 return getNotSCEV(getSCEV(BO->LHS));
8049
8050 // Model xor(and(x, C), C) as and(~x, C), if C is a low-bits mask.
8051 // This is a variant of the check for xor with -1, and it handles
8052 // the case where instcombine has trimmed non-demanded bits out
8053 // of an xor with -1.
8054 if (auto *LBO = dyn_cast<BinaryOperator>(BO->LHS))
8055 if (ConstantInt *LCI = dyn_cast<ConstantInt>(LBO->getOperand(1)))
8056 if (LBO->getOpcode() == Instruction::And &&
8057 LCI->getValue() == CI->getValue())
8058 if (const SCEVZeroExtendExpr *Z =
8060 Type *UTy = BO->LHS->getType();
8061 const SCEV *Z0 = Z->getOperand();
8062 Type *Z0Ty = Z0->getType();
8063 unsigned Z0TySize = getTypeSizeInBits(Z0Ty);
8064
8065 // If C is a low-bits mask, the zero extend is serving to
8066 // mask off the high bits. Complement the operand and
8067 // re-apply the zext.
8068 if (CI->getValue().isMask(Z0TySize))
8069 return getZeroExtendExpr(getNotSCEV(Z0), UTy);
8070
8071 // If C is a single bit, it may be in the sign-bit position
8072 // before the zero-extend. In this case, represent the xor
8073 // using an add, which is equivalent, and re-apply the zext.
8074 APInt Trunc = CI->getValue().trunc(Z0TySize);
8075 if (Trunc.zext(getTypeSizeInBits(UTy)) == CI->getValue() &&
8076 Trunc.isSignMask())
8077 return getZeroExtendExpr(getAddExpr(Z0, getConstant(Trunc)),
8078 UTy);
8079 }
8080 }
8081 break;
8082
8083 case Instruction::Shl:
8084 // Turn shift left of a constant amount into a multiply.
8085 if (ConstantInt *SA = dyn_cast<ConstantInt>(BO->RHS)) {
8086 uint32_t BitWidth = cast<IntegerType>(SA->getType())->getBitWidth();
8087
8088 // If the shift count is not less than the bitwidth, the result of
8089 // the shift is undefined. Don't try to analyze it, because the
8090 // resolution chosen here may differ from the resolution chosen in
8091 // other parts of the compiler.
8092 if (SA->getValue().uge(BitWidth))
8093 break;
8094
8095 // We can safely preserve the nuw flag in all cases. It's also safe to
8096 // turn a nuw nsw shl into a nuw nsw mul. However, nsw in isolation
8097 // requires special handling. It can be preserved as long as we're not
8098 // left shifting by bitwidth - 1.
8099 auto Flags = SCEV::FlagNone;
8100 if (BO->Op) {
8101 auto MulFlags = getNoWrapFlagsFromUB(BO->Op);
8102 if (any(MulFlags & SCEV::FlagNSW) &&
8103 (any(MulFlags & SCEV::FlagNUW) ||
8104 SA->getValue().ult(BitWidth - 1)))
8106 if (any(MulFlags & SCEV::FlagNUW))
8108 }
8109
8110 ConstantInt *X = ConstantInt::get(
8111 getContext(), APInt::getOneBitSet(BitWidth, SA->getZExtValue()));
8112 return getMulExpr(getSCEV(BO->LHS), getConstant(X), Flags);
8113 }
8114 break;
8115
8116 case Instruction::AShr:
8117 // AShr X, C, where C is a constant.
8118 ConstantInt *CI = dyn_cast<ConstantInt>(BO->RHS);
8119 if (!CI)
8120 break;
8121
8122 Type *OuterTy = BO->LHS->getType();
8124 // If the shift count is not less than the bitwidth, the result of
8125 // the shift is undefined. Don't try to analyze it, because the
8126 // resolution chosen here may differ from the resolution chosen in
8127 // other parts of the compiler.
8128 if (CI->getValue().uge(BitWidth))
8129 break;
8130
8131 if (CI->isZero())
8132 return getSCEV(BO->LHS); // shift by zero --> noop
8133
8134 uint64_t AShrAmt = CI->getZExtValue();
8135 Type *TruncTy = IntegerType::get(getContext(), BitWidth - AShrAmt);
8136
8137 Operator *L = dyn_cast<Operator>(BO->LHS);
8138 const SCEV *AddTruncateExpr = nullptr;
8139 ConstantInt *ShlAmtCI = nullptr;
8140 const SCEV *AddConstant = nullptr;
8141
8142 if (L && L->getOpcode() == Instruction::Add) {
8143 // X = Shl A, n
8144 // Y = Add X, c
8145 // Z = AShr Y, m
8146 // n, c and m are constants.
8147
8148 Operator *LShift = dyn_cast<Operator>(L->getOperand(0));
8149 ConstantInt *AddOperandCI = dyn_cast<ConstantInt>(L->getOperand(1));
8150 if (LShift && LShift->getOpcode() == Instruction::Shl) {
8151 if (AddOperandCI) {
8152 const SCEV *ShlOp0SCEV = getSCEV(LShift->getOperand(0));
8153 ShlAmtCI = dyn_cast<ConstantInt>(LShift->getOperand(1));
8154 // since we truncate to TruncTy, the AddConstant should be of the
8155 // same type, so create a new Constant with type same as TruncTy.
8156 // Also, the Add constant should be shifted right by AShr amount.
8157 APInt AddOperand = AddOperandCI->getValue().ashr(AShrAmt);
8158 AddConstant = getConstant(AddOperand.trunc(BitWidth - AShrAmt));
8159 // we model the expression as sext(add(trunc(A), c << n)), since the
8160 // sext(trunc) part is already handled below, we create a
8161 // AddExpr(TruncExp) which will be used later.
8162 AddTruncateExpr = getTruncateExpr(ShlOp0SCEV, TruncTy);
8163 }
8164 }
8165 } else if (L && L->getOpcode() == Instruction::Shl) {
8166 // X = Shl A, n
8167 // Y = AShr X, m
8168 // Both n and m are constant.
8169
8170 const SCEV *ShlOp0SCEV = getSCEV(L->getOperand(0));
8171 ShlAmtCI = dyn_cast<ConstantInt>(L->getOperand(1));
8172 AddTruncateExpr = getTruncateExpr(ShlOp0SCEV, TruncTy);
8173 }
8174
8175 if (AddTruncateExpr && ShlAmtCI) {
8176 // We can merge the two given cases into a single SCEV statement,
8177 // incase n = m, the mul expression will be 2^0, so it gets resolved to
8178 // a simpler case. The following code handles the two cases:
8179 //
8180 // 1) For a two-shift sext-inreg, i.e. n = m,
8181 // use sext(trunc(x)) as the SCEV expression.
8182 //
8183 // 2) When n > m, use sext(mul(trunc(x), 2^(n-m)))) as the SCEV
8184 // expression. We already checked that ShlAmt < BitWidth, so
8185 // the multiplier, 1 << (ShlAmt - AShrAmt), fits into TruncTy as
8186 // ShlAmt - AShrAmt < Amt.
8187 const APInt &ShlAmt = ShlAmtCI->getValue();
8188 if (ShlAmt.ult(BitWidth) && ShlAmt.uge(AShrAmt)) {
8189 APInt Mul = APInt::getOneBitSet(BitWidth - AShrAmt,
8190 ShlAmtCI->getZExtValue() - AShrAmt);
8191 const SCEV *CompositeExpr =
8192 getMulExpr(AddTruncateExpr, getConstant(Mul));
8193 if (L->getOpcode() != Instruction::Shl)
8194 CompositeExpr = getAddExpr(CompositeExpr, AddConstant);
8195
8196 return getSignExtendExpr(CompositeExpr, OuterTy);
8197 }
8198 }
8199 break;
8200 }
8201 }
8202
8203 switch (U->getOpcode()) {
8204 case Instruction::Trunc:
8205 return getTruncateExpr(getSCEV(U->getOperand(0)), U->getType());
8206
8207 case Instruction::ZExt:
8208 return getZeroExtendExpr(getSCEV(U->getOperand(0)), U->getType());
8209
8210 case Instruction::SExt:
8211 if (auto BO = MatchBinaryOp(U->getOperand(0), getDataLayout(), AC, DT,
8213 // The NSW flag of a subtract does not always survive the conversion to
8214 // A + (-1)*B. By pushing sign extension onto its operands we are much
8215 // more likely to preserve NSW and allow later AddRec optimisations.
8216 //
8217 // NOTE: This is effectively duplicating this logic from getSignExtend:
8218 // sext((A + B + ...)<nsw>) --> (sext(A) + sext(B) + ...)<nsw>
8219 // but by that point the NSW information has potentially been lost.
8220 if (BO->Opcode == Instruction::Sub && BO->IsNSW) {
8221 Type *Ty = U->getType();
8222 auto *V1 = getSignExtendExpr(getSCEV(BO->LHS), Ty);
8223 auto *V2 = getSignExtendExpr(getSCEV(BO->RHS), Ty);
8224 return getMinusSCEV(V1, V2, SCEV::FlagNSW);
8225 }
8226 }
8227 return getSignExtendExpr(getSCEV(U->getOperand(0)), U->getType());
8228
8229 case Instruction::BitCast:
8230 // BitCasts are no-op casts so we just eliminate the cast.
8231 if (isSCEVable(U->getType()) && isSCEVable(U->getOperand(0)->getType()))
8232 return getSCEV(U->getOperand(0));
8233 break;
8234
8235 case Instruction::PtrToAddr: {
8236 const SCEV *IntOp = getPtrToAddrExpr(getSCEV(U->getOperand(0)));
8237 if (isa<SCEVCouldNotCompute>(IntOp))
8238 return getUnknown(V);
8239 return IntOp;
8240 }
8241
8242 case Instruction::PtrToInt:
8243 // SCEV only models ptrtoaddr.
8244 return getUnknown(V);
8245
8246 case Instruction::IntToPtr:
8247 // Just don't deal with inttoptr casts.
8248 return getUnknown(V);
8249
8250 case Instruction::SDiv:
8251 // If both operands are non-negative, this is just an udiv.
8252 if (isKnownNonNegative(getSCEV(U->getOperand(0))) &&
8253 isKnownNonNegative(getSCEV(U->getOperand(1))))
8254 return getUDivExpr(getSCEV(U->getOperand(0)), getSCEV(U->getOperand(1)));
8255 break;
8256
8257 case Instruction::SRem:
8258 // If both operands are non-negative, this is just an urem.
8259 if (isKnownNonNegative(getSCEV(U->getOperand(0))) &&
8260 isKnownNonNegative(getSCEV(U->getOperand(1))))
8261 return getURemExpr(getSCEV(U->getOperand(0)), getSCEV(U->getOperand(1)));
8262 break;
8263
8264 case Instruction::GetElementPtr:
8265 return createNodeForGEP(cast<GEPOperator>(U));
8266
8267 case Instruction::PHI:
8268 return createNodeForPHI(cast<PHINode>(U));
8269
8270 case Instruction::Select:
8271 return createNodeForSelectOrPHI(U, U->getOperand(0), U->getOperand(1),
8272 U->getOperand(2));
8273
8274 case Instruction::Call:
8275 case Instruction::Invoke:
8276 if (Value *RV = cast<CallBase>(U)->getReturnedArgOperand())
8277 return getSCEV(RV);
8278
8279 if (auto *II = dyn_cast<IntrinsicInst>(U)) {
8280 switch (II->getIntrinsicID()) {
8281 case Intrinsic::abs:
8282 return getAbsExpr(
8283 getSCEV(II->getArgOperand(0)),
8284 /*IsNSW=*/cast<ConstantInt>(II->getArgOperand(1))->isOne());
8285 case Intrinsic::umax:
8286 LHS = getSCEV(II->getArgOperand(0));
8287 RHS = getSCEV(II->getArgOperand(1));
8288 return getUMaxExpr(LHS, RHS);
8289 case Intrinsic::umin:
8290 LHS = getSCEV(II->getArgOperand(0));
8291 RHS = getSCEV(II->getArgOperand(1));
8292 return getUMinExpr(LHS, RHS);
8293 case Intrinsic::smax:
8294 LHS = getSCEV(II->getArgOperand(0));
8295 RHS = getSCEV(II->getArgOperand(1));
8296 return getSMaxExpr(LHS, RHS);
8297 case Intrinsic::smin:
8298 LHS = getSCEV(II->getArgOperand(0));
8299 RHS = getSCEV(II->getArgOperand(1));
8300 return getSMinExpr(LHS, RHS);
8301 case Intrinsic::usub_sat: {
8302 const SCEV *X = getSCEV(II->getArgOperand(0));
8303 const SCEV *Y = getSCEV(II->getArgOperand(1));
8304 const SCEV *ClampedY = getUMinExpr(X, Y);
8305 return getMinusSCEV(X, ClampedY, SCEV::FlagNUW);
8306 }
8307 case Intrinsic::uadd_sat: {
8308 const SCEV *X = getSCEV(II->getArgOperand(0));
8309 const SCEV *Y = getSCEV(II->getArgOperand(1));
8310 const SCEV *ClampedX = getUMinExpr(X, getNotSCEV(Y));
8311 return getAddExpr(ClampedX, Y, SCEV::FlagNUW);
8312 }
8313 case Intrinsic::start_loop_iterations:
8314 case Intrinsic::annotation:
8315 case Intrinsic::ptr_annotation:
8316 // A start_loop_iterations or llvm.annotation or llvm.prt.annotation is
8317 // just eqivalent to the first operand for SCEV purposes.
8318 return getSCEV(II->getArgOperand(0));
8319 case Intrinsic::vscale:
8320 return getVScale(II->getType());
8321 default:
8322 break;
8323 }
8324 }
8325 break;
8326 }
8327
8328 return getUnknown(V);
8329}
8330
8331//===----------------------------------------------------------------------===//
8332// Iteration Count Computation Code
8333//
8334
8336 if (isa<SCEVCouldNotCompute>(ExitCount))
8337 return getCouldNotCompute();
8338
8339 auto *ExitCountType = ExitCount->getType();
8340 assert(ExitCountType->isIntegerTy());
8341 auto *EvalTy = Type::getIntNTy(ExitCountType->getContext(),
8342 1 + ExitCountType->getScalarSizeInBits());
8343 return getTripCountFromExitCount(ExitCount, EvalTy, nullptr);
8344}
8345
8347 Type *EvalTy,
8348 const Loop *L) {
8349 if (isa<SCEVCouldNotCompute>(ExitCount))
8350 return getCouldNotCompute();
8351
8352 unsigned ExitCountSize = getTypeSizeInBits(ExitCount->getType());
8353 unsigned EvalSize = EvalTy->getPrimitiveSizeInBits();
8354
8355 auto CanAddOneWithoutOverflow = [&]() {
8356 ConstantRange ExitCountRange =
8357 getRangeRef(ExitCount, RangeSignHint::HINT_RANGE_UNSIGNED);
8358 if (!ExitCountRange.contains(APInt::getMaxValue(ExitCountSize)))
8359 return true;
8360
8361 return L && isLoopEntryGuardedByCond(L, ICmpInst::ICMP_NE, ExitCount,
8362 getMinusOne(ExitCount->getType()));
8363 };
8364
8365 // If we need to zero extend the backedge count, check if we can add one to
8366 // it prior to zero extending without overflow. Provided this is safe, it
8367 // allows better simplification of the +1.
8368 if (EvalSize > ExitCountSize && CanAddOneWithoutOverflow())
8369 return getZeroExtendExpr(
8370 getAddExpr(ExitCount, getOne(ExitCount->getType())), EvalTy);
8371
8372 // Get the total trip count from the count by adding 1. This may wrap.
8373 return getAddExpr(getTruncateOrZeroExtend(ExitCount, EvalTy), getOne(EvalTy));
8374}
8375
8376static unsigned getConstantTripCount(const SCEVConstant *ExitCount) {
8377 if (!ExitCount)
8378 return 0;
8379
8380 ConstantInt *ExitConst = ExitCount->getValue();
8381
8382 // Guard against huge trip counts.
8383 if (ExitConst->getValue().getActiveBits() > 32)
8384 return 0;
8385
8386 // In case of integer overflow, this returns 0, which is correct.
8387 return ((unsigned)ExitConst->getZExtValue()) + 1;
8388}
8389
8391 auto *ExitCount = dyn_cast<SCEVConstant>(getBackedgeTakenCount(L, Exact));
8392 return getConstantTripCount(ExitCount);
8393}
8394
8395unsigned
8397 const BasicBlock *ExitingBlock) {
8398 assert(ExitingBlock && "Must pass a non-null exiting block!");
8399 assert(L->isLoopExiting(ExitingBlock) &&
8400 "Exiting block must actually branch out of the loop!");
8401 const SCEVConstant *ExitCount =
8402 dyn_cast<SCEVConstant>(getExitCount(L, ExitingBlock));
8403 return getConstantTripCount(ExitCount);
8404}
8405
8407 const Loop *L, SmallVectorImpl<const SCEVPredicate *> *Predicates) {
8408
8409 const auto *MaxExitCount =
8410 Predicates ? getPredicatedConstantMaxBackedgeTakenCount(L, *Predicates)
8412 return getConstantTripCount(dyn_cast<SCEVConstant>(MaxExitCount));
8413}
8414
8416 SmallVector<BasicBlock *, 8> ExitingBlocks;
8417 L->getExitingBlocks(ExitingBlocks);
8418
8419 // An exit with an uncomputable exit count makes the result 1.
8420 if (ExitingBlocks.empty() ||
8421 any_of(ExitingBlocks, [this, L](BasicBlock *ExitingBB) {
8422 return isa<SCEVCouldNotCompute>(getExitCount(L, ExitingBB));
8423 }))
8424 return 1;
8425
8426 LoopGuards Guards = LoopGuards::collect(L, *this);
8427 unsigned Res = 0;
8428 for (BasicBlock *ExitingBB : ExitingBlocks)
8429 Res = std::gcd(
8430 Res, getSmallConstantTripMultiple(getExitCount(L, ExitingBB), Guards));
8431 return Res;
8432}
8433
8434unsigned
8436 const LoopGuards &Guards) {
8437 assert(!isa<SCEVCouldNotCompute>(ExitCount) && "Must be computable!");
8438
8439 // Get the trip count
8440 const SCEV *TCExpr =
8441 getTripCountFromExitCount(applyLoopGuards(ExitCount, Guards));
8442
8443 APInt Multiple = getNonZeroConstantMultiple(TCExpr);
8444 // If a trip multiple is huge (>=2^32), the trip count is still divisible by
8445 // the greatest power of 2 divisor less than 2^32.
8446 return Multiple.getActiveBits() > 32
8447 ? 1U << std::min(31U, Multiple.countTrailingZeros())
8448 : (unsigned)Multiple.getZExtValue();
8449}
8450
8452 const SCEV *ExitCount) {
8453 if (isa<SCEVCouldNotCompute>(ExitCount))
8454 return 1;
8455
8456 return getSmallConstantTripMultiple(ExitCount, LoopGuards::collect(L, *this));
8457}
8458
8459/// Returns the largest constant divisor of the trip count of this loop as a
8460/// normal unsigned value, if possible. This means that the actual trip count is
8461/// always a multiple of the returned value (don't forget the trip count could
8462/// very well be zero as well!).
8463///
8464/// Returns 1 if the trip count is unknown or not guaranteed to be the
8465/// multiple of a constant (which is also the case if the trip count is simply
8466/// constant, use getSmallConstantTripCount for that case), Will also return 1
8467/// if the trip count is very large (>= 2^32).
8468///
8469/// As explained in the comments for getSmallConstantTripCount, this assumes
8470/// that control exits the loop via ExitingBlock.
8471unsigned
8473 const BasicBlock *ExitingBlock) {
8474 assert(ExitingBlock && "Must pass a non-null exiting block!");
8475 assert(L->isLoopExiting(ExitingBlock) &&
8476 "Exiting block must actually branch out of the loop!");
8477 const SCEV *ExitCount = getExitCount(L, ExitingBlock);
8478 return getSmallConstantTripMultiple(L, ExitCount);
8479}
8480
8482 const BasicBlock *ExitingBlock,
8483 ExitCountKind Kind) {
8484 switch (Kind) {
8485 case Exact:
8486 return getBackedgeTakenInfo(L).getExact(ExitingBlock, this);
8487 case SymbolicMaximum:
8488 return getBackedgeTakenInfo(L).getSymbolicMax(ExitingBlock, this);
8489 case ConstantMaximum:
8490 return getBackedgeTakenInfo(L).getConstantMax(ExitingBlock, this);
8491 };
8492 llvm_unreachable("Invalid ExitCountKind!");
8493}
8494
8496 const Loop *L, const BasicBlock *ExitingBlock,
8498 switch (Kind) {
8499 case Exact:
8500 return getPredicatedBackedgeTakenInfo(L).getExact(ExitingBlock, this,
8501 Predicates);
8502 case SymbolicMaximum:
8503 return getPredicatedBackedgeTakenInfo(L).getSymbolicMax(ExitingBlock, this,
8504 Predicates);
8505 case ConstantMaximum:
8506 return getPredicatedBackedgeTakenInfo(L).getConstantMax(ExitingBlock, this,
8507 Predicates);
8508 };
8509 llvm_unreachable("Invalid ExitCountKind!");
8510}
8511
8514 return getPredicatedBackedgeTakenInfo(L).getExact(L, this, &Preds);
8515}
8516
8518 ExitCountKind Kind) {
8519 switch (Kind) {
8520 case Exact:
8521 return getBackedgeTakenInfo(L).getExact(L, this);
8522 case ConstantMaximum:
8523 return getBackedgeTakenInfo(L).getConstantMax(this);
8524 case SymbolicMaximum:
8525 return getBackedgeTakenInfo(L).getSymbolicMax(L, this);
8526 };
8527 llvm_unreachable("Invalid ExitCountKind!");
8528}
8529
8532 return getPredicatedBackedgeTakenInfo(L).getSymbolicMax(L, this, &Preds);
8533}
8534
8537 return getPredicatedBackedgeTakenInfo(L).getConstantMax(this, &Preds);
8538}
8539
8541 return getBackedgeTakenInfo(L).isConstantMaxOrZero(this);
8542}
8543
8544/// Push PHI nodes in the header of the given loop onto the given Worklist.
8545static void PushLoopPHIs(const Loop *L,
8548 BasicBlock *Header = L->getHeader();
8549
8550 // Push all Loop-header PHIs onto the Worklist stack.
8551 for (PHINode &PN : Header->phis())
8552 if (Visited.insert(&PN).second)
8553 Worklist.push_back(&PN);
8554}
8555
8556ScalarEvolution::BackedgeTakenInfo &
8557ScalarEvolution::getPredicatedBackedgeTakenInfo(const Loop *L) {
8558 auto &BTI = getBackedgeTakenInfo(L);
8559 if (BTI.hasFullInfo())
8560 return BTI;
8561
8562 auto Pair = PredicatedBackedgeTakenCounts.try_emplace(L);
8563
8564 if (!Pair.second)
8565 return Pair.first->second;
8566
8567 BackedgeTakenInfo Result =
8568 computeBackedgeTakenCount(L, /*AllowPredicates=*/true);
8569
8570 return PredicatedBackedgeTakenCounts.find(L)->second = std::move(Result);
8571}
8572
8573ScalarEvolution::BackedgeTakenInfo &
8574ScalarEvolution::getBackedgeTakenInfo(const Loop *L) {
8575 // Initially insert an invalid entry for this loop. If the insertion
8576 // succeeds, proceed to actually compute a backedge-taken count and
8577 // update the value. The temporary CouldNotCompute value tells SCEV
8578 // code elsewhere that it shouldn't attempt to request a new
8579 // backedge-taken count, which could result in infinite recursion.
8580 std::pair<DenseMap<const Loop *, BackedgeTakenInfo>::iterator, bool> Pair =
8581 BackedgeTakenCounts.try_emplace(L);
8582 if (!Pair.second)
8583 return Pair.first->second;
8584
8585 // computeBackedgeTakenCount may allocate memory for its result. Inserting it
8586 // into the BackedgeTakenCounts map transfers ownership. Otherwise, the result
8587 // must be cleared in this scope.
8588 BackedgeTakenInfo Result = computeBackedgeTakenCount(L);
8589
8590 // Now that we know more about the trip count for this loop, forget any
8591 // existing SCEV values for PHI nodes in this loop since they are only
8592 // conservative estimates made without the benefit of trip count
8593 // information. This invalidation is not necessary for correctness, and is
8594 // only done to produce more precise results.
8595 if (Result.hasAnyInfo()) {
8596 // Invalidate any expression using an addrec in this loop.
8597 SmallVector<SCEVUse, 8> ToForget;
8598 auto LoopUsersIt = LoopUsers.find(L);
8599 if (LoopUsersIt != LoopUsers.end())
8600 append_range(ToForget, LoopUsersIt->second);
8601 forgetMemoizedResults(ToForget);
8602
8603 // Invalidate constant-evolved loop header phis.
8604 for (PHINode &PN : L->getHeader()->phis())
8605 ConstantEvolutionLoopExitValue.erase(&PN);
8606 }
8607
8608 // Re-lookup the insert position, since the call to
8609 // computeBackedgeTakenCount above could result in a
8610 // recusive call to getBackedgeTakenInfo (on a different
8611 // loop), which would invalidate the iterator computed
8612 // earlier.
8613 return BackedgeTakenCounts.find(L)->second = std::move(Result);
8614}
8615
8617 // This method is intended to forget all info about loops. It should
8618 // invalidate caches as if the following happened:
8619 // - The trip counts of all loops have changed arbitrarily
8620 // - Every llvm::Value has been updated in place to produce a different
8621 // result.
8622 BackedgeTakenCounts.clear();
8623 PredicatedBackedgeTakenCounts.clear();
8624 BECountUsers.clear();
8625 LoopPropertiesCache.clear();
8626 ConstantEvolutionLoopExitValue.clear();
8627 ValueExprMap.clear();
8628 ValuesAtScopes.clear();
8629 ValuesAtScopesUsers.clear();
8630 LoopDispositions.clear();
8631 BlockDispositions.clear();
8632 UnsignedRanges.clear();
8633 SignedRanges.clear();
8634 ExprValueMap.clear();
8635 HasRecMap.clear();
8636 ConstantMultipleCache.clear();
8637 PredicatedSCEVRewrites.clear();
8638 FoldCache.clear();
8639 FoldCacheUser.clear();
8640}
8641void ScalarEvolution::visitAndClearUsers(
8644 SmallVectorImpl<SCEVUse> &ToForget) {
8645 // Nothing can be invalidated if no value has a SCEV yet.
8646 if (ValueExprMap.empty()) {
8647 Worklist.clear();
8648 return;
8649 }
8650 while (!Worklist.empty()) {
8651 Instruction *I = Worklist.pop_back_val();
8652 if (!isSCEVable(I->getType()) && !isa<WithOverflowInst>(I))
8653 continue;
8654
8656 ValueExprMap.find_as(static_cast<Value *>(I));
8657 if (It != ValueExprMap.end()) {
8658 ToForget.push_back(It->second);
8659 eraseValueFromMap(It->first);
8660 if (PHINode *PN = dyn_cast<PHINode>(I))
8661 ConstantEvolutionLoopExitValue.erase(PN);
8662 }
8663
8664 PushDefUseChildren(I, Worklist, Visited);
8665 }
8666}
8667
8669 SmallVector<const Loop *, 16> LoopWorklist(1, L);
8672 SmallVector<SCEVUse, 16> ToForget;
8673
8674 // Iterate over all the loops and sub-loops to drop SCEV information.
8675 while (!LoopWorklist.empty()) {
8676 auto *CurrL = LoopWorklist.pop_back_val();
8677
8678 // Drop any stored trip count value.
8679 forgetBackedgeTakenCounts(CurrL, /* Predicated */ false);
8680 forgetBackedgeTakenCounts(CurrL, /* Predicated */ true);
8681
8682 // Drop information about predicated SCEV rewrites for this loop.
8683 PredicatedSCEVRewrites.remove_if(
8684 [&](const auto &Entry) { return Entry.first.second == CurrL; });
8685
8686 auto LoopUsersItr = LoopUsers.find(CurrL);
8687 if (LoopUsersItr != LoopUsers.end())
8688 llvm::append_range(ToForget, LoopUsersItr->second);
8689
8690 // Drop information about expressions based on loop-header PHIs.
8691 PushLoopPHIs(CurrL, Worklist, Visited);
8692 visitAndClearUsers(Worklist, Visited, ToForget);
8693
8694 LoopPropertiesCache.erase(CurrL);
8695 // Forget all contained loops too, to avoid dangling entries in the
8696 // ValuesAtScopes map.
8697 LoopWorklist.append(CurrL->begin(), CurrL->end());
8698 }
8699 forgetMemoizedResults(ToForget);
8700}
8701
8703 forgetLoop(L->getOutermostLoop());
8704}
8705
8708 if (!I) return;
8709
8710 // Drop information about expressions based on loop-header PHIs.
8713 SmallVector<SCEVUse, 8> ToForget;
8714 Worklist.push_back(I);
8715 Visited.insert(I);
8716 visitAndClearUsers(Worklist, Visited, ToForget);
8717
8718 forgetMemoizedResults(ToForget);
8719}
8720
8724 SmallVector<SCEVUse, 8> ToForget;
8725 for (Value *V : Values)
8726 if (auto *I = dyn_cast<Instruction>(V))
8727 if (Visited.insert(I).second)
8728 Worklist.push_back(I);
8729 visitAndClearUsers(Worklist, Visited, ToForget);
8730
8731 forgetMemoizedResults(ToForget);
8732}
8733
8735 // If SCEV looked through a trivial LCSSA phi node, we might have SCEV's
8736 // directly using a SCEVUnknown/SCEVAddRec defined in the loop. After an
8737 // extra predecessor is added, this is no longer valid. Find all Unknowns and
8738 // AddRecs defined in the loop and invalidate any SCEV's making use of them.
8739 auto InvalidateValue = [&](Value *Val) {
8740 if (!isSCEVable(Val->getType()))
8741 return;
8742 if (const SCEV *S = getExistingSCEV(Val)) {
8743 struct InvalidationRootCollector {
8744 Loop *L;
8746
8747 InvalidationRootCollector(Loop *L) : L(L) {}
8748
8749 bool follow(const SCEV *S) {
8750 if (auto *SU = dyn_cast<SCEVUnknown>(S)) {
8751 if (auto *I = dyn_cast<Instruction>(SU->getValue()))
8752 if (L->contains(I))
8753 Roots.push_back(S);
8754 } else if (auto *AddRec = dyn_cast<SCEVAddRecExpr>(S)) {
8755 if (L->contains(AddRec->getLoop()))
8756 Roots.push_back(S);
8757 }
8758 return true;
8759 }
8760 bool isDone() const { return false; }
8761 };
8762
8763 InvalidationRootCollector C(L);
8764 visitAll(S, C);
8765 forgetMemoizedResults(C.Roots);
8766 }
8767 };
8768
8769 InvalidateValue(V);
8770
8771 // If V has a non-SCEV-able type (e.g. {i64, i1} from a with.overflow
8772 // intrinsic), its users (e.g. extractvalue) may have stale SCEV
8773 // expressions referencing loop-internal values.
8774 if (!isSCEVable(V->getType()) &&
8775 any_of(V->incoming_values(), IsaPred<WithOverflowInst>))
8776 for (User *U : V->users())
8777 InvalidateValue(U);
8778 // Also perform the normal invalidation.
8779 forgetValue(V);
8780}
8781
8782void ScalarEvolution::forgetLoopDispositions() { LoopDispositions.clear(); }
8783
8785 // Unless a specific value is passed to invalidation, completely clear both
8786 // caches.
8787 if (!V) {
8788 BlockDispositions.clear();
8789 LoopDispositions.clear();
8790 return;
8791 }
8792
8793 if (!isSCEVable(V->getType()))
8794 return;
8795
8796 const SCEV *S = getExistingSCEV(V);
8797 if (!S)
8798 return;
8799
8800 // Invalidate the block and loop dispositions cached for S. Dispositions of
8801 // S's users may change if S's disposition changes (i.e. a user may change to
8802 // loop-invariant, if S changes to loop invariant), so also invalidate
8803 // dispositions of S's users recursively.
8804 SmallVector<SCEVUse, 8> Worklist = {S};
8806 while (!Worklist.empty()) {
8807 const SCEV *Curr = Worklist.pop_back_val();
8808 bool LoopDispoRemoved = LoopDispositions.erase(Curr);
8809 bool BlockDispoRemoved = BlockDispositions.erase(Curr);
8810 if (!LoopDispoRemoved && !BlockDispoRemoved)
8811 continue;
8812 auto Users = SCEVUsers.find(Curr);
8813 if (Users != SCEVUsers.end())
8814 for (const auto *User : Users->second)
8815 if (Seen.insert(User).second)
8816 Worklist.push_back(User);
8817 }
8818}
8819
8820/// Get the exact loop backedge taken count considering all loop exits. A
8821/// computable result can only be returned for loops with all exiting blocks
8822/// dominating the latch. howFarToZero assumes that the limit of each loop test
8823/// is never skipped. This is a valid assumption as long as the loop exits via
8824/// that test. For precise results, it is the caller's responsibility to specify
8825/// the relevant loop exiting block using getExact(ExitingBlock, SE).
8826const SCEV *ScalarEvolution::BackedgeTakenInfo::getExact(
8827 const Loop *L, ScalarEvolution *SE,
8829 // If any exits were not computable, the loop is not computable.
8830 if (!isComplete() || ExitNotTaken.empty())
8831 return SE->getCouldNotCompute();
8832
8833 const BasicBlock *Latch = L->getLoopLatch();
8834 // All exiting blocks we have collected must dominate the only backedge.
8835 if (!Latch)
8836 return SE->getCouldNotCompute();
8837
8838 // All exiting blocks we have gathered dominate loop's latch, so exact trip
8839 // count is simply a minimum out of all these calculated exit counts.
8841 for (const auto &ENT : ExitNotTaken) {
8842 const SCEV *BECount = ENT.ExactNotTaken;
8843 assert(BECount != SE->getCouldNotCompute() && "Bad exit SCEV!");
8844 assert(SE->DT.dominates(ENT.ExitingBlock, Latch) &&
8845 "We should only have known counts for exiting blocks that dominate "
8846 "latch!");
8847
8848 Ops.push_back(BECount);
8849
8850 if (Preds)
8851 append_range(*Preds, ENT.Predicates);
8852
8853 assert((Preds || ENT.hasAlwaysTruePredicate()) &&
8854 "Predicate should be always true!");
8855 }
8856
8857 // If an earlier exit exits on the first iteration (exit count zero), then
8858 // a later poison exit count should not propagate into the result. This are
8859 // exactly the semantics provided by umin_seq.
8860 return SE->getUMinFromMismatchedTypes(Ops, /* Sequential */ true);
8861}
8862
8863const ScalarEvolution::ExitNotTakenInfo *
8864ScalarEvolution::BackedgeTakenInfo::getExitNotTaken(
8865 const BasicBlock *ExitingBlock,
8866 SmallVectorImpl<const SCEVPredicate *> *Predicates) const {
8867 for (const auto &ENT : ExitNotTaken)
8868 if (ENT.ExitingBlock == ExitingBlock) {
8869 if (ENT.hasAlwaysTruePredicate())
8870 return &ENT;
8871 else if (Predicates) {
8872 append_range(*Predicates, ENT.Predicates);
8873 return &ENT;
8874 }
8875 }
8876
8877 return nullptr;
8878}
8879
8880/// getConstantMax - Get the constant max backedge taken count for the loop.
8881const SCEV *ScalarEvolution::BackedgeTakenInfo::getConstantMax(
8882 ScalarEvolution *SE,
8883 SmallVectorImpl<const SCEVPredicate *> *Predicates) const {
8884 if (!getConstantMax())
8885 return SE->getCouldNotCompute();
8886
8887 for (const auto &ENT : ExitNotTaken)
8888 if (!ENT.hasAlwaysTruePredicate()) {
8889 if (!Predicates)
8890 return SE->getCouldNotCompute();
8891 append_range(*Predicates, ENT.Predicates);
8892 }
8893
8894 assert((isa<SCEVCouldNotCompute>(getConstantMax()) ||
8895 isa<SCEVConstant>(getConstantMax())) &&
8896 "No point in having a non-constant max backedge taken count!");
8897 return getConstantMax();
8898}
8899
8900const SCEV *ScalarEvolution::BackedgeTakenInfo::getSymbolicMax(
8901 const Loop *L, ScalarEvolution *SE,
8902 SmallVectorImpl<const SCEVPredicate *> *Predicates) {
8903 if (!SymbolicMax) {
8904 // Form an expression for the maximum exit count possible for this loop. We
8905 // merge the max and exact information to approximate a version of
8906 // getConstantMaxBackedgeTakenCount which isn't restricted to just
8907 // constants.
8908 SmallVector<SCEVUse, 4> ExitCounts;
8909
8910 for (const auto &ENT : ExitNotTaken) {
8911 const SCEV *ExitCount = ENT.SymbolicMaxNotTaken;
8912 if (!isa<SCEVCouldNotCompute>(ExitCount)) {
8913 assert(SE->DT.dominates(ENT.ExitingBlock, L->getLoopLatch()) &&
8914 "We should only have known counts for exiting blocks that "
8915 "dominate latch!");
8916 ExitCounts.push_back(ExitCount);
8917 if (Predicates)
8918 append_range(*Predicates, ENT.Predicates);
8919
8920 assert((Predicates || ENT.hasAlwaysTruePredicate()) &&
8921 "Predicate should be always true!");
8922 }
8923 }
8924 if (ExitCounts.empty())
8925 SymbolicMax = SE->getCouldNotCompute();
8926 else
8927 SymbolicMax =
8928 SE->getUMinFromMismatchedTypes(ExitCounts, /*Sequential*/ true);
8929 }
8930 return SymbolicMax;
8931}
8932
8933bool ScalarEvolution::BackedgeTakenInfo::isConstantMaxOrZero(
8934 ScalarEvolution *SE) const {
8935 auto PredicateNotAlwaysTrue = [](const ExitNotTakenInfo &ENT) {
8936 return !ENT.hasAlwaysTruePredicate();
8937 };
8938 return MaxOrZero && !any_of(ExitNotTaken, PredicateNotAlwaysTrue);
8939}
8940
8943
8945 const SCEV *E, const SCEV *ConstantMaxNotTaken,
8946 const SCEV *SymbolicMaxNotTaken, bool MaxOrZero,
8950 // If we prove the max count is zero, so is the symbolic bound. This happens
8951 // in practice due to differences in a) how context sensitive we've chosen
8952 // to be and b) how we reason about bounds implied by UB.
8953 if (ConstantMaxNotTaken->isZero()) {
8954 this->ExactNotTaken = E = ConstantMaxNotTaken;
8955 this->SymbolicMaxNotTaken = SymbolicMaxNotTaken = ConstantMaxNotTaken;
8956 }
8957
8960 "Exact is not allowed to be less precise than Constant Max");
8963 "Exact is not allowed to be less precise than Symbolic Max");
8966 "Symbolic Max is not allowed to be less precise than Constant Max");
8969 "No point in having a non-constant max backedge taken count!");
8971 for (const auto PredList : PredLists)
8972 for (const auto *P : PredList) {
8973 if (SeenPreds.contains(P))
8974 continue;
8975 assert(!isa<SCEVUnionPredicate>(P) && "Only add leaf predicates here!");
8976 SeenPreds.insert(P);
8977 Predicates.push_back(P);
8978 }
8979 assert((isa<SCEVCouldNotCompute>(E) || !E->getType()->isPointerTy()) &&
8980 "Backedge count should be int");
8982 !ConstantMaxNotTaken->getType()->isPointerTy()) &&
8983 "Max backedge count should be int");
8984}
8985
8993
8994/// Allocate memory for BackedgeTakenInfo and copy the not-taken count of each
8995/// computable exit into a persistent ExitNotTakenInfo array.
8996ScalarEvolution::BackedgeTakenInfo::BackedgeTakenInfo(
8998 bool IsComplete, const SCEV *ConstantMax, bool MaxOrZero)
8999 : ConstantMax(ConstantMax), IsComplete(IsComplete), MaxOrZero(MaxOrZero) {
9000 using EdgeExitInfo = ScalarEvolution::BackedgeTakenInfo::EdgeExitInfo;
9001
9002 ExitNotTaken.reserve(ExitCounts.size());
9003 std::transform(ExitCounts.begin(), ExitCounts.end(),
9004 std::back_inserter(ExitNotTaken),
9005 [&](const EdgeExitInfo &EEI) {
9006 BasicBlock *ExitBB = EEI.first;
9007 const ExitLimit &EL = EEI.second;
9008 return ExitNotTakenInfo(ExitBB, EL.ExactNotTaken,
9009 EL.ConstantMaxNotTaken, EL.SymbolicMaxNotTaken,
9010 EL.Predicates);
9011 });
9012 assert((isa<SCEVCouldNotCompute>(ConstantMax) ||
9013 isa<SCEVConstant>(ConstantMax)) &&
9014 "No point in having a non-constant max backedge taken count!");
9015}
9016
9017/// Compute the number of times the backedge of the specified loop will execute.
9018ScalarEvolution::BackedgeTakenInfo
9019ScalarEvolution::computeBackedgeTakenCount(const Loop *L,
9020 bool AllowPredicates) {
9021 SmallVector<BasicBlock *, 8> ExitingBlocks;
9022 L->getExitingBlocks(ExitingBlocks);
9023
9024 using EdgeExitInfo = ScalarEvolution::BackedgeTakenInfo::EdgeExitInfo;
9025
9027 bool CouldComputeBECount = true;
9028 BasicBlock *Latch = L->getLoopLatch(); // may be NULL.
9029 const SCEV *MustExitMaxBECount = nullptr;
9030 const SCEV *MayExitMaxBECount = nullptr;
9031 bool MustExitMaxOrZero = false;
9032 bool IsOnlyExit = ExitingBlocks.size() == 1;
9033
9034 // Compute the ExitLimit for each loop exit. Use this to populate ExitCounts
9035 // and compute maxBECount.
9036 // Do a union of all the predicates here.
9037 for (BasicBlock *ExitBB : ExitingBlocks) {
9038 // We canonicalize untaken exits to br (constant), ignore them so that
9039 // proving an exit untaken doesn't negatively impact our ability to reason
9040 // about the loop as whole.
9041 if (auto *BI = dyn_cast<CondBrInst>(ExitBB->getTerminator()))
9042 if (auto *CI = dyn_cast<ConstantInt>(BI->getCondition())) {
9043 bool ExitIfTrue = !L->contains(BI->getSuccessor(0));
9044 if (ExitIfTrue == CI->isZero())
9045 continue;
9046 }
9047
9048 ExitLimit EL = computeExitLimit(L, ExitBB, IsOnlyExit, AllowPredicates);
9049
9050 assert((AllowPredicates || EL.Predicates.empty()) &&
9051 "Predicated exit limit when predicates are not allowed!");
9052
9053 // 1. For each exit that can be computed, add an entry to ExitCounts.
9054 // CouldComputeBECount is true only if all exits can be computed.
9055 if (EL.ExactNotTaken != getCouldNotCompute())
9056 ++NumExitCountsComputed;
9057 else
9058 // We couldn't compute an exact value for this exit, so
9059 // we won't be able to compute an exact value for the loop.
9060 CouldComputeBECount = false;
9061 // Remember exit count if either exact or symbolic is known. Because
9062 // Exact always implies symbolic, only check symbolic.
9063 if (EL.SymbolicMaxNotTaken != getCouldNotCompute())
9064 ExitCounts.emplace_back(ExitBB, EL);
9065 else {
9066 assert(EL.ExactNotTaken == getCouldNotCompute() &&
9067 "Exact is known but symbolic isn't?");
9068 ++NumExitCountsNotComputed;
9069 }
9070
9071 // 2. Derive the loop's MaxBECount from each exit's max number of
9072 // non-exiting iterations. Partition the loop exits into two kinds:
9073 // LoopMustExits and LoopMayExits.
9074 //
9075 // If the exit dominates the loop latch, it is a LoopMustExit otherwise it
9076 // is a LoopMayExit. If any computable LoopMustExit is found, then
9077 // MaxBECount is the minimum EL.ConstantMaxNotTaken of computable
9078 // LoopMustExits. Otherwise, MaxBECount is conservatively the maximum
9079 // EL.ConstantMaxNotTaken, where CouldNotCompute is considered greater than
9080 // any
9081 // computable EL.ConstantMaxNotTaken.
9082 if (EL.ConstantMaxNotTaken != getCouldNotCompute() && Latch &&
9083 DT.dominates(ExitBB, Latch)) {
9084 if (!MustExitMaxBECount) {
9085 MustExitMaxBECount = EL.ConstantMaxNotTaken;
9086 MustExitMaxOrZero = EL.MaxOrZero;
9087 } else {
9088 MustExitMaxBECount = getUMinFromMismatchedTypes(MustExitMaxBECount,
9089 EL.ConstantMaxNotTaken);
9090 }
9091 } else if (MayExitMaxBECount != getCouldNotCompute()) {
9092 if (!MayExitMaxBECount || EL.ConstantMaxNotTaken == getCouldNotCompute())
9093 MayExitMaxBECount = EL.ConstantMaxNotTaken;
9094 else {
9095 MayExitMaxBECount = getUMaxFromMismatchedTypes(MayExitMaxBECount,
9096 EL.ConstantMaxNotTaken);
9097 }
9098 }
9099 }
9100 const SCEV *MaxBECount = MustExitMaxBECount ? MustExitMaxBECount :
9101 (MayExitMaxBECount ? MayExitMaxBECount : getCouldNotCompute());
9102 // The loop backedge will be taken the maximum or zero times if there's
9103 // a single exit that must be taken the maximum or zero times.
9104 bool MaxOrZero = (MustExitMaxOrZero && ExitingBlocks.size() == 1);
9105
9106 // Remember which SCEVs are used in exit limits for invalidation purposes.
9107 // We only care about non-constant SCEVs here, so we can ignore
9108 // EL.ConstantMaxNotTaken
9109 // and MaxBECount, which must be SCEVConstant.
9110 for (const auto &Pair : ExitCounts) {
9111 if (!isa<SCEVConstant>(Pair.second.ExactNotTaken))
9112 BECountUsers[Pair.second.ExactNotTaken].insert({L, AllowPredicates});
9113 if (!isa<SCEVConstant>(Pair.second.SymbolicMaxNotTaken))
9114 BECountUsers[Pair.second.SymbolicMaxNotTaken].insert(
9115 {L, AllowPredicates});
9116 }
9117 return BackedgeTakenInfo(std::move(ExitCounts), CouldComputeBECount,
9118 MaxBECount, MaxOrZero);
9119}
9120
9121ScalarEvolution::ExitLimit
9122ScalarEvolution::computeExitLimit(const Loop *L, BasicBlock *ExitingBlock,
9123 bool IsOnlyExit, bool AllowPredicates) {
9124 assert(L->contains(ExitingBlock) && "Exit count for non-loop block?");
9125 // If our exiting block does not dominate the latch, then its connection with
9126 // loop's exit limit may be far from trivial.
9127 const BasicBlock *Latch = L->getLoopLatch();
9128 if (!Latch || !DT.dominates(ExitingBlock, Latch))
9129 return getCouldNotCompute();
9130
9131 Instruction *Term = ExitingBlock->getTerminator();
9132 if (CondBrInst *BI = dyn_cast<CondBrInst>(Term)) {
9133 bool ExitIfTrue = !L->contains(BI->getSuccessor(0));
9134 assert(ExitIfTrue == L->contains(BI->getSuccessor(1)) &&
9135 "It should have one successor in loop and one exit block!");
9136 // Proceed to the next level to examine the exit condition expression.
9137 return computeExitLimitFromCond(L, BI->getCondition(), ExitIfTrue,
9138 /*ControlsOnlyExit=*/IsOnlyExit,
9139 AllowPredicates);
9140 }
9141
9142 if (SwitchInst *SI = dyn_cast<SwitchInst>(Term)) {
9143 // For switch, make sure that there is a single exit from the loop.
9144 BasicBlock *Exit = nullptr;
9145 for (auto *SBB : successors(ExitingBlock))
9146 if (!L->contains(SBB)) {
9147 if (Exit) // Multiple exit successors.
9148 return getCouldNotCompute();
9149 Exit = SBB;
9150 }
9151 assert(Exit && "Exiting block must have at least one exit");
9152 return computeExitLimitFromSingleExitSwitch(
9153 L, SI, Exit, /*ControlsOnlyExit=*/IsOnlyExit);
9154 }
9155
9156 return getCouldNotCompute();
9157}
9158
9160 const Loop *L, Value *ExitCond, bool ExitIfTrue, bool ControlsOnlyExit,
9161 bool AllowPredicates) {
9162 ScalarEvolution::ExitLimitCacheTy Cache(L, ExitIfTrue, AllowPredicates);
9163 return computeExitLimitFromCondCached(Cache, L, ExitCond, ExitIfTrue,
9164 ControlsOnlyExit, AllowPredicates);
9165}
9166
9167std::optional<ScalarEvolution::ExitLimit>
9168ScalarEvolution::ExitLimitCache::find(const Loop *L, Value *ExitCond,
9169 bool ExitIfTrue, bool ControlsOnlyExit,
9170 bool AllowPredicates) {
9171 (void)this->L;
9172 (void)this->ExitIfTrue;
9173 (void)this->AllowPredicates;
9174
9175 assert(this->L == L && this->ExitIfTrue == ExitIfTrue &&
9176 this->AllowPredicates == AllowPredicates &&
9177 "Variance in assumed invariant key components!");
9178 auto Itr = TripCountMap.find({ExitCond, ControlsOnlyExit});
9179 if (Itr == TripCountMap.end())
9180 return std::nullopt;
9181 return Itr->second;
9182}
9183
9184void ScalarEvolution::ExitLimitCache::insert(const Loop *L, Value *ExitCond,
9185 bool ExitIfTrue,
9186 bool ControlsOnlyExit,
9187 bool AllowPredicates,
9188 const ExitLimit &EL) {
9189 assert(this->L == L && this->ExitIfTrue == ExitIfTrue &&
9190 this->AllowPredicates == AllowPredicates &&
9191 "Variance in assumed invariant key components!");
9192
9193 auto InsertResult = TripCountMap.insert({{ExitCond, ControlsOnlyExit}, EL});
9194 assert(InsertResult.second && "Expected successful insertion!");
9195 (void)InsertResult;
9196 (void)ExitIfTrue;
9197}
9198
9199ScalarEvolution::ExitLimit ScalarEvolution::computeExitLimitFromCondCached(
9200 ExitLimitCacheTy &Cache, const Loop *L, Value *ExitCond, bool ExitIfTrue,
9201 bool ControlsOnlyExit, bool AllowPredicates) {
9202
9203 if (auto MaybeEL = Cache.find(L, ExitCond, ExitIfTrue, ControlsOnlyExit,
9204 AllowPredicates))
9205 return *MaybeEL;
9206
9207 ExitLimit EL = computeExitLimitFromCondImpl(
9208 Cache, L, ExitCond, ExitIfTrue, ControlsOnlyExit, AllowPredicates);
9209 Cache.insert(L, ExitCond, ExitIfTrue, ControlsOnlyExit, AllowPredicates, EL);
9210 return EL;
9211}
9212
9213ScalarEvolution::ExitLimit ScalarEvolution::computeExitLimitFromCondImpl(
9214 ExitLimitCacheTy &Cache, const Loop *L, Value *ExitCond, bool ExitIfTrue,
9215 bool ControlsOnlyExit, bool AllowPredicates) {
9216 // Handle BinOp conditions (And, Or).
9217 if (auto LimitFromBinOp = computeExitLimitFromCondFromBinOp(
9218 Cache, L, ExitCond, ExitIfTrue, AllowPredicates))
9219 return *LimitFromBinOp;
9220
9221 // With an icmp, it may be feasible to compute an exact backedge-taken count.
9222 // Proceed to the next level to examine the icmp.
9223 if (ICmpInst *ExitCondICmp = dyn_cast<ICmpInst>(ExitCond)) {
9224 ExitLimit EL =
9225 computeExitLimitFromICmp(L, ExitCondICmp, ExitIfTrue, ControlsOnlyExit);
9226 if (EL.hasFullInfo() || !AllowPredicates)
9227 return EL;
9228
9229 // Try again, but use SCEV predicates this time.
9230 return computeExitLimitFromICmp(L, ExitCondICmp, ExitIfTrue,
9231 ControlsOnlyExit,
9232 /*AllowPredicates=*/true);
9233 }
9234
9235 // Check for a constant condition. These are normally stripped out by
9236 // SimplifyCFG, but ScalarEvolution may be used by a pass which wishes to
9237 // preserve the CFG and is temporarily leaving constant conditions
9238 // in place.
9239 if (ConstantInt *CI = dyn_cast<ConstantInt>(ExitCond)) {
9240 if (ExitIfTrue == !CI->getZExtValue())
9241 // The backedge is always taken.
9242 return getCouldNotCompute();
9243 // The backedge is never taken.
9244 return getZero(CI->getType());
9245 }
9246
9247 // If we're exiting based on the overflow flag of an x.with.overflow intrinsic
9248 // with a constant step, we can form an equivalent icmp predicate and figure
9249 // out how many iterations will be taken before we exit.
9250 const WithOverflowInst *WO;
9251 const APInt *C;
9252 if (match(ExitCond, m_ExtractValue<1>(m_WithOverflowInst(WO))) &&
9253 match(WO->getRHS(), m_APInt(C))) {
9254 ConstantRange NWR =
9256 WO->getNoWrapKind());
9257 CmpInst::Predicate Pred;
9258 APInt NewRHSC, Offset;
9259 NWR.getEquivalentICmp(Pred, NewRHSC, Offset);
9260 if (!ExitIfTrue)
9261 Pred = ICmpInst::getInversePredicate(Pred);
9262 auto *LHS = getSCEV(WO->getLHS());
9263 if (Offset != 0)
9265 auto EL = computeExitLimitFromICmp(L, Pred, LHS, getConstant(NewRHSC),
9266 ControlsOnlyExit, AllowPredicates);
9267 if (EL.hasAnyInfo())
9268 return EL;
9269 }
9270
9271 // If it's not an integer or pointer comparison then compute it the hard way.
9272 return computeExitCountExhaustively(L, ExitCond, ExitIfTrue);
9273}
9274
9275std::optional<ScalarEvolution::ExitLimit>
9276ScalarEvolution::computeExitLimitFromCondFromBinOp(ExitLimitCacheTy &Cache,
9277 const Loop *L,
9278 Value *ExitCond,
9279 bool ExitIfTrue,
9280 bool AllowPredicates) {
9281 // Check if the controlling expression for this loop is an And or Or.
9282 Value *Op0, *Op1;
9283 bool IsAnd;
9284 if (match(ExitCond, m_LogicalAnd(m_Value(Op0), m_Value(Op1))))
9285 IsAnd = true;
9286 else if (match(ExitCond, m_LogicalOr(m_Value(Op0), m_Value(Op1))))
9287 IsAnd = false;
9288 else
9289 return std::nullopt;
9290
9291 // A sub-condition of a non-trivial binop never solely controls the exit,
9292 // whether we exit always depends on both conditions.
9293 ExitLimit EL0 = computeExitLimitFromCondCached(
9294 Cache, L, Op0, ExitIfTrue, /*ControlsOnlyExit=*/false, AllowPredicates);
9295 ExitLimit EL1 = computeExitLimitFromCondCached(
9296 Cache, L, Op1, ExitIfTrue, /*ControlsOnlyExit=*/false, AllowPredicates);
9297
9298 // EitherMayExit is true in these two cases:
9299 // br (and Op0 Op1), loop, exit
9300 // br (or Op0 Op1), exit, loop
9301 bool EitherMayExit = IsAnd ^ ExitIfTrue;
9302
9303 const SCEV *BECount = getCouldNotCompute();
9304 const SCEV *ConstantMaxBECount = getCouldNotCompute();
9305 const SCEV *SymbolicMaxBECount = getCouldNotCompute();
9306 if (EitherMayExit) {
9307 bool UseSequentialUMin = !isa<BinaryOperator>(ExitCond);
9308 // Both conditions must be same for the loop to continue executing.
9309 // Choose the less conservative count.
9310 if (EL0.ExactNotTaken != getCouldNotCompute() &&
9311 EL1.ExactNotTaken != getCouldNotCompute()) {
9312 BECount = getUMinFromMismatchedTypes(EL0.ExactNotTaken, EL1.ExactNotTaken,
9313 UseSequentialUMin);
9314 }
9315 if (EL0.ConstantMaxNotTaken == getCouldNotCompute())
9316 ConstantMaxBECount = EL1.ConstantMaxNotTaken;
9317 else if (EL1.ConstantMaxNotTaken == getCouldNotCompute())
9318 ConstantMaxBECount = EL0.ConstantMaxNotTaken;
9319 else
9320 ConstantMaxBECount = getUMinFromMismatchedTypes(EL0.ConstantMaxNotTaken,
9321 EL1.ConstantMaxNotTaken);
9322 if (EL0.SymbolicMaxNotTaken == getCouldNotCompute())
9323 SymbolicMaxBECount = EL1.SymbolicMaxNotTaken;
9324 else if (EL1.SymbolicMaxNotTaken == getCouldNotCompute())
9325 SymbolicMaxBECount = EL0.SymbolicMaxNotTaken;
9326 else
9327 SymbolicMaxBECount = getUMinFromMismatchedTypes(
9328 EL0.SymbolicMaxNotTaken, EL1.SymbolicMaxNotTaken, UseSequentialUMin);
9329 } else {
9330 // Both conditions must be same at the same time for the loop to exit.
9331 // For now, be conservative.
9332 if (EL0.ExactNotTaken == EL1.ExactNotTaken)
9333 BECount = EL0.ExactNotTaken;
9334 }
9335
9336 // There are cases (e.g. PR26207) where computeExitLimitFromCond is able
9337 // to be more aggressive when computing BECount than when computing
9338 // ConstantMaxBECount. In these cases it is possible for EL0.ExactNotTaken
9339 // and
9340 // EL1.ExactNotTaken to match, but for EL0.ConstantMaxNotTaken and
9341 // EL1.ConstantMaxNotTaken to not.
9342 if (isa<SCEVCouldNotCompute>(ConstantMaxBECount) &&
9343 !isa<SCEVCouldNotCompute>(BECount))
9344 ConstantMaxBECount = getConstant(getUnsignedRangeMax(BECount));
9345 if (isa<SCEVCouldNotCompute>(SymbolicMaxBECount))
9346 SymbolicMaxBECount =
9347 isa<SCEVCouldNotCompute>(BECount) ? ConstantMaxBECount : BECount;
9348 return ExitLimit(BECount, ConstantMaxBECount, SymbolicMaxBECount, false,
9349 {ArrayRef(EL0.Predicates), ArrayRef(EL1.Predicates)});
9350}
9351
9352ScalarEvolution::ExitLimit ScalarEvolution::computeExitLimitFromICmp(
9353 const Loop *L, ICmpInst *ExitCond, bool ExitIfTrue, bool ControlsOnlyExit,
9354 bool AllowPredicates) {
9355 // If the condition was exit on true, convert the condition to exit on false
9356 CmpPredicate Pred;
9357 if (!ExitIfTrue)
9358 Pred = ExitCond->getCmpPredicate();
9359 else
9360 Pred = ExitCond->getInverseCmpPredicate();
9361 const ICmpInst::Predicate OriginalPred = Pred;
9362
9363 const SCEV *LHS = getSCEV(ExitCond->getOperand(0));
9364 const SCEV *RHS = getSCEV(ExitCond->getOperand(1));
9365
9366 ExitLimit EL = computeExitLimitFromICmp(L, Pred, LHS, RHS, ControlsOnlyExit,
9367 AllowPredicates);
9368 if (EL.hasAnyInfo())
9369 return EL;
9370
9371 auto *ExhaustiveCount =
9372 computeExitCountExhaustively(L, ExitCond, ExitIfTrue);
9373
9374 if (!isa<SCEVCouldNotCompute>(ExhaustiveCount))
9375 return ExhaustiveCount;
9376
9377 return computeShiftCompareExitLimit(ExitCond->getOperand(0),
9378 ExitCond->getOperand(1), L, OriginalPred);
9379}
9380ScalarEvolution::ExitLimit ScalarEvolution::computeExitLimitFromICmp(
9381 const Loop *L, CmpPredicate Pred, SCEVUse LHS, SCEVUse RHS,
9382 bool ControlsOnlyExit, bool AllowPredicates) {
9383
9384 // Try to evaluate any dependencies out of the loop.
9385 LHS = getSCEVAtScope(LHS, L);
9386 RHS = getSCEVAtScope(RHS, L);
9387
9388 // At this point, we would like to compute how many iterations of the
9389 // loop the predicate will return true for these inputs.
9390 if (isLoopInvariant(LHS, L) && !isLoopInvariant(RHS, L)) {
9391 // If there is a loop-invariant, force it into the RHS.
9392 std::swap(LHS, RHS);
9394 }
9395
9396 bool ControllingFiniteLoop = ControlsOnlyExit && loopHasNoAbnormalExits(L) &&
9398 // Simplify the operands before analyzing them.
9399 (void)SimplifyICmpOperands(Pred, LHS, RHS, /*Depth=*/0);
9400
9401 // If we have a comparison of a chrec against a constant, try to use value
9402 // ranges to answer this query.
9403 if (const SCEVConstant *RHSC = dyn_cast<SCEVConstant>(RHS))
9404 if (const SCEVAddRecExpr *AddRec = dyn_cast<SCEVAddRecExpr>(LHS))
9405 if (AddRec->getLoop() == L) {
9406 // Form the constant range.
9407 ConstantRange CompRange =
9408 ConstantRange::makeExactICmpRegion(Pred, RHSC->getAPInt());
9409
9410 const SCEV *Ret = AddRec->getNumIterationsInRange(CompRange, *this);
9411 if (!isa<SCEVCouldNotCompute>(Ret)) return Ret;
9412 }
9413
9414 // If this loop must exit based on this condition (or execute undefined
9415 // behaviour), see if we can improve wrap flags. This is essentially
9416 // a must execute style proof.
9417 if (ControllingFiniteLoop && isLoopInvariant(RHS, L)) {
9418 // If we can prove the test sequence produced must repeat the same values
9419 // on self-wrap of the IV, then we can infer that IV doesn't self wrap
9420 // because if it did, we'd have an infinite (undefined) loop.
9421 // TODO: We can peel off any functions which are invertible *in L*. Loop
9422 // invariant terms are effectively constants for our purposes here.
9423 SCEVUse InnerLHS = LHS;
9424 if (auto *ZExt = dyn_cast<SCEVZeroExtendExpr>(LHS))
9425 InnerLHS = ZExt->getOperand();
9426 if (const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(InnerLHS);
9427 AR && !AR->hasNoSelfWrap() && AR->getLoop() == L && AR->isAffine() &&
9428 isKnownToBeAPowerOfTwo(AR->getStepRecurrence(*this), /*OrZero=*/true,
9429 /*OrNegative=*/true)) {
9430 auto Flags = AR->getNoWrapFlags();
9431 Flags = setFlags(Flags, SCEV::FlagNW);
9434 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), Flags);
9435 }
9436
9437 // For a slt/ult condition with a positive step, can we prove nsw/nuw?
9438 // From no-self-wrap, this follows trivially from the fact that every
9439 // (un)signed-wrapped, but not self-wrapped value must be LT than the
9440 // last value before (un)signed wrap. Since we know that last value
9441 // didn't exit, nor will any smaller one.
9442 if (Pred == ICmpInst::ICMP_SLT || Pred == ICmpInst::ICMP_ULT) {
9443 auto WrapType = Pred == ICmpInst::ICMP_SLT ? SCEV::FlagNSW : SCEV::FlagNUW;
9444 if (const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(LHS);
9445 AR && AR->getLoop() == L && AR->isAffine() &&
9446 !AR->getNoWrapFlags(WrapType) && AR->hasNoSelfWrap() &&
9447 isKnownPositive(AR->getStepRecurrence(*this))) {
9448 auto Flags = AR->getNoWrapFlags();
9449 Flags = setFlags(Flags, WrapType);
9452 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), Flags);
9453 }
9454 }
9455 }
9456
9457 switch (Pred) {
9458 case ICmpInst::ICMP_NE: { // while (X != Y)
9459 // Convert to: while (X-Y != 0)
9460 if (LHS->getType()->isPointerTy()) {
9463 return LHS;
9464 }
9465 if (RHS->getType()->isPointerTy()) {
9468 return RHS;
9469 }
9470 ExitLimit EL = howFarToZero(getMinusSCEV(LHS, RHS), L, ControlsOnlyExit,
9471 AllowPredicates);
9472 if (EL.hasAnyInfo())
9473 return EL;
9474 break;
9475 }
9476 case ICmpInst::ICMP_EQ: { // while (X == Y)
9477 // Convert to: while (X-Y == 0)
9478 if (LHS->getType()->isPointerTy()) {
9481 return LHS;
9482 }
9483 if (RHS->getType()->isPointerTy()) {
9486 return RHS;
9487 }
9488 ExitLimit EL = howFarToNonZero(getMinusSCEV(LHS, RHS), L);
9489 if (EL.hasAnyInfo()) return EL;
9490 break;
9491 }
9492 case ICmpInst::ICMP_SLE:
9493 case ICmpInst::ICMP_ULE:
9494 // Since the loop is finite, an invariant RHS cannot include the boundary
9495 // value, otherwise it would loop forever.
9496 if (!EnableFiniteLoopControl || !ControllingFiniteLoop ||
9497 !isLoopInvariant(RHS, L)) {
9498 // Otherwise, perform the addition in a wider type, to avoid overflow.
9499 // If the LHS is an addrec with the appropriate nowrap flag, the
9500 // extension will be sunk into it and the exit count can be analyzed.
9501 auto *OldType = dyn_cast<IntegerType>(LHS->getType());
9502 if (!OldType)
9503 break;
9504 // Prefer doubling the bitwidth over adding a single bit to make it more
9505 // likely that we use a legal type.
9506 auto *NewType =
9507 Type::getIntNTy(OldType->getContext(), OldType->getBitWidth() * 2);
9508 if (ICmpInst::isSigned(Pred)) {
9509 LHS = getSignExtendExpr(LHS, NewType);
9510 RHS = getSignExtendExpr(RHS, NewType);
9511 } else {
9512 LHS = getZeroExtendExpr(LHS, NewType);
9513 RHS = getZeroExtendExpr(RHS, NewType);
9514 }
9515 }
9517 [[fallthrough]];
9518 case ICmpInst::ICMP_SLT:
9519 case ICmpInst::ICMP_ULT: { // while (X < Y)
9520 bool IsSigned = ICmpInst::isSigned(Pred);
9521 ExitLimit EL = howManyLessThans(LHS, RHS, L, IsSigned, /*Invert=*/false,
9522 ControlsOnlyExit, AllowPredicates);
9523 if (EL.hasAnyInfo())
9524 return EL;
9525 break;
9526 }
9527 case ICmpInst::ICMP_SGE:
9528 case ICmpInst::ICMP_UGE:
9529 // Since the loop is finite, an invariant RHS cannot include the boundary
9530 // value, otherwise it would loop forever.
9531 if (!EnableFiniteLoopControl || !ControllingFiniteLoop ||
9532 !isLoopInvariant(RHS, L))
9533 break;
9535 [[fallthrough]];
9536 case ICmpInst::ICMP_SGT:
9537 case ICmpInst::ICMP_UGT: { // while (X > Y)
9538 // "X > Y" is analyzed as the equivalent "~X < ~Y".
9539 bool IsSigned = ICmpInst::isSigned(Pred);
9540 ExitLimit EL = howManyLessThans(LHS, RHS, L, IsSigned, /*Invert=*/true,
9541 ControlsOnlyExit, AllowPredicates);
9542 if (EL.hasAnyInfo())
9543 return EL;
9544 break;
9545 }
9546 default:
9547 break;
9548 }
9549
9550 return getCouldNotCompute();
9551}
9552
9553ScalarEvolution::ExitLimit
9554ScalarEvolution::computeExitLimitFromSingleExitSwitch(const Loop *L,
9555 SwitchInst *Switch,
9556 BasicBlock *ExitingBlock,
9557 bool ControlsOnlyExit) {
9558 assert(!L->contains(ExitingBlock) && "Not an exiting block!");
9559
9560 // Give up if the exit is the default dest of a switch.
9561 if (Switch->getDefaultDest() == ExitingBlock)
9562 return getCouldNotCompute();
9563
9564 assert(L->contains(Switch->getDefaultDest()) &&
9565 "Default case must not exit the loop!");
9566 const SCEV *LHS = getSCEVAtScope(Switch->getCondition(), L);
9567 const SCEV *RHS = getConstant(Switch->findCaseDest(ExitingBlock));
9568
9569 // while (X != Y) --> while (X-Y != 0)
9570 ExitLimit EL = howFarToZero(getMinusSCEV(LHS, RHS), L, ControlsOnlyExit);
9571 if (EL.hasAnyInfo())
9572 return EL;
9573
9574 return getCouldNotCompute();
9575}
9576
9577static ConstantInt *
9579 ScalarEvolution &SE) {
9580 const SCEV *InVal = SE.getConstant(C);
9581 const SCEV *Val = AddRec->evaluateAtIteration(InVal, SE);
9583 "Evaluation of SCEV at constant didn't fold correctly?");
9584 return cast<SCEVConstant>(Val)->getValue();
9585}
9586
9587ScalarEvolution::ExitLimit ScalarEvolution::computeShiftCompareExitLimit(
9588 Value *LHS, Value *RHSV, const Loop *L, ICmpInst::Predicate Pred) {
9589 ConstantInt *RHS = dyn_cast<ConstantInt>(RHSV);
9590 if (!RHS)
9591 return getCouldNotCompute();
9592
9593 const BasicBlock *Latch = L->getLoopLatch();
9594 if (!Latch)
9595 return getCouldNotCompute();
9596
9597 const BasicBlock *Predecessor = L->getLoopPredecessor();
9598 if (!Predecessor)
9599 return getCouldNotCompute();
9600
9601 // Return true if V is of the form "LHS `shift_op` <positive constant>".
9602 // Return LHS in OutLHS, shift_op in OutOpCode, and the shift amount in
9603 // OutShiftAmt.
9604 auto MatchPositiveShift = [](Value *V, Value *&OutLHS,
9605 Instruction::BinaryOps &OutOpCode,
9606 unsigned &OutShiftAmt) {
9607 using namespace PatternMatch;
9608
9609 ConstantInt *ShiftAmt;
9610 if (match(V, m_LShr(m_Value(OutLHS), m_ConstantInt(ShiftAmt))))
9611 OutOpCode = Instruction::LShr;
9612 else if (match(V, m_AShr(m_Value(OutLHS), m_ConstantInt(ShiftAmt))))
9613 OutOpCode = Instruction::AShr;
9614 else if (match(V, m_Shl(m_Value(OutLHS), m_ConstantInt(ShiftAmt))))
9615 OutOpCode = Instruction::Shl;
9616 else
9617 return false;
9618
9619 uint64_t Amt = ShiftAmt->getValue().getLimitedValue();
9620 if (Amt == 0 || Amt >= OutLHS->getType()->getScalarSizeInBits())
9621 return false;
9622 OutShiftAmt = Amt;
9623 return true;
9624 };
9625
9626 // Recognize a "shift recurrence" either of the form %iv or of %iv.shifted in
9627 //
9628 // loop:
9629 // %iv = phi i32 [ %iv.shifted, %loop ], [ %val, %preheader ]
9630 // %iv.shifted = lshr i32 %iv, <positive constant>
9631 //
9632 // Return true on a successful match. Return the corresponding PHI node (%iv
9633 // above) in PNOut, the opcode of the shift operation in OpCodeOut, and the
9634 // shift amount in ShiftAmtOut.
9635 auto MatchShiftRecurrence = [&](Value *V, PHINode *&PNOut,
9636 Instruction::BinaryOps &OpCodeOut,
9637 unsigned &ShiftAmtOut) {
9638 std::optional<Instruction::BinaryOps> PostShiftOpCode;
9639
9640 {
9642 Value *V;
9643 unsigned Amt;
9644
9645 // If we encounter a shift instruction, "peel off" the shift operation,
9646 // and remember that we did so. Later when we inspect %iv's backedge
9647 // value, we will make sure that the backedge value uses the same
9648 // operation.
9649 //
9650 // Note: the peeled shift operation does not have to be the same
9651 // instruction as the one feeding into the PHI's backedge value. We only
9652 // really care about it being the same *kind* of shift instruction --
9653 // that's all that is required for our later inferences to hold.
9654 if (MatchPositiveShift(LHS, V, OpC, Amt)) {
9655 PostShiftOpCode = OpC;
9656 LHS = V;
9657 }
9658 }
9659
9660 PNOut = dyn_cast<PHINode>(LHS);
9661 if (!PNOut || PNOut->getParent() != L->getHeader())
9662 return false;
9663
9664 Value *BEValue = PNOut->getIncomingValueForBlock(Latch);
9665 Value *OpLHS;
9666
9667 return
9668 // The backedge value for the PHI node must be a shift by a positive
9669 // amount
9670 MatchPositiveShift(BEValue, OpLHS, OpCodeOut, ShiftAmtOut) &&
9671
9672 // of the PHI node itself
9673 OpLHS == PNOut &&
9674
9675 // and the kind of shift should be match the kind of shift we peeled
9676 // off, if any.
9677 (!PostShiftOpCode || *PostShiftOpCode == OpCodeOut);
9678 };
9679
9680 PHINode *PN;
9682 unsigned ShiftAmt;
9683 if (!MatchShiftRecurrence(LHS, PN, OpCode, ShiftAmt))
9684 return getCouldNotCompute();
9685
9686 const DataLayout &DL = getDataLayout();
9687
9688 // The key rationale for this optimization is that for some kinds of shift
9689 // recurrences, the value of the recurrence "stabilizes" to either 0 or -1
9690 // within a finite number of iterations. If the condition guarding the
9691 // backedge (in the sense that the backedge is taken if the condition is true)
9692 // is false for the value the shift recurrence stabilizes to, then we know
9693 // that the backedge is taken only a finite number of times.
9694
9695 ConstantInt *StableValue = nullptr;
9696 switch (OpCode) {
9697 default:
9698 llvm_unreachable("Impossible case!");
9699
9700 case Instruction::AShr: {
9701 // {K,ashr,<positive-constant>} stabilizes to signum(K) in at most
9702 // bitwidth(K) iterations.
9703 Value *FirstValue = PN->getIncomingValueForBlock(Predecessor);
9704 KnownBits Known = computeKnownBits(FirstValue, DL, &AC,
9705 Predecessor->getTerminator(), &DT);
9706 auto *Ty = cast<IntegerType>(RHS->getType());
9707 if (Known.isNonNegative())
9708 StableValue = ConstantInt::get(Ty, 0);
9709 else if (Known.isNegative())
9710 StableValue = ConstantInt::get(Ty, -1, true);
9711 else
9712 return getCouldNotCompute();
9713
9714 break;
9715 }
9716 case Instruction::LShr:
9717 case Instruction::Shl:
9718 // Both {K,lshr,<positive-constant>} and {K,shl,<positive-constant>}
9719 // stabilize to 0 in at most bitwidth(K) iterations.
9720 StableValue = ConstantInt::get(cast<IntegerType>(RHS->getType()), 0);
9721 break;
9722 }
9723
9724 auto *Result =
9725 ConstantFoldCompareInstOperands(Pred, StableValue, RHS, DL, &TLI);
9726 assert(Result->getType()->isIntegerTy(1) &&
9727 "Otherwise cannot be an operand to a branch instruction");
9728
9729 if (Result->isNullValue()) {
9730 unsigned BitWidth = getTypeSizeInBits(RHS->getType());
9731 unsigned MaxBTC = BitWidth;
9732
9733 // For right-shift recurrences (lshr/ashr with non-negative start), we can
9734 // compute a tighter max backedge-taken count from the range of the start
9735 // value. After k shifts of ShiftAmt, value = start >> (k * ShiftAmt).
9736 // The value reaches 0 (the stable value) when k * ShiftAmt >=
9737 // activeBits(start), so max BTC = ceil(activeBits(maxStart) / ShiftAmt).
9738 if (OpCode == Instruction::LShr || OpCode == Instruction::AShr) {
9739 Value *StartValue = PN->getIncomingValueForBlock(Predecessor);
9740 const SCEV *StartSCEV = getSCEV(StartValue);
9741 APInt MaxStart = getUnsignedRangeMax(StartSCEV);
9742 if (MaxStart.isStrictlyPositive()) {
9743 unsigned ActiveBits = MaxStart.getActiveBits();
9744 unsigned RangeBTC = divideCeil(ActiveBits, ShiftAmt);
9745 MaxBTC = std::min(MaxBTC, RangeBTC);
9746 }
9747 }
9748
9749 const SCEV *UpperBound =
9751 return ExitLimit(getCouldNotCompute(), UpperBound, UpperBound, false);
9752 }
9753
9754 return getCouldNotCompute();
9755}
9756
9757/// Return true if we can constant fold an instruction of the specified type,
9758/// assuming that all operands were constants.
9759static bool canConstantFold(const Instruction *I,
9760 const TargetLibraryInfo *TLI) {
9764 return true;
9765
9766 if (const CallInst *CI = dyn_cast<CallInst>(I))
9767 if (const Function *F = CI->getCalledFunction())
9768 return canConstantFoldCallTo(CI, F, TLI);
9769 return false;
9770}
9771
9772/// Determine whether this instruction can constant evolve within this loop
9773/// assuming its operands can all constant evolve.
9774static bool canConstantEvolve(Instruction *I, const Loop *L,
9775 const TargetLibraryInfo *TLI) {
9776 // An instruction outside of the loop can't be derived from a loop PHI.
9777 if (!L->contains(I)) return false;
9778
9779 if (isa<PHINode>(I)) {
9780 // We don't currently keep track of the control flow needed to evaluate
9781 // PHIs, so we cannot handle PHIs inside of loops.
9782 return L->getHeader() == I->getParent();
9783 }
9784
9785 // If we won't be able to constant fold this expression even if the operands
9786 // are constants, bail early.
9787 return canConstantFold(I, TLI);
9788}
9789
9790/// getConstantEvolvingPHIOperands - Implement getConstantEvolvingPHI by
9791/// recursing through each instruction operand until reaching a loop header phi.
9792static PHINode *
9795 const TargetLibraryInfo *TLI, unsigned Depth) {
9797 return nullptr;
9798
9799 // Otherwise, we can evaluate this instruction if all of its operands are
9800 // constant or derived from a PHI node themselves.
9801 PHINode *PHI = nullptr;
9802 for (Value *Op : UseInst->operands()) {
9803 if (isa<Constant>(Op)) continue;
9804
9806 if (!OpInst || !canConstantEvolve(OpInst, L, TLI))
9807 return nullptr;
9808
9809 PHINode *P = dyn_cast<PHINode>(OpInst);
9810 if (!P)
9811 // If this operand is already visited, reuse the prior result.
9812 // We may have P != PHI if this is the deepest point at which the
9813 // inconsistent paths meet.
9814 P = PHIMap.lookup(OpInst);
9815 if (!P) {
9816 // Recurse and memoize the results, whether a phi is found or not.
9817 // This recursive call invalidates pointers into PHIMap.
9818 P = getConstantEvolvingPHIOperands(OpInst, L, PHIMap, TLI, Depth + 1);
9819 PHIMap[OpInst] = P;
9820 }
9821 if (!P)
9822 return nullptr; // Not evolving from PHI
9823 if (PHI && PHI != P)
9824 return nullptr; // Evolving from multiple different PHIs.
9825 PHI = P;
9826 }
9827 // This is a expression evolving from a constant PHI!
9828 return PHI;
9829}
9830
9831/// getConstantEvolvingPHI - Given an LLVM value and a loop, return a PHI node
9832/// in the loop that V is derived from. We allow arbitrary operations along the
9833/// way, but the operands of an operation must either be constants or a value
9834/// derived from a constant PHI. If this expression does not fit with these
9835/// constraints, return null.
9837 const TargetLibraryInfo *TLI) {
9839 if (!I || !canConstantEvolve(I, L, TLI))
9840 return nullptr;
9841
9842 if (PHINode *PN = dyn_cast<PHINode>(I))
9843 return PN;
9844
9845 // Record non-constant instructions contained by the loop.
9847 return getConstantEvolvingPHIOperands(I, L, PHIMap, TLI, 0);
9848}
9849
9850/// EvaluateExpression - Given an expression that passes the
9851/// getConstantEvolvingPHI predicate, evaluate its value assuming the PHI node
9852/// in the loop has the value PHIVal. If we can't fold this expression for some
9853/// reason, return null.
9856 const DataLayout &DL,
9857 const TargetLibraryInfo *TLI) {
9858 // Convenient constant check, but redundant for recursive calls.
9859 if (Constant *C = dyn_cast<Constant>(V)) return C;
9861 if (!I) return nullptr;
9862
9863 if (Constant *C = Vals.lookup(I)) return C;
9864
9865 // An instruction inside the loop depends on a value outside the loop that we
9866 // weren't given a mapping for, or a value such as a call inside the loop.
9867 if (!canConstantEvolve(I, L, TLI))
9868 return nullptr;
9869
9870 // An unmapped PHI can be due to a branch or another loop inside this loop,
9871 // or due to this not being the initial iteration through a loop where we
9872 // couldn't compute the evolution of this particular PHI last time.
9873 if (isa<PHINode>(I)) return nullptr;
9874
9875 std::vector<Constant*> Operands(I->getNumOperands());
9876
9877 for (unsigned i = 0, e = I->getNumOperands(); i != e; ++i) {
9878 Instruction *Operand = dyn_cast<Instruction>(I->getOperand(i));
9879 if (!Operand) {
9880 Operands[i] = dyn_cast<Constant>(I->getOperand(i));
9881 if (!Operands[i]) return nullptr;
9882 continue;
9883 }
9884 Constant *C = EvaluateExpression(Operand, L, Vals, DL, TLI);
9885 Vals[Operand] = C;
9886 if (!C) return nullptr;
9887 Operands[i] = C;
9888 }
9889
9890 return ConstantFoldInstOperands(I, Operands, DL, TLI,
9891 /*AllowNonDeterministic=*/false);
9892}
9893
9894
9895// If every incoming value to PN except the one for BB is a specific Constant,
9896// return that, else return nullptr.
9898 Constant *IncomingVal = nullptr;
9899
9900 for (unsigned i = 0, e = PN->getNumIncomingValues(); i != e; ++i) {
9901 if (PN->getIncomingBlock(i) == BB)
9902 continue;
9903
9904 auto *CurrentVal = dyn_cast<Constant>(PN->getIncomingValue(i));
9905 if (!CurrentVal)
9906 return nullptr;
9907
9908 if (IncomingVal != CurrentVal) {
9909 if (IncomingVal)
9910 return nullptr;
9911 IncomingVal = CurrentVal;
9912 }
9913 }
9914
9915 return IncomingVal;
9916}
9917
9918/// getConstantEvolutionLoopExitValue - If we know that the specified Phi is
9919/// in the header of its containing loop, we know the loop executes a
9920/// constant number of times, and the PHI node is just a recurrence
9921/// involving constants, fold it.
9922Constant *
9923ScalarEvolution::getConstantEvolutionLoopExitValue(PHINode *PN,
9924 const APInt &BEs,
9925 const Loop *L) {
9926 auto [I, Inserted] = ConstantEvolutionLoopExitValue.try_emplace(PN);
9927 if (!Inserted)
9928 return I->second;
9929
9931 return nullptr; // Not going to evaluate it.
9932
9933 Constant *&RetVal = I->second;
9934
9935 DenseMap<Instruction *, Constant *> CurrentIterVals;
9936 BasicBlock *Header = L->getHeader();
9937 assert(PN->getParent() == Header && "Can't evaluate PHI not in loop header!");
9938
9939 BasicBlock *Latch = L->getLoopLatch();
9940 if (!Latch)
9941 return nullptr;
9942
9943 for (PHINode &PHI : Header->phis()) {
9944 if (auto *StartCST = getOtherIncomingValue(&PHI, Latch))
9945 CurrentIterVals[&PHI] = StartCST;
9946 }
9947 if (!CurrentIterVals.count(PN))
9948 return RetVal = nullptr;
9949
9950 Value *BEValue = PN->getIncomingValueForBlock(Latch);
9951
9952 // Execute the loop symbolically to determine the exit value.
9953 assert(BEs.getActiveBits() < CHAR_BIT * sizeof(unsigned) &&
9954 "BEs is <= MaxBruteForceIterations which is an 'unsigned'!");
9955
9956 unsigned NumIterations = BEs.getZExtValue(); // must be in range
9957 unsigned IterationNum = 0;
9958 const DataLayout &DL = getDataLayout();
9959 for (; ; ++IterationNum) {
9960 if (IterationNum == NumIterations)
9961 return RetVal = CurrentIterVals[PN]; // Got exit value!
9962
9963 // Compute the value of the PHIs for the next iteration.
9964 // EvaluateExpression adds non-phi values to the CurrentIterVals map.
9965 DenseMap<Instruction *, Constant *> NextIterVals;
9966 Constant *NextPHI =
9967 EvaluateExpression(BEValue, L, CurrentIterVals, DL, &TLI);
9968 if (!NextPHI)
9969 return nullptr; // Couldn't evaluate!
9970 NextIterVals[PN] = NextPHI;
9971
9972 bool StoppedEvolving = NextPHI == CurrentIterVals[PN];
9973
9974 // Also evaluate the other PHI nodes. However, we don't get to stop if we
9975 // cease to be able to evaluate one of them or if they stop evolving,
9976 // because that doesn't necessarily prevent us from computing PN.
9978 for (const auto &I : CurrentIterVals) {
9979 PHINode *PHI = dyn_cast<PHINode>(I.first);
9980 if (!PHI || PHI == PN || PHI->getParent() != Header) continue;
9981 PHIsToCompute.emplace_back(PHI, I.second);
9982 }
9983 // We use two distinct loops because EvaluateExpression may invalidate any
9984 // iterators into CurrentIterVals.
9985 for (const auto &I : PHIsToCompute) {
9986 PHINode *PHI = I.first;
9987 Constant *&NextPHI = NextIterVals[PHI];
9988 if (!NextPHI) { // Not already computed.
9989 Value *BEValue = PHI->getIncomingValueForBlock(Latch);
9990 NextPHI = EvaluateExpression(BEValue, L, CurrentIterVals, DL, &TLI);
9991 }
9992 if (NextPHI != I.second)
9993 StoppedEvolving = false;
9994 }
9995
9996 // If all entries in CurrentIterVals == NextIterVals then we can stop
9997 // iterating, the loop can't continue to change.
9998 if (StoppedEvolving)
9999 return RetVal = CurrentIterVals[PN];
10000
10001 CurrentIterVals.swap(NextIterVals);
10002 }
10003}
10004
10005const SCEV *ScalarEvolution::computeExitCountExhaustively(const Loop *L,
10006 Value *Cond,
10007 bool ExitWhen) {
10008 PHINode *PN = getConstantEvolvingPHI(Cond, L, &TLI);
10009 if (!PN) return getCouldNotCompute();
10010
10011 // If the loop is canonicalized, the PHI will have exactly two entries.
10012 // That's the only form we support here.
10013 if (PN->getNumIncomingValues() != 2) return getCouldNotCompute();
10014
10015 DenseMap<Instruction *, Constant *> CurrentIterVals;
10016 BasicBlock *Header = L->getHeader();
10017 assert(PN->getParent() == Header && "Can't evaluate PHI not in loop header!");
10018
10019 BasicBlock *Latch = L->getLoopLatch();
10020 assert(Latch && "Should follow from NumIncomingValues == 2!");
10021
10022 for (PHINode &PHI : Header->phis()) {
10023 if (auto *StartCST = getOtherIncomingValue(&PHI, Latch))
10024 CurrentIterVals[&PHI] = StartCST;
10025 }
10026 if (!CurrentIterVals.count(PN))
10027 return getCouldNotCompute();
10028
10029 // Okay, we find a PHI node that defines the trip count of this loop. Execute
10030 // the loop symbolically to determine when the condition gets a value of
10031 // "ExitWhen".
10032 unsigned MaxIterations = MaxBruteForceIterations; // Limit analysis.
10033 const DataLayout &DL = getDataLayout();
10034 for (unsigned IterationNum = 0; IterationNum != MaxIterations;++IterationNum){
10035 auto *CondVal = dyn_cast_or_null<ConstantInt>(
10036 EvaluateExpression(Cond, L, CurrentIterVals, DL, &TLI));
10037
10038 // Couldn't symbolically evaluate.
10039 if (!CondVal) return getCouldNotCompute();
10040
10041 if (CondVal->getValue() == uint64_t(ExitWhen)) {
10042 ++NumBruteForceTripCountsComputed;
10043 return getConstant(Type::getInt32Ty(getContext()), IterationNum);
10044 }
10045
10046 // Update all the PHI nodes for the next iteration.
10047 DenseMap<Instruction *, Constant *> NextIterVals;
10048
10049 // Create a list of which PHIs we need to compute. We want to do this before
10050 // calling EvaluateExpression on them because that may invalidate iterators
10051 // into CurrentIterVals.
10052 SmallVector<PHINode *, 8> PHIsToCompute;
10053 for (const auto &I : CurrentIterVals) {
10054 PHINode *PHI = dyn_cast<PHINode>(I.first);
10055 if (!PHI || PHI->getParent() != Header) continue;
10056 PHIsToCompute.push_back(PHI);
10057 }
10058 for (PHINode *PHI : PHIsToCompute) {
10059 Constant *&NextPHI = NextIterVals[PHI];
10060 if (NextPHI) continue; // Already computed!
10061
10062 Value *BEValue = PHI->getIncomingValueForBlock(Latch);
10063 NextPHI = EvaluateExpression(BEValue, L, CurrentIterVals, DL, &TLI);
10064 }
10065 CurrentIterVals.swap(NextIterVals);
10066 }
10067
10068 // Too many iterations were needed to evaluate.
10069 return getCouldNotCompute();
10070}
10071
10073 auto &Values = ValuesAtScopes[V];
10074 // Check to see if we've folded this expression at this loop before.
10075 for (auto &LS : Values)
10076 if (LS.first == L)
10077 return LS.second ? LS.second : SCEVUse(V);
10078
10079 Values.emplace_back(L, nullptr);
10080
10081 // Otherwise compute it.
10082 SCEVUse C = computeSCEVAtScope(V, L);
10083 for (auto &LS : reverse(ValuesAtScopes[V]))
10084 if (LS.first == L) {
10085 LS.second = C;
10086 // Record the dependency under the bare expression: invalidation walks
10087 // expressions, and any use flags on C do not change which expression
10088 // this is the value at scope of.
10089 if (!isa<SCEVConstant>(C))
10090 ValuesAtScopesUsers[C.getPointer()].push_back({L, V});
10091 break;
10092 }
10093 return C;
10094}
10095
10097 const BasicBlock *ExitingBlock) {
10098 SCEVUse ExitValue = getSCEVAtScope(V, L->getParentLoop());
10099 if (!isLoopInvariant(ExitValue, L)) {
10100 // If we failed to evaluate it in the outer scope, try to evaluate an
10101 // addrec for the specific exit.
10102 // TODO: Generalize this to other expressions.
10103 const SCEV *ExitCount = getExitCount(L, ExitingBlock);
10104 if (!isa<SCEVCouldNotCompute>(ExitCount))
10105 if (auto *AddRec = dyn_cast<SCEVAddRecExpr>(V))
10106 if (AddRec->getLoop() == L)
10107 ExitValue = AddRec->evaluateAtIteration(ExitCount, *this);
10108 }
10109 return ExitValue;
10110}
10111
10112/// This builds up a Constant using the ConstantExpr interface. That way, we
10113/// will return Constants for objects which aren't represented by a
10114/// SCEVConstant, because SCEVConstant is restricted to ConstantInt.
10115/// Returns NULL if the SCEV isn't representable as a Constant.
10117 switch (V->getSCEVType()) {
10118 case scCouldNotCompute:
10119 case scAddRecExpr:
10120 case scVScale:
10121 return nullptr;
10122 case scConstant:
10123 return cast<SCEVConstant>(V)->getValue();
10124 case scUnknown:
10126 case scPtrToAddr: {
10128 if (Constant *CastOp = BuildConstantFromSCEV(P2I->getOperand()))
10129 return ConstantExpr::getPtrToAddr(CastOp, P2I->getType());
10130
10131 return nullptr;
10132 }
10133 case scTruncate: {
10135 if (Constant *CastOp = BuildConstantFromSCEV(ST->getOperand()))
10136 return ConstantExpr::getTrunc(CastOp, ST->getType());
10137 return nullptr;
10138 }
10139 case scAddExpr: {
10140 const SCEVAddExpr *SA = cast<SCEVAddExpr>(V);
10141 Constant *C = nullptr;
10142 for (const SCEV *Op : SA->operands()) {
10144 if (!OpC)
10145 return nullptr;
10146 if (!C) {
10147 C = OpC;
10148 continue;
10149 }
10150 assert(!C->getType()->isPointerTy() &&
10151 "Can only have one pointer, and it must be last");
10152 if (OpC->getType()->isPointerTy()) {
10153 // The offsets have been converted to bytes. We can add bytes using
10154 // an i8 GEP.
10155 C = ConstantExpr::getPtrAdd(OpC, C);
10156 } else {
10157 C = ConstantExpr::getAdd(C, OpC);
10158 }
10159 }
10160 return C;
10161 }
10162 case scMulExpr:
10163 case scSignExtend:
10164 case scZeroExtend:
10165 case scUDivExpr:
10166 case scSMaxExpr:
10167 case scUMaxExpr:
10168 case scSMinExpr:
10169 case scUMinExpr:
10171 return nullptr;
10172 }
10173 llvm_unreachable("Unknown SCEV kind!");
10174}
10175
10176const SCEV *ScalarEvolution::getWithOperands(const SCEV *S,
10177 SmallVectorImpl<SCEVUse> &NewOps) {
10178 switch (S->getSCEVType()) {
10179 case scTruncate:
10180 case scZeroExtend:
10181 case scSignExtend:
10182 case scPtrToAddr:
10183 return getCastExpr(S->getSCEVType(), NewOps[0], S->getType());
10184 case scAddRecExpr: {
10185 auto *AddRec = cast<SCEVAddRecExpr>(S);
10186 return getAddRecExpr(NewOps, AddRec->getLoop(), AddRec->getNoWrapFlags());
10187 }
10188 case scAddExpr:
10189 return getAddExpr(NewOps, cast<SCEVAddExpr>(S)->getNoWrapFlags());
10190 case scMulExpr:
10191 return getMulExpr(NewOps, cast<SCEVMulExpr>(S)->getNoWrapFlags());
10192 case scUDivExpr:
10193 return getUDivExpr(NewOps[0], NewOps[1]);
10194 case scUMaxExpr:
10195 case scSMaxExpr:
10196 case scUMinExpr:
10197 case scSMinExpr:
10198 return getMinMaxExpr(S->getSCEVType(), NewOps);
10200 return getSequentialMinMaxExpr(S->getSCEVType(), NewOps);
10201 case scConstant:
10202 case scVScale:
10203 case scUnknown:
10204 return S;
10205 case scCouldNotCompute:
10206 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
10207 }
10208 llvm_unreachable("Unknown SCEV kind!");
10209}
10210
10211SCEVUse ScalarEvolution::computeSCEVAtScope(const SCEV *V, const Loop *L) {
10212 switch (V->getSCEVType()) {
10213 case scConstant:
10214 case scVScale:
10215 return V;
10216 case scAddRecExpr: {
10217 // If this is a loop recurrence for a loop that does not contain L, then we
10218 // are dealing with the final value computed by the loop.
10219 const SCEVAddRecExpr *AddRec = cast<SCEVAddRecExpr>(V);
10220 // First, attempt to evaluate each operand.
10221 // Avoid performing the look-up in the common case where the specified
10222 // expression has no loop-variant portions.
10223 for (unsigned i = 0, e = AddRec->getNumOperands(); i != e; ++i) {
10224 SCEVUse OpAtScope = getSCEVAtScope(AddRec->getOperand(i), L);
10225 if (OpAtScope == AddRec->getOperand(i))
10226 continue;
10227
10228 // Okay, at least one of these operands is loop variant but might be
10229 // foldable. Build a new instance of the folded commutative expression.
10231 NewOps.reserve(AddRec->getNumOperands());
10232 append_range(NewOps, AddRec->operands().take_front(i));
10233 NewOps.push_back(OpAtScope);
10234 for (++i; i != e; ++i)
10235 NewOps.push_back(getSCEVAtScope(AddRec->getOperand(i), L));
10236
10237 const SCEV *FoldedRec = getAddRecExpr(
10238 NewOps, AddRec->getLoop(), AddRec->getNoWrapFlags(SCEV::FlagNW));
10239 AddRec = dyn_cast<SCEVAddRecExpr>(FoldedRec);
10240 // The addrec may be folded to a nonrecurrence, for example, if the
10241 // induction variable is multiplied by zero after constant folding. Go
10242 // ahead and return the folded value.
10243 if (!AddRec)
10244 return FoldedRec;
10245 break;
10246 }
10247
10248 // If the scope is outside the addrec's loop, evaluate it by using the
10249 // loop exit value of the addrec.
10250 if (!AddRec->getLoop()->contains(L)) {
10251 SCEVUse ExitValue = AddRec->getExitValue(*this);
10252 if (isa<SCEVCouldNotCompute>(ExitValue))
10253 return AddRec;
10254 return ExitValue;
10255 }
10256
10257 return AddRec;
10258 }
10259 case scTruncate:
10260 case scZeroExtend:
10261 case scSignExtend:
10262 case scPtrToAddr:
10263 case scAddExpr:
10264 case scMulExpr:
10265 case scUDivExpr:
10266 case scUMaxExpr:
10267 case scSMaxExpr:
10268 case scUMinExpr:
10269 case scSMinExpr:
10270 case scSequentialUMinExpr: {
10271 ArrayRef<SCEVUse> Ops = V->operands();
10272 // Avoid performing the look-up in the common case where the specified
10273 // expression has no loop-variant portions.
10274 for (unsigned i = 0, e = Ops.size(); i != e; ++i) {
10275 SCEVUse OpAtScope = getSCEVAtScope(Ops[i].getPointer(), L);
10276 if (OpAtScope != Ops[i].getPointer()) {
10277 // Okay, at least one of these operands is loop variant but might be
10278 // foldable. Build a new instance of the folded commutative expression.
10280 NewOps.reserve(Ops.size());
10281 append_range(NewOps, Ops.take_front(i));
10282 NewOps.push_back(OpAtScope);
10283
10284 for (++i; i != e; ++i) {
10285 OpAtScope = getSCEVAtScope(Ops[i].getPointer(), L);
10286 NewOps.push_back(OpAtScope);
10287 }
10288
10289 return getWithOperands(V, NewOps);
10290 }
10291 }
10292 // If we got here, all operands are loop invariant.
10293 return V;
10294 }
10295 case scUnknown: {
10296 // If this instruction is evolved from a constant-evolving PHI, compute the
10297 // exit value from the loop without using SCEVs.
10298 const SCEVUnknown *SU = cast<SCEVUnknown>(V);
10300 if (!I)
10301 return V; // This is some other type of SCEVUnknown, just return it.
10302
10303 if (PHINode *PN = dyn_cast<PHINode>(I)) {
10304 const Loop *CurrLoop = this->LI[I->getParent()];
10305 // Looking for loop exit value.
10306 if (CurrLoop && CurrLoop->getParentLoop() == L &&
10307 PN->getParent() == CurrLoop->getHeader()) {
10308 // Okay, there is no closed form solution for the PHI node. Check
10309 // to see if the loop that contains it has a known backedge-taken
10310 // count. If so, we may be able to force computation of the exit
10311 // value.
10312 const SCEV *BackedgeTakenCount = getBackedgeTakenCount(CurrLoop);
10313 // This trivial case can show up in some degenerate cases where
10314 // the incoming IR has not yet been fully simplified.
10315 if (BackedgeTakenCount->isZero()) {
10316 Value *InitValue = nullptr;
10317 bool MultipleInitValues = false;
10318 for (unsigned i = 0; i < PN->getNumIncomingValues(); i++) {
10319 if (!CurrLoop->contains(PN->getIncomingBlock(i))) {
10320 if (!InitValue)
10321 InitValue = PN->getIncomingValue(i);
10322 else if (InitValue != PN->getIncomingValue(i)) {
10323 MultipleInitValues = true;
10324 break;
10325 }
10326 }
10327 }
10328 if (!MultipleInitValues && InitValue)
10329 return getSCEV(InitValue);
10330 }
10331 // Do we have a loop invariant value flowing around the backedge
10332 // for a loop which must execute the backedge?
10333 if (!isa<SCEVCouldNotCompute>(BackedgeTakenCount) &&
10334 isKnownNonZero(BackedgeTakenCount) &&
10335 PN->getNumIncomingValues() == 2) {
10336
10337 unsigned InLoopPred =
10338 CurrLoop->contains(PN->getIncomingBlock(0)) ? 0 : 1;
10339 Value *BackedgeVal = PN->getIncomingValue(InLoopPred);
10340 if (CurrLoop->isLoopInvariant(BackedgeVal))
10341 return getSCEV(BackedgeVal);
10342 }
10343 if (auto *BTCC = dyn_cast<SCEVConstant>(BackedgeTakenCount)) {
10344 // Okay, we know how many times the containing loop executes. If
10345 // this is a constant evolving PHI node, get the final value at
10346 // the specified iteration number.
10347 Constant *RV =
10348 getConstantEvolutionLoopExitValue(PN, BTCC->getAPInt(), CurrLoop);
10349 if (RV)
10350 return getSCEV(RV);
10351 }
10352 }
10353 }
10354
10355 // Okay, this is an expression that we cannot symbolically evaluate
10356 // into a SCEV. Check to see if it's possible to symbolically evaluate
10357 // the arguments into constants, and if so, try to constant propagate the
10358 // result. This is particularly useful for computing loop exit values.
10359 if (!canConstantFold(I, &TLI))
10360 return V; // This is some other type of SCEVUnknown, just return it.
10361
10362 SmallVector<Constant *, 4> Operands;
10363 Operands.reserve(I->getNumOperands());
10364 bool MadeImprovement = false;
10365 for (Value *Op : I->operands()) {
10366 if (Constant *C = dyn_cast<Constant>(Op)) {
10367 Operands.push_back(C);
10368 continue;
10369 }
10370
10371 // If any of the operands is non-constant and if they are
10372 // non-integer and non-pointer, don't even try to analyze them
10373 // with scev techniques.
10374 if (!isSCEVable(Op->getType()))
10375 return V;
10376
10377 const SCEV *OrigV = getSCEV(Op);
10378 const SCEV *OpV = getSCEVAtScope(OrigV, L);
10379 MadeImprovement |= OrigV != OpV;
10380
10382 if (!C)
10383 return V;
10384 assert(C->getType() == Op->getType() && "Type mismatch");
10385 Operands.push_back(C);
10386 }
10387
10388 // Check to see if getSCEVAtScope actually made an improvement.
10389 if (!MadeImprovement)
10390 return V; // This is some other type of SCEVUnknown, just return it.
10391
10392 Constant *C = nullptr;
10393 const DataLayout &DL = getDataLayout();
10394 C = ConstantFoldInstOperands(I, Operands, DL, &TLI,
10395 /*AllowNonDeterministic=*/false);
10396 if (!C)
10397 return V;
10398 return getSCEV(C);
10399 }
10400 case scCouldNotCompute:
10401 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
10402 }
10403 llvm_unreachable("Unknown SCEV type!");
10404}
10405
10407 return getSCEVAtScope(getSCEV(V), L);
10408}
10409
10410const SCEV *ScalarEvolution::stripInjectiveFunctions(const SCEV *S) const {
10412 return stripInjectiveFunctions(ZExt->getOperand());
10414 return stripInjectiveFunctions(SExt->getOperand());
10415 return S;
10416}
10417
10418/// Finds the minimum unsigned root of the following equation:
10419///
10420/// A * X = B (mod N)
10421///
10422/// where N = 2^BW and BW is the common bit width of A and B. The signedness of
10423/// A and B isn't important.
10424///
10425/// If the equation does not have a solution, SCEVCouldNotCompute is returned.
10426static const SCEV *
10429 ScalarEvolution &SE, const Loop *L) {
10430 uint32_t BW = A.getBitWidth();
10431 assert(BW == SE.getTypeSizeInBits(B->getType()));
10432 assert(A != 0 && "A must be non-zero.");
10433
10434 // 1. D = gcd(A, N)
10435 //
10436 // The gcd of A and N may have only one prime factor: 2. The number of
10437 // trailing zeros in A is its multiplicity
10438 uint32_t Mult2 = A.countr_zero();
10439 // D = 2^Mult2
10440
10441 // 2. Check if B is divisible by D.
10442 //
10443 // B is divisible by D if and only if the multiplicity of prime factor 2 for B
10444 // is not less than multiplicity of this prime factor for D.
10445 unsigned MinTZ = SE.getMinTrailingZeros(B);
10446 // Try again with the terminator of the loop predecessor for context-specific
10447 // result, if MinTZ s too small.
10448 if (MinTZ < Mult2 && L->getLoopPredecessor())
10449 MinTZ = SE.getMinTrailingZeros(B, L->getLoopPredecessor()->getTerminator());
10450 if (MinTZ < Mult2) {
10451 // Check if we can prove there's no remainder using URem.
10452 const SCEV *URem =
10453 SE.getURemExpr(B, SE.getConstant(APInt::getOneBitSet(BW, Mult2)));
10454 const SCEV *Zero = SE.getZero(B->getType());
10455 if (!SE.isKnownPredicate(CmpInst::ICMP_EQ, URem, Zero)) {
10456 // Try to add a predicate ensuring B is a multiple of 1 << Mult2.
10457 if (!Predicates)
10458 return SE.getCouldNotCompute();
10459
10460 // Avoid adding a predicate that is known to be false.
10461 if (SE.isKnownPredicate(CmpInst::ICMP_NE, URem, Zero))
10462 return SE.getCouldNotCompute();
10463 Predicates->push_back(SE.getEqualPredicate(URem, Zero));
10464 }
10465 }
10466
10467 // 3. Compute I: the multiplicative inverse of (A / D) in arithmetic
10468 // modulo (N / D).
10469 //
10470 // If D == 1, (N / D) == N == 2^BW, so we need one extra bit to represent
10471 // (N / D) in general. The inverse itself always fits into BW bits, though,
10472 // so we immediately truncate it.
10473 APInt AD = A.lshr(Mult2).trunc(BW - Mult2); // AD = A / D
10474 APInt I = AD.multiplicativeInverse().zext(BW);
10475
10476 // 4. Compute the minimum unsigned root of the equation:
10477 // I * (B / D) mod (N / D)
10478 // To simplify the computation, we factor out the divide by D:
10479 // (I * B mod N) / D
10480 const SCEV *D = SE.getConstant(APInt::getOneBitSet(BW, Mult2));
10481 return SE.getUDivExactExpr(SE.getMulExpr(B, SE.getConstant(I)), D);
10482}
10483
10484/// For a given quadratic addrec, generate coefficients of the corresponding
10485/// quadratic equation, multiplied by a common value to ensure that they are
10486/// integers.
10487/// The returned value is a tuple { A, B, C, M, BitWidth }, where
10488/// Ax^2 + Bx + C is the quadratic function, M is the value that A, B and C
10489/// were multiplied by, and BitWidth is the bit width of the original addrec
10490/// coefficients.
10491/// This function returns std::nullopt if the addrec coefficients are not
10492/// compile- time constants.
10493static std::optional<std::tuple<APInt, APInt, APInt, APInt, unsigned>>
10495 assert(AddRec->getNumOperands() == 3 && "This is not a quadratic chrec!");
10496 const SCEVConstant *LC = dyn_cast<SCEVConstant>(AddRec->getOperand(0));
10497 const SCEVConstant *MC = dyn_cast<SCEVConstant>(AddRec->getOperand(1));
10498 const SCEVConstant *NC = dyn_cast<SCEVConstant>(AddRec->getOperand(2));
10499 LLVM_DEBUG(dbgs() << __func__ << ": analyzing quadratic addrec: "
10500 << *AddRec << '\n');
10501
10502 // We currently can only solve this if the coefficients are constants.
10503 if (!LC || !MC || !NC) {
10504 LLVM_DEBUG(dbgs() << __func__ << ": coefficients are not constant\n");
10505 return std::nullopt;
10506 }
10507
10508 APInt L = LC->getAPInt();
10509 APInt M = MC->getAPInt();
10510 APInt N = NC->getAPInt();
10511 assert(!N.isZero() && "This is not a quadratic addrec");
10512
10513 unsigned BitWidth = LC->getAPInt().getBitWidth();
10514 unsigned NewWidth = BitWidth + 1;
10515 LLVM_DEBUG(dbgs() << __func__ << ": addrec coeff bw: "
10516 << BitWidth << '\n');
10517 // The sign-extension (as opposed to a zero-extension) here matches the
10518 // extension used in SolveQuadraticEquationWrap (with the same motivation).
10519 N = N.sext(NewWidth);
10520 M = M.sext(NewWidth);
10521 L = L.sext(NewWidth);
10522
10523 // The increments are M, M+N, M+2N, ..., so the accumulated values are
10524 // L+M, (L+M)+(M+N), (L+M)+(M+N)+(M+2N), ..., that is,
10525 // L+M, L+2M+N, L+3M+3N, ...
10526 // After n iterations the accumulated value Acc is L + nM + n(n-1)/2 N.
10527 //
10528 // The equation Acc = 0 is then
10529 // L + nM + n(n-1)/2 N = 0, or 2L + 2M n + n(n-1) N = 0.
10530 // In a quadratic form it becomes:
10531 // N n^2 + (2M-N) n + 2L = 0.
10532
10533 APInt A = N;
10534 APInt B = 2 * M - A;
10535 APInt C = 2 * L;
10536 APInt T = APInt(NewWidth, 2);
10537 LLVM_DEBUG(dbgs() << __func__ << ": equation " << A << "x^2 + " << B
10538 << "x + " << C << ", coeff bw: " << NewWidth
10539 << ", multiplied by " << T << '\n');
10540 return std::make_tuple(A, B, C, T, BitWidth);
10541}
10542
10543/// Helper function to compare optional APInts:
10544/// (a) if X and Y both exist, return min(X, Y),
10545/// (b) if neither X nor Y exist, return std::nullopt,
10546/// (c) if exactly one of X and Y exists, return that value.
10547static std::optional<APInt> MinOptional(std::optional<APInt> X,
10548 std::optional<APInt> Y) {
10549 if (X && Y) {
10550 unsigned W = std::max(X->getBitWidth(), Y->getBitWidth());
10551 APInt XW = X->sext(W);
10552 APInt YW = Y->sext(W);
10553 return XW.slt(YW) ? *X : *Y;
10554 }
10555 if (!X && !Y)
10556 return std::nullopt;
10557 return X ? *X : *Y;
10558}
10559
10560/// Helper function to truncate an optional APInt to a given BitWidth.
10561/// When solving addrec-related equations, it is preferable to return a value
10562/// that has the same bit width as the original addrec's coefficients. If the
10563/// solution fits in the original bit width, truncate it (except for i1).
10564/// Returning a value of a different bit width may inhibit some optimizations.
10565///
10566/// In general, a solution to a quadratic equation generated from an addrec
10567/// may require BW+1 bits, where BW is the bit width of the addrec's
10568/// coefficients. The reason is that the coefficients of the quadratic
10569/// equation are BW+1 bits wide (to avoid truncation when converting from
10570/// the addrec to the equation).
10571static std::optional<APInt> TruncIfPossible(std::optional<APInt> X,
10572 unsigned BitWidth) {
10573 if (!X)
10574 return std::nullopt;
10575 unsigned W = X->getBitWidth();
10577 return X->trunc(BitWidth);
10578 return X;
10579}
10580
10581/// Let c(n) be the value of the quadratic chrec {L,+,M,+,N} after n
10582/// iterations. The values L, M, N are assumed to be signed, and they
10583/// should all have the same bit widths.
10584/// Find the least n >= 0 such that c(n) = 0 in the arithmetic modulo 2^BW,
10585/// where BW is the bit width of the addrec's coefficients.
10586/// If the calculated value is a BW-bit integer (for BW > 1), it will be
10587/// returned as such, otherwise the bit width of the returned value may
10588/// be greater than BW.
10589///
10590/// This function returns std::nullopt if
10591/// (a) the addrec coefficients are not constant, or
10592/// (b) SolveQuadraticEquationWrap was unable to find a solution. For cases
10593/// like x^2 = 5, no integer solutions exist, in other cases an integer
10594/// solution may exist, but SolveQuadraticEquationWrap may fail to find it.
10595static std::optional<APInt>
10597 APInt A, B, C, M;
10598 unsigned BitWidth;
10599 auto T = GetQuadraticEquation(AddRec);
10600 if (!T)
10601 return std::nullopt;
10602
10603 std::tie(A, B, C, M, BitWidth) = *T;
10604 LLVM_DEBUG(dbgs() << __func__ << ": solving for unsigned overflow\n");
10605 std::optional<APInt> X =
10607 if (!X)
10608 return std::nullopt;
10609
10610 ConstantInt *CX = ConstantInt::get(SE.getContext(), *X);
10611 ConstantInt *V = EvaluateConstantChrecAtConstant(AddRec, CX, SE);
10612 if (!V->isZero())
10613 return std::nullopt;
10614
10615 return TruncIfPossible(X, BitWidth);
10616}
10617
10618/// Let c(n) be the value of the quadratic chrec {0,+,M,+,N} after n
10619/// iterations. The values M, N are assumed to be signed, and they
10620/// should all have the same bit widths.
10621/// Find the least n such that c(n) does not belong to the given range,
10622/// while c(n-1) does.
10623///
10624/// This function returns std::nullopt if
10625/// (a) the addrec coefficients are not constant, or
10626/// (b) SolveQuadraticEquationWrap was unable to find a solution for the
10627/// bounds of the range.
10628static std::optional<APInt>
10630 const ConstantRange &Range, ScalarEvolution &SE) {
10631 assert(AddRec->getOperand(0)->isZero() &&
10632 "Starting value of addrec should be 0");
10633 LLVM_DEBUG(dbgs() << __func__ << ": solving boundary crossing for range "
10634 << Range << ", addrec " << *AddRec << '\n');
10635 // This case is handled in getNumIterationsInRange. Here we can assume that
10636 // we start in the range.
10637 assert(Range.contains(APInt(SE.getTypeSizeInBits(AddRec->getType()), 0)) &&
10638 "Addrec's initial value should be in range");
10639
10640 APInt A, B, C, M;
10641 unsigned BitWidth;
10642 auto T = GetQuadraticEquation(AddRec);
10643 if (!T)
10644 return std::nullopt;
10645
10646 // Be careful about the return value: there can be two reasons for not
10647 // returning an actual number. First, if no solutions to the equations
10648 // were found, and second, if the solutions don't leave the given range.
10649 // The first case means that the actual solution is "unknown", the second
10650 // means that it's known, but not valid. If the solution is unknown, we
10651 // cannot make any conclusions.
10652 // Return a pair: the optional solution and a flag indicating if the
10653 // solution was found.
10654 auto SolveForBoundary =
10655 [&](APInt Bound) -> std::pair<std::optional<APInt>, bool> {
10656 // Solve for signed overflow and unsigned overflow, pick the lower
10657 // solution.
10658 LLVM_DEBUG(dbgs() << "SolveQuadraticAddRecRange: checking boundary "
10659 << Bound << " (before multiplying by " << M << ")\n");
10660 Bound *= M; // The quadratic equation multiplier.
10661
10662 std::optional<APInt> SO;
10663 if (BitWidth > 1) {
10664 LLVM_DEBUG(dbgs() << "SolveQuadraticAddRecRange: solving for "
10665 "signed overflow\n");
10667 }
10668 LLVM_DEBUG(dbgs() << "SolveQuadraticAddRecRange: solving for "
10669 "unsigned overflow\n");
10670 std::optional<APInt> UO =
10672
10673 auto LeavesRange = [&] (const APInt &X) {
10674 ConstantInt *C0 = ConstantInt::get(SE.getContext(), X);
10675 ConstantInt *V0 = EvaluateConstantChrecAtConstant(AddRec, C0, SE);
10676 if (Range.contains(V0->getValue()))
10677 return false;
10678 // X should be at least 1, so X-1 is non-negative.
10679 ConstantInt *C1 = ConstantInt::get(SE.getContext(), X-1);
10681 if (Range.contains(V1->getValue()))
10682 return true;
10683 return false;
10684 };
10685
10686 // If SolveQuadraticEquationWrap returns std::nullopt, it means that there
10687 // can be a solution, but the function failed to find it. We cannot treat it
10688 // as "no solution".
10689 if (!SO || !UO)
10690 return {std::nullopt, false};
10691
10692 // Check the smaller value first to see if it leaves the range.
10693 // At this point, both SO and UO must have values.
10694 std::optional<APInt> Min = MinOptional(SO, UO);
10695 if (LeavesRange(*Min))
10696 return { Min, true };
10697 std::optional<APInt> Max = Min == SO ? UO : SO;
10698 if (LeavesRange(*Max))
10699 return { Max, true };
10700
10701 // Solutions were found, but were eliminated, hence the "true".
10702 return {std::nullopt, true};
10703 };
10704
10705 std::tie(A, B, C, M, BitWidth) = *T;
10706 // Lower bound is inclusive, subtract 1 to represent the exiting value.
10707 APInt Lower = Range.getLower().sext(A.getBitWidth()) - 1;
10708 APInt Upper = Range.getUpper().sext(A.getBitWidth());
10709 auto SL = SolveForBoundary(Lower);
10710 auto SU = SolveForBoundary(Upper);
10711 // If any of the solutions was unknown, no meaninigful conclusions can
10712 // be made.
10713 if (!SL.second || !SU.second)
10714 return std::nullopt;
10715
10716 // Claim: The correct solution is not some value between Min and Max.
10717 //
10718 // Justification: Assuming that Min and Max are different values, one of
10719 // them is when the first signed overflow happens, the other is when the
10720 // first unsigned overflow happens. Crossing the range boundary is only
10721 // possible via an overflow (treating 0 as a special case of it, modeling
10722 // an overflow as crossing k*2^W for some k).
10723 //
10724 // The interesting case here is when Min was eliminated as an invalid
10725 // solution, but Max was not. The argument is that if there was another
10726 // overflow between Min and Max, it would also have been eliminated if
10727 // it was considered.
10728 //
10729 // For a given boundary, it is possible to have two overflows of the same
10730 // type (signed/unsigned) without having the other type in between: this
10731 // can happen when the vertex of the parabola is between the iterations
10732 // corresponding to the overflows. This is only possible when the two
10733 // overflows cross k*2^W for the same k. In such case, if the second one
10734 // left the range (and was the first one to do so), the first overflow
10735 // would have to enter the range, which would mean that either we had left
10736 // the range before or that we started outside of it. Both of these cases
10737 // are contradictions.
10738 //
10739 // Claim: In the case where SolveForBoundary returns std::nullopt, the correct
10740 // solution is not some value between the Max for this boundary and the
10741 // Min of the other boundary.
10742 //
10743 // Justification: Assume that we had such Max_A and Min_B corresponding
10744 // to range boundaries A and B and such that Max_A < Min_B. If there was
10745 // a solution between Max_A and Min_B, it would have to be caused by an
10746 // overflow corresponding to either A or B. It cannot correspond to B,
10747 // since Min_B is the first occurrence of such an overflow. If it
10748 // corresponded to A, it would have to be either a signed or an unsigned
10749 // overflow that is larger than both eliminated overflows for A. But
10750 // between the eliminated overflows and this overflow, the values would
10751 // cover the entire value space, thus crossing the other boundary, which
10752 // is a contradiction.
10753
10754 return TruncIfPossible(MinOptional(SL.first, SU.first), BitWidth);
10755}
10756
10757ScalarEvolution::ExitLimit ScalarEvolution::howFarToZero(const SCEV *V,
10758 const Loop *L,
10759 bool ControlsOnlyExit,
10760 bool AllowPredicates) {
10761
10762 // This is only used for loops with a "x != y" exit test. The exit condition
10763 // is now expressed as a single expression, V = x-y. So the exit test is
10764 // effectively V != 0. We know and take advantage of the fact that this
10765 // expression only being used in a comparison by zero context.
10766
10768 // If the value is a constant
10769 if (const SCEVConstant *C = dyn_cast<SCEVConstant>(V)) {
10770 // If the value is already zero, the branch will execute zero times.
10771 if (C->getValue()->isZero()) return C;
10772 return getCouldNotCompute(); // Otherwise it will loop infinitely.
10773 }
10774
10775 const SCEVAddRecExpr *AddRec =
10776 dyn_cast<SCEVAddRecExpr>(stripInjectiveFunctions(V));
10777
10778 if (!AddRec && AllowPredicates)
10779 // Try to make this an AddRec using runtime tests, in the first X
10780 // iterations of this loop, where X is the SCEV expression found by the
10781 // algorithm below.
10782 AddRec = convertSCEVToAddRecWithPredicates(V, L, Predicates);
10783
10784 if (!AddRec || AddRec->getLoop() != L)
10785 return getCouldNotCompute();
10786
10787 // If this is a quadratic (3-term) AddRec {L,+,M,+,N}, find the roots of
10788 // the quadratic equation to solve it.
10789 if (AddRec->isQuadratic() && AddRec->getType()->isIntegerTy()) {
10790 // We can only use this value if the chrec ends up with an exact zero
10791 // value at this index. When solving for "X*X != 5", for example, we
10792 // should not accept a root of 2.
10793 if (auto S = SolveQuadraticAddRecExact(AddRec, *this)) {
10794 const auto *R = cast<SCEVConstant>(getConstant(*S));
10795 return ExitLimit(R, R, R, false, Predicates);
10796 }
10797 return getCouldNotCompute();
10798 }
10799
10800 // Otherwise we can only handle this if it is affine.
10801 if (!AddRec->isAffine())
10802 return getCouldNotCompute();
10803
10804 // If this is an affine expression, the execution count of this branch is
10805 // the minimum unsigned root of the following equation:
10806 //
10807 // Start + Step*N = 0 (mod 2^BW)
10808 //
10809 // equivalent to:
10810 //
10811 // Step*N = -Start (mod 2^BW)
10812 //
10813 // where BW is the common bit width of Start and Step.
10814
10815 // Get the initial value for the loop.
10816 const SCEV *Start = getSCEVAtScope(AddRec->getStart(), L->getParentLoop());
10817 const SCEV *Step = getSCEVAtScope(AddRec->getOperand(1), L->getParentLoop());
10818
10819 if (!isLoopInvariant(Step, L))
10820 return getCouldNotCompute();
10821
10822 LoopGuards Guards = LoopGuards::collect(L, *this);
10823 // Specialize step for this loop so we get context sensitive facts below.
10824 const SCEV *StepWLG = applyLoopGuards(Step, Guards);
10825
10826 // For positive steps (counting up until unsigned overflow):
10827 // N = -Start/Step (as unsigned)
10828 // For negative steps (counting down to zero):
10829 // N = Start/-Step
10830 // First compute the unsigned distance from zero in the direction of Step.
10831 bool CountDown = isKnownNegative(StepWLG);
10832 if (!CountDown && !isKnownNonNegative(StepWLG))
10833 return getCouldNotCompute();
10834
10835 const SCEV *Distance = CountDown ? Start : getNegativeSCEV(Start);
10836 // Handle unitary steps, which cannot wraparound.
10837 // 1*N = -Start; -1*N = Start (mod 2^BW), so:
10838 // N = Distance (as unsigned)
10839
10840 if (match(Step, m_CombineOr(m_scev_One(), m_scev_AllOnes()))) {
10841 APInt MaxBECount = getUnsignedRangeMax(applyLoopGuards(Distance, Guards));
10842 MaxBECount = APIntOps::umin(MaxBECount, getUnsignedRangeMax(Distance));
10843
10844 // When a loop like "for (int i = 0; i != n; ++i) { /* body */ }" is rotated,
10845 // we end up with a loop whose backedge-taken count is n - 1. Detect this
10846 // case, and see if we can improve the bound.
10847 //
10848 // Explicitly handling this here is necessary because getUnsignedRange
10849 // isn't context-sensitive; it doesn't know that we only care about the
10850 // range inside the loop.
10851 const SCEV *Zero = getZero(Distance->getType());
10852 const SCEV *One = getOne(Distance->getType());
10853 const SCEV *DistancePlusOne = getAddExpr(Distance, One);
10854 if (isLoopEntryGuardedByCond(L, ICmpInst::ICMP_NE, DistancePlusOne, Zero)) {
10855 // If Distance + 1 doesn't overflow, we can compute the maximum distance
10856 // as "unsigned_max(Distance + 1) - 1". Also apply the loop guards to
10857 // Distance + 1; the range of Distance itself may be a wrapped set even
10858 // when the guards bound Distance + 1 tightly.
10859 APInt Max = APIntOps::umin(
10860 getUnsignedRangeMax(applyLoopGuards(DistancePlusOne, Guards)),
10861 getUnsignedRangeMax(DistancePlusOne));
10862 MaxBECount = APIntOps::umin(MaxBECount, Max - 1);
10863 }
10864 return ExitLimit(Distance, getConstant(MaxBECount), Distance, false,
10865 Predicates);
10866 }
10867
10868 // If the condition controls loop exit (the loop exits only if the expression
10869 // is true) and the addition is no-wrap we can use unsigned divide to
10870 // compute the backedge count. In this case, the step may not divide the
10871 // distance, but we don't care because if the condition is "missed" the loop
10872 // will have undefined behavior due to wrapping.
10873 if (ControlsOnlyExit && AddRec->hasNoSelfWrap() &&
10874 loopHasNoAbnormalExits(AddRec->getLoop())) {
10875
10876 // If the stride is zero and the start is non-zero, the loop must be
10877 // infinite. In C++, most loops are finite by assumption, in which case the
10878 // step being zero implies UB must execute if the loop is entered.
10879 if (!(loopIsFiniteByAssumption(L) && isKnownNonZero(Start)) &&
10880 !isKnownNonZero(StepWLG))
10881 return getCouldNotCompute();
10882
10883 const SCEV *Exact =
10884 getUDivExpr(Distance, CountDown ? getNegativeSCEV(Step) : Step);
10885 const SCEV *ConstantMax = getCouldNotCompute();
10886 if (Exact != getCouldNotCompute()) {
10887 APInt MaxInt = getUnsignedRangeMax(applyLoopGuards(Exact, Guards));
10888 ConstantMax =
10890 }
10891 const SCEV *SymbolicMax =
10892 isa<SCEVCouldNotCompute>(Exact) ? ConstantMax : Exact;
10893 return ExitLimit(Exact, ConstantMax, SymbolicMax, false, Predicates);
10894 }
10895
10896 // Solve the general equation.
10897 const SCEVConstant *StepC = dyn_cast<SCEVConstant>(Step);
10898 if (!StepC || StepC->getValue()->isZero())
10899 return getCouldNotCompute();
10900 const SCEV *E = SolveLinEquationWithOverflow(
10901 StepC->getAPInt(), getNegativeSCEV(Start),
10902 AllowPredicates ? &Predicates : nullptr, *this, L);
10903
10904 const SCEV *M = E;
10905 if (E != getCouldNotCompute()) {
10906 APInt MaxWithGuards = getUnsignedRangeMax(applyLoopGuards(E, Guards));
10907 M = getConstant(APIntOps::umin(MaxWithGuards, getUnsignedRangeMax(E)));
10908 }
10909 auto *S = isa<SCEVCouldNotCompute>(E) ? M : E;
10910 return ExitLimit(E, M, S, false, Predicates);
10911}
10912
10913ScalarEvolution::ExitLimit
10914ScalarEvolution::howFarToNonZero(const SCEV *V, const Loop *L) {
10915 // Loops that look like: while (X == 0) are very strange indeed. We don't
10916 // handle them yet except for the trivial case. This could be expanded in the
10917 // future as needed.
10918
10919 // If the value is a constant, check to see if it is known to be non-zero
10920 // already. If so, the backedge will execute zero times.
10921 if (const SCEVConstant *C = dyn_cast<SCEVConstant>(V)) {
10922 if (!C->getValue()->isZero())
10923 return getZero(C->getType());
10924 return getCouldNotCompute(); // Otherwise it will loop infinitely.
10925 }
10926
10927 // We could implement others, but I really doubt anyone writes loops like
10928 // this, and if they did, they would already be constant folded.
10929 return getCouldNotCompute();
10930}
10931
10932std::pair<const BasicBlock *, const BasicBlock *>
10933ScalarEvolution::getPredecessorWithUniqueSuccessorForBB(const BasicBlock *BB)
10934 const {
10935 // If the block has a unique predecessor, then there is no path from the
10936 // predecessor to the block that does not go through the direct edge
10937 // from the predecessor to the block.
10938 if (const BasicBlock *Pred = BB->getSinglePredecessor())
10939 return {Pred, BB};
10940
10941 // A loop's header is defined to be a block that dominates the loop.
10942 // If the header has a unique predecessor outside the loop, it must be
10943 // a block that has exactly one successor that can reach the loop.
10944 if (const Loop *L = LI.getLoopFor(BB))
10945 return {L->getLoopPredecessor(), L->getHeader()};
10946
10947 return {nullptr, BB};
10948}
10949
10950/// SCEV structural equivalence is usually sufficient for testing whether two
10951/// expressions are equal, however for the purposes of looking for a condition
10952/// guarding a loop, it can be useful to be a little more general, since a
10953/// front-end may have replicated the controlling expression.
10954static bool HasSameValue(const SCEV *A, const SCEV *B) {
10955 // Quick check to see if they are the same SCEV.
10956 if (A == B) return true;
10957
10958 auto ComputesEqualValues = [](const Instruction *A, const Instruction *B) {
10959 // Not all instructions that are "identical" compute the same value. For
10960 // instance, two distinct alloca instructions allocating the same type are
10961 // identical and do not read memory; but compute distinct values.
10962 return A->isIdenticalTo(B) && (isa<BinaryOperator>(A) || isa<GetElementPtrInst>(A));
10963 };
10964
10965 // Otherwise, if they're both SCEVUnknown, it's possible that they hold
10966 // two different instructions with the same value. Check for this case.
10967 if (const SCEVUnknown *AU = dyn_cast<SCEVUnknown>(A))
10968 if (const SCEVUnknown *BU = dyn_cast<SCEVUnknown>(B))
10969 if (const Instruction *AI = dyn_cast<Instruction>(AU->getValue()))
10970 if (const Instruction *BI = dyn_cast<Instruction>(BU->getValue()))
10971 if (ComputesEqualValues(AI, BI))
10972 return true;
10973
10974 // Otherwise assume they may have a different value.
10975 return false;
10976}
10977
10978static bool MatchBinarySub(const SCEV *S, SCEVUse &LHS, SCEVUse &RHS) {
10979 const SCEV *Op0, *Op1;
10980 if (!match(S, m_scev_Add(m_SCEV(Op0), m_SCEV(Op1))))
10981 return false;
10982 if (match(Op0, m_scev_Mul(m_scev_AllOnes(), m_SCEV(RHS)))) {
10983 LHS = Op1;
10984 return true;
10985 }
10986 if (match(Op1, m_scev_Mul(m_scev_AllOnes(), m_SCEV(RHS)))) {
10987 LHS = Op0;
10988 return true;
10989 }
10990 return false;
10991}
10992
10994 SCEVUse &RHS, unsigned Depth) {
10995 bool Changed = false;
10996 // Simplifies ICMP to trivial true or false by turning it into '0 == 0' or
10997 // '0 != 0'.
10998 auto TrivialCase = [&](bool TriviallyTrue) {
11000 Pred = TriviallyTrue ? ICmpInst::ICMP_EQ : ICmpInst::ICMP_NE;
11001 return true;
11002 };
11003 // If we hit the max recursion limit bail out.
11004 if (Depth >= 3)
11005 return false;
11006
11007 const SCEV *NewLHS, *NewRHS;
11008 if (match(LHS, m_scev_c_Mul(m_SCEV(NewLHS), m_SCEVVScale())) &&
11009 match(RHS, m_scev_c_Mul(m_SCEV(NewRHS), m_SCEVVScale()))) {
11010 const SCEVMulExpr *LMul = cast<SCEVMulExpr>(LHS);
11011 const SCEVMulExpr *RMul = cast<SCEVMulExpr>(RHS);
11012
11013 // (X * vscale) pred (Y * vscale) ==> X pred Y
11014 // when both multiples are NSW.
11015 // (X * vscale) uicmp/eq/ne (Y * vscale) ==> X uicmp/eq/ne Y
11016 // when both multiples are NUW.
11017 if ((LMul->hasNoSignedWrap() && RMul->hasNoSignedWrap()) ||
11018 (LMul->hasNoUnsignedWrap() && RMul->hasNoUnsignedWrap() &&
11019 !ICmpInst::isSigned(Pred))) {
11020 LHS = NewLHS;
11021 RHS = NewRHS;
11022 Changed = true;
11023 }
11024 }
11025
11026 // Canonicalize a constant to the right side.
11027 if (const SCEVConstant *LHSC = dyn_cast<SCEVConstant>(LHS)) {
11028 // Check for both operands constant.
11029 if (const SCEVConstant *RHSC = dyn_cast<SCEVConstant>(RHS)) {
11030 if (!ICmpInst::compare(LHSC->getAPInt(), RHSC->getAPInt(), Pred))
11031 return TrivialCase(false);
11032 return TrivialCase(true);
11033 }
11034 // Otherwise swap the operands to put the constant on the right.
11035 std::swap(LHS, RHS);
11037 Changed = true;
11038 }
11039
11040 // (K + A) pred (K + B) --> A pred B
11041 // For equality, no flags are needed.
11042 // For signed, both adds must be NSW. For unsigned, both must be NUW.
11043 {
11044 const SCEVConstant *C = nullptr;
11045 if (match(LHS, m_scev_Add(m_SCEVConstant(C), m_SCEV(NewLHS))) &&
11046 match(RHS, m_scev_Add(m_scev_Specific(C), m_SCEV(NewRHS)))) {
11047 const auto *LAdd = cast<SCEVAddExpr>(LHS);
11048 const auto *RAdd = cast<SCEVAddExpr>(RHS);
11049 if (ICmpInst::isEquality(Pred) ||
11050 (ICmpInst::isSigned(Pred) && LAdd->hasNoSignedWrap() &&
11051 RAdd->hasNoSignedWrap()) ||
11052 (ICmpInst::isUnsigned(Pred) && LAdd->hasNoUnsignedWrap() &&
11053 RAdd->hasNoUnsignedWrap())) {
11054 LHS = NewLHS;
11055 RHS = NewRHS;
11056 Changed = true;
11057 }
11058 }
11059 }
11060
11061 // (C * A) pred (C * B) --> A pred B
11062 // For equality predicates, both muls must be NUW or both must be NSW
11063 // (either suffices to make multiplication by C injective; C == 0 is
11064 // impossible because SCEV folds 0 * X to 0).
11065 // For signed ordering, C must be positive and both muls must be NSW.
11066 // For unsigned ordering, both muls must be NUW.
11067 {
11068 const SCEVConstant *C = nullptr;
11069 if (match(LHS, m_scev_Mul(m_SCEVConstant(C), m_SCEV(NewLHS))) &&
11070 match(RHS, m_scev_Mul(m_scev_Specific(C), m_SCEV(NewRHS)))) {
11071 const auto *LMul = cast<SCEVMulExpr>(LHS);
11072 const auto *RMul = cast<SCEVMulExpr>(RHS);
11073 bool BothNUW = LMul->hasNoUnsignedWrap() && RMul->hasNoUnsignedWrap();
11074 bool BothNSW = LMul->hasNoSignedWrap() && RMul->hasNoSignedWrap();
11075 if ((ICmpInst::isEquality(Pred) && (BothNUW || BothNSW)) ||
11076 (ICmpInst::isSigned(Pred) && BothNSW &&
11077 C->getAPInt().isStrictlyPositive()) ||
11078 (ICmpInst::isUnsigned(Pred) && BothNUW)) {
11079 LHS = NewLHS;
11080 RHS = NewRHS;
11081 Changed = true;
11082 }
11083 }
11084 }
11085
11086 // If we're comparing an addrec with a value which is loop-invariant in the
11087 // addrec's loop, put the addrec on the left. Also make a dominance check,
11088 // as both operands could be addrecs loop-invariant in each other's loop.
11089 if (const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(RHS)) {
11090 const Loop *L = AR->getLoop();
11091 if (isLoopInvariant(LHS, L) && properlyDominates(LHS, L->getHeader())) {
11092 std::swap(LHS, RHS);
11094 Changed = true;
11095 }
11096 }
11097
11098 // If there's a constant operand, canonicalize comparisons with boundary
11099 // cases, and canonicalize *-or-equal comparisons to regular comparisons.
11100 if (const SCEVConstant *RC = dyn_cast<SCEVConstant>(RHS)) {
11101 const APInt &RA = RC->getAPInt();
11102
11103 bool SimplifiedByConstantRange = false;
11104
11105 if (!ICmpInst::isEquality(Pred)) {
11107 if (ExactCR.isFullSet())
11108 return TrivialCase(true);
11109 if (ExactCR.isEmptySet())
11110 return TrivialCase(false);
11111
11112 APInt NewRHS;
11113 CmpInst::Predicate NewPred;
11114 if (ExactCR.getEquivalentICmp(NewPred, NewRHS) &&
11115 ICmpInst::isEquality(NewPred)) {
11116 // We were able to convert an inequality to an equality.
11117 Pred = NewPred;
11118 RHS = getConstant(NewRHS);
11119 Changed = SimplifiedByConstantRange = true;
11120 }
11121 }
11122
11123 if (!SimplifiedByConstantRange) {
11124 switch (Pred) {
11125 default:
11126 break;
11127 case ICmpInst::ICMP_EQ:
11128 case ICmpInst::ICMP_NE:
11129 // Fold ((-1) * %a) + %b == 0 (equivalent to %b-%a == 0) into %a == %b.
11130 if (RA.isZero() && MatchBinarySub(LHS, LHS, RHS))
11131 Changed = true;
11132 break;
11133
11134 // The "Should have been caught earlier!" messages refer to the fact
11135 // that the ExactCR.isFullSet() or ExactCR.isEmptySet() check above
11136 // should have fired on the corresponding cases, and canonicalized the
11137 // check to trivial case.
11138
11139 case ICmpInst::ICMP_UGE:
11140 assert(!RA.isMinValue() && "Should have been caught earlier!");
11141 Pred = ICmpInst::ICMP_UGT;
11142 RHS = getConstant(RA - 1);
11143 Changed = true;
11144 break;
11145 case ICmpInst::ICMP_ULE:
11146 assert(!RA.isMaxValue() && "Should have been caught earlier!");
11147 Pred = ICmpInst::ICMP_ULT;
11148 RHS = getConstant(RA + 1);
11149 Changed = true;
11150 break;
11151 case ICmpInst::ICMP_SGE:
11152 assert(!RA.isMinSignedValue() && "Should have been caught earlier!");
11153 Pred = ICmpInst::ICMP_SGT;
11154 RHS = getConstant(RA - 1);
11155 Changed = true;
11156 break;
11157 case ICmpInst::ICMP_SLE:
11158 assert(!RA.isMaxSignedValue() && "Should have been caught earlier!");
11159 Pred = ICmpInst::ICMP_SLT;
11160 RHS = getConstant(RA + 1);
11161 Changed = true;
11162 break;
11163 }
11164 }
11165 }
11166
11167 // a /u b == 0 => a < b
11168 // a /u b != 0 => a >= b
11169 if (ICmpInst::isEquality(Pred) && RHS->isZero() &&
11170 match(LHS, m_scev_UDiv(m_SCEV(LHS), m_SCEV(RHS)))) {
11172 Changed = true;
11173 }
11174
11175 // Check for obvious equality.
11176 if (HasSameValue(LHS, RHS)) {
11177 if (ICmpInst::isTrueWhenEqual(Pred))
11178 return TrivialCase(true);
11180 return TrivialCase(false);
11181 }
11182
11183 // If possible, canonicalize GE/LE comparisons to GT/LT comparisons, by
11184 // adding or subtracting 1 from one of the operands.
11185 switch (Pred) {
11186 case ICmpInst::ICMP_SLE:
11187 if (!getSignedRangeMax(RHS).isMaxSignedValue()) {
11188 RHS = getAddExpr(getConstant(RHS->getType(), 1, true), RHS,
11190 Pred = ICmpInst::ICMP_SLT;
11191 Changed = true;
11192 } else if (!getSignedRangeMin(LHS).isMinSignedValue()) {
11193 LHS = getAddExpr(getConstant(RHS->getType(), (uint64_t)-1, true), LHS,
11195 Pred = ICmpInst::ICMP_SLT;
11196 Changed = true;
11197 }
11198 break;
11199 case ICmpInst::ICMP_SGE:
11200 if (!getSignedRangeMin(RHS).isMinSignedValue()) {
11201 RHS = getAddExpr(getConstant(RHS->getType(), (uint64_t)-1, true), RHS,
11203 Pred = ICmpInst::ICMP_SGT;
11204 Changed = true;
11205 } else if (!getSignedRangeMax(LHS).isMaxSignedValue()) {
11206 LHS = getAddExpr(getConstant(RHS->getType(), 1, true), LHS,
11208 Pred = ICmpInst::ICMP_SGT;
11209 Changed = true;
11210 }
11211 break;
11212 case ICmpInst::ICMP_ULE:
11213 if (!getUnsignedRangeMax(RHS).isMaxValue()) {
11214 RHS = getAddExpr(getConstant(RHS->getType(), 1, true), RHS,
11216 Pred = ICmpInst::ICMP_ULT;
11217 Changed = true;
11218 } else if (!getUnsignedRangeMin(LHS).isMinValue()) {
11219 LHS = getAddExpr(getConstant(RHS->getType(), (uint64_t)-1, true), LHS);
11220 Pred = ICmpInst::ICMP_ULT;
11221 Changed = true;
11222 }
11223 break;
11224 case ICmpInst::ICMP_UGE:
11225 // If RHS is an op we can fold the -1, try that first.
11226 // Otherwise prefer LHS to preserve the nuw flag.
11227 if ((isa<SCEVConstant>(RHS) ||
11229 isa<SCEVConstant>(cast<SCEVNAryExpr>(RHS)->getOperand(0)))) &&
11230 !getUnsignedRangeMin(RHS).isMinValue()) {
11231 RHS = getAddExpr(getConstant(RHS->getType(), (uint64_t)-1, true), RHS);
11232 Pred = ICmpInst::ICMP_UGT;
11233 Changed = true;
11234 } else if (!getUnsignedRangeMax(LHS).isMaxValue()) {
11235 LHS = getAddExpr(getConstant(RHS->getType(), 1, true), LHS,
11237 Pred = ICmpInst::ICMP_UGT;
11238 Changed = true;
11239 } else if (!getUnsignedRangeMin(RHS).isMinValue()) {
11240 RHS = getAddExpr(getConstant(RHS->getType(), (uint64_t)-1, true), RHS);
11241 Pred = ICmpInst::ICMP_UGT;
11242 Changed = true;
11243 }
11244 break;
11245 default:
11246 break;
11247 }
11248
11249 // TODO: More simplifications are possible here.
11250
11251 // Recursively simplify until we either hit a recursion limit or nothing
11252 // changes.
11253 if (Changed)
11254 (void)SimplifyICmpOperands(Pred, LHS, RHS, Depth + 1);
11255
11256 return Changed;
11257}
11258
11260 return getSignedRangeMax(S).isNegative();
11261}
11262
11266
11268 return !getSignedRangeMin(S).isNegative();
11269}
11270
11274
11276 // Query push down for cases where the unsigned range is
11277 // less than sufficient.
11278 if (const auto *SExt = dyn_cast<SCEVSignExtendExpr>(S))
11279 return isKnownNonZero(SExt->getOperand(0));
11280 return getUnsignedRangeMin(S) != 0;
11281}
11282
11284 bool OrNegative) {
11285 auto NonRecursive = [OrNegative](const SCEV *S) {
11286 if (auto *C = dyn_cast<SCEVConstant>(S))
11287 return C->getAPInt().isPowerOf2() ||
11288 (OrNegative && C->getAPInt().isNegatedPowerOf2());
11289
11290 // vscale is a power-of-two.
11291 return isa<SCEVVScale>(S);
11292 };
11293
11294 if (NonRecursive(S))
11295 return true;
11296
11297 auto *Mul = dyn_cast<SCEVMulExpr>(S);
11298 if (!Mul)
11299 return false;
11300 return all_of(Mul->operands(), NonRecursive) && (OrZero || isKnownNonZero(S));
11301}
11302
11304 const SCEV *S, uint64_t M,
11306 if (M == 0)
11307 return false;
11308 if (M == 1)
11309 return true;
11310
11311 // For a constant, check that "S % M == 0".
11312 if (auto *Cst = dyn_cast<SCEVConstant>(S)) {
11313 APInt C = Cst->getAPInt();
11314 return C.urem(M) == 0;
11315 }
11316
11317 // Basic tests have failed.
11318 // Check "S % M == 0" at compile time and record runtime Assumptions.
11319 auto *STy = dyn_cast<IntegerType>(S->getType());
11320 const SCEV *SmodM =
11321 getURemExpr(S, getConstant(ConstantInt::get(STy, M, false)));
11322 const SCEV *Zero = getZero(STy);
11323
11324 // Check whether "S % M == 0" is known at compile time.
11325 if (isKnownPredicate(ICmpInst::ICMP_EQ, SmodM, Zero))
11326 return true;
11327
11328 // Check whether "S % M != 0" is known at compile time.
11329 if (isKnownPredicate(ICmpInst::ICMP_NE, SmodM, Zero))
11330 return false;
11331
11332 if (!Predicates)
11333 return false;
11334
11335 // Look through Add and AddRec expressions with nuw to improve the
11336 // precision of added predicates. S is a multiple of M if S starts with a
11337 // multiple of M and at every iteration step S only adds multiples of M.
11340 all_of(S->operands(),
11341 [&](SCEVUse Op) { return isKnownMultipleOf(Op, M, Predicates); }))
11342 return true;
11343
11344 // Similarly, look through Mul with nuw, where any operand being a
11345 // known-multiple is sufficient.
11346 if (auto *Mul = dyn_cast<SCEVMulExpr>(S))
11347 if (Mul->hasNoUnsignedWrap() && any_of(S->operands(), [&](SCEVUse Op) {
11348 return isKnownMultipleOf(Op, M, Predicates);
11349 }))
11350 return true;
11351
11352 // Similarly, look through MinMax, with no wrapping arithmetic to consider.
11353 if (isa<SCEVMinMaxExpr>(S) && all_of(S->operands(), [&](SCEVUse Op) {
11354 return isKnownMultipleOf(Op, M, Predicates);
11355 }))
11356 return true;
11357
11359
11360 // Detect redundant predicates.
11361 for (auto *A : *Predicates)
11362 if (A->implies(P, *this))
11363 return true;
11364
11365 // Only record non-redundant predicates.
11366 Predicates->push_back(P);
11367 return true;
11368}
11369
11371 return ((isKnownNonNegative(S1) && isKnownNonNegative(S2)) ||
11373}
11374
11375std::pair<const SCEV *, const SCEV *>
11377 // Compute SCEV on entry of loop L.
11378 const SCEV *Start = SCEVInitRewriter::rewrite(S, L, *this);
11379 if (Start == getCouldNotCompute())
11380 return { Start, Start };
11381 // Compute post increment SCEV for loop L.
11382 const SCEV *PostInc = SCEVPostIncRewriter::rewrite(S, L, *this);
11383 assert(PostInc != getCouldNotCompute() && "Unexpected could not compute");
11384 return { Start, PostInc };
11385}
11386
11388 SCEVUse RHS) {
11389 // First collect all loops.
11391 getUsedLoops(LHS, LoopsUsed);
11392 getUsedLoops(RHS, LoopsUsed);
11393
11394 if (LoopsUsed.empty())
11395 return false;
11396
11397 // Domination relationship must be a linear order on collected loops.
11398#ifndef NDEBUG
11399 for (const auto *L1 : LoopsUsed)
11400 for (const auto *L2 : LoopsUsed)
11401 assert((DT.dominates(L1->getHeader(), L2->getHeader()) ||
11402 DT.dominates(L2->getHeader(), L1->getHeader())) &&
11403 "Domination relationship is not a linear order");
11404#endif
11405
11406 const Loop *MDL =
11407 *llvm::max_element(LoopsUsed, [&](const Loop *L1, const Loop *L2) {
11408 return DT.properlyDominates(L1->getHeader(), L2->getHeader());
11409 });
11410
11411 // Get init and post increment value for LHS.
11412 auto SplitLHS = SplitIntoInitAndPostInc(MDL, LHS);
11413 // if LHS contains unknown non-invariant SCEV then bail out.
11414 if (SplitLHS.first == getCouldNotCompute())
11415 return false;
11416 assert (SplitLHS.second != getCouldNotCompute() && "Unexpected CNC");
11417 // Get init and post increment value for RHS.
11418 auto SplitRHS = SplitIntoInitAndPostInc(MDL, RHS);
11419 // if RHS contains unknown non-invariant SCEV then bail out.
11420 if (SplitRHS.first == getCouldNotCompute())
11421 return false;
11422 assert (SplitRHS.second != getCouldNotCompute() && "Unexpected CNC");
11423 // It is possible that init SCEV contains an invariant load but it does
11424 // not dominate MDL and is not available at MDL loop entry, so we should
11425 // check it here.
11426 if (!isAvailableAtLoopEntry(SplitLHS.first, MDL) ||
11427 !isAvailableAtLoopEntry(SplitRHS.first, MDL))
11428 return false;
11429
11430 // It seems backedge guard check is faster than entry one so in some cases
11431 // it can speed up whole estimation by short circuit
11432 return isLoopBackedgeGuardedByCond(MDL, Pred, SplitLHS.second,
11433 SplitRHS.second) &&
11434 isLoopEntryGuardedByCond(MDL, Pred, SplitLHS.first, SplitRHS.first);
11435}
11436
11438 SCEVUse RHS) {
11439 // Canonicalize the inputs first.
11440 (void)SimplifyICmpOperands(Pred, LHS, RHS);
11441
11442 return isKnownViaInduction(Pred, LHS, RHS) ||
11443 isKnownPredicateViaSplitting(Pred, LHS, RHS) ||
11444 isKnownViaNonRecursiveReasoning(Pred, LHS, RHS);
11445}
11446
11448 const SCEV *LHS,
11449 const SCEV *RHS) {
11450 if (isKnownPredicate(Pred, LHS, RHS))
11451 return true;
11453 return false;
11454 return std::nullopt;
11455}
11456
11458 const SCEV *RHS,
11459 const Instruction *CtxI) {
11460 // TODO: Analyze guards and assumes from Context's block.
11461 return isKnownPredicate(Pred, LHS, RHS) ||
11462 isBasicBlockEntryGuardedByCond(CtxI->getParent(), Pred, LHS, RHS);
11463}
11464
11465std::optional<bool>
11467 const SCEV *RHS, const Instruction *CtxI) {
11468 std::optional<bool> KnownWithoutContext = evaluatePredicate(Pred, LHS, RHS);
11469 if (KnownWithoutContext)
11470 return KnownWithoutContext;
11471
11472 if (isBasicBlockEntryGuardedByCond(CtxI->getParent(), Pred, LHS, RHS))
11473 return true;
11475 CtxI->getParent(), ICmpInst::getInverseCmpPredicate(Pred), LHS, RHS))
11476 return false;
11477 return std::nullopt;
11478}
11479
11481 const SCEVAddRecExpr *LHS,
11482 const SCEV *RHS) {
11483 const Loop *L = LHS->getLoop();
11484 return isLoopEntryGuardedByCond(L, Pred, LHS->getStart(), RHS) &&
11485 isLoopBackedgeGuardedByCond(L, Pred, LHS->getPostIncExpr(*this), RHS);
11486}
11487
11488std::optional<ScalarEvolution::MonotonicPredicateType>
11490 ICmpInst::Predicate Pred) {
11491 auto Result = getMonotonicPredicateTypeImpl(LHS, Pred);
11492
11493#ifndef NDEBUG
11494 // Verify an invariant: inverting the predicate should turn a monotonically
11495 // increasing change to a monotonically decreasing one, and vice versa.
11496 if (Result) {
11497 auto ResultSwapped =
11498 getMonotonicPredicateTypeImpl(LHS, ICmpInst::getSwappedPredicate(Pred));
11499
11500 assert(*ResultSwapped != *Result &&
11501 "monotonicity should flip as we flip the predicate");
11502 }
11503#endif
11504
11505 return Result;
11506}
11507
11508std::optional<ScalarEvolution::MonotonicPredicateType>
11509ScalarEvolution::getMonotonicPredicateTypeImpl(const SCEVAddRecExpr *LHS,
11510 ICmpInst::Predicate Pred) {
11511 // A zero step value for LHS means the induction variable is essentially a
11512 // loop invariant value. We don't really depend on the predicate actually
11513 // flipping from false to true (for increasing predicates, and the other way
11514 // around for decreasing predicates), all we care about is that *if* the
11515 // predicate changes then it only changes from false to true.
11516 //
11517 // A zero step value in itself is not very useful, but there may be places
11518 // where SCEV can prove X >= 0 but not prove X > 0, so it is helpful to be
11519 // as general as possible.
11520
11521 // Only handle LE/LT/GE/GT predicates.
11522 if (!ICmpInst::isRelational(Pred))
11523 return std::nullopt;
11524
11525 bool IsGreater = ICmpInst::isGE(Pred) || ICmpInst::isGT(Pred);
11526 assert((IsGreater || ICmpInst::isLE(Pred) || ICmpInst::isLT(Pred)) &&
11527 "Should be greater or less!");
11528
11529 // Check that AR does not wrap.
11530 if (ICmpInst::isUnsigned(Pred)) {
11531 if (!LHS->hasNoUnsignedWrap())
11532 return std::nullopt;
11534 }
11535 assert(ICmpInst::isSigned(Pred) &&
11536 "Relational predicate is either signed or unsigned!");
11537 if (!LHS->hasNoSignedWrap())
11538 return std::nullopt;
11539
11540 const SCEV *Step = LHS->getStepRecurrence(*this);
11541
11542 if (isKnownNonNegative(Step))
11544
11545 if (isKnownNonPositive(Step))
11547
11548 return std::nullopt;
11549}
11550
11551std::optional<ScalarEvolution::LoopInvariantPredicate>
11553 const SCEV *RHS, const Loop *L,
11554 const Instruction *CtxI) {
11555 // If there is a loop-invariant, force it into the RHS, otherwise bail out.
11556 if (!isLoopInvariant(RHS, L)) {
11557 if (!isLoopInvariant(LHS, L))
11558 return std::nullopt;
11559
11560 std::swap(LHS, RHS);
11562 }
11563
11564 const SCEVAddRecExpr *ArLHS = dyn_cast<SCEVAddRecExpr>(LHS);
11565 if (!ArLHS || ArLHS->getLoop() != L)
11566 return std::nullopt;
11567
11568 auto MonotonicType = getMonotonicPredicateType(ArLHS, Pred);
11569 if (!MonotonicType)
11570 return std::nullopt;
11571 // If the predicate "ArLHS `Pred` RHS" monotonically increases from false to
11572 // true as the loop iterates, and the backedge is control dependent on
11573 // "ArLHS `Pred` RHS" == true then we can reason as follows:
11574 //
11575 // * if the predicate was false in the first iteration then the predicate
11576 // is never evaluated again, since the loop exits without taking the
11577 // backedge.
11578 // * if the predicate was true in the first iteration then it will
11579 // continue to be true for all future iterations since it is
11580 // monotonically increasing.
11581 //
11582 // For both the above possibilities, we can replace the loop varying
11583 // predicate with its value on the first iteration of the loop (which is
11584 // loop invariant).
11585 //
11586 // A similar reasoning applies for a monotonically decreasing predicate, by
11587 // replacing true with false and false with true in the above two bullets.
11589 auto P = Increasing ? Pred : ICmpInst::getInverseCmpPredicate(Pred);
11590
11591 if (isLoopBackedgeGuardedByCond(L, P, LHS, RHS))
11593 RHS);
11594
11595 if (!CtxI)
11596 return std::nullopt;
11597 // Try to prove via context.
11598 // TODO: Support other cases.
11599 switch (Pred) {
11600 default:
11601 break;
11602 case ICmpInst::ICMP_ULE:
11603 case ICmpInst::ICMP_ULT: {
11604 assert(ArLHS->hasNoUnsignedWrap() && "Is a requirement of monotonicity!");
11605 // Given preconditions
11606 // (1) ArLHS does not cross the border of positive and negative parts of
11607 // range because of:
11608 // - Positive step; (TODO: lift this limitation)
11609 // - nuw - does not cross zero boundary;
11610 // - nsw - does not cross SINT_MAX boundary;
11611 // (2) ArLHS <s RHS
11612 // (3) RHS >=s 0
11613 // we can replace the loop variant ArLHS <u RHS condition with loop
11614 // invariant Start(ArLHS) <u RHS.
11615 //
11616 // Because of (1) there are two options:
11617 // - ArLHS is always negative. It means that ArLHS <u RHS is always false;
11618 // - ArLHS is always non-negative. Because of (3) RHS is also non-negative.
11619 // It means that ArLHS <s RHS <=> ArLHS <u RHS.
11620 // Because of (2) ArLHS <u RHS is trivially true.
11621 // All together it means that ArLHS <u RHS <=> Start(ArLHS) >=s 0.
11622 // We can strengthen this to Start(ArLHS) <u RHS.
11623 auto SignFlippedPred = ICmpInst::getFlippedSignednessPredicate(Pred);
11624 if (ArLHS->hasNoSignedWrap() && ArLHS->isAffine() &&
11625 isKnownPositive(ArLHS->getStepRecurrence(*this)) &&
11626 isKnownNonNegative(RHS) &&
11627 isKnownPredicateAt(SignFlippedPred, ArLHS, RHS, CtxI))
11629 RHS);
11630 }
11631 }
11632
11633 return std::nullopt;
11634}
11635
11636std::optional<ScalarEvolution::LoopInvariantPredicate>
11638 CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, const Loop *L,
11639 const Instruction *CtxI, const SCEV *MaxIter) {
11641 Pred, LHS, RHS, L, CtxI, MaxIter))
11642 return LIP;
11643 if (auto *UMin = dyn_cast<SCEVUMinExpr>(MaxIter))
11644 // Number of iterations expressed as UMIN isn't always great for expressing
11645 // the value on the last iteration. If the straightforward approach didn't
11646 // work, try the following trick: if the a predicate is invariant for X, it
11647 // is also invariant for umin(X, ...). So try to find something that works
11648 // among subexpressions of MaxIter expressed as umin.
11649 for (SCEVUse Op : UMin->operands())
11651 Pred, LHS, RHS, L, CtxI, Op))
11652 return LIP;
11653 return std::nullopt;
11654}
11655
11656std::optional<ScalarEvolution::LoopInvariantPredicate>
11658 CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, const Loop *L,
11659 const Instruction *CtxI, const SCEV *MaxIter) {
11660 // Try to prove the following set of facts:
11661 // - The predicate is monotonic in the iteration space.
11662 // - If the check does not fail on the 1st iteration:
11663 // - No overflow will happen during first MaxIter iterations;
11664 // - It will not fail on the MaxIter'th iteration.
11665 // If the check does fail on the 1st iteration, we leave the loop and no
11666 // other checks matter.
11667
11668 // If there is a loop-invariant, force it into the RHS, otherwise bail out.
11669 if (!isLoopInvariant(RHS, L)) {
11670 if (!isLoopInvariant(LHS, L))
11671 return std::nullopt;
11672
11673 std::swap(LHS, RHS);
11675 }
11676
11677 auto *AR = dyn_cast<SCEVAddRecExpr>(LHS);
11678 if (!AR || AR->getLoop() != L)
11679 return std::nullopt;
11680
11681 // Even if both are valid, we need to consistently chose the unsigned or the
11682 // signed predicate below, not mixtures of both. For now, prefer the unsigned
11683 // predicate.
11684 Pred = Pred.dropSameSign();
11685
11686 // The predicate must be relational (i.e. <, <=, >=, >).
11687 if (!ICmpInst::isRelational(Pred))
11688 return std::nullopt;
11689
11690 // TODO: Support steps other than +/- 1.
11691 const SCEV *Step = AR->getStepRecurrence(*this);
11692 auto *One = getOne(Step->getType());
11693 auto *MinusOne = getNegativeSCEV(One);
11694 if (Step != One && Step != MinusOne)
11695 return std::nullopt;
11696
11697 // Type mismatch here means that MaxIter is potentially larger than max
11698 // unsigned value in start type, which mean we cannot prove no wrap for the
11699 // indvar.
11700 if (AR->getType() != MaxIter->getType())
11701 return std::nullopt;
11702
11703 // Value of IV on suggested last iteration.
11704 const SCEV *Last = AR->evaluateAtIteration(MaxIter, *this);
11705 // Does it still meet the requirement?
11706 if (!isLoopBackedgeGuardedByCond(L, Pred, Last, RHS))
11707 return std::nullopt;
11708 // Because step is +/- 1 and MaxIter has same type as Start (i.e. it does
11709 // not exceed max unsigned value of this type), this effectively proves
11710 // that there is no wrap during the iteration. To prove that there is no
11711 // signed/unsigned wrap, we need to check that
11712 // Start <= Last for step = 1 or Start >= Last for step = -1.
11713 ICmpInst::Predicate NoOverflowPred =
11715 if (Step == MinusOne)
11716 NoOverflowPred = ICmpInst::getSwappedPredicate(NoOverflowPred);
11717 const SCEV *Start = AR->getStart();
11718 if (!isKnownPredicateAt(NoOverflowPred, Start, Last, CtxI))
11719 return std::nullopt;
11720
11721 // Everything is fine.
11722 return ScalarEvolution::LoopInvariantPredicate(Pred, Start, RHS);
11723}
11724
11725bool ScalarEvolution::isKnownPredicateViaConstantRanges(CmpPredicate Pred,
11726 SCEVUse LHS,
11727 SCEVUse RHS) {
11728 if (HasSameValue(LHS, RHS))
11729 return ICmpInst::isTrueWhenEqual(Pred);
11730
11731 auto CheckRange = [&](bool IsSigned) {
11732 auto RangeLHS = IsSigned ? getSignedRange(LHS) : getUnsignedRange(LHS);
11733 auto RangeRHS = IsSigned ? getSignedRange(RHS) : getUnsignedRange(RHS);
11734 return RangeLHS.icmp(Pred, RangeRHS);
11735 };
11736
11737 // The check at the top of the function catches the case where the values are
11738 // known to be equal.
11739 if (Pred == CmpInst::ICMP_EQ)
11740 return false;
11741
11742 if (Pred == CmpInst::ICMP_NE) {
11743 if (CheckRange(true) || CheckRange(false))
11744 return true;
11745 auto *Diff = getMinusSCEV(LHS, RHS);
11746 return !isa<SCEVCouldNotCompute>(Diff) && isKnownNonZero(Diff);
11747 }
11748
11749 return CheckRange(CmpInst::isSigned(Pred));
11750}
11751
11752bool ScalarEvolution::isKnownPredicateViaNoOverflow(CmpPredicate Pred,
11754 // Match X to (A + C1)<ExpectedFlags> and Y to (A + C2)<ExpectedFlags>, where
11755 // C1 and C2 are constant integers. If either X or Y are not add expressions,
11756 // consider them as X + 0 and Y + 0 respectively. C1 and C2 are returned via
11757 // OutC1 and OutC2.
11758 auto MatchBinaryAddToConst = [this](SCEVUse X, SCEVUse Y, APInt &OutC1,
11759 APInt &OutC2, SCEVFlags ExpectedFlags) {
11760 SCEVUse XNonConstOp, XConstOp;
11761 SCEVUse YNonConstOp, YConstOp;
11762 SCEVFlags XFlagsPresent;
11763 SCEVFlags YFlagsPresent;
11764
11765 if (!splitBinaryAdd(X, XConstOp, XNonConstOp, XFlagsPresent)) {
11766 XConstOp = getZero(X->getType());
11767 XNonConstOp = X;
11768 XFlagsPresent = ExpectedFlags;
11769 }
11770 if (!isa<SCEVConstant>(XConstOp))
11771 return false;
11772
11773 if (!splitBinaryAdd(Y, YConstOp, YNonConstOp, YFlagsPresent)) {
11774 YConstOp = getZero(Y->getType());
11775 YNonConstOp = Y;
11776 YFlagsPresent = ExpectedFlags;
11777 }
11778
11779 if (YNonConstOp != XNonConstOp)
11780 return false;
11781
11782 if (!isa<SCEVConstant>(YConstOp))
11783 return false;
11784
11785 // When matching ADDs with NUW flags (and unsigned predicates), only the
11786 // second ADD (with the larger constant) requires NUW.
11787 if ((YFlagsPresent & ExpectedFlags) != ExpectedFlags)
11788 return false;
11789 if (ExpectedFlags != SCEV::FlagNUW &&
11790 (XFlagsPresent & ExpectedFlags) != ExpectedFlags) {
11791 return false;
11792 }
11793
11794 OutC1 = cast<SCEVConstant>(XConstOp)->getAPInt();
11795 OutC2 = cast<SCEVConstant>(YConstOp)->getAPInt();
11796
11797 return true;
11798 };
11799
11800 APInt C1;
11801 APInt C2;
11802
11803 switch (Pred) {
11804 default:
11805 break;
11806
11807 case ICmpInst::ICMP_SGE:
11808 std::swap(LHS, RHS);
11809 [[fallthrough]];
11810 case ICmpInst::ICMP_SLE:
11811 // (X + C1)<nsw> s<= (X + C2)<nsw> if C1 s<= C2.
11812 if (MatchBinaryAddToConst(LHS, RHS, C1, C2, SCEV::FlagNSW) && C1.sle(C2))
11813 return true;
11814
11815 break;
11816
11817 case ICmpInst::ICMP_SGT:
11818 std::swap(LHS, RHS);
11819 [[fallthrough]];
11820 case ICmpInst::ICMP_SLT:
11821 // (X + C1)<nsw> s< (X + C2)<nsw> if C1 s< C2.
11822 if (MatchBinaryAddToConst(LHS, RHS, C1, C2, SCEV::FlagNSW) && C1.slt(C2))
11823 return true;
11824
11825 break;
11826
11827 case ICmpInst::ICMP_UGE:
11828 std::swap(LHS, RHS);
11829 [[fallthrough]];
11830 case ICmpInst::ICMP_ULE:
11831 // (X + C1) u<= (X + C2)<nuw> for C1 u<= C2.
11832 if (MatchBinaryAddToConst(LHS, RHS, C1, C2, SCEV::FlagNUW) && C1.ule(C2))
11833 return true;
11834
11835 break;
11836
11837 case ICmpInst::ICMP_UGT:
11838 std::swap(LHS, RHS);
11839 [[fallthrough]];
11840 case ICmpInst::ICMP_ULT:
11841 // (X + C1) u< (X + C2)<nuw> if C1 u< C2.
11842 if (MatchBinaryAddToConst(LHS, RHS, C1, C2, SCEV::FlagNUW) && C1.ult(C2))
11843 return true;
11844 break;
11845 }
11846
11847 return false;
11848}
11849
11850bool ScalarEvolution::isKnownPredicateViaSplitting(CmpPredicate Pred,
11852 if (Pred != ICmpInst::ICMP_ULT || ProvingSplitPredicate)
11853 return false;
11854
11855 // Allowing arbitrary number of activations of isKnownPredicateViaSplitting on
11856 // the stack can result in exponential time complexity.
11857 SaveAndRestore Restore(ProvingSplitPredicate, true);
11858
11859 // If L >= 0 then I `ult` L <=> I >= 0 && I `slt` L
11860 //
11861 // To prove L >= 0 we use isKnownNonNegative whereas to prove I >= 0 we use
11862 // isKnownPredicate. isKnownPredicate is more powerful, but also more
11863 // expensive; and using isKnownNonNegative(RHS) is sufficient for most of the
11864 // interesting cases seen in practice. We can consider "upgrading" L >= 0 to
11865 // use isKnownPredicate later if needed.
11866 return isKnownNonNegative(RHS) &&
11869}
11870
11871bool ScalarEvolution::isImpliedViaGuard(const BasicBlock *BB, CmpPredicate Pred,
11872 const SCEV *LHS, const SCEV *RHS) {
11873 // No need to even try if we know the module has no guards.
11874 if (!HasGuards)
11875 return false;
11876
11877 return any_of(*BB, [&](const Instruction &I) {
11878 using namespace llvm::PatternMatch;
11879
11880 Value *Condition;
11882 m_Value(Condition))) &&
11883 isImpliedCond(Pred, LHS, RHS, Condition, false);
11884 });
11885}
11886
11887/// isLoopBackedgeGuardedByCond - Test whether the backedge of the loop is
11888/// protected by a conditional between LHS and RHS. This is used to
11889/// to eliminate casts.
11891 CmpPredicate Pred,
11892 const SCEV *LHS,
11893 const SCEV *RHS) {
11894 // Interpret a null as meaning no loop, where there is obviously no guard
11895 // (interprocedural conditions notwithstanding). Do not bother about
11896 // unreachable loops.
11897 if (!L || !DT.isReachableFromEntry(L->getHeader()))
11898 return true;
11899
11900 if (VerifyIR)
11901 assert(!verifyFunction(*L->getHeader()->getParent(), &dbgs()) &&
11902 "This cannot be done on broken IR!");
11903
11904
11905 if (isKnownViaNonRecursiveReasoning(Pred, LHS, RHS))
11906 return true;
11907
11908 BasicBlock *Latch = L->getLoopLatch();
11909 if (!Latch)
11910 return false;
11911
11912 CondBrInst *LoopContinuePredicate =
11914 if (LoopContinuePredicate &&
11915 isImpliedCond(Pred, LHS, RHS, LoopContinuePredicate->getCondition(),
11916 LoopContinuePredicate->getSuccessor(0) != L->getHeader()))
11917 return true;
11918
11919 // We don't want more than one activation of the following loops on the stack
11920 // -- that can lead to O(n!) time complexity.
11921 if (WalkingBEDominatingConds)
11922 return false;
11923
11924 SaveAndRestore ClearOnExit(WalkingBEDominatingConds, true);
11925
11926 // See if we can exploit a trip count to prove the predicate.
11927 const auto &BETakenInfo = getBackedgeTakenInfo(L);
11928 const SCEV *LatchBECount = BETakenInfo.getExact(Latch, this);
11929 if (LatchBECount != getCouldNotCompute()) {
11930 // We know that Latch branches back to the loop header exactly
11931 // LatchBECount times. This means the backdege condition at Latch is
11932 // equivalent to "{0,+,1} u< LatchBECount".
11933 Type *Ty = LatchBECount->getType();
11934 auto NoWrapFlags = SCEVFlags(SCEV::FlagNUW | SCEV::FlagNW);
11935 const SCEV *LoopCounter =
11936 getAddRecExpr(getZero(Ty), getOne(Ty), L, NoWrapFlags);
11937 if (isImpliedCond(Pred, LHS, RHS, ICmpInst::ICMP_ULT, LoopCounter,
11938 LatchBECount))
11939 return true;
11940 }
11941
11942 // Check conditions due to any @llvm.assume intrinsics.
11943 for (auto &AssumeVH : AC.assumptions()) {
11944 if (!AssumeVH)
11945 continue;
11946 auto *CI = cast<CallInst>(AssumeVH);
11947 if (!DT.dominates(CI, Latch->getTerminator()))
11948 continue;
11949
11950 if (isImpliedCond(Pred, LHS, RHS, CI->getArgOperand(0), false))
11951 return true;
11952 }
11953
11954 if (isImpliedViaGuard(Latch, Pred, LHS, RHS))
11955 return true;
11956
11957 for (DomTreeNode *DTN = DT[Latch], *HeaderDTN = DT[L->getHeader()];
11958 DTN != HeaderDTN; DTN = DTN->getIDom()) {
11959 assert(DTN && "should reach the loop header before reaching the root!");
11960
11961 BasicBlock *BB = DTN->getBlock();
11962 if (isImpliedViaGuard(BB, Pred, LHS, RHS))
11963 return true;
11964
11965 BasicBlock *PBB = BB->getSinglePredecessor();
11966 if (!PBB)
11967 continue;
11968
11970 if (!ContBr || ContBr->getSuccessor(0) == ContBr->getSuccessor(1))
11971 continue;
11972
11973 // If we have an edge `E` within the loop body that dominates the only
11974 // latch, the condition guarding `E` also guards the backedge. This
11975 // reasoning works only for loops with a single latch.
11976 // We're constructively (and conservatively) enumerating edges within the
11977 // loop body that dominate the latch. The dominator tree better agree
11978 // with us on this:
11979 assert(DT.dominates(BasicBlockEdge(PBB, BB), Latch) && "should be!");
11980 if (isImpliedCond(Pred, LHS, RHS, ContBr->getCondition(),
11981 BB != ContBr->getSuccessor(0)))
11982 return true;
11983 }
11984
11985 return false;
11986}
11987
11989 CmpPredicate Pred,
11990 const SCEV *LHS,
11991 const SCEV *RHS) {
11992 // Do not bother proving facts for unreachable code.
11993 if (!DT.isReachableFromEntry(BB))
11994 return true;
11995 if (VerifyIR)
11996 assert(!verifyFunction(*BB->getParent(), &dbgs()) &&
11997 "This cannot be done on broken IR!");
11998
11999 // If we cannot prove strict comparison (e.g. a > b), maybe we can prove
12000 // the facts (a >= b && a != b) separately. A typical situation is when the
12001 // non-strict comparison is known from ranges and non-equality is known from
12002 // dominating predicates. If we are proving strict comparison, we always try
12003 // to prove non-equality and non-strict comparison separately.
12004 CmpPredicate NonStrictPredicate = ICmpInst::getNonStrictCmpPredicate(Pred);
12005 const bool ProvingStrictComparison =
12006 Pred != NonStrictPredicate.dropSameSign();
12007 bool ProvedNonStrictComparison = false;
12008 bool ProvedNonEquality = false;
12009
12010 auto SplitAndProve = [&](std::function<bool(CmpPredicate)> Fn) -> bool {
12011 if (!ProvedNonStrictComparison)
12012 ProvedNonStrictComparison = Fn(NonStrictPredicate);
12013 if (!ProvedNonEquality)
12014 ProvedNonEquality = Fn(ICmpInst::ICMP_NE);
12015 if (ProvedNonStrictComparison && ProvedNonEquality)
12016 return true;
12017 return false;
12018 };
12019
12020 if (ProvingStrictComparison) {
12021 auto ProofFn = [&](CmpPredicate P) {
12022 return isKnownViaNonRecursiveReasoning(P, LHS, RHS);
12023 };
12024 if (SplitAndProve(ProofFn))
12025 return true;
12026 }
12027
12028 // Try to prove (Pred, LHS, RHS) using isImpliedCond.
12029 auto ProveViaCond = [&](const Value *Condition, bool Inverse) {
12030 const Instruction *CtxI = &BB->front();
12031 if (isImpliedCond(Pred, LHS, RHS, Condition, Inverse, CtxI))
12032 return true;
12033 if (ProvingStrictComparison) {
12034 auto ProofFn = [&](CmpPredicate P) {
12035 return isImpliedCond(P, LHS, RHS, Condition, Inverse, CtxI);
12036 };
12037 if (SplitAndProve(ProofFn))
12038 return true;
12039 }
12040 return false;
12041 };
12042
12043 // Starting at the block's predecessor, climb up the predecessor chain, as long
12044 // as there are predecessors that can be found that have unique successors
12045 // leading to the original block.
12046 const Loop *ContainingLoop = LI.getLoopFor(BB);
12047 const BasicBlock *PredBB;
12048 if (ContainingLoop && ContainingLoop->getHeader() == BB)
12049 PredBB = ContainingLoop->getLoopPredecessor();
12050 else
12051 PredBB = BB->getSinglePredecessor();
12052 for (std::pair<const BasicBlock *, const BasicBlock *> Pair(PredBB, BB);
12053 Pair.first; Pair = getPredecessorWithUniqueSuccessorForBB(Pair.first)) {
12054 const CondBrInst *BlockEntryPredicate =
12055 dyn_cast<CondBrInst>(Pair.first->getTerminator());
12056 if (!BlockEntryPredicate)
12057 continue;
12058
12059 if (ProveViaCond(BlockEntryPredicate->getCondition(),
12060 BlockEntryPredicate->getSuccessor(0) != Pair.second))
12061 return true;
12062 }
12063
12064 // Check conditions due to any @llvm.assume intrinsics.
12065 for (auto &AssumeVH : AC.assumptions()) {
12066 if (!AssumeVH)
12067 continue;
12068 auto *CI = cast<CallInst>(AssumeVH);
12069 if (!DT.dominates(CI, BB))
12070 continue;
12071
12072 if (ProveViaCond(CI->getArgOperand(0), false))
12073 return true;
12074 }
12075
12076 // Check conditions due to any @llvm.experimental.guard intrinsics.
12077 auto *GuardDecl = Intrinsic::getDeclarationIfExists(
12078 F.getParent(), Intrinsic::experimental_guard);
12079 if (GuardDecl)
12080 for (const auto *GU : GuardDecl->users())
12081 if (const auto *Guard = dyn_cast<IntrinsicInst>(GU))
12082 if (Guard->getFunction() == BB->getParent() && DT.dominates(Guard, BB))
12083 if (ProveViaCond(Guard->getArgOperand(0), false))
12084 return true;
12085 return false;
12086}
12087
12089 const SCEV *LHS,
12090 const SCEV *RHS) {
12091 // Interpret a null as meaning no loop, where there is obviously no guard
12092 // (interprocedural conditions notwithstanding).
12093 if (!L)
12094 return false;
12095
12096 // Both LHS and RHS must be available at loop entry.
12098 "LHS is not available at Loop Entry");
12100 "RHS is not available at Loop Entry");
12101
12102 if (isKnownViaNonRecursiveReasoning(Pred, LHS, RHS))
12103 return true;
12104
12105 return isBasicBlockEntryGuardedByCond(L->getHeader(), Pred, LHS, RHS);
12106}
12107
12108bool ScalarEvolution::isImpliedCond(CmpPredicate Pred, const SCEV *LHS,
12109 const SCEV *RHS,
12110 const Value *FoundCondValue, bool Inverse,
12111 const Instruction *CtxI) {
12112 // False conditions implies anything. Do not bother analyzing it further.
12113 if (FoundCondValue ==
12114 ConstantInt::getBool(FoundCondValue->getContext(), Inverse))
12115 return true;
12116
12117 if (!PendingLoopPredicates.insert(FoundCondValue).second)
12118 return false;
12119
12120 llvm::scope_exit ClearOnExit(
12121 [&]() { PendingLoopPredicates.erase(FoundCondValue); });
12122
12123 // Recursively handle And and Or conditions.
12124 const Value *Op0, *Op1;
12125 if (match(FoundCondValue, m_LogicalAnd(m_Value(Op0), m_Value(Op1)))) {
12126 if (!Inverse)
12127 return isImpliedCond(Pred, LHS, RHS, Op0, Inverse, CtxI) ||
12128 isImpliedCond(Pred, LHS, RHS, Op1, Inverse, CtxI);
12129 } else if (match(FoundCondValue, m_LogicalOr(m_Value(Op0), m_Value(Op1)))) {
12130 if (Inverse)
12131 return isImpliedCond(Pred, LHS, RHS, Op0, Inverse, CtxI) ||
12132 isImpliedCond(Pred, LHS, RHS, Op1, Inverse, CtxI);
12133 }
12134
12135 const ICmpInst *ICI = dyn_cast<ICmpInst>(FoundCondValue);
12136 if (!ICI) return false;
12137
12138 // Now that we found a conditional branch that dominates the loop or controls
12139 // the loop latch. Check to see if it is the comparison we are looking for.
12140 CmpPredicate FoundPred;
12141 if (Inverse)
12142 FoundPred = ICI->getInverseCmpPredicate();
12143 else
12144 FoundPred = ICI->getCmpPredicate();
12145
12146 const SCEV *FoundLHS = getSCEV(ICI->getOperand(0));
12147 const SCEV *FoundRHS = getSCEV(ICI->getOperand(1));
12148
12149 return isImpliedCond(Pred, LHS, RHS, FoundPred, FoundLHS, FoundRHS, CtxI);
12150}
12151
12152bool ScalarEvolution::isImpliedCond(CmpPredicate Pred, const SCEV *LHS,
12153 const SCEV *RHS, CmpPredicate FoundPred,
12154 const SCEV *FoundLHS, const SCEV *FoundRHS,
12155 const Instruction *CtxI) {
12156 // Balance the types.
12157 if (getTypeSizeInBits(LHS->getType()) <
12158 getTypeSizeInBits(FoundLHS->getType())) {
12159 // For unsigned and equality predicates, try to prove that both found
12160 // operands fit into narrow unsigned range. If so, try to prove facts in
12161 // narrow types.
12162 if (!CmpInst::isSigned(FoundPred) && !FoundLHS->getType()->isPointerTy() &&
12163 !FoundRHS->getType()->isPointerTy()) {
12164 auto *NarrowType = LHS->getType();
12165 auto *WideType = FoundLHS->getType();
12166 auto BitWidth = getTypeSizeInBits(NarrowType);
12167 const SCEV *MaxValue = getZeroExtendExpr(
12169 if (isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_ULE, FoundLHS,
12170 MaxValue) &&
12171 isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_ULE, FoundRHS,
12172 MaxValue)) {
12173 const SCEV *TruncFoundLHS = getTruncateExpr(FoundLHS, NarrowType);
12174 const SCEV *TruncFoundRHS = getTruncateExpr(FoundRHS, NarrowType);
12175 // We cannot preserve samesign after truncation.
12176 if (isImpliedCondBalancedTypes(Pred, LHS, RHS, FoundPred.dropSameSign(),
12177 TruncFoundLHS, TruncFoundRHS, CtxI))
12178 return true;
12179 }
12180 }
12181
12182 if (LHS->getType()->isPointerTy() || RHS->getType()->isPointerTy())
12183 return false;
12184 if (CmpInst::isSigned(Pred)) {
12185 LHS = getSignExtendExpr(LHS, FoundLHS->getType());
12186 RHS = getSignExtendExpr(RHS, FoundLHS->getType());
12187 } else {
12188 LHS = getZeroExtendExpr(LHS, FoundLHS->getType());
12189 RHS = getZeroExtendExpr(RHS, FoundLHS->getType());
12190 }
12191 } else if (getTypeSizeInBits(LHS->getType()) >
12192 getTypeSizeInBits(FoundLHS->getType())) {
12193 if (FoundLHS->getType()->isPointerTy() || FoundRHS->getType()->isPointerTy())
12194 return false;
12195 if (CmpInst::isSigned(FoundPred)) {
12196 FoundLHS = getSignExtendExpr(FoundLHS, LHS->getType());
12197 FoundRHS = getSignExtendExpr(FoundRHS, LHS->getType());
12198 } else {
12199 FoundLHS = getZeroExtendExpr(FoundLHS, LHS->getType());
12200 FoundRHS = getZeroExtendExpr(FoundRHS, LHS->getType());
12201 }
12202 }
12203 return isImpliedCondBalancedTypes(Pred, LHS, RHS, FoundPred, FoundLHS,
12204 FoundRHS, CtxI);
12205}
12206
12207bool ScalarEvolution::isImpliedCondBalancedTypes(
12208 CmpPredicate Pred, SCEVUse LHS, SCEVUse RHS, CmpPredicate FoundPred,
12209 SCEVUse FoundLHS, SCEVUse FoundRHS, const Instruction *CtxI) {
12211 getTypeSizeInBits(FoundLHS->getType()) &&
12212 "Types should be balanced!");
12213 // Canonicalize the query to match the way instcombine will have
12214 // canonicalized the comparison.
12215 if (SimplifyICmpOperands(Pred, LHS, RHS))
12216 if (LHS == RHS)
12217 return CmpInst::isTrueWhenEqual(Pred);
12218 if (SimplifyICmpOperands(FoundPred, FoundLHS, FoundRHS))
12219 if (FoundLHS == FoundRHS)
12220 return CmpInst::isFalseWhenEqual(FoundPred);
12221
12222 // Check to see if we can make the LHS or RHS match.
12223 if (LHS == FoundRHS || RHS == FoundLHS) {
12224 if (isa<SCEVConstant>(RHS)) {
12225 std::swap(FoundLHS, FoundRHS);
12226 FoundPred = ICmpInst::getSwappedCmpPredicate(FoundPred);
12227 } else {
12228 std::swap(LHS, RHS);
12230 }
12231 }
12232
12233 // Check whether the found predicate is the same as the desired predicate.
12234 if (auto P = CmpPredicate::getMatching(FoundPred, Pred))
12235 return isImpliedCondOperands(*P, LHS, RHS, FoundLHS, FoundRHS, CtxI);
12236
12237 // Check whether swapping the found predicate makes it the same as the
12238 // desired predicate.
12239 if (auto P = CmpPredicate::getMatching(
12240 ICmpInst::getSwappedCmpPredicate(FoundPred), Pred)) {
12241 // We can write the implication
12242 // 0. LHS Pred RHS <- FoundLHS SwapPred FoundRHS
12243 // using one of the following ways:
12244 // 1. LHS Pred RHS <- FoundRHS Pred FoundLHS
12245 // 2. RHS SwapPred LHS <- FoundLHS SwapPred FoundRHS
12246 // Both require swapping the operands of one condition. Don't do this if it
12247 // would break canonical constant/addrec ordering.
12249 return isImpliedCondOperands(ICmpInst::getSwappedCmpPredicate(*P), RHS,
12250 LHS, FoundLHS, FoundRHS, CtxI);
12251 if (!isa<SCEVConstant>(FoundRHS) && !isa<SCEVAddRecExpr>(FoundLHS))
12252 return isImpliedCondOperands(*P, LHS, RHS, FoundRHS, FoundLHS, CtxI);
12253
12254 return false;
12255 }
12256
12257 auto IsSignFlippedPredicate = [](CmpInst::Predicate P1,
12259 assert(P1 != P2 && "Handled earlier!");
12260 return CmpInst::isRelational(P2) &&
12262 };
12263 if (IsSignFlippedPredicate(Pred, FoundPred)) {
12264 // Unsigned comparison is the same as signed comparison when both the
12265 // operands are non-negative or negative.
12266 if (haveSameSign(FoundLHS, FoundRHS))
12267 return isImpliedCondOperands(Pred, LHS, RHS, FoundLHS, FoundRHS, CtxI);
12268 // Create local copies that we can freely swap and canonicalize our
12269 // conditions to "le/lt".
12270 CmpPredicate CanonicalPred = Pred, CanonicalFoundPred = FoundPred;
12271 const SCEV *CanonicalLHS = LHS, *CanonicalRHS = RHS,
12272 *CanonicalFoundLHS = FoundLHS, *CanonicalFoundRHS = FoundRHS;
12273 if (ICmpInst::isGT(CanonicalPred) || ICmpInst::isGE(CanonicalPred)) {
12274 CanonicalPred = ICmpInst::getSwappedCmpPredicate(CanonicalPred);
12275 CanonicalFoundPred = ICmpInst::getSwappedCmpPredicate(CanonicalFoundPred);
12276 std::swap(CanonicalLHS, CanonicalRHS);
12277 std::swap(CanonicalFoundLHS, CanonicalFoundRHS);
12278 }
12279 assert((ICmpInst::isLT(CanonicalPred) || ICmpInst::isLE(CanonicalPred)) &&
12280 "Must be!");
12281 assert((ICmpInst::isLT(CanonicalFoundPred) ||
12282 ICmpInst::isLE(CanonicalFoundPred)) &&
12283 "Must be!");
12284 if (ICmpInst::isSigned(CanonicalPred) && isKnownNonNegative(CanonicalRHS))
12285 // Use implication:
12286 // x <u y && y >=s 0 --> x <s y.
12287 // If we can prove the left part, the right part is also proven.
12288 return isImpliedCondOperands(CanonicalFoundPred, CanonicalLHS,
12289 CanonicalRHS, CanonicalFoundLHS,
12290 CanonicalFoundRHS);
12291 if (ICmpInst::isUnsigned(CanonicalPred) && isKnownNegative(CanonicalRHS))
12292 // Use implication:
12293 // x <s y && y <s 0 --> x <u y.
12294 // If we can prove the left part, the right part is also proven.
12295 return isImpliedCondOperands(CanonicalFoundPred, CanonicalLHS,
12296 CanonicalRHS, CanonicalFoundLHS,
12297 CanonicalFoundRHS);
12298 }
12299
12300 // Check if we can make progress by sharpening ranges.
12301 if (FoundPred == ICmpInst::ICMP_NE &&
12302 (isa<SCEVConstant>(FoundLHS) || isa<SCEVConstant>(FoundRHS))) {
12303
12304 const SCEVConstant *C = nullptr;
12305 const SCEV *V = nullptr;
12306
12307 if (isa<SCEVConstant>(FoundLHS)) {
12308 C = cast<SCEVConstant>(FoundLHS);
12309 V = FoundRHS;
12310 } else {
12311 C = cast<SCEVConstant>(FoundRHS);
12312 V = FoundLHS;
12313 }
12314
12315 // The guarding predicate tells us that C != V. If the known range
12316 // of V is [C, t), we can sharpen the range to [C + 1, t). The
12317 // range we consider has to correspond to same signedness as the
12318 // predicate we're interested in folding.
12319
12320 APInt Min = ICmpInst::isSigned(Pred) ?
12322
12323 if (Min == C->getAPInt()) {
12324 // Given (V >= Min && V != Min) we conclude V >= (Min + 1).
12325 // This is true even if (Min + 1) wraps around -- in case of
12326 // wraparound, (Min + 1) < Min, so (V >= Min => V >= (Min + 1)).
12327
12328 APInt SharperMin = Min + 1;
12329
12330 switch (Pred) {
12331 case ICmpInst::ICMP_SGE:
12332 case ICmpInst::ICMP_UGE:
12333 // We know V `Pred` SharperMin. If this implies LHS `Pred`
12334 // RHS, we're done.
12335 if (isImpliedCondOperands(Pred, LHS, RHS, V, getConstant(SharperMin),
12336 CtxI))
12337 return true;
12338 [[fallthrough]];
12339
12340 case ICmpInst::ICMP_SGT:
12341 case ICmpInst::ICMP_UGT:
12342 // We know from the range information that (V `Pred` Min ||
12343 // V == Min). We know from the guarding condition that !(V
12344 // == Min). This gives us
12345 //
12346 // V `Pred` Min || V == Min && !(V == Min)
12347 // => V `Pred` Min
12348 //
12349 // If V `Pred` Min implies LHS `Pred` RHS, we're done.
12350
12351 if (isImpliedCondOperands(Pred, LHS, RHS, V, getConstant(Min), CtxI))
12352 return true;
12353 break;
12354
12355 // `LHS < RHS` and `LHS <= RHS` are handled in the same way as `RHS > LHS` and `RHS >= LHS` respectively.
12356 case ICmpInst::ICMP_SLE:
12357 case ICmpInst::ICMP_ULE:
12358 if (isImpliedCondOperands(ICmpInst::getSwappedCmpPredicate(Pred), RHS,
12359 LHS, V, getConstant(SharperMin), CtxI))
12360 return true;
12361 [[fallthrough]];
12362
12363 case ICmpInst::ICMP_SLT:
12364 case ICmpInst::ICMP_ULT:
12365 if (isImpliedCondOperands(ICmpInst::getSwappedCmpPredicate(Pred), RHS,
12366 LHS, V, getConstant(Min), CtxI))
12367 return true;
12368 break;
12369
12370 default:
12371 // No change
12372 break;
12373 }
12374 }
12375 }
12376
12377 // Check whether the actual condition is beyond sufficient.
12378 if (FoundPred == ICmpInst::ICMP_EQ)
12379 if (ICmpInst::isTrueWhenEqual(Pred))
12380 if (isImpliedCondOperands(Pred, LHS, RHS, FoundLHS, FoundRHS, CtxI))
12381 return true;
12382 if (Pred == ICmpInst::ICMP_NE)
12383 if (!ICmpInst::isTrueWhenEqual(FoundPred))
12384 if (isImpliedCondOperands(FoundPred, LHS, RHS, FoundLHS, FoundRHS, CtxI))
12385 return true;
12386
12387 if (isImpliedCondOperandsViaRanges(Pred, LHS, RHS, FoundPred, FoundLHS, FoundRHS))
12388 return true;
12389
12390 // Otherwise assume the worst.
12391 return false;
12392}
12393
12394bool ScalarEvolution::splitBinaryAdd(SCEVUse Expr, SCEVUse &L, SCEVUse &R,
12395 SCEVFlags &Flags) {
12396 if (!match(Expr, m_scev_Add(m_SCEV(L), m_SCEV(R))))
12397 return false;
12398
12399 Flags = cast<SCEVAddExpr>(Expr)->getNoWrapFlags();
12400 return true;
12401}
12402
12403std::optional<APInt>
12405 // We avoid subtracting expressions here because this function is usually
12406 // fairly deep in the call stack (i.e. is called many times).
12407
12408 unsigned BW = getTypeSizeInBits(More->getType());
12409 APInt Diff(BW, 0);
12410 APInt DiffMul(BW, 1);
12411 // Try various simplifications to reduce the difference to a constant. Limit
12412 // the number of allowed simplifications to keep compile-time low.
12413 for (unsigned I = 0; I < 8; ++I) {
12414 if (More == Less)
12415 return Diff;
12416
12417 // Reduce addrecs with identical steps to their start value.
12419 const auto *LAR = cast<SCEVAddRecExpr>(Less);
12420 const auto *MAR = cast<SCEVAddRecExpr>(More);
12421
12422 if (LAR->getLoop() != MAR->getLoop())
12423 return std::nullopt;
12424
12425 // We look at affine expressions only; not for correctness but to keep
12426 // getStepRecurrence cheap.
12427 if (!LAR->isAffine() || !MAR->isAffine())
12428 return std::nullopt;
12429
12430 if (LAR->getStepRecurrence(*this) != MAR->getStepRecurrence(*this))
12431 return std::nullopt;
12432
12433 Less = LAR->getStart();
12434 More = MAR->getStart();
12435 continue;
12436 }
12437
12438 // Try to match a common constant multiply.
12439 auto MatchConstMul =
12440 [](const SCEV *S) -> std::optional<std::pair<const SCEV *, APInt>> {
12441 const APInt *C;
12442 const SCEV *Op;
12443 if (match(S, m_scev_Mul(m_scev_APInt(C), m_SCEV(Op))))
12444 return {{Op, *C}};
12445 return std::nullopt;
12446 };
12447 if (auto MatchedMore = MatchConstMul(More)) {
12448 if (auto MatchedLess = MatchConstMul(Less)) {
12449 if (MatchedMore->second == MatchedLess->second) {
12450 More = MatchedMore->first;
12451 Less = MatchedLess->first;
12452 DiffMul *= MatchedMore->second;
12453 continue;
12454 }
12455 }
12456 }
12457
12458 // Try to cancel out common factors in two add expressions.
12460 auto Add = [&](const SCEV *S, int Mul) {
12461 if (auto *C = dyn_cast<SCEVConstant>(S)) {
12462 if (Mul == 1) {
12463 Diff += C->getAPInt() * DiffMul;
12464 } else {
12465 assert(Mul == -1);
12466 Diff -= C->getAPInt() * DiffMul;
12467 }
12468 } else
12469 Multiplicity[S] += Mul;
12470 };
12471 auto Decompose = [&](const SCEV *S, int Mul) {
12472 if (isa<SCEVAddExpr>(S)) {
12473 for (const SCEV *Op : S->operands())
12474 Add(Op, Mul);
12475 } else
12476 Add(S, Mul);
12477 };
12478 Decompose(More, 1);
12479 Decompose(Less, -1);
12480
12481 // Check whether all the non-constants cancel out, or reduce to new
12482 // More/Less values.
12483 const SCEV *NewMore = nullptr, *NewLess = nullptr;
12484 for (const auto &[S, Mul] : Multiplicity) {
12485 if (Mul == 0)
12486 continue;
12487 if (Mul == 1) {
12488 if (NewMore)
12489 return std::nullopt;
12490 NewMore = S;
12491 } else if (Mul == -1) {
12492 if (NewLess)
12493 return std::nullopt;
12494 NewLess = S;
12495 } else
12496 return std::nullopt;
12497 }
12498
12499 // Values stayed the same, no point in trying further.
12500 if (NewMore == More || NewLess == Less)
12501 return std::nullopt;
12502
12503 More = NewMore;
12504 Less = NewLess;
12505
12506 // Reduced to constant.
12507 if (!More && !Less)
12508 return Diff;
12509
12510 // Left with variable on only one side, bail out.
12511 if (!More || !Less)
12512 return std::nullopt;
12513 }
12514
12515 // Did not reduce to constant.
12516 return std::nullopt;
12517}
12518
12519bool ScalarEvolution::isImpliedCondOperandsViaAddRecStart(
12520 CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, const SCEV *FoundLHS,
12521 const SCEV *FoundRHS, const Instruction *CtxI) {
12522 // Try to recognize the following pattern:
12523 //
12524 // FoundRHS = ...
12525 // ...
12526 // loop:
12527 // FoundLHS = {Start,+,W}
12528 // context_bb: // Basic block from the same loop
12529 // known(Pred, FoundLHS, FoundRHS)
12530 //
12531 // If some predicate is known in the context of a loop, it is also known on
12532 // each iteration of this loop, including the first iteration. Therefore, in
12533 // this case, `FoundLHS Pred FoundRHS` implies `Start Pred FoundRHS`. Try to
12534 // prove the original pred using this fact.
12535 if (!CtxI)
12536 return false;
12537 const BasicBlock *ContextBB = CtxI->getParent();
12538 // Make sure AR varies in the context block.
12539 if (auto *AR = dyn_cast<SCEVAddRecExpr>(FoundLHS)) {
12540 const Loop *L = AR->getLoop();
12541 const auto *Latch = L->getLoopLatch();
12542 // Make sure that context belongs to the loop and executes on 1st iteration
12543 // (if it ever executes at all).
12544 if (!L->contains(ContextBB) || !Latch || !DT.dominates(ContextBB, Latch))
12545 return false;
12546 if (!isAvailableAtLoopEntry(FoundRHS, AR->getLoop()))
12547 return false;
12548 return isImpliedCondOperands(Pred, LHS, RHS, AR->getStart(), FoundRHS);
12549 }
12550
12551 if (auto *AR = dyn_cast<SCEVAddRecExpr>(FoundRHS)) {
12552 const Loop *L = AR->getLoop();
12553 const auto *Latch = L->getLoopLatch();
12554 // Make sure that context belongs to the loop and executes on 1st iteration
12555 // (if it ever executes at all).
12556 if (!L->contains(ContextBB) || !Latch || !DT.dominates(ContextBB, Latch))
12557 return false;
12558 if (!isAvailableAtLoopEntry(FoundLHS, AR->getLoop()))
12559 return false;
12560 return isImpliedCondOperands(Pred, LHS, RHS, FoundLHS, AR->getStart());
12561 }
12562
12563 return false;
12564}
12565
12566bool ScalarEvolution::isImpliedCondOperandsViaNoOverflow(CmpPredicate Pred,
12567 const SCEV *LHS,
12568 const SCEV *RHS,
12569 const SCEV *FoundLHS,
12570 const SCEV *FoundRHS) {
12571 if (Pred != CmpInst::ICMP_SLT && Pred != CmpInst::ICMP_ULT)
12572 return false;
12573
12574 const auto *AddRecLHS = dyn_cast<SCEVAddRecExpr>(LHS);
12575 if (!AddRecLHS)
12576 return false;
12577
12578 const auto *AddRecFoundLHS = dyn_cast<SCEVAddRecExpr>(FoundLHS);
12579 if (!AddRecFoundLHS)
12580 return false;
12581
12582 // We'd like to let SCEV reason about control dependencies, so we constrain
12583 // both the inequalities to be about add recurrences on the same loop. This
12584 // way we can use isLoopEntryGuardedByCond later.
12585
12586 const Loop *L = AddRecFoundLHS->getLoop();
12587 if (L != AddRecLHS->getLoop())
12588 return false;
12589
12590 // FoundLHS u< FoundRHS u< -C => (FoundLHS + C) u< (FoundRHS + C) ... (1)
12591 //
12592 // FoundLHS s< FoundRHS s< INT_MIN - C => (FoundLHS + C) s< (FoundRHS + C)
12593 // ... (2)
12594 //
12595 // Informal proof for (2), assuming (1) [*]:
12596 //
12597 // We'll also assume (A s< B) <=> ((A + INT_MIN) u< (B + INT_MIN)) ... (3)[**]
12598 //
12599 // Then
12600 //
12601 // FoundLHS s< FoundRHS s< INT_MIN - C
12602 // <=> (FoundLHS + INT_MIN) u< (FoundRHS + INT_MIN) u< -C [ using (3) ]
12603 // <=> (FoundLHS + INT_MIN + C) u< (FoundRHS + INT_MIN + C) [ using (1) ]
12604 // <=> (FoundLHS + INT_MIN + C + INT_MIN) s<
12605 // (FoundRHS + INT_MIN + C + INT_MIN) [ using (3) ]
12606 // <=> FoundLHS + C s< FoundRHS + C
12607 //
12608 // [*]: (1) can be proved by ruling out overflow.
12609 //
12610 // [**]: This can be proved by analyzing all the four possibilities:
12611 // (A s< 0, B s< 0), (A s< 0, B s>= 0), (A s>= 0, B s< 0) and
12612 // (A s>= 0, B s>= 0).
12613 //
12614 // Note:
12615 // Despite (2), "FoundRHS s< INT_MIN - C" does not mean that "FoundRHS + C"
12616 // will not sign underflow. For instance, say FoundLHS = (i8 -128), FoundRHS
12617 // = (i8 -127) and C = (i8 -100). Then INT_MIN - C = (i8 -28), and FoundRHS
12618 // s< (INT_MIN - C). Lack of sign overflow / underflow in "FoundRHS + C" is
12619 // neither necessary nor sufficient to prove "(FoundLHS + C) s< (FoundRHS +
12620 // C)".
12621
12622 std::optional<APInt> LDiff = computeConstantDifference(LHS, FoundLHS);
12623 if (!LDiff)
12624 return false;
12625 std::optional<APInt> RDiff = computeConstantDifference(RHS, FoundRHS);
12626 if (!RDiff || *LDiff != *RDiff)
12627 return false;
12628
12629 if (LDiff->isMinValue())
12630 return true;
12631
12632 APInt FoundRHSLimit;
12633
12634 if (Pred == CmpInst::ICMP_ULT) {
12635 FoundRHSLimit = -(*RDiff);
12636 } else {
12637 assert(Pred == CmpInst::ICMP_SLT && "Checked above!");
12638 FoundRHSLimit = APInt::getSignedMinValue(getTypeSizeInBits(RHS->getType())) - *RDiff;
12639 }
12640
12641 // Try to prove (1) or (2), as needed.
12642 return isAvailableAtLoopEntry(FoundRHS, L) &&
12643 isLoopEntryGuardedByCond(L, Pred, FoundRHS,
12644 getConstant(FoundRHSLimit));
12645}
12646
12647bool ScalarEvolution::isImpliedViaMerge(CmpPredicate Pred, const SCEV *LHS,
12648 const SCEV *RHS, const SCEV *FoundLHS,
12649 const SCEV *FoundRHS, unsigned Depth) {
12650 const PHINode *LPhi = nullptr, *RPhi = nullptr;
12651
12652 llvm::scope_exit ClearOnExit([&]() {
12653 if (LPhi) {
12654 bool Erased = PendingMerges.erase(LPhi);
12655 assert(Erased && "Failed to erase LPhi!");
12656 (void)Erased;
12657 }
12658 if (RPhi) {
12659 bool Erased = PendingMerges.erase(RPhi);
12660 assert(Erased && "Failed to erase RPhi!");
12661 (void)Erased;
12662 }
12663 });
12664
12665 // Find respective Phis and check that they are not being pending.
12666 if (const SCEVUnknown *LU = dyn_cast<SCEVUnknown>(LHS))
12667 if (auto *Phi = dyn_cast<PHINode>(LU->getValue())) {
12668 if (!PendingMerges.insert(Phi).second)
12669 return false;
12670 LPhi = Phi;
12671 }
12672 if (const SCEVUnknown *RU = dyn_cast<SCEVUnknown>(RHS))
12673 if (auto *Phi = dyn_cast<PHINode>(RU->getValue())) {
12674 // If we detect a loop of Phi nodes being processed by this method, for
12675 // example:
12676 //
12677 // %a = phi i32 [ %some1, %preheader ], [ %b, %latch ]
12678 // %b = phi i32 [ %some2, %preheader ], [ %a, %latch ]
12679 //
12680 // we don't want to deal with a case that complex, so return conservative
12681 // answer false.
12682 if (!PendingMerges.insert(Phi).second)
12683 return false;
12684 RPhi = Phi;
12685 }
12686
12687 // If none of LHS, RHS is a Phi, nothing to do here.
12688 if (!LPhi && !RPhi)
12689 return false;
12690
12691 // If there is a SCEVUnknown Phi we are interested in, make it left.
12692 if (!LPhi) {
12693 std::swap(LHS, RHS);
12694 std::swap(FoundLHS, FoundRHS);
12695 std::swap(LPhi, RPhi);
12697 }
12698
12699 assert(LPhi && "LPhi should definitely be a SCEVUnknown Phi!");
12700 const BasicBlock *LBB = LPhi->getParent();
12701 const SCEVAddRecExpr *RAR = dyn_cast<SCEVAddRecExpr>(RHS);
12702
12703 auto ProvedEasily = [&](const SCEV *S1, const SCEV *S2) {
12704 return isKnownViaNonRecursiveReasoning(Pred, S1, S2) ||
12705 isImpliedCondOperandsViaRanges(Pred, S1, S2, Pred, FoundLHS, FoundRHS) ||
12706 isImpliedViaOperations(Pred, S1, S2, FoundLHS, FoundRHS, Depth);
12707 };
12708
12709 if (RPhi && RPhi->getParent() == LBB) {
12710 // Case one: RHS is also a SCEVUnknown Phi from the same basic block.
12711 // If we compare two Phis from the same block, and for each entry block
12712 // the predicate is true for incoming values from this block, then the
12713 // predicate is also true for the Phis.
12714 for (const BasicBlock *IncBB : predecessors(LBB)) {
12715 const SCEV *L = getSCEV(LPhi->getIncomingValueForBlock(IncBB));
12716 const SCEV *R = getSCEV(RPhi->getIncomingValueForBlock(IncBB));
12717 if (!ProvedEasily(L, R))
12718 return false;
12719 }
12720 } else if (RAR && RAR->getLoop()->getHeader() == LBB) {
12721 // Case two: RHS is also a Phi from the same basic block, and it is an
12722 // AddRec. It means that there is a loop which has both AddRec and Unknown
12723 // PHIs, for it we can compare incoming values of AddRec from above the loop
12724 // and latch with their respective incoming values of LPhi.
12725 // TODO: Generalize to handle loops with many inputs in a header.
12726 if (LPhi->getNumIncomingValues() != 2) return false;
12727
12728 auto *RLoop = RAR->getLoop();
12729 auto *Predecessor = RLoop->getLoopPredecessor();
12730 assert(Predecessor && "Loop with AddRec with no predecessor?");
12731 const SCEV *L1 = getSCEV(LPhi->getIncomingValueForBlock(Predecessor));
12732 if (!ProvedEasily(L1, RAR->getStart()))
12733 return false;
12734 auto *Latch = RLoop->getLoopLatch();
12735 assert(Latch && "Loop with AddRec with no latch?");
12736 const SCEV *L2 = getSCEV(LPhi->getIncomingValueForBlock(Latch));
12737 if (!ProvedEasily(L2, RAR->getPostIncExpr(*this)))
12738 return false;
12739 } else {
12740 // In all other cases go over inputs of LHS and compare each of them to RHS,
12741 // the predicate is true for (LHS, RHS) if it is true for all such pairs.
12742 // At this point RHS is either a non-Phi, or it is a Phi from some block
12743 // different from LBB.
12744 for (const BasicBlock *IncBB : predecessors(LBB)) {
12745 // Check that RHS is available in this block.
12746 if (!dominates(RHS, IncBB))
12747 return false;
12748 const SCEV *L = getSCEV(LPhi->getIncomingValueForBlock(IncBB));
12749 // Make sure L does not refer to a value from a potentially previous
12750 // iteration of a loop.
12751 if (!properlyDominates(L, LBB))
12752 return false;
12753 // Addrecs are considered to properly dominate their loop, so are missed
12754 // by the previous check. Discard any values that have computable
12755 // evolution in this loop.
12756 if (auto *Loop = LI.getLoopFor(LBB))
12758 return false;
12759 if (!ProvedEasily(L, RHS))
12760 return false;
12761 }
12762 }
12763 return true;
12764}
12765
12766bool ScalarEvolution::isImpliedCondOperandsViaShift(CmpPredicate Pred,
12767 const SCEV *LHS,
12768 const SCEV *RHS,
12769 const SCEV *FoundLHS,
12770 const SCEV *FoundRHS) {
12771 // We want to imply LHS < RHS from LHS < (RHS >> shiftvalue). First, make
12772 // sure that we are dealing with same LHS.
12773 if (RHS == FoundRHS) {
12774 std::swap(LHS, RHS);
12775 std::swap(FoundLHS, FoundRHS);
12777 }
12778 if (LHS != FoundLHS)
12779 return false;
12780
12781 auto *SUFoundRHS = dyn_cast<SCEVUnknown>(FoundRHS);
12782 if (!SUFoundRHS)
12783 return false;
12784
12785 Value *Shiftee, *ShiftValue;
12786
12787 using namespace PatternMatch;
12788 if (match(SUFoundRHS->getValue(),
12789 m_LShr(m_Value(Shiftee), m_Value(ShiftValue)))) {
12790 auto *ShifteeS = getSCEV(Shiftee);
12791 // Prove one of the following:
12792 // LHS <u (shiftee >> shiftvalue) && shiftee <=u RHS ---> LHS <u RHS
12793 // LHS <=u (shiftee >> shiftvalue) && shiftee <=u RHS ---> LHS <=u RHS
12794 // LHS <s (shiftee >> shiftvalue) && shiftee <=s RHS && shiftee >=s 0
12795 // ---> LHS <s RHS
12796 // LHS <=s (shiftee >> shiftvalue) && shiftee <=s RHS && shiftee >=s 0
12797 // ---> LHS <=s RHS
12798 if (Pred == ICmpInst::ICMP_ULT || Pred == ICmpInst::ICMP_ULE)
12799 return isKnownPredicate(ICmpInst::ICMP_ULE, ShifteeS, RHS);
12800 if (Pred == ICmpInst::ICMP_SLT || Pred == ICmpInst::ICMP_SLE)
12801 if (isKnownNonNegative(ShifteeS))
12802 return isKnownPredicate(ICmpInst::ICMP_SLE, ShifteeS, RHS);
12803 }
12804
12805 return false;
12806}
12807
12808bool ScalarEvolution::isImpliedCondOperandsViaMatchingDiff(
12809 CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, const SCEV *FoundLHS,
12810 const SCEV *FoundRHS) {
12811 // Only valid for equality predicates: (A == B) implies (C == D) when
12812 // the SCEV difference A - B equals C - D (they check the same
12813 // underlying relationship at every iteration).
12814 if (!ICmpInst::isEquality(Pred))
12815 return false;
12816
12817 // Restrict to cases involving loop recurrences - that's where this
12818 // pattern arises (correlated IV comparisons). This avoids calling
12819 // getMinusSCEV on arbitrary non-loop expressions.
12821 (!isa<SCEVAddRecExpr>(FoundLHS) && !isa<SCEVAddRecExpr>(FoundRHS)))
12822 return false;
12823
12824 // AddRecs from different loops can never produce matching differences.
12825 const SCEVAddRecExpr *QueryAddRec = dyn_cast<SCEVAddRecExpr>(LHS);
12826 if (!QueryAddRec)
12827 QueryAddRec = cast<SCEVAddRecExpr>(RHS);
12828 const SCEVAddRecExpr *FoundAddRec = dyn_cast<SCEVAddRecExpr>(FoundLHS);
12829 if (!FoundAddRec)
12830 FoundAddRec = cast<SCEVAddRecExpr>(FoundRHS);
12831 if (QueryAddRec->getLoop() != FoundAddRec->getLoop())
12832 return false;
12833
12834 // If the strides differ, the differences can never match.
12835 if (QueryAddRec->getStepRecurrence(*this) !=
12836 FoundAddRec->getStepRecurrence(*this))
12837 return false;
12838
12839 // Compute differences. For pointer-typed operands sharing the same base,
12840 // getMinusSCEV strips the common base and returns an integer SCEV.
12841 // For example, {base,+,8} - (base+8*n) = {-8n,+,8}
12842 const SCEV *FoundDiff = getMinusSCEV(FoundLHS, FoundRHS);
12843 if (isa<SCEVCouldNotCompute>(FoundDiff))
12844 return false;
12845
12846 const SCEV *Diff = getMinusSCEV(LHS, RHS);
12847 if (isa<SCEVCouldNotCompute>(Diff))
12848 return false;
12849
12850 return Diff == FoundDiff;
12851}
12852
12853bool ScalarEvolution::isImpliedCondOperands(CmpPredicate Pred, const SCEV *LHS,
12854 const SCEV *RHS,
12855 const SCEV *FoundLHS,
12856 const SCEV *FoundRHS,
12857 const Instruction *CtxI) {
12858 return isImpliedCondOperandsViaRanges(Pred, LHS, RHS, Pred, FoundLHS,
12859 FoundRHS) ||
12860 isImpliedCondOperandsViaNoOverflow(Pred, LHS, RHS, FoundLHS,
12861 FoundRHS) ||
12862 isImpliedCondOperandsViaShift(Pred, LHS, RHS, FoundLHS, FoundRHS) ||
12863 isImpliedCondOperandsViaAddRecStart(Pred, LHS, RHS, FoundLHS, FoundRHS,
12864 CtxI) ||
12865 isImpliedCondOperandsViaMatchingDiff(Pred, LHS, RHS, FoundLHS,
12866 FoundRHS) ||
12867 isImpliedCondOperandsHelper(Pred, LHS, RHS, FoundLHS, FoundRHS);
12868}
12869
12870/// Is MaybeMinMaxExpr an (U|S)(Min|Max) of Candidate and some other values?
12871template <typename MinMaxExprType>
12872static bool IsMinMaxConsistingOf(const SCEV *MaybeMinMaxExpr,
12873 const SCEV *Candidate) {
12874 const MinMaxExprType *MinMaxExpr = dyn_cast<MinMaxExprType>(MaybeMinMaxExpr);
12875 if (!MinMaxExpr)
12876 return false;
12877
12878 return is_contained(MinMaxExpr->operands(), Candidate);
12879}
12880
12882 CmpPredicate Pred, const SCEV *LHS,
12883 const SCEV *RHS) {
12884 // If both sides are affine addrecs for the same loop, with equal
12885 // steps, and we know the recurrences don't wrap, then we only
12886 // need to check the predicate on the starting values.
12887
12888 if (!ICmpInst::isRelational(Pred))
12889 return false;
12890
12891 const SCEV *LStart, *RStart, *Step;
12892 const Loop *L;
12893 if (!match(LHS,
12894 m_scev_AffineAddRec(m_SCEV(LStart), m_SCEV(Step), m_Loop(L))) ||
12896 m_SpecificLoop(L))))
12897 return false;
12901 if (!LAR->getNoWrapFlags(NW) || !RAR->getNoWrapFlags(NW))
12902 return false;
12903
12904 return SE.isKnownPredicate(Pred, LStart, RStart);
12905}
12906
12907/// Is LHS `Pred` RHS true because one of them is an AddRec that is known not to
12908/// go below its own start value?
12910 CmpPredicate Pred,
12911 const SCEV *LHS,
12912 const SCEV *RHS) {
12913 // Normalize to (AddRec Pred Start).
12916 std::swap(LHS, RHS);
12917 }
12918
12919 // The recurrence is equal to Start in the first iteration, so only the
12920 // non-strict predicate holds.
12921 if (Pred != ICmpInst::ICMP_UGE && Pred != ICmpInst::ICMP_SGE)
12922 return false;
12923
12924 const auto *AR = dyn_cast<SCEVAddRecExpr>(LHS);
12925 if (!AR || AR->getStart() != RHS)
12926 return false;
12927
12928 return SE.getMonotonicPredicateType(AR, Pred) ==
12930}
12931
12932/// Is LHS `Pred` RHS true on the virtue of LHS or RHS being a Min or Max
12933/// expression?
12935 const SCEV *LHS, const SCEV *RHS) {
12936 switch (Pred) {
12937 default:
12938 return false;
12939
12940 case ICmpInst::ICMP_SGE:
12941 std::swap(LHS, RHS);
12942 [[fallthrough]];
12943 case ICmpInst::ICMP_SLE:
12944 return
12945 // min(A, ...) <= A
12947 // A <= max(A, ...)
12949
12950 case ICmpInst::ICMP_UGE:
12951 std::swap(LHS, RHS);
12952 [[fallthrough]];
12953 case ICmpInst::ICMP_ULE:
12954 return
12955 // min(A, ...) <= A
12956 // FIXME: what about umin_seq?
12958 // A <= max(A, ...)
12960
12961 case ICmpInst::ICMP_UGT:
12962 std::swap(LHS, RHS);
12963 [[fallthrough]];
12964 case ICmpInst::ICMP_ULT:
12965 // umin(Ops) u<= each Op, so proving Op u< RHS for any Op proves
12966 // umin(Ops) u< RHS.
12967 //
12968 // Use computeConstantDifference instead of the more powerful
12969 // isKnownPredicate to keep this check cheap: isKnownPredicateViaMinOrMax
12970 // is called from isKnownViaNonRecursiveReasoning, so recursing into
12971 // the full predicate prover would be expensive.
12972 if (const auto *Min = dyn_cast<SCEVUMinExpr>(LHS)) {
12973 for (SCEVUse Op : Min->operands()) {
12974 std::optional<APInt> Diff = SE.computeConstantDifference(RHS, Op);
12975 // When Op and RHS share a common base differing by a
12976 // constant offset D (RHS - Op = D), Op u< RHS holds iff D != 0 and
12977 // RHS >= D (unsigned), i.e. the subtraction doesn't underflow.
12978 if (Diff && !Diff->isZero() && SE.getUnsignedRangeMin(RHS).uge(*Diff))
12979 return true;
12980 }
12981 }
12982 return false;
12983 }
12984
12985 llvm_unreachable("covered switch fell through?!");
12986}
12987
12988bool ScalarEvolution::isImpliedViaOperations(CmpPredicate Pred, const SCEV *LHS,
12989 const SCEV *RHS,
12990 const SCEV *FoundLHS,
12991 const SCEV *FoundRHS,
12992 unsigned Depth) {
12995 "LHS and RHS have different sizes?");
12996 assert(getTypeSizeInBits(FoundLHS->getType()) ==
12997 getTypeSizeInBits(FoundRHS->getType()) &&
12998 "FoundLHS and FoundRHS have different sizes?");
12999 // We want to avoid hurting the compile time with analysis of too big trees.
13001 return false;
13002
13003 // We only want to work with GT comparison so far.
13004 if (ICmpInst::isLT(Pred)) {
13006 std::swap(LHS, RHS);
13007 std::swap(FoundLHS, FoundRHS);
13008 }
13009
13011
13012 // For unsigned, try to reduce it to corresponding signed comparison.
13013 if (P == ICmpInst::ICMP_UGT)
13014 // We can replace unsigned predicate with its signed counterpart if all
13015 // involved values are non-negative.
13016 // TODO: We could have better support for unsigned.
13017 if (isKnownNonNegative(FoundLHS) && isKnownNonNegative(FoundRHS)) {
13018 // Knowing that both FoundLHS and FoundRHS are non-negative, and knowing
13019 // FoundLHS >u FoundRHS, we also know that FoundLHS >s FoundRHS. Let us
13020 // use this fact to prove that LHS and RHS are non-negative.
13021 const SCEV *MinusOne = getMinusOne(LHS->getType());
13022 if (isImpliedCondOperands(ICmpInst::ICMP_SGT, LHS, MinusOne, FoundLHS,
13023 FoundRHS) &&
13024 isImpliedCondOperands(ICmpInst::ICMP_SGT, RHS, MinusOne, FoundLHS,
13025 FoundRHS))
13027 }
13028
13029 if (P != ICmpInst::ICMP_SGT)
13030 return false;
13031
13032 auto GetOpFromSExt = [&](const SCEV *S) -> const SCEV * {
13033 if (auto *Ext = dyn_cast<SCEVSignExtendExpr>(S))
13034 return Ext->getOperand();
13035 // TODO: If S is a SCEVConstant then you can cheaply "strip" the sext off
13036 // the constant in some cases.
13037 return S;
13038 };
13039
13040 // Acquire values from extensions.
13041 auto *OrigLHS = LHS;
13042 auto *OrigFoundLHS = FoundLHS;
13043 LHS = GetOpFromSExt(LHS);
13044 FoundLHS = GetOpFromSExt(FoundLHS);
13045
13046 // Is the SGT predicate can be proved trivially or using the found context.
13047 auto IsSGTViaContext = [&](const SCEV *S1, const SCEV *S2) {
13048 return isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_SGT, S1, S2) ||
13049 isImpliedViaOperations(ICmpInst::ICMP_SGT, S1, S2, OrigFoundLHS,
13050 FoundRHS, Depth + 1);
13051 };
13052
13053 if (auto *LHSAddExpr = dyn_cast<SCEVAddExpr>(LHS)) {
13054 // We want to avoid creation of any new non-constant SCEV. Since we are
13055 // going to compare the operands to RHS, we should be certain that we don't
13056 // need any size extensions for this. So let's decline all cases when the
13057 // sizes of types of LHS and RHS do not match.
13058 // TODO: Maybe try to get RHS from sext to catch more cases?
13060 return false;
13061
13062 // Should not overflow.
13063 if (!LHSAddExpr->hasNoSignedWrap())
13064 return false;
13065
13066 SCEVUse LL = LHSAddExpr->getOperand(0);
13067 SCEVUse LR = LHSAddExpr->getOperand(1);
13068 auto *MinusOne = getMinusOne(RHS->getType());
13069
13070 // Checks that S1 >= 0 && S2 > RHS, trivially or using the found context.
13071 auto IsSumGreaterThanRHS = [&](const SCEV *S1, const SCEV *S2) {
13072 return IsSGTViaContext(S1, MinusOne) && IsSGTViaContext(S2, RHS);
13073 };
13074 // Try to prove the following rule:
13075 // (LHS = LL + LR) && (LL >= 0) && (LR > RHS) => (LHS > RHS).
13076 // (LHS = LL + LR) && (LR >= 0) && (LL > RHS) => (LHS > RHS).
13077 if (IsSumGreaterThanRHS(LL, LR) || IsSumGreaterThanRHS(LR, LL))
13078 return true;
13079 } else if (auto *LHSUnknownExpr = dyn_cast<SCEVUnknown>(LHS)) {
13080 Value *LL, *LR;
13081 // FIXME: Once we have SDiv implemented, we can get rid of this matching.
13082
13083 using namespace llvm::PatternMatch;
13084
13085 if (match(LHSUnknownExpr->getValue(), m_SDiv(m_Value(LL), m_Value(LR)))) {
13086 // Rules for division.
13087 // We are going to perform some comparisons with Denominator and its
13088 // derivative expressions. In general case, creating a SCEV for it may
13089 // lead to a complex analysis of the entire graph, and in particular it
13090 // can request trip count recalculation for the same loop. This would
13091 // cache as SCEVCouldNotCompute to avoid the infinite recursion. To avoid
13092 // this, we only want to create SCEVs that are constants in this section.
13093 // So we bail if Denominator is not a constant.
13094 if (!isa<ConstantInt>(LR))
13095 return false;
13096
13097 auto *Denominator = cast<SCEVConstant>(getSCEV(LR));
13098
13099 // We want to make sure that LHS = FoundLHS / Denominator. If it is so,
13100 // then a SCEV for the numerator already exists and matches with FoundLHS.
13101 auto *Numerator = getExistingSCEV(LL);
13102 if (!Numerator || Numerator->getType() != FoundLHS->getType())
13103 return false;
13104
13105 // Make sure that the numerator matches with FoundLHS and the denominator
13106 // is positive.
13107 if (!HasSameValue(Numerator, FoundLHS) || !isKnownPositive(Denominator))
13108 return false;
13109
13110 auto *DTy = Denominator->getType();
13111 auto *FRHSTy = FoundRHS->getType();
13112 if (DTy->isPointerTy() != FRHSTy->isPointerTy())
13113 // One of types is a pointer and another one is not. We cannot extend
13114 // them properly to a wider type, so let us just reject this case.
13115 // TODO: Usage of getEffectiveSCEVType for DTy, FRHSTy etc should help
13116 // to avoid this check.
13117 return false;
13118
13119 // Given that:
13120 // FoundLHS > FoundRHS, LHS = FoundLHS / Denominator, Denominator > 0.
13121 auto *WTy = getWiderType(DTy, FRHSTy);
13122 auto *DenominatorExt = getNoopOrSignExtend(Denominator, WTy);
13123 auto *FoundRHSExt = getNoopOrSignExtend(FoundRHS, WTy);
13124
13125 // Try to prove the following rule:
13126 // (FoundRHS > Denominator - 2) && (RHS <= 0) => (LHS > RHS).
13127 // For example, given that FoundLHS > 2. It means that FoundLHS is at
13128 // least 3. If we divide it by Denominator < 4, we will have at least 1.
13129 auto *DenomMinusTwo = getMinusSCEV(DenominatorExt, getConstant(WTy, 2));
13130 if (isKnownNonPositive(RHS) &&
13131 IsSGTViaContext(FoundRHSExt, DenomMinusTwo))
13132 return true;
13133
13134 // Try to prove the following rule:
13135 // (FoundRHS > -1 - Denominator) && (RHS < 0) => (LHS > RHS).
13136 // For example, given that FoundLHS > -3. Then FoundLHS is at least -2.
13137 // If we divide it by Denominator > 2, then:
13138 // 1. If FoundLHS is negative, then the result is 0.
13139 // 2. If FoundLHS is non-negative, then the result is non-negative.
13140 // Anyways, the result is non-negative.
13141 auto *MinusOne = getMinusOne(WTy);
13142 auto *NegDenomMinusOne = getMinusSCEV(MinusOne, DenominatorExt);
13143 if (isKnownNegative(RHS) &&
13144 IsSGTViaContext(FoundRHSExt, NegDenomMinusOne))
13145 return true;
13146 }
13147 }
13148
13149 // If our expression contained SCEVUnknown Phis, and we split it down and now
13150 // need to prove something for them, try to prove the predicate for every
13151 // possible incoming values of those Phis.
13152 if (isImpliedViaMerge(Pred, OrigLHS, RHS, OrigFoundLHS, FoundRHS, Depth + 1))
13153 return true;
13154
13155 return false;
13156}
13157
13159 const SCEV *RHS) {
13160 // zext x u<= sext x, sext x s<= zext x
13161 const SCEV *Op;
13162 switch (Pred) {
13163 case ICmpInst::ICMP_SGE:
13164 std::swap(LHS, RHS);
13165 [[fallthrough]];
13166 case ICmpInst::ICMP_SLE: {
13167 // If operand >=s 0 then ZExt == SExt. If operand <s 0 then SExt <s ZExt.
13168 return match(LHS, m_scev_SExt(m_SCEV(Op))) &&
13170 }
13171 case ICmpInst::ICMP_UGE:
13172 std::swap(LHS, RHS);
13173 [[fallthrough]];
13174 case ICmpInst::ICMP_ULE: {
13175 // If operand >=u 0 then ZExt == SExt. If operand <u 0 then ZExt <u SExt.
13176 return match(LHS, m_scev_ZExt(m_SCEV(Op))) &&
13178 }
13179 default:
13180 return false;
13181 };
13182 llvm_unreachable("unhandled case");
13183}
13184
13185bool ScalarEvolution::isKnownViaNonRecursiveReasoning(CmpPredicate Pred,
13186 SCEVUse LHS,
13187 SCEVUse RHS) {
13188 return isKnownPredicateExtendIdiom(Pred, LHS, RHS) ||
13189 isKnownPredicateViaConstantRanges(Pred, LHS, RHS) ||
13190 IsKnownPredicateViaMinOrMax(*this, Pred, LHS, RHS) ||
13191 IsKnownPredicateViaAddRecStart(*this, Pred, LHS, RHS) ||
13193 isKnownPredicateViaNoOverflow(Pred, LHS, RHS);
13194}
13195
13196bool ScalarEvolution::isImpliedCondOperandsHelper(CmpPredicate Pred,
13197 const SCEV *LHS,
13198 const SCEV *RHS,
13199 const SCEV *FoundLHS,
13200 const SCEV *FoundRHS) {
13201 switch (Pred) {
13202 default:
13203 llvm_unreachable("Unexpected CmpPredicate value!");
13204 case ICmpInst::ICMP_EQ:
13205 case ICmpInst::ICMP_NE:
13206 if (HasSameValue(LHS, FoundLHS) && HasSameValue(RHS, FoundRHS))
13207 return true;
13208 break;
13209 case ICmpInst::ICMP_SLT:
13210 case ICmpInst::ICMP_SLE:
13211 if (isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_SLE, LHS, FoundLHS) &&
13212 isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_SGE, RHS, FoundRHS))
13213 return true;
13214 break;
13215 case ICmpInst::ICMP_SGT:
13216 case ICmpInst::ICMP_SGE:
13217 if (isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_SGE, LHS, FoundLHS) &&
13218 isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_SLE, RHS, FoundRHS))
13219 return true;
13220 break;
13221 case ICmpInst::ICMP_ULT:
13222 case ICmpInst::ICMP_ULE:
13223 if (isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_ULE, LHS, FoundLHS) &&
13224 isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_UGE, RHS, FoundRHS))
13225 return true;
13226 break;
13227 case ICmpInst::ICMP_UGT:
13228 case ICmpInst::ICMP_UGE:
13229 if (isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_UGE, LHS, FoundLHS) &&
13230 isKnownViaNonRecursiveReasoning(ICmpInst::ICMP_ULE, RHS, FoundRHS))
13231 return true;
13232 break;
13233 }
13234
13235 // Maybe it can be proved via operations?
13236 if (isImpliedViaOperations(Pred, LHS, RHS, FoundLHS, FoundRHS))
13237 return true;
13238
13239 return false;
13240}
13241
13242bool ScalarEvolution::isImpliedCondOperandsViaRanges(
13243 CmpPredicate Pred, const SCEV *LHS, const SCEV *RHS, CmpPredicate FoundPred,
13244 const SCEV *FoundLHS, const SCEV *FoundRHS) {
13245 if (!isa<SCEVConstant>(RHS) || !isa<SCEVConstant>(FoundRHS))
13246 // The restriction on `FoundRHS` be lifted easily -- it exists only to
13247 // reduce the compile time impact of this optimization.
13248 return false;
13249
13250 std::optional<APInt> Addend = computeConstantDifference(LHS, FoundLHS);
13251 if (!Addend)
13252 return false;
13253
13254 const APInt &ConstFoundRHS = cast<SCEVConstant>(FoundRHS)->getAPInt();
13255
13256 // `FoundLHSRange` is the range we know `FoundLHS` to be in by virtue of the
13257 // antecedent "`FoundLHS` `FoundPred` `FoundRHS`".
13258 ConstantRange FoundLHSRange =
13259 ConstantRange::makeExactICmpRegion(FoundPred, ConstFoundRHS);
13260
13261 // Since `LHS` is `FoundLHS` + `Addend`, we can compute a range for `LHS`:
13262 ConstantRange LHSRange = FoundLHSRange.add(ConstantRange(*Addend));
13263
13264 // We can also compute the range of values for `LHS` that satisfy the
13265 // consequent, "`LHS` `Pred` `RHS`":
13266 const APInt &ConstRHS = cast<SCEVConstant>(RHS)->getAPInt();
13267 // The antecedent implies the consequent if every value of `LHS` that
13268 // satisfies the antecedent also satisfies the consequent.
13269 return LHSRange.icmp(Pred, ConstRHS);
13270}
13271
13272bool ScalarEvolution::canIVOverflowOnLT(const SCEV *RHS, const SCEV *Stride,
13273 bool IsSigned, bool Invert) {
13274 assert(isKnownPositive(Stride) && "Positive stride expected!");
13275
13276 unsigned BitWidth = getTypeSizeInBits(RHS->getType());
13277 const SCEV *One = getOne(Stride->getType());
13278
13279 if (IsSigned) {
13280 APInt MaxRHS = getRangeMax(RHS, /*IsSigned=*/true, Invert);
13281 APInt MaxValue = APInt::getSignedMaxValue(BitWidth);
13282 APInt MaxStrideMinusOne = getSignedRangeMax(getMinusSCEV(Stride, One));
13283
13284 // SMaxRHS + SMaxStrideMinusOne > SMaxValue => overflow!
13285 return (std::move(MaxValue) - MaxStrideMinusOne).slt(MaxRHS);
13286 }
13287
13288 APInt MaxRHS = getRangeMax(RHS, /*IsSigned=*/false, Invert);
13289 APInt MaxValue = APInt::getMaxValue(BitWidth);
13290 APInt MaxStrideMinusOne = getUnsignedRangeMax(getMinusSCEV(Stride, One));
13291
13292 // UMaxRHS + UMaxStrideMinusOne > UMaxValue => overflow!
13293 return (std::move(MaxValue) - MaxStrideMinusOne).ult(MaxRHS);
13294}
13295
13297 // umin(N, 1) + floor((N - umin(N, 1)) / D)
13298 // This is equivalent to "1 + floor((N - 1) / D)" for N != 0. The umin
13299 // expression fixes the case of N=0.
13300 const SCEV *MinNOne = getUMinExpr(N, getOne(N->getType()));
13301 const SCEV *NMinusOne = getMinusSCEV(N, MinNOne);
13302 return getAddExpr(MinNOne, getUDivExpr(NMinusOne, D));
13303}
13304
13305const SCEV *
13306ScalarEvolution::computeMaxBECountForLT(const SCEV *Start, const SCEV *Stride,
13307 const SCEV *End, unsigned BitWidth,
13308 bool IsSigned, bool Invert) {
13309 // The logic in this function assumes we can represent a positive stride.
13310 // If we can't, the backedge-taken count must be zero.
13311 if (IsSigned && BitWidth == 1)
13312 return getZero(Stride->getType());
13313
13314 // This code below only been closely audited for negative strides in the
13315 // unsigned comparison case, it may be correct for signed comparison, but
13316 // that needs to be established.
13317 if (IsSigned && isKnownNegative(Stride))
13318 return getCouldNotCompute();
13319
13320 // Calculate the maximum backedge count based on the range of values
13321 // permitted by Start, End, and Stride. If Invert is true, both Start and End
13322 // need inverting. Stride was already negated by the caller.
13323 APInt MinStart = getRangeMin(Start, IsSigned, Invert);
13324
13325 APInt MinStride =
13326 IsSigned ? getSignedRangeMin(Stride) : getUnsignedRangeMin(Stride);
13327
13328 // We assume either the stride is positive, or the backedge-taken count
13329 // is zero. So force StrideForMaxBECount to be at least one.
13330 APInt One(BitWidth, 1);
13331 APInt StrideForMaxBECount = IsSigned ? APIntOps::smax(One, MinStride)
13332 : APIntOps::umax(One, MinStride);
13333
13334 APInt MaxValue = IsSigned ? APInt::getSignedMaxValue(BitWidth)
13335 : APInt::getMaxValue(BitWidth);
13336 APInt Limit = MaxValue - (StrideForMaxBECount - 1);
13337
13338 // Although End can be a MAX expression we estimate MaxEnd considering only
13339 // the case End = RHS of the loop termination condition. This is safe because
13340 // in the other case (End - Start) is zero, leading to a zero maximum backedge
13341 // taken count.
13342 APInt MaxEnd = getRangeMax(End, IsSigned, Invert);
13343 MaxEnd =
13344 IsSigned ? APIntOps::smin(MaxEnd, Limit) : APIntOps::umin(MaxEnd, Limit);
13345
13346 // MaxBECount = ceil((max(MaxEnd, MinStart) - MinStart) / Stride)
13347 MaxEnd = IsSigned ? APIntOps::smax(MaxEnd, MinStart)
13348 : APIntOps::umax(MaxEnd, MinStart);
13349
13350 APInt Delta = MaxEnd - MinStart;
13351
13352 // Try to refine Delta in case End - Start (or Start - End if Invert) gives a
13353 // tighter bound after folding.
13354 const SCEV *DeltaExpr =
13355 Invert ? getMinusSCEV(Start, End) : getMinusSCEV(End, Start);
13356 Delta = APIntOps::umin(Delta, getUnsignedRangeMax(DeltaExpr));
13357
13358 return getUDivCeilSCEV(getConstant(Delta), getConstant(StrideForMaxBECount));
13359}
13360
13362ScalarEvolution::howManyLessThans(const SCEV *LHS, const SCEV *RHS,
13363 const Loop *L, bool IsSigned, bool Invert,
13364 bool ControlsOnlyExit, bool AllowPredicates) {
13366
13367 // Loop guards for L, collected on demand.
13368 std::optional<LoopGuards> CachedGuards;
13369 auto getGuards = [&]() -> const LoopGuards & {
13370 if (!CachedGuards)
13371 CachedGuards.emplace(LoopGuards::collect(L, *this));
13372 return *CachedGuards;
13373 };
13374
13375 // FIXME: Extend the non-invariant RHS analysis to greater-than comparisons.
13376 if (Invert && !isLoopInvariant(RHS, L))
13377 return getCouldNotCompute();
13378
13379 const SCEVAddRecExpr *IV = dyn_cast<SCEVAddRecExpr>(LHS);
13380 bool PredicatedIV = false;
13381 // FIXME: Generalize the NUW inference below to decreasing IVs.
13382 if (!IV && !Invert) {
13383 if (auto *ZExt = dyn_cast<SCEVZeroExtendExpr>(LHS)) {
13384 const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(ZExt->getOperand());
13385 if (AR && AR->getLoop() == L && AR->isAffine()) {
13386 auto canProveNUW = [&]() {
13387 // We can use the comparison to infer no-wrap flags only if it fully
13388 // controls the loop exit.
13389 if (!ControlsOnlyExit)
13390 return false;
13391
13392 if (!isLoopInvariant(RHS, L))
13393 return false;
13394
13395 if (!isKnownNonZero(AR->getStepRecurrence(*this)))
13396 // We need the sequence defined by AR to strictly increase in the
13397 // unsigned integer domain for the logic below to hold.
13398 return false;
13399
13400 const unsigned InnerBitWidth = getTypeSizeInBits(AR->getType());
13401 const unsigned OuterBitWidth = getTypeSizeInBits(RHS->getType());
13402 // If RHS <=u Limit, then there must exist a value V in the sequence
13403 // defined by AR (e.g. {Start,+,Step}) such that V >u RHS, and
13404 // V <=u UINT_MAX. Thus, we must exit the loop before unsigned
13405 // overflow occurs. This limit also implies that a signed comparison
13406 // (in the wide bitwidth) is equivalent to an unsigned comparison as
13407 // the high bits on both sides must be zero.
13408 APInt StrideMax = getUnsignedRangeMax(AR->getStepRecurrence(*this));
13409 APInt Limit = APInt::getMaxValue(InnerBitWidth) - (StrideMax - 1);
13410 Limit = Limit.zext(OuterBitWidth);
13411 return getUnsignedRangeMax(applyLoopGuards(RHS, getGuards()))
13412 .ule(Limit);
13413 };
13414 auto Flags = AR->getNoWrapFlags();
13415 if (!hasFlags(Flags, SCEV::FlagNUW) && canProveNUW())
13416 Flags = setFlags(Flags, SCEV::FlagNUW);
13417
13418 setNoWrapFlags(const_cast<SCEVAddRecExpr *>(AR), Flags);
13419 if (AR->hasNoUnsignedWrap()) {
13420 // Emulate what getZeroExtendExpr would have done during construction
13421 // if we'd been able to infer the fact just above at that time.
13422 const SCEV *Step = AR->getStepRecurrence(*this);
13423 Type *Ty = ZExt->getType();
13424 const SCEV *S = getAddRecExpr(
13426 getZeroExtendExpr(Step, Ty, 0), L, AR->getNoWrapFlags());
13428 }
13429 }
13430 }
13431 }
13432
13433 if (!IV && AllowPredicates) {
13434 // Try to make this an AddRec using runtime tests, in the first X
13435 // iterations of this loop, where X is the SCEV expression found by the
13436 // algorithm below.
13437 IV = convertSCEVToAddRecWithPredicates(LHS, L, Predicates);
13438 PredicatedIV = true;
13439 }
13440
13441 // Avoid weird loops
13442 if (!IV || IV->getLoop() != L || !IV->isAffine())
13443 return getCouldNotCompute();
13444
13445 // A precondition of this method is that the condition being analyzed
13446 // reaches an exiting branch which dominates the latch. Given that, we can
13447 // assume that an increment which violates the nowrap specification and
13448 // produces poison must cause undefined behavior when the resulting poison
13449 // value is branched upon and thus we can conclude that the backedge is
13450 // taken no more often than would be required to produce that poison value.
13451 // Note that a well defined loop can exit on the iteration which violates
13452 // the nowrap specification if there is another exit (either explicit or
13453 // implicit/exceptional) which causes the loop to execute before the
13454 // exiting instruction we're analyzing would trigger UB.
13455 auto WrapType = IsSigned ? SCEV::FlagNSW : SCEV::FlagNUW;
13456 bool NoWrap = ControlsOnlyExit && any(IV->getNoWrapFlags(WrapType));
13457 // Reverse the ordering for greater-than comparisons.
13459 if (Invert)
13461
13462 // The step of ~IV is the negated step of IV.
13463 const SCEV *Stride = IV->getStepRecurrence(*this);
13464 if (Invert)
13465 Stride = getNegativeSCEV(Stride);
13466 const SCEV *GuardedStride = Stride;
13467
13468 // Whether the IV may reach the maximum (or minimum if inverted) value
13469 // before the exit is taken.
13470 bool IVMayOverflow = true;
13471
13472 bool PositiveStride = isKnownPositive(Stride);
13473 // A dominating guard may prove the stride positive.
13474 if (!PositiveStride) {
13475 const SCEV *LoopGuardedStride = applyLoopGuards(Stride, getGuards());
13476 if (isKnownPositive(LoopGuardedStride)) {
13477 GuardedStride = LoopGuardedStride;
13478 PositiveStride = true;
13479 // Encode the context-sensitive stride > 0 fact into the expression
13480 Stride = getUMaxExpr(Stride, getOne(Stride->getType()));
13481 }
13482 }
13483
13484 // Avoid negative or zero stride values.
13485 if (!PositiveStride) {
13486 // FIXME: Generalize the unknown-stride analysis to decreasing IVs.
13487 if (Invert)
13488 return getCouldNotCompute();
13489
13490 // We can compute the correct backedge taken count for loops with unknown
13491 // strides if we can prove that the loop is not an infinite loop with side
13492 // effects. Here's the loop structure we are trying to handle -
13493 //
13494 // i = start
13495 // do {
13496 // A[i] = i;
13497 // i += s;
13498 // } while (i < end);
13499 //
13500 // The backedge taken count for such loops is evaluated as -
13501 // (max(end, start + stride) - start - 1) /u stride
13502 //
13503 // The additional preconditions that we need to check to prove correctness
13504 // of the above formula is as follows -
13505 //
13506 // a) IV is either nuw or nsw depending upon signedness (indicated by the
13507 // NoWrap flag).
13508 // b) the loop is guaranteed to be finite (e.g. is mustprogress and has
13509 // b) the loop is guaranteed to be finite (e.g. is mustprogress and has
13510 // no side effects within the loop) or a predicate is added to ensure
13511 // stride is positive.
13512 // c) loop has a single static exit (with no abnormal exits)
13513 //
13514 // Precondition a) implies that if the stride is negative, this is a single
13515 // trip loop. The backedge taken count formula reduces to zero in this case.
13516 //
13517 // Precondition b) and c) combine to imply that if rhs is invariant in L,
13518 // then a zero stride means the backedge can't be taken without executing
13519 // undefined behavior.
13520 //
13521 // The positive stride case is the same as isKnownPositive(Stride) returning
13522 // true (original behavior of the function).
13523 //
13524 if (PredicatedIV || !NoWrap || !loopHasNoAbnormalExits(L))
13525 return getCouldNotCompute();
13526
13527 if (!loopIsFiniteByAssumption(L)) {
13528 // If the loop may be infinite, add a predicate ensuring Stride is
13529 // positive, to guarantee forward progress.
13530 if (!AllowPredicates || !isLoopInvariant(Stride, L))
13531 return getCouldNotCompute();
13532
13533 const SCEV *Zero = getZero(Stride->getType());
13534 const SCEVPredicate *P =
13536 Predicates.push_back(P);
13537 // When the predicate holds (Stride > 0), umax(Stride, 1) == Stride,
13538 // so the result is unchanged. To prevent div by zero.
13539 Stride = getUMaxExpr(Stride, getOne(Stride->getType()));
13540 } else if (!isKnownNonZero(Stride)) {
13541 // If we have a step of zero, and RHS isn't invariant in L, we don't know
13542 // if it might eventually be greater than start and if so, on which
13543 // iteration. We can't even produce a useful upper bound.
13544 if (!isLoopInvariant(RHS, L))
13545 return getCouldNotCompute();
13546
13547 // We allow a potentially zero stride, but we need to divide by stride
13548 // below. Since the loop can't be infinite and this check must control
13549 // the sole exit, we can infer the exit must be taken on the first
13550 // iteration (e.g. backedge count = 0) if the stride is zero. Given that,
13551 // we know the numerator in the divides below must be zero, so we can
13552 // pick an arbitrary non-zero value for the denominator (e.g. stride)
13553 // and produce the right result.
13554 // FIXME: Handle the case where Stride is poison?
13555 auto wouldZeroStrideBeUB = [&]() {
13556 // Proof by contradiction. Suppose the stride were zero. If we can
13557 // prove that the backedge *is* taken on the first iteration, then since
13558 // we know this condition controls the sole exit, we must have an
13559 // infinite loop. We can't have a (well defined) infinite loop per
13560 // check just above.
13561 // Note: The (Start - Stride) term is used to get the start' term from
13562 // (start' + stride,+,stride). Remember that we only care about the
13563 // result of this expression when stride == 0 at runtime.
13564 auto *StartIfZero = getMinusSCEV(IV->getStart(), Stride);
13565 return isLoopEntryGuardedByCond(L, Cond, StartIfZero, RHS);
13566 };
13567 if (!wouldZeroStrideBeUB()) {
13568 Stride = getUMaxExpr(Stride, getOne(Stride->getType()));
13569 }
13570 }
13571 } else {
13572 // Avoid proven overflow cases: this will ensure that the backedge taken
13573 // count will not generate any unsigned overflow.
13574 IVMayOverflow = canIVOverflowOnLT(RHS, GuardedStride, IsSigned, Invert);
13575 if (IVMayOverflow && !NoWrap)
13576 return getCouldNotCompute();
13577 }
13578
13579 // On all paths just preceeding, we established the following invariant:
13580 // IV can be assumed not to overflow up to and including the exiting
13581 // iteration. We proved this in one of two ways:
13582 // 1) We can show overflow doesn't occur before the exiting iteration
13583 // 1a) canIVOverflowOnLT, and b) step of one
13584 // 2) We can show that if overflow occurs, the loop must execute UB
13585 // before any possible exit.
13586 // Note that we have not yet proved RHS invariant (in general).
13587
13588 const SCEV *Start = IV->getStart();
13589
13590 // Preserve pointer-typed Start/RHS to pass to isLoopEntryGuardedByCond.
13591 // If we convert to integers, isLoopEntryGuardedByCond will miss some cases.
13592 // Use integer-typed versions for actual computation; we can't subtract
13593 // pointers in general.
13594 const SCEV *OrigStart = Start;
13595 const SCEV *OrigRHS = RHS;
13596 if (Start->getType()->isPointerTy()) {
13597 Start = getPtrToAddrExpr(Start);
13598 if (isa<SCEVCouldNotCompute>(Start))
13599 return Start;
13600 }
13601 if (RHS->getType()->isPointerTy()) {
13604 return RHS;
13605 }
13606
13607 const SCEV *End = nullptr, *BECount = getCouldNotCompute(),
13608 *BECountIfBackedgeTaken = getCouldNotCompute();
13609 if (!isLoopInvariant(RHS, L)) {
13610 assert(!Invert && "RHS must be loop-invariant for Invert");
13611 const auto *RHSAddRec = dyn_cast<SCEVAddRecExpr>(RHS);
13612 if (PositiveStride && RHSAddRec != nullptr && RHSAddRec->getLoop() == L &&
13613 any(RHSAddRec->getNoWrapFlags())) {
13614 // The structure of loop we are trying to calculate backedge count of:
13615 //
13616 // left = left_start
13617 // right = right_start
13618 //
13619 // while(left < right){
13620 // ... do something here ...
13621 // left += s1; // stride of left is s1 (s1 > 0)
13622 // right += s2; // stride of right is s2 (s2 < 0)
13623 // }
13624 //
13625
13626 const SCEV *RHSStart = RHSAddRec->getStart();
13627 const SCEV *RHSStride = RHSAddRec->getStepRecurrence(*this);
13628
13629 // If Stride - RHSStride is positive and does not overflow, we can write
13630 // backedge count as ->
13631 // ceil((End - Start) /u (Stride - RHSStride))
13632 // Where, End = max(RHSStart, Start)
13633
13634 // Check if RHSStride < 0 and Stride - RHSStride will not overflow.
13635 if (isKnownNegative(RHSStride) &&
13636 willNotOverflow(Instruction::Sub, /*Signed=*/true, Stride,
13637 RHSStride)) {
13638
13639 const SCEV *Denominator = getMinusSCEV(Stride, RHSStride);
13640 if (isKnownPositive(Denominator)) {
13641 End = IsSigned ? getSMaxExpr(RHSStart, Start)
13642 : getUMaxExpr(RHSStart, Start);
13643
13644 // We can do this because End >= Start, as End = max(RHSStart, Start)
13645 const SCEV *Delta = getMinusSCEV(End, Start);
13646
13647 BECount = getUDivCeilSCEV(Delta, Denominator);
13648 BECountIfBackedgeTaken =
13649 getUDivCeilSCEV(getMinusSCEV(RHSStart, Start), Denominator);
13650 }
13651 }
13652 }
13653 } else {
13654 // Let End = max(RHS,Start). We use the expression (End-Start)/Stride to
13655 // describe the backedge count: if the backedge is taken at least once then
13656 // End is RHS, and if not End is Start so we get a backedge count of zero.
13657 // Inverted, End is min(RHS, Start).
13658 //
13659 // AddingStrideMinusOneMayOverflow has the following preconditions:
13660 //
13661 // 1. Start <= End, signed if IsSigned (inverted: End <= Start)
13662 // 2. The index variable doesn't overflow.
13663 //
13664 // Therefore, we know N exists such that
13665 // (Start + Stride * N) >= End, and computing "(Start + Stride * N)"
13666 // doesn't overflow.
13667 //
13668 // Using this information, try to prove whether the addition in
13669 // "(End - Start) + (Stride - 1)" has unsigned overflow.
13670 //
13671 // If the IV cannot overflow, RHS is at least Stride - 1 below the maximum
13672 // value, so the distance End - Start is at most UMAX - (Stride - 1) and
13673 // the (Stride - 1) addition below cannot overflow.
13674 const SCEV *One = getOne(Stride->getType());
13675 bool AddingStrideMinusOneMayOverflow = IVMayOverflow && [&] {
13676 if (isKnownToBeAPowerOfTwo(Stride)) {
13677 // Suppose Stride is a power of two, and Start/End are unsigned
13678 // integers. Let UMAX be the largest representable unsigned
13679 // integer.
13680 //
13681 // By the preconditions of this function, we know
13682 // "(Start + Stride * N) >= End", and this doesn't overflow.
13683 // As a formula:
13684 //
13685 // End <= (Start + Stride * N) <= UMAX
13686 //
13687 // Subtracting Start from all the terms:
13688 //
13689 // End - Start <= Stride * N <= UMAX - Start
13690 //
13691 // Since Start is unsigned, UMAX - Start <= UMAX. Therefore:
13692 //
13693 // End - Start <= Stride * N <= UMAX
13694 //
13695 // Stride * N is a multiple of Stride. Therefore,
13696 //
13697 // End - Start <= Stride * N <= UMAX - (UMAX mod Stride)
13698 //
13699 // Since Stride is a power of two, UMAX + 1 is divisible by
13700 // Stride. Therefore, UMAX mod Stride == Stride - 1. So we can
13701 // write:
13702 //
13703 // End - Start <= Stride * N <= UMAX - Stride - 1
13704 //
13705 // Dropping the middle term:
13706 //
13707 // End - Start <= UMAX - Stride - 1
13708 //
13709 // Adding Stride - 1 to both sides:
13710 //
13711 // (End - Start) + (Stride - 1) <= UMAX
13712 //
13713 // In other words, the addition doesn't have unsigned overflow.
13714 //
13715 // A similar proof works if we treat Start/End as signed values.
13716 // Just rewrite steps before "End - Start <= Stride * N <= UMAX"
13717 // to use signed max instead of unsigned max. Note that we're
13718 // trying to prove a lack of unsigned overflow in either case.
13719 // Inverted: "Start - End <= Stride * N <= Start - MIN <= UMAX", same.
13720 return false;
13721 }
13722 if (!Invert && (Start == Stride || Start == getMinusSCEV(Stride, One))) {
13723 // If Start is equal to Stride, (End - Start) + (Stride - 1) == End
13724 // - 1. If !IsSigned, 0 <u Stride == Start <=u End; so 0 <u End - 1
13725 // <u End. If IsSigned, 0 <s Stride == Start <=s End; so 0 <s End -
13726 // 1 <s End.
13727 //
13728 // If Start is equal to Stride - 1, (End - Start) + Stride - 1 ==
13729 // End.
13730 //
13731 // Both need Start to be the smaller value, so neither applies inverted.
13732 return false;
13733 }
13734 return true;
13735 }();
13736
13737 // If inverted, the analyzed values are complements: "~V - Offset" is "~(V +
13738 // Offset)" and "~To - ~From" is "From - To".
13739 auto StepBack = [&](const SCEV *V, const SCEV *Offset) -> const SCEV * {
13740 if (Invert)
13741 return getAddExpr(V, Offset);
13742 return getMinusSCEV(V, Offset);
13743 };
13744 auto Distance = [&](const SCEV *From, const SCEV *To) {
13745 return Invert ? getMinusSCEV(From, To) : getMinusSCEV(To, From);
13746 };
13747
13748 const SCEV *OrigPrevStart = StepBack(OrigStart, Stride);
13749 assert(isAvailableAtLoopEntry(OrigPrevStart, L) && "Must be!");
13750 assert(isAvailableAtLoopEntry(OrigStart, L) && "Must be!");
13751 assert(isAvailableAtLoopEntry(OrigRHS, L) && "Must be!");
13752 // Can we prove Start - Stride < RHS, and either Start - Stride < Start or
13753 // (via !AddingStrideMinusOneMayOverflow) that (RHS - Start) + (Stride - 1)
13754 // does not overflow?
13755 if ((!AddingStrideMinusOneMayOverflow ||
13756 isLoopEntryGuardedByCond(L, Cond, OrigPrevStart, OrigStart)) &&
13757 isLoopEntryGuardedByCond(L, Cond, OrigPrevStart, OrigRHS)) {
13758 // In this case, we can use a refined formula for computing backedge
13759 // taken count. The general formula remains:
13760 // "End-Start /uceiling Stride"
13761 // We want to use the alternate formula:
13762 // "((RHS - 1) - (Start - Stride)) /u Stride"
13763 // Let's do a quick case analysis to show these are equivalent under
13764 // our preconditions. When inverted, the proof uses complemented Start,
13765 // RHS and End; Stride remains positive.
13766 // * For RHS <= Start (End is Start), the backedge-taken count must be
13767 // zero. Together with the precondition "Start - Stride < RHS", we have
13768 // "Start - Stride < RHS <= Start". Subtracting Start - Stride from
13769 // all sides we get "0 < RHS - (Start - Stride) <= Stride".
13770 // Subtracting 1 we get "0 <= (RHS - 1) - (Start - Stride) < Stride".
13771 // So dividing that by Stride gives zero.
13772 //
13773 // * For RHS > Start (End is RHS), the backedge count must be
13774 // "RHS-Start /uceil Stride", so it is sufficient to show that the
13775 // numerator "((RHS - 1) - (Start - Stride))" does not overflow.
13776 //
13777 // If "Start - Stride < Start" holds, we have
13778 // "RHS > Start > Start - Stride". As such
13779 // "RHS - (Start - Stride) - 1" does not overflow, which is the
13780 // reassociated numerator.
13781 //
13782 // Otherwise !AddingStrideMinusOneMayOverflow guarantees that
13783 // "(End - Start) + (Stride - 1)" does not overflow unsigned. Here
13784 // "End" is "RHS", as "RHS > Start", so this is the reassociated
13785 // numerator. Neither sub-term wraps unsigned: "RHS - Start"
13786 // due to "RHS > Start", and "Stride - 1", as Stride is non-zero.
13787 const SCEV *Numerator =
13788 getMinusSCEV(Distance(StepBack(Start, Stride), RHS), One);
13789 BECount = getUDivExpr(Numerator, Stride);
13790 }
13791
13792 if (isa<SCEVCouldNotCompute>(BECount)) {
13793 auto canProveRHSIsAtOrBeyondStart = [&]() {
13794 // Inverted, the claim is "Start >= RHS". Reverse the comparisons below
13795 // by swapping their operands rather than their predicates:
13796 // isLoopEntryGuardedByCond is sensitive to operand order and loses the
13797 // proof if the IV bound moves to the other side.
13798 auto SwapIfInverted = [&](const SCEV *A, const SCEV *B) {
13799 return Invert ? std::pair(B, A) : std::pair(A, B);
13800 };
13801
13802 auto CondGE = IsSigned ? ICmpInst::ICMP_SGE : ICmpInst::ICMP_UGE;
13803 const SCEV *GuardedRHS = applyLoopGuards(OrigRHS, getGuards());
13804 const SCEV *GuardedStart = applyLoopGuards(OrigStart, getGuards());
13805 if (Invert)
13806 std::swap(GuardedRHS, GuardedStart);
13807
13808 auto [GELHS, GERHS] = SwapIfInverted(OrigRHS, OrigStart);
13809 if (isLoopEntryGuardedByCond(L, CondGE, GELHS, GERHS) ||
13810 isKnownPredicate(CondGE, GuardedRHS, GuardedStart))
13811 return true;
13812
13813 // (RHS > Start - 1) implies RHS >= Start.
13814 // * "RHS >= Start" is trivially equivalent to "RHS > Start - 1" if
13815 // "Start - 1" doesn't overflow.
13816 // * For signed comparison, if Start - 1 does overflow, it's equal
13817 // to INT_MAX, and "RHS >s INT_MAX" is trivially false.
13818 // * For unsigned comparison, if Start - 1 does overflow, it's equal
13819 // to UINT_MAX, and "RHS >u UINT_MAX" is trivially false.
13820 //
13821 // FIXME: Should isLoopEntryGuardedByCond do this for us?
13822 auto CondGT = IsSigned ? ICmpInst::ICMP_SGT : ICmpInst::ICMP_UGT;
13823 auto [GTLHS, GTRHS] = SwapIfInverted(OrigRHS, StepBack(OrigStart, One));
13824 return isLoopEntryGuardedByCond(L, CondGT, GTLHS, GTRHS);
13825 };
13826
13827 // If we know that RHS >= Start in the context of loop, then we know
13828 // that max(RHS, Start) = RHS at this point.
13829 if (canProveRHSIsAtOrBeyondStart()) {
13830 End = RHS;
13831 } else {
13832 // If RHS < Start, the backedge will be taken zero times. So in
13833 // general, we can write the backedge-taken count as:
13834 //
13835 // RHS >= Start ? ceil(RHS - Start) / Stride : 0
13836 //
13837 // We convert it to the following to make it more convenient for SCEV:
13838 //
13839 // ceil(max(RHS, Start) - Start) / Stride
13840 //
13841 // Inverted, this is ceil(Start - min(RHS, Start)) / Stride.
13842 if (Invert)
13843 End = IsSigned ? getSMinExpr(RHS, Start) : getUMinExpr(RHS, Start);
13844 else
13845 End = IsSigned ? getSMaxExpr(RHS, Start) : getUMaxExpr(RHS, Start);
13846
13847 // See what would happen if we assume the backedge is taken. This is
13848 // used to compute MaxBECount.
13849 BECountIfBackedgeTaken = getUDivCeilSCEV(Distance(Start, RHS), Stride);
13850 }
13851
13852 const SCEV *Delta = Distance(Start, End);
13853 if (!AddingStrideMinusOneMayOverflow) {
13854 // floor((D + (S - 1)) / S)
13855 // We prefer this formulation if it's legal because it's fewer
13856 // operations.
13857 BECount =
13858 getUDivExpr(getAddExpr(Delta, getMinusSCEV(Stride, One)), Stride);
13859 } else {
13860 BECount = getUDivCeilSCEV(Delta, Stride);
13861 }
13862 }
13863 }
13864
13865 const SCEV *ConstantMaxBECount;
13866 bool MaxOrZero = false;
13867 if (isa<SCEVConstant>(BECount)) {
13868 ConstantMaxBECount = BECount;
13869 } else {
13870 ConstantMaxBECount = computeMaxBECountForLT(
13871 Start, Stride, RHS, getTypeSizeInBits(LHS->getType()), IsSigned,
13872 Invert);
13873 // If we know exactly how many times the backedge will be taken if it's
13874 // taken at least once, then the backedge count will either be that or
13875 // zero. If that count exceeds the range-based bound, the backedge can
13876 // never be taken.
13877 const APInt *IfTaken, *RangeMax;
13878 if (match(BECountIfBackedgeTaken, m_scev_APInt(IfTaken))) {
13879 if (match(ConstantMaxBECount, m_scev_APInt(RangeMax)) &&
13880 IfTaken->ugt(*RangeMax)) {
13881 ConstantMaxBECount = getZero(BECountIfBackedgeTaken->getType());
13882 } else {
13883 ConstantMaxBECount = BECountIfBackedgeTaken;
13884 MaxOrZero = true;
13885 }
13886 }
13887 }
13888
13889 if (isa<SCEVCouldNotCompute>(ConstantMaxBECount) &&
13890 !isa<SCEVCouldNotCompute>(BECount))
13891 ConstantMaxBECount = getConstant(getUnsignedRangeMax(BECount));
13892
13893 const SCEV *SymbolicMaxBECount =
13894 isa<SCEVCouldNotCompute>(BECount) ? ConstantMaxBECount : BECount;
13895 return ExitLimit(BECount, ConstantMaxBECount, SymbolicMaxBECount, MaxOrZero,
13896 Predicates);
13897}
13898
13900 ScalarEvolution &SE) const {
13901 if (Range.isFullSet()) // Infinite loop.
13902 return SE.getCouldNotCompute();
13903
13904 // If the start is a non-zero constant, shift the range to simplify things.
13905 if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(getStart()))
13906 if (!SC->getValue()->isZero()) {
13908 Operands[0] = SE.getZero(SC->getType());
13909 const SCEV *Shifted = SE.getAddRecExpr(Operands, getLoop(),
13911 if (const auto *ShiftedAddRec = dyn_cast<SCEVAddRecExpr>(Shifted))
13912 return ShiftedAddRec->getNumIterationsInRange(
13913 Range.subtract(SC->getAPInt()), SE);
13914 // This is strange and shouldn't happen.
13915 return SE.getCouldNotCompute();
13916 }
13917
13918 // The only time we can solve this is when we have all constant indices.
13919 // Otherwise, we cannot determine the overflow conditions.
13921 return SE.getCouldNotCompute();
13922
13923 // Okay at this point we know that all elements of the chrec are constants and
13924 // that the start element is zero.
13925
13926 // First check to see if the range contains zero. If not, the first
13927 // iteration exits.
13928 unsigned BitWidth = SE.getTypeSizeInBits(getType());
13929 if (!Range.contains(APInt(BitWidth, 0)))
13930 return SE.getZero(getType());
13931
13932 if (isAffine()) {
13933 // If this is an affine expression then we have this situation:
13934 // Solve {0,+,A} in Range === Ax in Range
13935
13936 // We know that zero is in the range. If A is positive then we know that
13937 // the upper value of the range must be the first possible exit value.
13938 // If A is negative then the lower of the range is the last possible loop
13939 // value. Also note that we already checked for a full range.
13940 APInt A = cast<SCEVConstant>(getOperand(1))->getAPInt();
13941 APInt End = A.sge(1) ? (Range.getUpper() - 1) : Range.getLower();
13942
13943 // The exit value should be (End+A)/A.
13944 APInt ExitVal = (End + A).udiv(A);
13945 ConstantInt *ExitValue = ConstantInt::get(SE.getContext(), ExitVal);
13946
13947 // Evaluate at the exit value. If we really did fall out of the valid
13948 // range, then we computed our trip count, otherwise wrap around or other
13949 // things must have happened.
13950 ConstantInt *Val = EvaluateConstantChrecAtConstant(this, ExitValue, SE);
13951 if (Range.contains(Val->getValue()))
13952 return SE.getCouldNotCompute(); // Something strange happened
13953
13954 // Ensure that the previous value is in the range.
13955 assert(Range.contains(
13957 ConstantInt::get(SE.getContext(), ExitVal - 1), SE)->getValue()) &&
13958 "Linear scev computation is off in a bad way!");
13959 return SE.getConstant(ExitValue);
13960 }
13961
13962 if (isQuadratic()) {
13963 if (auto S = SolveQuadraticAddRecRange(this, Range, SE))
13964 return SE.getConstant(*S);
13965 }
13966
13967 return SE.getCouldNotCompute();
13968}
13969
13970const SCEVAddRecExpr *
13972 assert(getNumOperands() > 1 && "AddRec with zero step?");
13973 // There is a temptation to just call getAddExpr(this, getStepRecurrence(SE)),
13974 // but in this case we cannot guarantee that the value returned will be an
13975 // AddRec because SCEV does not have a fixed point where it stops
13976 // simplification: it is legal to return ({rec1} + {rec2}). For example, it
13977 // may happen if we reach arithmetic depth limit while simplifying. So we
13978 // construct the returned value explicitly.
13980 // If this is {A,+,B,+,C,...,+,N}, then its step is {B,+,C,+,...,+,N}, and
13981 // (this + Step) is {A+B,+,B+C,+...,+,N}.
13982 for (unsigned i = 0, e = getNumOperands() - 1; i < e; ++i)
13983 Ops.push_back(SE.getAddExpr(getOperand(i), getOperand(i + 1)));
13984 // We know that the last operand is not a constant zero (otherwise it would
13985 // have been popped out earlier). This guarantees us that if the result has
13986 // the same last operand, then it will also not be popped out, meaning that
13987 // the returned value will be an AddRec.
13988 const SCEV *Last = getOperand(getNumOperands() - 1);
13989 assert(!Last->isZero() && "Recurrency with zero step?");
13990 Ops.push_back(Last);
13992}
13993
13994// Return true when S contains at least an undef value.
13996 return SCEVExprContains(
13997 S, [](const SCEV *S) { return match(S, m_scev_UndefOrPoison()); });
13998}
13999
14000// Return true when S contains a value that is a nullptr.
14002 return SCEVExprContains(S, [](const SCEV *S) {
14003 if (const auto *SU = dyn_cast<SCEVUnknown>(S))
14004 return SU->getValue() == nullptr;
14005 return false;
14006 });
14007}
14008
14009/// Return the size of an element read or written by Inst.
14011 if (!isa<LoadInst, StoreInst>(Inst))
14012 return nullptr;
14014 return getSizeOfExpr(ETy, getLoadStoreType(Inst));
14015}
14016
14017//===----------------------------------------------------------------------===//
14018// SCEVCallbackVH Class Implementation
14019//===----------------------------------------------------------------------===//
14020
14022 assert(SE && "SCEVCallbackVH called with a null ScalarEvolution!");
14023 if (PHINode *PN = dyn_cast<PHINode>(getValPtr()))
14024 SE->ConstantEvolutionLoopExitValue.erase(PN);
14025 SE->eraseValueFromMap(getValPtr());
14026 // this now dangles!
14027}
14028
14029void ScalarEvolution::SCEVCallbackVH::allUsesReplacedWith(Value *V) {
14030 assert(SE && "SCEVCallbackVH called with a null ScalarEvolution!");
14031
14032 // Forget all the expressions associated with users of the old value,
14033 // so that future queries will recompute the expressions using the new
14034 // value.
14035 SE->forgetValue(getValPtr());
14036 // this now dangles!
14037}
14038
14039ScalarEvolution::SCEVCallbackVH::SCEVCallbackVH(Value *V, ScalarEvolution *se)
14040 : CallbackVH(V), SE(se) {}
14041
14042//===----------------------------------------------------------------------===//
14043// ScalarEvolution Class Implementation
14044//===----------------------------------------------------------------------===//
14045
14048 LoopInfo &LI)
14049 : F(F), DL(F.getDataLayout()), TLI(TLI), AC(AC), DT(DT), LI(LI),
14050 CouldNotCompute(new SCEVCouldNotCompute()), ValuesAtScopes(64),
14051 LoopDispositions(64), BlockDispositions(64) {
14052 // To use guards for proving predicates, we need to scan every instruction in
14053 // relevant basic blocks, and not just terminators. Doing this is a waste of
14054 // time if the IR does not actually contain any calls to
14055 // @llvm.experimental.guard, so do a quick check and remember this beforehand.
14056 //
14057 // This pessimizes the case where a pass that preserves ScalarEvolution wants
14058 // to _add_ guards to the module when there weren't any before, and wants
14059 // ScalarEvolution to optimize based on those guards. For now we prefer to be
14060 // efficient in lieu of being smart in that rather obscure case.
14061
14062 auto *GuardDecl = Intrinsic::getDeclarationIfExists(
14063 F.getParent(), Intrinsic::experimental_guard);
14064 HasGuards = GuardDecl && !GuardDecl->use_empty();
14065}
14066
14068 : F(Arg.F), DL(Arg.DL), HasGuards(Arg.HasGuards), TLI(Arg.TLI), AC(Arg.AC),
14069 DT(Arg.DT), LI(Arg.LI), CouldNotCompute(std::move(Arg.CouldNotCompute)),
14070 ValueExprMap(std::move(Arg.ValueExprMap)),
14071 PendingLoopPredicates(std::move(Arg.PendingLoopPredicates)),
14072 PendingMerges(std::move(Arg.PendingMerges)),
14073 ConstantMultipleCache(std::move(Arg.ConstantMultipleCache)),
14074 BackedgeTakenCounts(std::move(Arg.BackedgeTakenCounts)),
14075 PredicatedBackedgeTakenCounts(
14076 std::move(Arg.PredicatedBackedgeTakenCounts)),
14077 BECountUsers(std::move(Arg.BECountUsers)),
14078 ConstantEvolutionLoopExitValue(
14079 std::move(Arg.ConstantEvolutionLoopExitValue)),
14080 ValuesAtScopes(std::move(Arg.ValuesAtScopes)),
14081 ValuesAtScopesUsers(std::move(Arg.ValuesAtScopesUsers)),
14082 LoopDispositions(std::move(Arg.LoopDispositions)),
14083 LoopPropertiesCache(std::move(Arg.LoopPropertiesCache)),
14084 BlockDispositions(std::move(Arg.BlockDispositions)),
14085 SCEVUsers(std::move(Arg.SCEVUsers)),
14086 UnsignedRanges(std::move(Arg.UnsignedRanges)),
14087 SignedRanges(std::move(Arg.SignedRanges)),
14088 UniqueSCEVs(std::move(Arg.UniqueSCEVs)),
14089 UniquePreds(std::move(Arg.UniquePreds)),
14090 SCEVAllocator(std::move(Arg.SCEVAllocator)),
14091 ConstantSCEVs(std::move(Arg.ConstantSCEVs)),
14092 LoopUsers(std::move(Arg.LoopUsers)),
14093 PredicatedSCEVRewrites(std::move(Arg.PredicatedSCEVRewrites)),
14094 FirstUnknown(Arg.FirstUnknown) {
14095 Arg.FirstUnknown = nullptr;
14096}
14097
14099 // Iterate through all the SCEVUnknown instances and call their
14100 // destructors, so that they release their references to their values.
14101 for (SCEVUnknown *U = FirstUnknown; U;) {
14102 SCEVUnknown *Tmp = U;
14103 U = U->Next;
14104 Tmp->~SCEVUnknown();
14105 }
14106 FirstUnknown = nullptr;
14107
14108 ExprValueMap.clear();
14109 ValueExprMap.clear();
14110 HasRecMap.clear();
14111 BackedgeTakenCounts.clear();
14112 PredicatedBackedgeTakenCounts.clear();
14113
14114 assert(PendingLoopPredicates.empty() && "isImpliedCond garbage");
14115 assert(PendingMerges.empty() && "isImpliedViaMerge garbage");
14116 assert(!WalkingBEDominatingConds && "isLoopBackedgeGuardedByCond garbage!");
14117 assert(!ProvingSplitPredicate && "ProvingSplitPredicate garbage!");
14118}
14119
14123
14124/// When printing a top-level SCEV for trip counts, it's helpful to include
14125/// a type for constants which are otherwise hard to disambiguate.
14126static void PrintSCEVWithTypeHint(raw_ostream &OS, const SCEV* S) {
14127 if (isa<SCEVConstant>(S))
14128 OS << *S->getType() << " ";
14129 OS << *S;
14130}
14131
14133 const Loop *L) {
14134 // Print all inner loops first
14135 for (Loop *I : *L)
14136 PrintLoopInfo(OS, SE, I);
14137
14138 OS << "Loop ";
14139 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14140 OS << ": ";
14141
14142 SmallVector<BasicBlock *, 8> ExitingBlocks;
14143 L->getExitingBlocks(ExitingBlocks);
14144 if (ExitingBlocks.size() != 1)
14145 OS << "<multiple exits> ";
14146
14147 auto *BTC = SE->getBackedgeTakenCount(L);
14148 if (!isa<SCEVCouldNotCompute>(BTC)) {
14149 OS << "backedge-taken count is ";
14150 PrintSCEVWithTypeHint(OS, BTC);
14151 } else
14152 OS << "Unpredictable backedge-taken count.";
14153 OS << "\n";
14154
14155 if (ExitingBlocks.size() > 1)
14156 for (BasicBlock *ExitingBlock : ExitingBlocks) {
14157 OS << " exit count for " << ExitingBlock->getName() << ": ";
14158 const SCEV *EC = SE->getExitCount(L, ExitingBlock);
14159 PrintSCEVWithTypeHint(OS, EC);
14160 if (isa<SCEVCouldNotCompute>(EC)) {
14161 // Retry with predicates.
14163 EC = SE->getPredicatedExitCount(L, ExitingBlock, &Predicates);
14164 if (!isa<SCEVCouldNotCompute>(EC)) {
14165 OS << "\n predicated exit count for " << ExitingBlock->getName()
14166 << ": ";
14167 PrintSCEVWithTypeHint(OS, EC);
14168 OS << "\n Predicates:\n";
14169 for (const auto *P : Predicates)
14170 P->print(OS, 4);
14171 }
14172 }
14173 OS << "\n";
14174 }
14175
14176 OS << "Loop ";
14177 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14178 OS << ": ";
14179
14180 auto *ConstantBTC = SE->getConstantMaxBackedgeTakenCount(L);
14181 if (!isa<SCEVCouldNotCompute>(ConstantBTC)) {
14182 OS << "constant max backedge-taken count is ";
14183 PrintSCEVWithTypeHint(OS, ConstantBTC);
14185 OS << ", actual taken count either this or zero.";
14186 } else {
14187 OS << "Unpredictable constant max backedge-taken count. ";
14188 }
14189
14190 OS << "\n"
14191 "Loop ";
14192 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14193 OS << ": ";
14194
14195 auto *SymbolicBTC = SE->getSymbolicMaxBackedgeTakenCount(L);
14196 if (!isa<SCEVCouldNotCompute>(SymbolicBTC)) {
14197 OS << "symbolic max backedge-taken count is ";
14198 PrintSCEVWithTypeHint(OS, SymbolicBTC);
14200 OS << ", actual taken count either this or zero.";
14201 } else {
14202 OS << "Unpredictable symbolic max backedge-taken count. ";
14203 }
14204 OS << "\n";
14205
14206 if (ExitingBlocks.size() > 1)
14207 for (BasicBlock *ExitingBlock : ExitingBlocks) {
14208 OS << " symbolic max exit count for " << ExitingBlock->getName() << ": ";
14209 auto *ExitBTC = SE->getExitCount(L, ExitingBlock,
14211 PrintSCEVWithTypeHint(OS, ExitBTC);
14212 if (isa<SCEVCouldNotCompute>(ExitBTC)) {
14213 // Retry with predicates.
14215 ExitBTC = SE->getPredicatedExitCount(L, ExitingBlock, &Predicates,
14217 if (!isa<SCEVCouldNotCompute>(ExitBTC)) {
14218 OS << "\n predicated symbolic max exit count for "
14219 << ExitingBlock->getName() << ": ";
14220 PrintSCEVWithTypeHint(OS, ExitBTC);
14221 OS << "\n Predicates:\n";
14222 for (const auto *P : Predicates)
14223 P->print(OS, 4);
14224 }
14225 }
14226 OS << "\n";
14227 }
14228
14230 auto *PBT = SE->getPredicatedBackedgeTakenCount(L, Preds);
14231 if (PBT != BTC) {
14232 OS << "Loop ";
14233 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14234 OS << ": ";
14235 if (!isa<SCEVCouldNotCompute>(PBT)) {
14236 OS << "Predicated backedge-taken count is ";
14237 PrintSCEVWithTypeHint(OS, PBT);
14238 } else
14239 OS << "Unpredictable predicated backedge-taken count.";
14240 OS << "\n";
14241 OS << " Predicates:\n";
14242 for (const auto *P : Preds)
14243 P->print(OS, 4);
14244 }
14245 Preds.clear();
14246
14247 auto *PredConstantMax =
14249 if (PredConstantMax != ConstantBTC) {
14250 OS << "Loop ";
14251 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14252 OS << ": ";
14253 if (!isa<SCEVCouldNotCompute>(PredConstantMax)) {
14254 OS << "Predicated constant max backedge-taken count is ";
14255 PrintSCEVWithTypeHint(OS, PredConstantMax);
14256 } else
14257 OS << "Unpredictable predicated constant max backedge-taken count.";
14258 OS << "\n";
14259 OS << " Predicates:\n";
14260 for (const auto *P : Preds)
14261 P->print(OS, 4);
14262 }
14263 Preds.clear();
14264
14265 auto *PredSymbolicMax =
14267 if (SymbolicBTC != PredSymbolicMax) {
14268 OS << "Loop ";
14269 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14270 OS << ": ";
14271 if (!isa<SCEVCouldNotCompute>(PredSymbolicMax)) {
14272 OS << "Predicated symbolic max backedge-taken count is ";
14273 PrintSCEVWithTypeHint(OS, PredSymbolicMax);
14274 } else
14275 OS << "Unpredictable predicated symbolic max backedge-taken count.";
14276 OS << "\n";
14277 OS << " Predicates:\n";
14278 for (const auto *P : Preds)
14279 P->print(OS, 4);
14280 }
14281
14283 OS << "Loop ";
14284 L->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14285 OS << ": ";
14286 OS << "Trip multiple is " << SE->getSmallConstantTripMultiple(L) << "\n";
14287 }
14288}
14289
14290namespace llvm {
14291// Note: these overloaded operators need to be in the llvm namespace for them
14292// to be resolved correctly. If we put them outside the llvm namespace, the
14293//
14294// OS << ": " << SE.getLoopDisposition(SV, InnerL);
14295//
14296// code below "breaks" and start printing raw enum values as opposed to the
14297// string values.
14300 switch (LD) {
14302 OS << "Variant";
14303 break;
14305 OS << "Invariant";
14306 break;
14308 OS << "Uniform";
14309 break;
14311 OS << "Computable";
14312 break;
14313 }
14314 return OS;
14315}
14316
14319 switch (BD) {
14321 OS << "DoesNotDominate";
14322 break;
14324 OS << "Dominates";
14325 break;
14327 OS << "ProperlyDominates";
14328 break;
14329 }
14330 return OS;
14331}
14332} // namespace llvm
14333
14335 // ScalarEvolution's implementation of the print method is to print
14336 // out SCEV values of all instructions that are interesting. Doing
14337 // this potentially causes it to create new SCEV objects though,
14338 // which technically conflicts with the const qualifier. This isn't
14339 // observable from outside the class though, so casting away the
14340 // const isn't dangerous.
14341 ScalarEvolution &SE = *const_cast<ScalarEvolution *>(this);
14342
14343 if (ClassifyExpressions) {
14344 OS << "Classifying expressions for: ";
14345 F.printAsOperand(OS, /*PrintType=*/false);
14346 OS << "\n";
14347 for (Instruction &I : instructions(F))
14348 if (isSCEVable(I.getType()) && !isa<CmpInst>(I)) {
14349 OS << I << '\n';
14350 OS << " --> ";
14351 const SCEV *SV = SE.getSCEV(&I);
14352 SV->print(OS);
14353 if (!isa<SCEVCouldNotCompute>(SV)) {
14354 OS << " U: ";
14355 SE.getUnsignedRange(SV).print(OS);
14356 OS << " S: ";
14357 SE.getSignedRange(SV).print(OS);
14358 }
14359
14360 const Loop *L = LI.getLoopFor(I.getParent());
14361
14362 SCEVUse AtUse = SE.getSCEVAtScope(SV, L);
14363 if (AtUse != SV) {
14364 OS << " --> ";
14365 OS << AtUse;
14366 if (!isa<SCEVCouldNotCompute>(AtUse)) {
14367 OS << " U: ";
14368 SE.getUnsignedRange(AtUse).print(OS);
14369 OS << " S: ";
14370 SE.getSignedRange(AtUse).print(OS);
14371 }
14372 }
14373
14374 if (L) {
14375 OS << "\t\t" "Exits: ";
14376 SCEVUse ExitValue = SE.getSCEVAtScope(SV, L->getParentLoop());
14377 if (!SE.isLoopInvariant(ExitValue, L)) {
14378 OS << "<<Unknown>>";
14379 } else {
14380 OS << ExitValue;
14381 }
14382
14383 ListSeparator LS(", ", "\t\tLoopDispositions: { ");
14384 for (const auto *Iter = L; Iter; Iter = Iter->getParentLoop()) {
14385 OS << LS;
14386 Iter->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14387 OS << ": " << SE.getLoopDisposition(SV, Iter);
14388 }
14389
14390 for (const auto *InnerL : depth_first(L)) {
14391 if (InnerL == L)
14392 continue;
14393 OS << LS;
14394 InnerL->getHeader()->printAsOperand(OS, /*PrintType=*/false);
14395 OS << ": " << SE.getLoopDisposition(SV, InnerL);
14396 }
14397
14398 OS << " }";
14399 }
14400
14401 OS << "\n";
14402 }
14403 }
14404
14405 OS << "Determining loop execution counts for: ";
14406 F.printAsOperand(OS, /*PrintType=*/false);
14407 OS << "\n";
14408 for (Loop *I : LI)
14409 PrintLoopInfo(OS, &SE, I);
14410}
14411
14414 auto &Values = LoopDispositions[S];
14415 for (auto &V : Values) {
14416 if (V.getPointer() == L)
14417 return V.getInt();
14418 }
14419 Values.emplace_back(L, LoopVariant);
14420 LoopDisposition D = computeLoopDisposition(S, L);
14421 auto &Values2 = LoopDispositions[S];
14422 for (auto &V : llvm::reverse(Values2)) {
14423 if (V.getPointer() == L) {
14424 V.setInt(D);
14425 break;
14426 }
14427 }
14428 return D;
14429}
14430
14432ScalarEvolution::computeLoopDisposition(const SCEV *S, const Loop *L) {
14433 switch (S->getSCEVType()) {
14434 case scConstant:
14435 case scVScale:
14436 return LoopInvariant;
14437 case scAddRecExpr: {
14438 const SCEVAddRecExpr *AR = cast<SCEVAddRecExpr>(S);
14439
14440 // If L is the addrec's loop, it's computable.
14441 if (AR->getLoop() == L)
14442 return LoopComputable;
14443
14444 // Add recurrences are never invariant in the function-body (null loop).
14445 if (!L)
14446 return LoopVariant;
14447
14448 // Everything that is not defined at loop entry is variant.
14449 if (DT.dominates(L->getHeader(), AR->getLoop()->getHeader())) {
14450 if (L->contains(AR->getLoop()) &&
14451 llvm::all_of(AR->operands(),
14452 [&](const SCEV *Op) { return isLoopUniform(Op, L); }))
14453 return LoopUniform;
14454
14455 return LoopVariant;
14456 }
14457 assert(!L->contains(AR->getLoop()) && "Containing loop's header does not"
14458 " dominate the contained loop's header?");
14459
14460 // This recurrence is invariant w.r.t. L if AR's loop contains L.
14461 if (AR->getLoop()->contains(L))
14462 return LoopInvariant;
14463
14464 // This recurrence is variant w.r.t. L if any of its operands
14465 // are variant.
14466 for (SCEVUse Op : AR->operands())
14467 if (!isLoopInvariant(Op, L))
14468 return LoopVariant;
14469
14470 // Otherwise it's loop-invariant.
14471 return LoopInvariant;
14472 }
14473 case scTruncate:
14474 case scZeroExtend:
14475 case scSignExtend:
14476 case scPtrToAddr:
14477 case scAddExpr:
14478 case scMulExpr:
14479 case scUDivExpr:
14480 case scUMaxExpr:
14481 case scSMaxExpr:
14482 case scUMinExpr:
14483 case scSMinExpr:
14484 case scSequentialUMinExpr: {
14485 bool HasVarying = false;
14486 bool HasUniform = false;
14487 for (SCEVUse Op : S->operands()) {
14489 if (D == LoopVariant)
14490 return LoopVariant;
14491 if (D == LoopComputable)
14492 HasVarying = true;
14493 if (D == LoopUniform)
14494 HasUniform = true;
14495 }
14496 return HasVarying ? (HasUniform ? LoopVariant : LoopComputable)
14497 : (HasUniform ? LoopUniform : LoopInvariant);
14498 }
14499 case scUnknown:
14500 // All non-instruction values are loop invariant. All instructions are loop
14501 // invariant if they are not contained in the specified loop.
14502 // Instructions are never considered invariant in the function body
14503 // (null loop) because they are defined within the "loop".
14505 return (L && !L->contains(I)) ? LoopInvariant : LoopVariant;
14506 return LoopInvariant;
14507 case scCouldNotCompute:
14508 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
14509 }
14510 llvm_unreachable("Unknown SCEV kind!");
14511}
14512
14513bool ScalarEvolution::isLoopUniform(const SCEV *S, const Loop *L) {
14515 return D == LoopUniform || D == LoopInvariant;
14516}
14517
14519 return getLoopDisposition(S, L) == LoopInvariant;
14520}
14521
14523 return getLoopDisposition(S, L) == LoopComputable;
14524}
14525
14528 auto &Values = BlockDispositions[S];
14529 for (auto &V : Values) {
14530 if (V.getPointer() == BB)
14531 return V.getInt();
14532 }
14533 Values.emplace_back(BB, DoesNotDominateBlock);
14534 BlockDisposition D = computeBlockDisposition(S, BB);
14535 auto &Values2 = BlockDispositions[S];
14536 for (auto &V : llvm::reverse(Values2)) {
14537 if (V.getPointer() == BB) {
14538 V.setInt(D);
14539 break;
14540 }
14541 }
14542 return D;
14543}
14544
14546ScalarEvolution::computeBlockDisposition(const SCEV *S, const BasicBlock *BB) {
14547 switch (S->getSCEVType()) {
14548 case scConstant:
14549 case scVScale:
14551 case scAddRecExpr: {
14552 // This uses a "dominates" query instead of "properly dominates" query
14553 // to test for proper dominance too, because the instruction which
14554 // produces the addrec's value is a PHI, and a PHI effectively properly
14555 // dominates its entire containing block.
14556 const SCEVAddRecExpr *AR = cast<SCEVAddRecExpr>(S);
14557 if (!DT.dominates(AR->getLoop()->getHeader(), BB))
14558 return DoesNotDominateBlock;
14559
14560 // Fall through into SCEVNAryExpr handling.
14561 [[fallthrough]];
14562 }
14563 case scTruncate:
14564 case scZeroExtend:
14565 case scSignExtend:
14566 case scPtrToAddr:
14567 case scAddExpr:
14568 case scMulExpr:
14569 case scUDivExpr:
14570 case scUMaxExpr:
14571 case scSMaxExpr:
14572 case scUMinExpr:
14573 case scSMinExpr:
14574 case scSequentialUMinExpr: {
14575 bool Proper = true;
14576 for (const SCEV *NAryOp : S->operands()) {
14578 if (D == DoesNotDominateBlock)
14579 return DoesNotDominateBlock;
14580 if (D == DominatesBlock)
14581 Proper = false;
14582 }
14583 return Proper ? ProperlyDominatesBlock : DominatesBlock;
14584 }
14585 case scUnknown:
14586 if (Instruction *I =
14588 if (I->getParent() == BB)
14589 return DominatesBlock;
14590 if (DT.properlyDominates(I->getParent(), BB))
14592 return DoesNotDominateBlock;
14593 }
14595 case scCouldNotCompute:
14596 llvm_unreachable("Attempt to use a SCEVCouldNotCompute object!");
14597 }
14598 llvm_unreachable("Unknown SCEV kind!");
14599}
14600
14601bool ScalarEvolution::dominates(const SCEV *S, const BasicBlock *BB) {
14602 return getBlockDisposition(S, BB) >= DominatesBlock;
14603}
14604
14607}
14608
14609void ScalarEvolution::forgetBackedgeTakenCounts(const Loop *L,
14610 bool Predicated) {
14611 auto &BECounts =
14612 Predicated ? PredicatedBackedgeTakenCounts : BackedgeTakenCounts;
14613 auto It = BECounts.find(L);
14614 if (It != BECounts.end()) {
14615 for (const ExitNotTakenInfo &ENT : It->second.ExitNotTaken) {
14616 for (const SCEV *S : {ENT.ExactNotTaken, ENT.SymbolicMaxNotTaken}) {
14617 if (!isa<SCEVConstant>(S)) {
14618 auto UserIt = BECountUsers.find(S);
14619 assert(UserIt != BECountUsers.end());
14620 UserIt->second.erase({L, Predicated});
14621 }
14622 }
14623 }
14624 BECounts.erase(It);
14625 }
14626}
14627
14628void ScalarEvolution::forgetMemoizedResults(ArrayRef<SCEVUse> SCEVs) {
14629 SmallPtrSet<const SCEV *, 8> ToForget(llvm::from_range, SCEVs);
14630 SmallVector<SCEVUse, 8> Worklist(ToForget.begin(), ToForget.end());
14631
14632 while (!Worklist.empty()) {
14633 const SCEV *Curr = Worklist.pop_back_val();
14634 auto Users = SCEVUsers.find(Curr);
14635 if (Users != SCEVUsers.end())
14636 for (const auto *User : Users->second)
14637 if (ToForget.insert(User).second)
14638 Worklist.push_back(User);
14639 }
14640
14641 for (const auto *S : ToForget)
14642 forgetMemoizedResultsImpl(S);
14643
14644 PredicatedSCEVRewrites.remove_if(
14645 [&](const auto &Entry) { return ToForget.count(Entry.first.first); });
14646}
14647
14648void ScalarEvolution::forgetMemoizedResultsImpl(const SCEV *S) {
14649 LoopDispositions.erase(S);
14650 BlockDispositions.erase(S);
14651 UnsignedRanges.erase(S);
14652 SignedRanges.erase(S);
14653 HasRecMap.erase(S);
14654 ConstantMultipleCache.erase(S);
14655
14656 if (auto *AR = dyn_cast<SCEVAddRecExpr>(S)) {
14657 UnsignedWrapViaInductionTried.erase(AR);
14658 SignedWrapViaInductionTried.erase(AR);
14659 }
14660
14661 auto ExprIt = ExprValueMap.find(S);
14662 if (ExprIt != ExprValueMap.end()) {
14663 for (Value *V : ExprIt->second) {
14664 auto ValueIt = ValueExprMap.find_as(V);
14665 if (ValueIt != ValueExprMap.end())
14666 ValueExprMap.erase(ValueIt);
14667 }
14668 ExprValueMap.erase(ExprIt);
14669 }
14670
14671 auto ScopeIt = ValuesAtScopes.find(S);
14672 if (ScopeIt != ValuesAtScopes.end()) {
14673 for (const auto &Pair : ScopeIt->second)
14674 if (!isa_and_nonnull<SCEVConstant>(Pair.second))
14675 llvm::erase(ValuesAtScopesUsers[Pair.second.getPointer()],
14676 std::make_pair(Pair.first, S));
14677 ValuesAtScopes.erase(ScopeIt);
14678 }
14679
14680 auto ScopeUserIt = ValuesAtScopesUsers.find(S);
14681 if (ScopeUserIt != ValuesAtScopesUsers.end()) {
14682 for (const auto &Pair : ScopeUserIt->second)
14683 // The recorded value at scope is a use of S, which may carry no-wrap
14684 // flags that are not part of this key.
14685 llvm::erase_if(ValuesAtScopes[Pair.second], [&](const auto &LS) {
14686 return LS.first == Pair.first && LS.second.getPointer() == S;
14687 });
14688 ValuesAtScopesUsers.erase(ScopeUserIt);
14689 }
14690
14691 auto BEUsersIt = BECountUsers.find(S);
14692 if (BEUsersIt != BECountUsers.end()) {
14693 // Work on a copy, as forgetBackedgeTakenCounts() will modify the original.
14694 auto Copy = BEUsersIt->second;
14695 for (const auto &Pair : Copy)
14696 forgetBackedgeTakenCounts(Pair.getPointer(), Pair.getInt());
14697 BECountUsers.erase(BEUsersIt);
14698 }
14699
14700 auto FoldUser = FoldCacheUser.find(S);
14701 if (FoldUser != FoldCacheUser.end())
14702 for (auto &KV : FoldUser->second)
14703 FoldCache.erase(KV);
14704 FoldCacheUser.erase(S);
14705}
14706
14707void
14708ScalarEvolution::getUsedLoops(const SCEV *S,
14709 SmallPtrSetImpl<const Loop *> &LoopsUsed) {
14710 struct FindUsedLoops {
14711 FindUsedLoops(SmallPtrSetImpl<const Loop *> &LoopsUsed)
14712 : LoopsUsed(LoopsUsed) {}
14713 SmallPtrSetImpl<const Loop *> &LoopsUsed;
14714 bool follow(const SCEV *S) {
14715 if (auto *AR = dyn_cast<SCEVAddRecExpr>(S))
14716 LoopsUsed.insert(AR->getLoop());
14717 return true;
14718 }
14719
14720 bool isDone() const { return false; }
14721 };
14722
14723 FindUsedLoops F(LoopsUsed);
14724 SCEVTraversal<FindUsedLoops>(F).visitAll(S);
14725}
14726
14727void ScalarEvolution::getReachableBlocks(
14730 Worklist.push_back(&F.getEntryBlock());
14731 while (!Worklist.empty()) {
14732 BasicBlock *BB = Worklist.pop_back_val();
14733 if (!Reachable.insert(BB).second)
14734 continue;
14735
14736 Value *Cond;
14737 BasicBlock *TrueBB, *FalseBB;
14738 if (match(BB->getTerminator(), m_Br(m_Value(Cond), m_BasicBlock(TrueBB),
14739 m_BasicBlock(FalseBB)))) {
14740 if (auto *C = dyn_cast<ConstantInt>(Cond)) {
14741 Worklist.push_back(C->isOne() ? TrueBB : FalseBB);
14742 continue;
14743 }
14744
14745 if (auto *Cmp = dyn_cast<ICmpInst>(Cond)) {
14746 const SCEV *L = getSCEV(Cmp->getOperand(0));
14747 const SCEV *R = getSCEV(Cmp->getOperand(1));
14748 if (isKnownPredicateViaConstantRanges(Cmp->getCmpPredicate(), L, R)) {
14749 Worklist.push_back(TrueBB);
14750 continue;
14751 }
14752 if (isKnownPredicateViaConstantRanges(Cmp->getInverseCmpPredicate(), L,
14753 R)) {
14754 Worklist.push_back(FalseBB);
14755 continue;
14756 }
14757 }
14758 }
14759
14760 append_range(Worklist, successors(BB));
14761 }
14762}
14763
14765 ScalarEvolution &SE = *const_cast<ScalarEvolution *>(this);
14766 ScalarEvolution SE2(F, TLI, AC, DT, LI);
14767
14768 SmallVector<Loop *, 8> LoopStack(LI.begin(), LI.end());
14769
14770 // Map's SCEV expressions from one ScalarEvolution "universe" to another.
14771 struct SCEVMapper : public SCEVRewriteVisitor<SCEVMapper> {
14772 SCEVMapper(ScalarEvolution &SE) : SCEVRewriteVisitor<SCEVMapper>(SE) {}
14773
14774 const SCEV *visitConstant(const SCEVConstant *Constant) {
14775 return SE.getConstant(Constant->getAPInt());
14776 }
14777
14778 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
14779 return SE.getUnknown(Expr->getValue());
14780 }
14781
14782 const SCEV *visitCouldNotCompute(const SCEVCouldNotCompute *Expr) {
14783 return SE.getCouldNotCompute();
14784 }
14785 };
14786
14787 SCEVMapper SCM(SE2);
14788 SmallPtrSet<BasicBlock *, 16> ReachableBlocks;
14789 SE2.getReachableBlocks(ReachableBlocks, F);
14790
14791 auto GetDelta = [&](const SCEV *Old, const SCEV *New) -> const SCEV * {
14792 if (containsUndefs(Old) || containsUndefs(New)) {
14793 // SCEV treats "undef" as an unknown but consistent value (i.e. it does
14794 // not propagate undef aggressively). This means we can (and do) fail
14795 // verification in cases where a transform makes a value go from "undef"
14796 // to "undef+1" (say). The transform is fine, since in both cases the
14797 // result is "undef", but SCEV thinks the value increased by 1.
14798 return nullptr;
14799 }
14800
14801 // Unless VerifySCEVStrict is set, we only compare constant deltas.
14802 const SCEV *Delta = SE2.getMinusSCEV(Old, New);
14803 if (!VerifySCEVStrict && !isa<SCEVConstant>(Delta))
14804 return nullptr;
14805
14806 return Delta;
14807 };
14808
14809 while (!LoopStack.empty()) {
14810 auto *L = LoopStack.pop_back_val();
14811 llvm::append_range(LoopStack, *L);
14812
14813 // Only verify BECounts in reachable loops. For an unreachable loop,
14814 // any BECount is legal.
14815 if (!ReachableBlocks.contains(L->getHeader()))
14816 continue;
14817
14818 // Only verify cached BECounts. Computing new BECounts may change the
14819 // results of subsequent SCEV uses.
14820 auto It = BackedgeTakenCounts.find(L);
14821 if (It == BackedgeTakenCounts.end())
14822 continue;
14823
14824 auto *CurBECount =
14825 SCM.visit(It->second.getExact(L, const_cast<ScalarEvolution *>(this)));
14826 auto *NewBECount = SE2.getBackedgeTakenCount(L);
14827
14828 if (CurBECount == SE2.getCouldNotCompute() ||
14829 NewBECount == SE2.getCouldNotCompute()) {
14830 // NB! This situation is legal, but is very suspicious -- whatever pass
14831 // change the loop to make a trip count go from could not compute to
14832 // computable or vice-versa *should have* invalidated SCEV. However, we
14833 // choose not to assert here (for now) since we don't want false
14834 // positives.
14835 continue;
14836 }
14837
14838 if (SE.getTypeSizeInBits(CurBECount->getType()) >
14839 SE.getTypeSizeInBits(NewBECount->getType()))
14840 NewBECount = SE2.getZeroExtendExpr(NewBECount, CurBECount->getType());
14841 else if (SE.getTypeSizeInBits(CurBECount->getType()) <
14842 SE.getTypeSizeInBits(NewBECount->getType()))
14843 CurBECount = SE2.getZeroExtendExpr(CurBECount, NewBECount->getType());
14844
14845 const SCEV *Delta = GetDelta(CurBECount, NewBECount);
14846 if (Delta && !Delta->isZero()) {
14847 dbgs() << "Trip Count for " << *L << " Changed!\n";
14848 dbgs() << "Old: " << *CurBECount << "\n";
14849 dbgs() << "New: " << *NewBECount << "\n";
14850 dbgs() << "Delta: " << *Delta << "\n";
14851 std::abort();
14852 }
14853 }
14854
14855 // Collect all valid loops currently in LoopInfo.
14856 SmallPtrSet<Loop *, 32> ValidLoops;
14857 SmallVector<Loop *, 32> Worklist(LI.begin(), LI.end());
14858 while (!Worklist.empty()) {
14859 Loop *L = Worklist.pop_back_val();
14860 if (ValidLoops.insert(L).second)
14861 Worklist.append(L->begin(), L->end());
14862 }
14863 for (const auto &KV : ValueExprMap) {
14864#ifndef NDEBUG
14865 // Check for SCEV expressions referencing invalid/deleted loops.
14866 if (auto *AR = dyn_cast<SCEVAddRecExpr>(KV.second)) {
14867 assert(ValidLoops.contains(AR->getLoop()) &&
14868 "AddRec references invalid loop");
14869 }
14870#endif
14871
14872 // Check that the value is also part of the reverse map.
14873 auto It = ExprValueMap.find(KV.second);
14874 if (It == ExprValueMap.end() || !It->second.contains(KV.first)) {
14875 dbgs() << "Value " << *KV.first
14876 << " is in ValueExprMap but not in ExprValueMap\n";
14877 std::abort();
14878 }
14879
14880 if (auto *I = dyn_cast<Instruction>(&*KV.first)) {
14881 if (!ReachableBlocks.contains(I->getParent()))
14882 continue;
14883 const SCEV *OldSCEV = SCM.visit(KV.second);
14884 const SCEV *NewSCEV = SE2.getSCEV(I);
14885 const SCEV *Delta = GetDelta(OldSCEV, NewSCEV);
14886 if (Delta && !Delta->isZero()) {
14887 dbgs() << "SCEV for value " << *I << " changed!\n"
14888 << "Old: " << *OldSCEV << "\n"
14889 << "New: " << *NewSCEV << "\n"
14890 << "Delta: " << *Delta << "\n";
14891 std::abort();
14892 }
14893 }
14894 }
14895
14896 for (const auto &KV : ExprValueMap) {
14897 for (Value *V : KV.second) {
14898 const SCEV *S = ValueExprMap.lookup(V);
14899 if (!S) {
14900 dbgs() << "Value " << *V
14901 << " is in ExprValueMap but not in ValueExprMap\n";
14902 std::abort();
14903 }
14904 if (S != KV.first) {
14905 dbgs() << "Value " << *V << " mapped to " << *S << " rather than "
14906 << *KV.first << "\n";
14907 std::abort();
14908 }
14909 }
14910 }
14911
14912 // Verify integrity of SCEV users.
14913 for (const auto &S : UniqueSCEVs) {
14914 for (SCEVUse Op : S.operands()) {
14915 // We do not store dependencies of constants.
14916 if (isa<SCEVConstant>(Op))
14917 continue;
14918 auto It = SCEVUsers.find(Op);
14919 if (It != SCEVUsers.end() && It->second.count(&S))
14920 continue;
14921 dbgs() << "Use of operand " << *Op << " by user " << S
14922 << " is not being tracked!\n";
14923 std::abort();
14924 }
14925 }
14926
14927 // Verify integrity of ValuesAtScopes users.
14928 for (const auto &ValueAndVec : ValuesAtScopes) {
14929 const SCEV *Value = ValueAndVec.first;
14930 for (const auto &LoopAndValueAtScope : ValueAndVec.second) {
14931 const Loop *L = LoopAndValueAtScope.first;
14932 SCEVUse ValueAtScope = LoopAndValueAtScope.second;
14933 if (!isa<SCEVConstant>(ValueAtScope)) {
14934 auto It = ValuesAtScopesUsers.find(ValueAtScope.getPointer());
14935 if (It != ValuesAtScopesUsers.end() &&
14936 is_contained(It->second, std::make_pair(L, Value)))
14937 continue;
14938 dbgs() << "Value: " << *Value << ", Loop: " << *L << ", ValueAtScope: "
14939 << *ValueAtScope << " missing in ValuesAtScopesUsers\n";
14940 std::abort();
14941 }
14942 }
14943 }
14944
14945 for (const auto &ValueAtScopeAndVec : ValuesAtScopesUsers) {
14946 const SCEV *ValueAtScope = ValueAtScopeAndVec.first;
14947 for (const auto &LoopAndValue : ValueAtScopeAndVec.second) {
14948 const Loop *L = LoopAndValue.first;
14949 const SCEV *Value = LoopAndValue.second;
14951 auto It = ValuesAtScopes.find(Value);
14952 // The recorded value at scope may carry no-wrap flags that are not part
14953 // of the key it is recorded under.
14954 if (It != ValuesAtScopes.end() && any_of(It->second, [&](const auto &LS) {
14955 return LS.first == L && LS.second.getPointer() == ValueAtScope;
14956 }))
14957 continue;
14958 dbgs() << "Value: " << *Value << ", Loop: " << *L << ", ValueAtScope: "
14959 << *ValueAtScope << " missing in ValuesAtScopes\n";
14960 std::abort();
14961 }
14962 }
14963
14964 // Verify integrity of BECountUsers.
14965 auto VerifyBECountUsers = [&](bool Predicated) {
14966 auto &BECounts =
14967 Predicated ? PredicatedBackedgeTakenCounts : BackedgeTakenCounts;
14968 for (const auto &LoopAndBEInfo : BECounts) {
14969 for (const ExitNotTakenInfo &ENT : LoopAndBEInfo.second.ExitNotTaken) {
14970 for (const SCEV *S : {ENT.ExactNotTaken, ENT.SymbolicMaxNotTaken}) {
14971 if (!isa<SCEVConstant>(S)) {
14972 auto UserIt = BECountUsers.find(S);
14973 if (UserIt != BECountUsers.end() &&
14974 UserIt->second.contains({ LoopAndBEInfo.first, Predicated }))
14975 continue;
14976 dbgs() << "Value " << *S << " for loop " << *LoopAndBEInfo.first
14977 << " missing from BECountUsers\n";
14978 std::abort();
14979 }
14980 }
14981 }
14982 }
14983 };
14984 VerifyBECountUsers(/* Predicated */ false);
14985 VerifyBECountUsers(/* Predicated */ true);
14986
14987 // Verify intergity of loop disposition cache.
14988 for (auto &[S, Values] : LoopDispositions) {
14989 for (auto [Loop, CachedDisposition] : Values) {
14990 const auto RecomputedDisposition = SE2.getLoopDisposition(S, Loop);
14991 if (CachedDisposition != RecomputedDisposition) {
14992 dbgs() << "Cached disposition of " << *S << " for loop " << *Loop
14993 << " is incorrect: cached " << CachedDisposition << ", actual "
14994 << RecomputedDisposition << "\n";
14995 std::abort();
14996 }
14997 }
14998 }
14999
15000 // Verify integrity of the block disposition cache.
15001 for (auto &[S, Values] : BlockDispositions) {
15002 for (auto [BB, CachedDisposition] : Values) {
15003 const auto RecomputedDisposition = SE2.getBlockDisposition(S, BB);
15004 if (CachedDisposition != RecomputedDisposition) {
15005 dbgs() << "Cached disposition of " << *S << " for block %"
15006 << BB->getName() << " is incorrect: cached " << CachedDisposition
15007 << ", actual " << RecomputedDisposition << "\n";
15008 std::abort();
15009 }
15010 }
15011 }
15012
15013 // Verify FoldCache/FoldCacheUser caches.
15014 for (auto [FoldID, Expr] : FoldCache) {
15015 auto I = FoldCacheUser.find(Expr);
15016 if (I == FoldCacheUser.end()) {
15017 dbgs() << "Missing entry in FoldCacheUser for cached expression " << *Expr
15018 << "!\n";
15019 std::abort();
15020 }
15021 if (!is_contained(I->second, FoldID)) {
15022 dbgs() << "Missing FoldID in cached users of " << *Expr << "!\n";
15023 std::abort();
15024 }
15025 }
15026 for (auto [Expr, IDs] : FoldCacheUser) {
15027 for (auto &FoldID : IDs) {
15028 const SCEV *S = FoldCache.lookup(FoldID);
15029 if (!S) {
15030 dbgs() << "Missing entry in FoldCache for expression " << *Expr
15031 << "!\n";
15032 std::abort();
15033 }
15034 if (S != Expr) {
15035 dbgs() << "Entry in FoldCache doesn't match FoldCacheUser: " << *S
15036 << " != " << *Expr << "!\n";
15037 std::abort();
15038 }
15039 }
15040 }
15041
15042 // Verify that ConstantMultipleCache computations are correct. We check that
15043 // cached multiples and recomputed multiples are multiples of each other to
15044 // verify correctness. It is possible that a recomputed multiple is different
15045 // from the cached multiple due to strengthened no wrap flags or changes in
15046 // KnownBits computations.
15047 for (auto [S, Multiple] : ConstantMultipleCache) {
15048 APInt RecomputedMultiple = SE2.getConstantMultiple(S);
15049 if ((Multiple != 0 && RecomputedMultiple != 0 &&
15050 Multiple.urem(RecomputedMultiple) != 0 &&
15051 RecomputedMultiple.urem(Multiple) != 0)) {
15052 dbgs() << "Incorrect cached computation in ConstantMultipleCache for "
15053 << *S << " : Computed " << RecomputedMultiple
15054 << " but cache contains " << Multiple << "!\n";
15055 std::abort();
15056 }
15057 }
15058}
15059
15061 Function &F, const PreservedAnalyses &PA,
15062 FunctionAnalysisManager::Invalidator &Inv) {
15063 // Invalidate the ScalarEvolution object whenever it isn't preserved or one
15064 // of its dependencies is invalidated.
15065 auto PAC = PA.getChecker<ScalarEvolutionAnalysis>();
15066 return !(PAC.preserved() || PAC.preservedSet<AllAnalysesOn<Function>>()) ||
15067 Inv.invalidate<AssumptionAnalysis>(F, PA) ||
15068 Inv.invalidate<DominatorTreeAnalysis>(F, PA) ||
15069 Inv.invalidate<LoopAnalysis>(F, PA);
15070}
15071
15072AnalysisKey ScalarEvolutionAnalysis::Key;
15073
15076 auto &TLI = AM.getResult<TargetLibraryAnalysis>(F);
15077 auto &AC = AM.getResult<AssumptionAnalysis>(F);
15078 auto &DT = AM.getResult<DominatorTreeAnalysis>(F);
15079 auto &LI = AM.getResult<LoopAnalysis>(F);
15080 return ScalarEvolution(F, TLI, AC, DT, LI);
15081}
15082
15088
15091 // For compatibility with opt's -analyze feature under legacy pass manager
15092 // which was not ported to NPM. This keeps tests using
15093 // update_analyze_test_checks.py working.
15094 OS << "Printing analysis 'Scalar Evolution Analysis' for function '"
15095 << F.getName() << "':\n";
15097 return PreservedAnalyses::all();
15098}
15099
15101 "Scalar Evolution Analysis", false, true)
15107 "Scalar Evolution Analysis", false, true)
15108
15109char ScalarEvolutionWrapperPass::ID = 0;
15110
15112
15114 SE.reset(new ScalarEvolution(
15116 getAnalysis<AssumptionCacheTracker>().getAssumptionCache(F),
15118 getAnalysis<LoopInfoWrapperPass>().getLoopInfo()));
15119 return false;
15120}
15121
15123
15125 SE->print(OS);
15126}
15127
15129 if (!VerifySCEV)
15130 return;
15131
15132 SE->verify();
15133}
15134
15142
15144 const SCEV *RHS) {
15145 return getComparePredicate(ICmpInst::ICMP_EQ, LHS, RHS);
15146}
15147
15148const SCEVPredicate *
15150 const SCEV *LHS, const SCEV *RHS) {
15152 assert(LHS->getType() == RHS->getType() &&
15153 "Type mismatch between LHS and RHS");
15154 // Unique this node based on the arguments
15155 ID.AddInteger(SCEVPredicate::P_Compare);
15156 ID.AddInteger(Pred);
15157 ID.AddPointer(LHS);
15158 ID.AddPointer(RHS);
15160 if (const auto *S = UniquePreds.lookup(ID, Token))
15161 return S;
15162 SCEVComparePredicate *Eq = new (SCEVAllocator)
15163 SCEVComparePredicate(ID.Intern(SCEVAllocator), Pred, LHS, RHS);
15164 UniquePreds.insert(Eq, Token);
15165 return Eq;
15166}
15167
15169 const SCEVAddRecExpr *AR,
15172 // Unique this node based on the arguments
15174 ID.AddPointer(AR);
15175 ID.AddInteger(AddedFlags);
15177 if (const auto *S = UniquePreds.lookup(ID, Token))
15178 return S;
15179 auto *OF = new (SCEVAllocator)
15180 SCEVWrapPredicate(ID.Intern(SCEVAllocator), AR, AddedFlags);
15181 UniquePreds.insert(OF, Token);
15182 return OF;
15183}
15184
15185namespace {
15186
15187class SCEVPredicateRewriter : public SCEVRewriteVisitor<SCEVPredicateRewriter> {
15188public:
15189
15190 /// Rewrites \p S in the context of a loop L and the SCEV predication
15191 /// infrastructure.
15192 ///
15193 /// If \p Pred is non-null, the SCEV expression is rewritten to respect the
15194 /// equivalences present in \p Pred.
15195 ///
15196 /// If \p NewPreds is non-null, rewrite is free to add further predicates to
15197 /// \p NewPreds such that the result will be an AddRecExpr.
15198 static const SCEV *rewrite(const SCEV *S, const Loop *L, ScalarEvolution &SE,
15200 const SCEVPredicate *Pred) {
15201 SCEVPredicateRewriter Rewriter(L, SE, NewPreds, Pred);
15202 return Rewriter.visit(S);
15203 }
15204
15205 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
15206 if (Pred) {
15207 if (auto *U = dyn_cast<SCEVUnionPredicate>(Pred)) {
15208 for (const auto *Pred : U->getPredicates())
15209 if (const auto *IPred = dyn_cast<SCEVComparePredicate>(Pred))
15210 if (IPred->getLHS() == Expr &&
15211 IPred->getPredicate() == ICmpInst::ICMP_EQ)
15212 return IPred->getRHS();
15213 } else if (const auto *IPred = dyn_cast<SCEVComparePredicate>(Pred)) {
15214 if (IPred->getLHS() == Expr &&
15215 IPred->getPredicate() == ICmpInst::ICMP_EQ)
15216 return IPred->getRHS();
15217 }
15218 }
15219 return convertToAddRecWithPreds(Expr);
15220 }
15221
15222 const SCEV *visitZeroExtendExpr(const SCEVZeroExtendExpr *Expr) {
15223 const SCEV *Operand = visit(Expr->getOperand());
15224 const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(Operand);
15225 if (AR && AR->getLoop() == L && AR->isAffine()) {
15226 // This couldn't be folded because the operand didn't have the nuw
15227 // flag. Add the nusw flag as an assumption that we could make.
15228 const SCEV *Step = AR->getStepRecurrence(SE);
15229 Type *Ty = Expr->getType();
15230 if (addOverflowAssumption(AR, SCEVWrapPredicate::IncrementNUSW))
15231 return SE.getAddRecExpr(SE.getZeroExtendExpr(AR->getStart(), Ty),
15232 SE.getSignExtendExpr(Step, Ty), L,
15233 AR->getNoWrapFlags());
15234 }
15235 return SE.getZeroExtendExpr(Operand, Expr->getType());
15236 }
15237
15238 const SCEV *visitSignExtendExpr(const SCEVSignExtendExpr *Expr) {
15239 const SCEV *Operand = visit(Expr->getOperand());
15240 const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(Operand);
15241 if (AR && AR->getLoop() == L && AR->isAffine()) {
15242 // This couldn't be folded because the operand didn't have the nsw
15243 // flag. Add the nssw flag as an assumption that we could make.
15244 const SCEV *Step = AR->getStepRecurrence(SE);
15245 Type *Ty = Expr->getType();
15246 if (addOverflowAssumption(AR, SCEVWrapPredicate::IncrementNSSW))
15247 return SE.getAddRecExpr(SE.getSignExtendExpr(AR->getStart(), Ty),
15248 SE.getSignExtendExpr(Step, Ty), L,
15249 AR->getNoWrapFlags());
15250 }
15251 return SE.getSignExtendExpr(Operand, Expr->getType());
15252 }
15253
15254private:
15255 explicit SCEVPredicateRewriter(
15256 const Loop *L, ScalarEvolution &SE,
15257 SmallVectorImpl<const SCEVPredicate *> *NewPreds,
15258 const SCEVPredicate *Pred)
15259 : SCEVRewriteVisitor(SE), NewPreds(NewPreds), Pred(Pred), L(L) {}
15260
15261 bool addOverflowAssumption(const SCEVPredicate *P) {
15262 if (!NewPreds) {
15263 // Check if we've already made this assumption.
15264 return Pred && Pred->implies(P, SE);
15265 }
15266 NewPreds->push_back(P);
15267 return true;
15268 }
15269
15270 bool addOverflowAssumption(const SCEVAddRecExpr *AR,
15272 auto *A = SE.getWrapPredicate(AR, AddedFlags);
15273 return addOverflowAssumption(A);
15274 }
15275
15276 // If \p Expr represents a PHINode, we try to see if it can be represented
15277 // as an AddRec, possibly under a predicate (PHISCEVPred). If it is possible
15278 // to add this predicate as a runtime overflow check, we return the AddRec.
15279 // If \p Expr does not meet these conditions (is not a PHI node, or we
15280 // couldn't create an AddRec for it, or couldn't add the predicate), we just
15281 // return \p Expr.
15282 const SCEV *convertToAddRecWithPreds(const SCEVUnknown *Expr) {
15283 if (!isa<PHINode>(Expr->getValue()))
15284 return Expr;
15285 std::optional<
15286 std::pair<const SCEV *, SmallVector<const SCEVPredicate *, 3>>>
15287 PredicatedRewrite = SE.createAddRecFromPHIWithCasts(Expr);
15288 if (!PredicatedRewrite)
15289 return Expr;
15290 for (const auto *P : PredicatedRewrite->second){
15291 // Wrap predicates from outer loops are not supported.
15292 if (auto *WP = dyn_cast<const SCEVWrapPredicate>(P)) {
15293 if (L != WP->getExpr()->getLoop())
15294 return Expr;
15295 }
15296 if (!addOverflowAssumption(P))
15297 return Expr;
15298 }
15299 return PredicatedRewrite->first;
15300 }
15301
15302 SmallVectorImpl<const SCEVPredicate *> *NewPreds;
15303 const SCEVPredicate *Pred;
15304 const Loop *L;
15305};
15306
15307} // end anonymous namespace
15308
15309const SCEV *
15311 const SCEVPredicate &Preds) {
15312 return SCEVPredicateRewriter::rewrite(S, L, *this, nullptr, &Preds);
15313}
15314
15316 const SCEV *S, const Loop *L,
15319 S = SCEVPredicateRewriter::rewrite(S, L, *this, &TransformPreds, nullptr);
15320 auto *AddRec = dyn_cast<SCEVAddRecExpr>(S);
15321
15322 if (!AddRec)
15323 return nullptr;
15324
15325 // Check if any of the transformed predicates is known to be false. In that
15326 // case, it doesn't make sense to convert to a predicated AddRec, as the
15327 // versioned loop will never execute.
15328 for (const SCEVPredicate *Pred : TransformPreds) {
15329 auto *WrapPred = dyn_cast<SCEVWrapPredicate>(Pred);
15330 if (!WrapPred || WrapPred->getFlags() != SCEVWrapPredicate::IncrementNSSW)
15331 continue;
15332
15333 const SCEVAddRecExpr *AddRecToCheck = WrapPred->getExpr();
15334 const SCEV *ExitCount = getBackedgeTakenCount(AddRecToCheck->getLoop());
15335 if (isa<SCEVCouldNotCompute>(ExitCount))
15336 continue;
15337
15338 const SCEV *Step = AddRecToCheck->getStepRecurrence(*this);
15339 if (!Step->isOne())
15340 continue;
15341
15342 ExitCount = getTruncateOrSignExtend(ExitCount, Step->getType());
15343 const SCEV *Add = getAddExpr(AddRecToCheck->getStart(), ExitCount);
15344 if (isKnownPredicate(CmpInst::ICMP_SLT, Add, AddRecToCheck->getStart()))
15345 return nullptr;
15346 }
15347
15348 // Since the transformation was successful, we can now transfer the SCEV
15349 // predicates.
15350 Preds.append(TransformPreds.begin(), TransformPreds.end());
15351
15352 return AddRec;
15353}
15354
15355/// SCEV predicates
15359
15361 const ICmpInst::Predicate Pred,
15362 const SCEV *LHS, const SCEV *RHS)
15363 : SCEVPredicate(ID, P_Compare), Pred(Pred), LHS(LHS), RHS(RHS) {
15364 assert(LHS->getType() == RHS->getType() && "LHS and RHS types don't match");
15365 assert(LHS != RHS && "LHS and RHS are the same SCEV");
15366}
15367
15369 ScalarEvolution &SE) const {
15370 const auto *Op = dyn_cast<SCEVComparePredicate>(N);
15371
15372 if (!Op)
15373 return false;
15374
15375 if (Pred != ICmpInst::ICMP_EQ)
15376 return false;
15377
15378 return Op->LHS == LHS && Op->RHS == RHS;
15379}
15380
15381bool SCEVComparePredicate::isAlwaysTrue() const { return false; }
15382
15384 if (Pred == ICmpInst::ICMP_EQ)
15385 OS.indent(Depth) << "Equal predicate: " << *LHS << " == " << *RHS << "\n";
15386 else
15387 OS.indent(Depth) << "Compare predicate: " << *LHS << " " << Pred << ") "
15388 << *RHS << "\n";
15389
15390}
15391
15393 const SCEVAddRecExpr *AR,
15394 IncrementWrapFlags Flags)
15395 : SCEVPredicate(ID, P_Wrap), AR(AR), Flags(Flags) {}
15396
15397const SCEVAddRecExpr *SCEVWrapPredicate::getExpr() const { return AR; }
15398
15400 ScalarEvolution &SE) const {
15401 const auto *Op = dyn_cast<SCEVWrapPredicate>(N);
15402 if (!Op || setFlags(Flags, Op->Flags) != Flags)
15403 return false;
15404
15405 if (Op->AR == AR)
15406 return true;
15407
15408 if (Flags != SCEVWrapPredicate::IncrementNSSW &&
15410 return false;
15411
15412 const SCEV *Start = AR->getStart();
15413 const SCEV *OpStart = Op->AR->getStart();
15414 if (Start->getType()->isPointerTy() != OpStart->getType()->isPointerTy())
15415 return false;
15416
15417 // Reject pointers to different address spaces.
15418 if (Start->getType()->isPointerTy() && Start->getType() != OpStart->getType())
15419 return false;
15420
15421 // NUSW/NSSW on a wider-type AddRec does not imply the same on a
15422 // narrower-type AddRec.
15423 if (SE.getTypeSizeInBits(AR->getType()) >
15424 SE.getTypeSizeInBits(Op->AR->getType()))
15425 return false;
15426
15427 const SCEV *Step = AR->getStepRecurrence(SE);
15428 const SCEV *OpStep = Op->AR->getStepRecurrence(SE);
15429 if (!SE.isKnownPositive(Step) || !SE.isKnownPositive(OpStep))
15430 return false;
15431
15432 // If both steps are positive, this implies N, if N's start and step are
15433 // ULE/SLE (for NSUW/NSSW) than this'.
15434 Type *WiderTy = SE.getWiderType(Step->getType(), OpStep->getType());
15435 Step = SE.getNoopOrZeroExtend(Step, WiderTy);
15436 OpStep = SE.getNoopOrZeroExtend(OpStep, WiderTy);
15437
15438 bool IsNUW = Flags == SCEVWrapPredicate::IncrementNUSW;
15439 OpStart = IsNUW ? SE.getNoopOrZeroExtend(OpStart, WiderTy)
15440 : SE.getNoopOrSignExtend(OpStart, WiderTy);
15441 Start = IsNUW ? SE.getNoopOrZeroExtend(Start, WiderTy)
15442 : SE.getNoopOrSignExtend(Start, WiderTy);
15444 return SE.isKnownPredicate(Pred, OpStep, Step) &&
15445 SE.isKnownPredicate(Pred, OpStart, Start);
15446}
15447
15449 SCEVFlags ScevFlags = AR->getNoWrapFlags();
15450 IncrementWrapFlags IFlags = Flags;
15451
15452 if (ScalarEvolution::setFlags(ScevFlags, SCEV::FlagNSW) == ScevFlags)
15453 IFlags = clearFlags(IFlags, IncrementNSSW);
15454
15455 return IFlags == IncrementAnyWrap;
15456}
15457
15458void SCEVWrapPredicate::print(raw_ostream &OS, unsigned Depth) const {
15459 OS.indent(Depth) << *getExpr() << " Added Flags: ";
15461 OS << "<nusw>";
15463 OS << "<nssw>";
15464 OS << "\n";
15465}
15466
15467/// Union predicates don't get cached so create a dummy set ID for it.
15469 ScalarEvolution &SE)
15471 for (const auto *P : Preds)
15472 add(P, SE);
15473}
15474
15476 return all_of(Preds,
15477 [](const SCEVPredicate *I) { return I->isAlwaysTrue(); });
15478}
15479
15481 ScalarEvolution &SE) const {
15482 if (const auto *Set = dyn_cast<SCEVUnionPredicate>(N))
15483 return all_of(Set->Preds, [this, &SE](const SCEVPredicate *I) {
15484 return this->implies(I, SE);
15485 });
15486
15487 if (any_of(Preds,
15488 [N, &SE](const SCEVPredicate *I) { return I->implies(N, SE); }))
15489 return true;
15490
15491 // A wrap predicate may be implied by a wrap predicate in Preds after applying
15492 // equal predicates.
15493 const auto *NWrap = dyn_cast<SCEVWrapPredicate>(N);
15494 if (!NWrap)
15495 return false;
15496 const Loop *L = NWrap->getExpr()->getLoop();
15497 return any_of(Preds, [&](const SCEVPredicate *I) {
15498 const auto *IWrap = dyn_cast<SCEVWrapPredicate>(I);
15499 if (!IWrap)
15500 return false;
15501 const auto *RewrittenAR = dyn_cast<SCEVAddRecExpr>(
15502 SE.rewriteUsingPredicate(IWrap->getExpr(), L, *this));
15503 return RewrittenAR &&
15504 SE.getWrapPredicate(RewrittenAR, IWrap->getFlags())->implies(N, SE);
15505 });
15506}
15507
15509 for (const auto *Pred : Preds)
15510 Pred->print(OS, Depth);
15511}
15512
15513void SCEVUnionPredicate::add(const SCEVPredicate *N, ScalarEvolution &SE) {
15514 if (const auto *Set = dyn_cast<SCEVUnionPredicate>(N)) {
15515 for (const auto *Pred : Set->Preds)
15516 add(Pred, SE);
15517 return;
15518 }
15519
15520 // Implication checks are quadratic in the number of predicates. Stop doing
15521 // them if there are many predicates, as they should be too expensive to use
15522 // anyway at that point.
15523 bool CheckImplies = Preds.size() < 16;
15524
15525 // Only add predicate if it is not already implied by this union predicate.
15526 if (CheckImplies && implies(N, SE))
15527 return;
15528
15529 // Build a new vector containing the current predicates, except the ones that
15530 // are implied by the new predicate N.
15532 for (auto *P : Preds) {
15533 if (CheckImplies && N->implies(P, SE))
15534 continue;
15535 PrunedPreds.push_back(P);
15536 }
15537 Preds = std::move(PrunedPreds);
15538 Preds.push_back(N);
15539}
15540
15542 Loop &L)
15543 : SE(SE), L(L) {
15545 Preds = std::make_unique<SCEVUnionPredicate>(Empty, SE);
15546}
15547
15549 for (const SCEV *Op : Ops)
15550 // We do not expect that forgetting cached data for SCEVConstants will ever
15551 // open any prospects for sharpening or introduce any correctness issues,
15552 // so we don't bother storing their dependencies.
15553 if (!isa<SCEVConstant>(Op))
15554 SCEVUsers[Op].insert(User);
15555}
15556
15558 const SCEV *Expr = SE.getSCEV(V);
15559 return getPredicatedSCEV(Expr);
15560}
15561
15563 RewriteEntry &Entry = RewriteMap[Expr];
15564
15565 // If we already have an entry and the version matches, return it.
15566 if (Entry.second && Generation == Entry.first)
15567 return Entry.second;
15568
15569 // We found an entry but it's stale. Rewrite the stale entry
15570 // according to the current predicate.
15571 if (Entry.second)
15572 Expr = Entry.second;
15573
15574 const SCEV *NewSCEV = SE.rewriteUsingPredicate(Expr, &L, *Preds);
15575 Entry = {Generation, NewSCEV};
15576
15577 return NewSCEV;
15578}
15579
15581 if (!BackedgeCount) {
15583 BackedgeCount = SE.getPredicatedBackedgeTakenCount(&L, Preds);
15584 for (const auto *P : Preds)
15585 addPredicate(*P);
15586 }
15587 return BackedgeCount;
15588}
15589
15591 if (!SymbolicMaxBackedgeCount) {
15593 SymbolicMaxBackedgeCount =
15594 SE.getPredicatedSymbolicMaxBackedgeTakenCount(&L, Preds);
15595 for (const auto *P : Preds)
15596 addPredicate(*P);
15597 }
15598 return SymbolicMaxBackedgeCount;
15599}
15600
15602 if (!SmallConstantMaxTripCount) {
15604 SmallConstantMaxTripCount = SE.getSmallConstantMaxTripCount(&L, &Preds);
15605 for (const auto *P : Preds)
15606 addPredicate(*P);
15607 }
15608 return *SmallConstantMaxTripCount;
15609}
15610
15612 if (Preds->implies(&Pred, SE))
15613 return;
15614
15615 SmallVector<const SCEVPredicate *, 4> NewPreds(Preds->getPredicates());
15616 NewPreds.push_back(&Pred);
15617 Preds = std::make_unique<SCEVUnionPredicate>(NewPreds, SE);
15618 updateGeneration();
15619}
15620
15623 for (const SCEVPredicate *P : Preds)
15624 addPredicate(*P);
15625}
15626
15628 return *Preds;
15629}
15630
15631void PredicatedScalarEvolution::updateGeneration() {
15632 // If the generation number wrapped recompute everything.
15633 if (++Generation == 0) {
15634 for (auto &II : RewriteMap) {
15635 const SCEV *Rewritten = II.second.second;
15636 II.second = {Generation, SE.rewriteUsingPredicate(Rewritten, &L, *Preds)};
15637 }
15638 }
15639}
15640
15643 const SCEV *Expr = this->getSCEV(V);
15645 auto *New = SE.convertSCEVToAddRecWithPredicates(Expr, &L, NewPreds);
15646
15647 if (!New)
15648 return nullptr;
15649
15650 if (ExtraPreds) {
15651 ExtraPreds->append(NewPreds);
15652 return New;
15653 }
15654
15655 addPredicates(NewPreds);
15656
15657 RewriteMap[SE.getSCEV(V)] = {Generation, New};
15658 return New;
15659}
15660
15663 : RewriteMap(Init.RewriteMap), SE(Init.SE), L(Init.L),
15664 Preds(std::make_unique<SCEVUnionPredicate>(Init.Preds->getPredicates(),
15665 SE)),
15666 Generation(Init.Generation), BackedgeCount(Init.BackedgeCount) {}
15667
15669 // For each block.
15670 for (auto *BB : L.getBlocks())
15671 for (auto &I : *BB) {
15672 if (!SE.isSCEVable(I.getType()))
15673 continue;
15674
15675 auto *Expr = SE.getSCEV(&I);
15676 auto II = RewriteMap.find(Expr);
15677
15678 if (II == RewriteMap.end())
15679 continue;
15680
15681 // Don't print things that are not interesting.
15682 if (II->second.second == Expr)
15683 continue;
15684
15685 OS.indent(Depth) << "[PSE]" << I << ":\n";
15686 OS.indent(Depth + 2) << *Expr << "\n";
15687 OS.indent(Depth + 2) << "--> " << *II->second.second << "\n";
15688 }
15689}
15690
15693 BasicBlock *Header = L->getHeader();
15694 BasicBlock *Pred = L->getLoopPredecessor();
15695 LoopGuards Guards(SE);
15696 if (!Pred)
15697 return Guards;
15699 collectFromBlock(SE, Guards, Header, Pred, VisitedBlocks);
15700 return Guards;
15701}
15702
15703void ScalarEvolution::LoopGuards::collectFromPHI(
15707 unsigned Depth) {
15708 if (!SE.isSCEVable(Phi.getType()))
15709 return;
15710
15711 using MinMaxPattern = std::pair<const SCEVConstant *, SCEVTypes>;
15712 auto GetMinMaxConst = [&](unsigned IncomingIdx) -> MinMaxPattern {
15713 const BasicBlock *InBlock = Phi.getIncomingBlock(IncomingIdx);
15714 if (!VisitedBlocks.insert(InBlock).second)
15715 return {nullptr, scCouldNotCompute};
15716
15717 // Avoid analyzing unreachable blocks so that we don't get trapped
15718 // traversing cycles with ill-formed dominance or infinite cycles
15719 if (!SE.DT.isReachableFromEntry(InBlock))
15720 return {nullptr, scCouldNotCompute};
15721
15722 auto [G, Inserted] = IncomingGuards.try_emplace(InBlock, LoopGuards(SE));
15723 if (Inserted)
15724 collectFromBlock(SE, G->second, Phi.getParent(), InBlock, VisitedBlocks,
15725 Depth + 1);
15726 auto &RewriteMap = G->second.RewriteMap;
15727 if (RewriteMap.empty())
15728 return {nullptr, scCouldNotCompute};
15729 auto S = RewriteMap.find(SE.getSCEV(Phi.getIncomingValue(IncomingIdx)));
15730 if (S == RewriteMap.end())
15731 return {nullptr, scCouldNotCompute};
15732 auto *SM = dyn_cast_if_present<SCEVMinMaxExpr>(S->second);
15733 if (!SM)
15734 return {nullptr, scCouldNotCompute};
15735 if (const SCEVConstant *C0 = dyn_cast<SCEVConstant>(SM->getOperand(0)))
15736 return {C0, SM->getSCEVType()};
15737 return {nullptr, scCouldNotCompute};
15738 };
15739 auto MergeMinMaxConst = [](MinMaxPattern P1,
15740 MinMaxPattern P2) -> MinMaxPattern {
15741 auto [C1, T1] = P1;
15742 auto [C2, T2] = P2;
15743 if (!C1 || !C2 || T1 != T2)
15744 return {nullptr, scCouldNotCompute};
15745 switch (T1) {
15746 case scUMaxExpr:
15747 return {C1->getAPInt().ult(C2->getAPInt()) ? C1 : C2, T1};
15748 case scSMaxExpr:
15749 return {C1->getAPInt().slt(C2->getAPInt()) ? C1 : C2, T1};
15750 case scUMinExpr:
15751 return {C1->getAPInt().ugt(C2->getAPInt()) ? C1 : C2, T1};
15752 case scSMinExpr:
15753 return {C1->getAPInt().sgt(C2->getAPInt()) ? C1 : C2, T1};
15754 default:
15755 llvm_unreachable("Trying to merge non-MinMaxExpr SCEVs.");
15756 }
15757 };
15758 auto P = GetMinMaxConst(0);
15759 for (unsigned int In = 1; In < Phi.getNumIncomingValues(); In++) {
15760 if (!P.first)
15761 break;
15762 P = MergeMinMaxConst(P, GetMinMaxConst(In));
15763 }
15764 if (P.first) {
15765 const SCEV *LHS = SE.getSCEV(const_cast<PHINode *>(&Phi));
15766 SmallVector<SCEVUse, 2> Ops({P.first, LHS});
15767 const SCEV *RHS = SE.getMinMaxExpr(P.second, Ops);
15768 Guards.RewriteMap.insert({LHS, RHS});
15769 }
15770}
15771
15772// Return a new SCEV that modifies \p Expr to the closest number divides by
15773// \p Divisor and less or equal than Expr. For now, only handle constant
15774// Expr.
15776 const APInt &DivisorVal,
15777 ScalarEvolution &SE) {
15778 const APInt *ExprVal;
15779 if (!match(Expr, m_scev_APInt(ExprVal)) || ExprVal->isNegative() ||
15780 DivisorVal.isNonPositive())
15781 return Expr;
15782 APInt Rem = ExprVal->urem(DivisorVal);
15783 // return the SCEV: Expr - Expr % Divisor
15784 return SE.getConstant(*ExprVal - Rem);
15785}
15786
15787// Return a new SCEV that modifies \p Expr to the closest number divides by
15788// \p Divisor and greater or equal than Expr. For now, only handle constant
15789// Expr.
15790static const SCEV *getNextSCEVDivisibleByDivisor(const SCEV *Expr,
15791 const APInt &DivisorVal,
15792 ScalarEvolution &SE) {
15793 const APInt *ExprVal;
15794 if (!match(Expr, m_scev_APInt(ExprVal)) || ExprVal->isNegative() ||
15795 DivisorVal.isNonPositive())
15796 return Expr;
15797 APInt Rem = ExprVal->urem(DivisorVal);
15798 if (Rem.isZero())
15799 return Expr;
15800 // return the SCEV: Expr + Divisor - Expr % Divisor
15801 return SE.getConstant(*ExprVal + DivisorVal - Rem);
15802}
15803
15805 ICmpInst::Predicate Predicate, const SCEV *LHS, const SCEV *RHS,
15808 // If we have LHS == 0, check if LHS is computing a property of some unknown
15809 // SCEV %v which we can rewrite %v to express explicitly.
15811 return false;
15812 // If LHS is A % B, i.e. A % B == 0, rewrite A to (A /u B) * B to
15813 // explicitly express that.
15814 const SCEVUnknown *URemLHS = nullptr;
15815 const SCEV *URemRHS = nullptr;
15816 if (!match(LHS, m_scev_URem(m_SCEVUnknown(URemLHS), m_SCEV(URemRHS), SE)))
15817 return false;
15818
15819 const SCEV *Multiple =
15820 SE.getMulExpr(SE.getUDivExpr(URemLHS, URemRHS), URemRHS);
15821 DivInfo[URemLHS] = Multiple;
15822 if (auto *C = dyn_cast<SCEVConstant>(URemRHS))
15823 Multiples[URemLHS] = C->getAPInt();
15824 return true;
15825}
15826
15827// Check if the condition is a divisibility guard (A % B == 0).
15828static bool isDivisibilityGuard(const SCEV *LHS, const SCEV *RHS,
15829 ScalarEvolution &SE) {
15830 const SCEV *X, *Y;
15831 return match(LHS, m_scev_URem(m_SCEV(X), m_SCEV(Y), SE)) && RHS->isZero();
15832}
15833
15834// Apply divisibility by \p Divisor on MinMaxExpr with constant values,
15835// recursively. This is done by aligning up/down the constant value to the
15836// Divisor.
15837static const SCEV *applyDivisibilityOnMinMaxExpr(const SCEV *MinMaxExpr,
15838 APInt Divisor,
15839 ScalarEvolution &SE) {
15840 // Return true if \p Expr is a MinMax SCEV expression with a non-negative
15841 // constant operand. If so, return in \p SCTy the SCEV type and in \p RHS
15842 // the non-constant operand and in \p LHS the constant operand.
15843 auto IsMinMaxSCEVWithNonNegativeConstant =
15844 [&](const SCEV *Expr, SCEVTypes &SCTy, const SCEV *&LHS,
15845 const SCEV *&RHS) {
15846 if (auto *MinMax = dyn_cast<SCEVMinMaxExpr>(Expr)) {
15847 if (MinMax->getNumOperands() != 2)
15848 return false;
15849 if (auto *C = dyn_cast<SCEVConstant>(MinMax->getOperand(0))) {
15850 if (C->getAPInt().isNegative())
15851 return false;
15852 SCTy = MinMax->getSCEVType();
15853 LHS = MinMax->getOperand(0);
15854 RHS = MinMax->getOperand(1);
15855 return true;
15856 }
15857 }
15858 return false;
15859 };
15860
15861 const SCEV *MinMaxLHS = nullptr, *MinMaxRHS = nullptr;
15862 SCEVTypes SCTy;
15863 if (!IsMinMaxSCEVWithNonNegativeConstant(MinMaxExpr, SCTy, MinMaxLHS,
15864 MinMaxRHS))
15865 return MinMaxExpr;
15866 auto IsMin = isa<SCEVSMinExpr>(MinMaxExpr) || isa<SCEVUMinExpr>(MinMaxExpr);
15867 assert(SE.isKnownNonNegative(MinMaxLHS) && "Expected non-negative operand!");
15868 auto *DivisibleExpr =
15869 IsMin ? getPreviousSCEVDivisibleByDivisor(MinMaxLHS, Divisor, SE)
15870 : getNextSCEVDivisibleByDivisor(MinMaxLHS, Divisor, SE);
15872 applyDivisibilityOnMinMaxExpr(MinMaxRHS, Divisor, SE), DivisibleExpr};
15873 return SE.getMinMaxExpr(SCTy, Ops);
15874}
15875
15876void ScalarEvolution::LoopGuards::collectFromBlock(
15877 ScalarEvolution &SE, ScalarEvolution::LoopGuards &Guards,
15878 const BasicBlock *Block, const BasicBlock *Pred,
15879 SmallPtrSetImpl<const BasicBlock *> &VisitedBlocks, unsigned Depth) {
15880
15882
15883 SmallVector<SCEVUse> ExprsToRewrite;
15884 auto CollectCondition = [&](ICmpInst::Predicate Predicate, const SCEV *LHS,
15885 const SCEV *RHS,
15886 DenseMap<const SCEV *, const SCEV *> &RewriteMap,
15887 const LoopGuards &DivGuards) {
15888 // WARNING: It is generally unsound to apply any wrap flags to the proposed
15889 // replacement SCEV which isn't directly implied by the structure of that
15890 // SCEV. In particular, using contextual facts to imply flags is *NOT*
15891 // legal. See the scoping rules for flags in the header to understand why.
15892
15893 // Puts rewrite rule \p From -> \p To into the rewrite map. Also if \p From
15894 // and \p FromRewritten are the same (i.e. there has been no rewrite
15895 // registered for \p From), then puts this value in the list of rewritten
15896 // expressions.
15897 auto AddRewrite = [&](const SCEV *From, const SCEV *FromRewritten,
15898 const SCEV *To) {
15899 if (From == FromRewritten)
15900 ExprsToRewrite.push_back(From);
15901 RewriteMap[From] = To;
15902 };
15903
15904 // Checks whether \p S has already been rewritten. In that case returns the
15905 // existing rewrite because we want to chain further rewrites onto the
15906 // already rewritten value. Otherwise returns \p S.
15907 auto GetMaybeRewritten = [&](const SCEV *S) {
15908 return RewriteMap.lookup_or(S, S);
15909 };
15910
15911 // Check for a condition of the form (-C1 + X < C2). InstCombine will
15912 // create this form when combining two checks of the form (X u< C2 + C1) and
15913 // (X >=u C1).
15914 auto MatchRangeCheckIdiom = [&](ICmpInst::Predicate Pred,
15915 const SCEV *MatchLHS,
15916 const SCEV *MatchRHS) {
15917 const SCEVConstant *C1;
15918 const SCEVUnknown *LHSUnknown;
15919 auto *C2 = dyn_cast<SCEVConstant>(MatchRHS);
15920 if (!match(MatchLHS,
15921 m_scev_Add(m_SCEVConstant(C1), m_SCEVUnknown(LHSUnknown))) ||
15922 !C2)
15923 return false;
15924
15925 auto ExactRegion =
15926 ConstantRange::makeExactICmpRegion(Pred, C2->getAPInt())
15927 .sub(C1->getAPInt());
15928
15929 // Tighten the raw range with what we already know about LHSUnknown
15930 // from prior guards recorded in RewriteMap, or from SCEV's own range
15931 // analysis.
15932 const SCEV *RewrittenLHS = GetMaybeRewritten(LHSUnknown);
15933 ExactRegion = ExactRegion.intersectWith(SE.getUnsignedRange(RewrittenLHS),
15935
15936 // Bail if the guard is inconsistent with prior facts, or if the range
15937 // is still not a monotonic non-wrapping interval after tightening.
15938 if (ExactRegion.isEmptySet() || ExactRegion.isWrappedSet() ||
15939 ExactRegion.isFullSet())
15940 return false;
15941
15942 const SCEV *RegionMin = SE.getConstant(ExactRegion.getUnsignedMin());
15943 const SCEV *RegionMax = SE.getConstant(ExactRegion.getUnsignedMax());
15944 const SCEV *ClampedLHS =
15945 SE.getUMaxExpr(RegionMin, SE.getUMinExpr(RewrittenLHS, RegionMax));
15946 AddRewrite(LHSUnknown, RewrittenLHS, ClampedLHS);
15947 return true;
15948 };
15949 if (MatchRangeCheckIdiom(Predicate, LHS, RHS))
15950 return;
15951
15952 // Do not apply information for constants or if RHS contains an AddRec.
15954 return;
15955
15956 // If RHS is SCEVUnknown, make sure the information is applied to it.
15958 std::swap(LHS, RHS);
15960 }
15961
15962 const SCEV *RewrittenLHS = GetMaybeRewritten(LHS);
15963 // Apply divisibility information when computing the constant multiple.
15964 const APInt &DividesBy =
15965 SE.getConstantMultiple(DivGuards.rewrite(RewrittenLHS));
15966
15967 // Collect rewrites for LHS and its transitive operands based on the
15968 // condition.
15969 // For min/max expressions, also apply the guard to its operands:
15970 // 'min(a, b) >= c' -> '(a >= c) and (b >= c)',
15971 // 'min(a, b) > c' -> '(a > c) and (b > c)',
15972 // 'max(a, b) <= c' -> '(a <= c) and (b <= c)',
15973 // 'max(a, b) < c' -> '(a < c) and (b < c)'.
15974
15975 // We cannot express strict predicates in SCEV, so instead we replace them
15976 // with non-strict ones against plus or minus one of RHS depending on the
15977 // predicate.
15978 const SCEV *One = SE.getOne(RHS->getType());
15979 switch (Predicate) {
15980 case CmpInst::ICMP_ULT:
15981 if (RHS->getType()->isPointerTy())
15982 return;
15983 RHS = SE.getUMaxExpr(RHS, One);
15984 [[fallthrough]];
15985 case CmpInst::ICMP_SLT: {
15986 RHS = SE.getMinusSCEV(RHS, One);
15987 RHS = getPreviousSCEVDivisibleByDivisor(RHS, DividesBy, SE);
15988 break;
15989 }
15990 case CmpInst::ICMP_UGT:
15991 case CmpInst::ICMP_SGT:
15992 RHS = SE.getAddExpr(RHS, One);
15993 RHS = getNextSCEVDivisibleByDivisor(RHS, DividesBy, SE);
15994 break;
15995 case CmpInst::ICMP_ULE:
15996 case CmpInst::ICMP_SLE:
15997 RHS = getPreviousSCEVDivisibleByDivisor(RHS, DividesBy, SE);
15998 break;
15999 case CmpInst::ICMP_UGE:
16000 case CmpInst::ICMP_SGE:
16001 RHS = getNextSCEVDivisibleByDivisor(RHS, DividesBy, SE);
16002 break;
16003 default:
16004 break;
16005 }
16006
16007 SmallVector<SCEVUse, 16> Worklist(1, LHS);
16008 SmallPtrSet<const SCEV *, 16> Visited;
16009
16010 auto EnqueueOperands = [&Worklist](const SCEVNAryExpr *S) {
16011 append_range(Worklist, S->operands());
16012 };
16013
16014 while (!Worklist.empty()) {
16015 const SCEV *From = Worklist.pop_back_val();
16016 if (isa<SCEVConstant>(From))
16017 continue;
16018 if (!Visited.insert(From).second)
16019 continue;
16020 const SCEV *FromRewritten = GetMaybeRewritten(From);
16021 const SCEV *To = nullptr;
16022
16023 switch (Predicate) {
16024 case CmpInst::ICMP_ULT:
16025 case CmpInst::ICMP_ULE:
16026 To = SE.getUMinExpr(FromRewritten, RHS);
16027 if (auto *UMax = dyn_cast<SCEVUMaxExpr>(FromRewritten))
16028 EnqueueOperands(UMax);
16029 break;
16030 case CmpInst::ICMP_SLT:
16031 case CmpInst::ICMP_SLE:
16032 To = SE.getSMinExpr(FromRewritten, RHS);
16033 if (auto *SMax = dyn_cast<SCEVSMaxExpr>(FromRewritten))
16034 EnqueueOperands(SMax);
16035 break;
16036 case CmpInst::ICMP_UGT:
16037 case CmpInst::ICMP_UGE:
16038 To = SE.getUMaxExpr(FromRewritten, RHS);
16039 if (auto *UMin = dyn_cast<SCEVUMinExpr>(FromRewritten))
16040 EnqueueOperands(UMin);
16041 break;
16042 case CmpInst::ICMP_SGT:
16043 case CmpInst::ICMP_SGE:
16044 To = SE.getSMaxExpr(FromRewritten, RHS);
16045 if (auto *SMin = dyn_cast<SCEVSMinExpr>(FromRewritten))
16046 EnqueueOperands(SMin);
16047 break;
16048 case CmpInst::ICMP_EQ:
16050 To = RHS;
16051 break;
16052 case CmpInst::ICMP_NE:
16053 if (match(RHS, m_scev_Zero())) {
16054 const SCEV *OneAlignedUp =
16055 getNextSCEVDivisibleByDivisor(One, DividesBy, SE);
16056 To = SE.getUMaxExpr(FromRewritten, OneAlignedUp);
16057 } else {
16058 // LHS != RHS can be rewritten as (LHS - RHS) = UMax(1, LHS - RHS),
16059 // but creating the subtraction eagerly is expensive. Track the
16060 // inequalities in a separate map, and materialize the rewrite lazily
16061 // when encountering a suitable subtraction while re-writing.
16062 if (LHS->getType()->isPointerTy()) {
16063 LHS = SE.getPtrToAddrExpr(LHS);
16064 RHS = SE.getPtrToAddrExpr(RHS);
16066 break;
16067 }
16068 const SCEVConstant *C;
16069 const SCEV *A, *B;
16072 RHS = A;
16073 LHS = B;
16074 }
16075 if (LHS > RHS)
16076 std::swap(LHS, RHS);
16077 Guards.NotEqual.insert({LHS, RHS});
16078 continue;
16079 }
16080 break;
16081 default:
16082 break;
16083 }
16084
16085 if (To)
16086 AddRewrite(From, FromRewritten, To);
16087 }
16088 };
16089
16091 // First, collect information from assumptions dominating the loop.
16092 for (auto &AssumeVH : SE.AC.assumptions()) {
16093 if (!AssumeVH)
16094 continue;
16095 auto *AssumeI = cast<CallInst>(AssumeVH);
16096 if (!SE.DT.dominates(AssumeI, Block))
16097 continue;
16098 Terms.emplace_back(AssumeI->getOperand(0), true);
16099 }
16100
16101 // Second, collect information from llvm.experimental.guards dominating the loop.
16102 auto *GuardDecl = Intrinsic::getDeclarationIfExists(
16103 SE.F.getParent(), Intrinsic::experimental_guard);
16104 if (GuardDecl)
16105 for (const auto *GU : GuardDecl->users())
16106 if (const auto *Guard = dyn_cast<IntrinsicInst>(GU))
16107 if (Guard->getFunction() == Block->getParent() &&
16108 SE.DT.dominates(Guard, Block))
16109 Terms.emplace_back(Guard->getArgOperand(0), true);
16110
16111 // Third, collect conditions from dominating branches. Starting at the loop
16112 // predecessor, climb up the predecessor chain, as long as there are
16113 // predecessors that can be found that have unique successors leading to the
16114 // original header.
16115 // TODO: share this logic with isLoopEntryGuardedByCond.
16116 unsigned NumCollectedConditions = 0;
16118 std::pair<const BasicBlock *, const BasicBlock *> Pair(Pred, Block);
16119 for (; Pair.first;
16120 Pair = SE.getPredecessorWithUniqueSuccessorForBB(Pair.first)) {
16121 VisitedBlocks.insert(Pair.second);
16122 const CondBrInst *LoopEntryPredicate =
16123 dyn_cast<CondBrInst>(Pair.first->getTerminator());
16124 if (!LoopEntryPredicate)
16125 continue;
16126
16127 Terms.emplace_back(LoopEntryPredicate->getCondition(),
16128 LoopEntryPredicate->getSuccessor(0) == Pair.second);
16129 NumCollectedConditions++;
16130
16131 // If we are recursively collecting guards stop after 2
16132 // conditions to limit compile-time impact for now.
16133 if (Depth > 0 && NumCollectedConditions == 2)
16134 break;
16135 }
16136 // Finally, if we stopped climbing the predecessor chain because
16137 // there wasn't a unique one to continue, try to collect conditions
16138 // for PHINodes by recursively following all of their incoming
16139 // blocks and try to merge the found conditions to build a new one
16140 // for the Phi.
16141 if (Pair.second->hasNPredecessorsOrMore(2) &&
16143 SmallDenseMap<const BasicBlock *, LoopGuards> IncomingGuards;
16144 for (auto &Phi : Pair.second->phis())
16145 collectFromPHI(SE, Guards, Phi, VisitedBlocks, IncomingGuards, Depth);
16146 }
16147
16148 // Now apply the information from the collected conditions to
16149 // Guards.RewriteMap. Conditions are processed in reverse order, so the
16150 // earliest conditions is processed first, except guards with divisibility
16151 // information, which are moved to the back. This ensures the SCEVs with the
16152 // shortest dependency chains are constructed first.
16154 GuardsToProcess;
16155 for (auto [Term, EnterIfTrue] : reverse(Terms)) {
16156 SmallVector<Value *, 8> Worklist;
16157 SmallPtrSet<Value *, 8> Visited;
16158 Worklist.push_back(Term);
16159 while (!Worklist.empty()) {
16160 Value *Cond = Worklist.pop_back_val();
16161 if (!Visited.insert(Cond).second)
16162 continue;
16163
16164 if (auto *Cmp = dyn_cast<ICmpInst>(Cond)) {
16165 auto Predicate =
16166 EnterIfTrue ? Cmp->getPredicate() : Cmp->getInversePredicate();
16167 const auto *LHS = SE.getSCEV(Cmp->getOperand(0));
16168 const auto *RHS = SE.getSCEV(Cmp->getOperand(1));
16169 // If LHS is a constant, apply information to the other expression.
16170 // TODO: If LHS is not a constant, check if using CompareSCEVComplexity
16171 // can improve results.
16172 if (isa<SCEVConstant>(LHS)) {
16173 std::swap(LHS, RHS);
16175 }
16176 GuardsToProcess.emplace_back(Predicate, LHS, RHS);
16177 continue;
16178 }
16179
16180 Value *L, *R;
16181 if (EnterIfTrue ? match(Cond, m_LogicalAnd(m_Value(L), m_Value(R)))
16182 : match(Cond, m_LogicalOr(m_Value(L), m_Value(R)))) {
16183 Worklist.push_back(L);
16184 Worklist.push_back(R);
16185 }
16186 }
16187 }
16188
16189 // Process divisibility guards in reverse order to populate DivGuards early.
16190 DenseMap<const SCEV *, APInt> Multiples;
16191 LoopGuards DivGuards(SE);
16192 for (const auto &[Predicate, LHS, RHS] : GuardsToProcess) {
16193 if (!isDivisibilityGuard(LHS, RHS, SE))
16194 continue;
16195 collectDivisibilityInformation(Predicate, LHS, RHS, DivGuards.RewriteMap,
16196 Multiples, SE);
16197 }
16198
16199 for (const auto &[Predicate, LHS, RHS] : GuardsToProcess)
16200 CollectCondition(Predicate, LHS, RHS, Guards.RewriteMap, DivGuards);
16201
16202 // Apply divisibility information last. This ensures it is applied to the
16203 // outermost expression after other rewrites for the given value.
16204 for (const auto &[K, Divisor] : Multiples) {
16205 const SCEV *DivisorSCEV = SE.getConstant(Divisor);
16206 Guards.RewriteMap[K] =
16208 Guards.rewrite(K), Divisor, SE),
16209 DivisorSCEV),
16210 DivisorSCEV);
16211 ExprsToRewrite.push_back(K);
16212 }
16213
16214 // Let the rewriter preserve NUW/NSW flags if the unsigned/signed ranges of
16215 // the replacement expressions are contained in the ranges of the replaced
16216 // expressions.
16217 Guards.PreserveNUW = true;
16218 Guards.PreserveNSW = true;
16219 for (const SCEV *Expr : ExprsToRewrite) {
16220 const SCEV *RewriteTo = Guards.RewriteMap[Expr];
16221 Guards.PreserveNUW &=
16222 SE.getUnsignedRange(Expr).contains(SE.getUnsignedRange(RewriteTo));
16223 Guards.PreserveNSW &=
16224 SE.getSignedRange(Expr).contains(SE.getSignedRange(RewriteTo));
16225 }
16226
16227 // Now that all rewrite information is collect, rewrite the collected
16228 // expressions with the information in the map. This applies information to
16229 // sub-expressions.
16230 if (ExprsToRewrite.size() > 1) {
16231 for (const SCEV *Expr : ExprsToRewrite) {
16232 const SCEV *RewriteTo = Guards.RewriteMap[Expr];
16233 Guards.RewriteMap.erase(Expr);
16234 Guards.RewriteMap.insert({Expr, Guards.rewrite(RewriteTo)});
16235 }
16236 }
16237}
16238
16240 /// A rewriter to replace SCEV expressions in Map with the corresponding entry
16241 /// in the map. It skips AddRecExpr because we cannot guarantee that the
16242 /// replacement is loop invariant in the loop of the AddRec.
16243 class SCEVLoopGuardRewriter
16244 : public SCEVRewriteVisitor<SCEVLoopGuardRewriter> {
16247
16248 SCEVFlags FlagMask = SCEV::FlagNone;
16249
16250 public:
16251 SCEVLoopGuardRewriter(ScalarEvolution &SE,
16252 const ScalarEvolution::LoopGuards &Guards)
16253 : SCEVRewriteVisitor(SE), Map(Guards.RewriteMap),
16254 NotEqual(Guards.NotEqual) {
16255 if (Guards.PreserveNUW)
16256 FlagMask = ScalarEvolution::setFlags(FlagMask, SCEV::FlagNUW);
16257 if (Guards.PreserveNSW)
16258 FlagMask = ScalarEvolution::setFlags(FlagMask, SCEV::FlagNSW);
16259 }
16260
16261 const SCEV *visitAddRecExpr(const SCEVAddRecExpr *Expr) { return Expr; }
16262
16263 const SCEV *visitUnknown(const SCEVUnknown *Expr) {
16264 return Map.lookup_or(Expr, Expr);
16265 }
16266
16267 const SCEV *visitPtrToAddrExpr(const SCEVPtrToAddrExpr *Expr) {
16268 if (const SCEV *S = Map.lookup(Expr))
16269 return S;
16271 Expr);
16272 }
16273
16274 const SCEV *visitZeroExtendExpr(const SCEVZeroExtendExpr *Expr) {
16275 if (const SCEV *S = Map.lookup(Expr))
16276 return S;
16277
16278 // If we didn't find the extact ZExt expr in the map, check if there's
16279 // an entry for a smaller ZExt we can use instead.
16280 Type *Ty = Expr->getType();
16281 const SCEV *Op = Expr->getOperand(0);
16282 unsigned Bitwidth = Ty->getScalarSizeInBits() / 2;
16283 while (Bitwidth % 8 == 0 && Bitwidth >= 8 &&
16284 Bitwidth > Op->getType()->getScalarSizeInBits()) {
16285 Type *NarrowTy = IntegerType::get(SE.getContext(), Bitwidth);
16286 auto *NarrowExt = SE.getZeroExtendExpr(Op, NarrowTy);
16287 if (const SCEV *S = Map.lookup(NarrowExt))
16288 return SE.getZeroExtendExpr(S, Ty);
16289 Bitwidth = Bitwidth / 2;
16290 }
16291
16293 Expr);
16294 }
16295
16296 const SCEV *visitSignExtendExpr(const SCEVSignExtendExpr *Expr) {
16297 if (const SCEV *S = Map.lookup(Expr))
16298 return S;
16300 Expr);
16301 }
16302
16303 const SCEV *visitUMinExpr(const SCEVUMinExpr *Expr) {
16304 if (const SCEV *S = Map.lookup(Expr))
16305 return S;
16307 }
16308
16309 const SCEV *visitSMinExpr(const SCEVSMinExpr *Expr) {
16310 if (const SCEV *S = Map.lookup(Expr))
16311 return S;
16313 }
16314
16315 const SCEV *visitAddExpr(const SCEVAddExpr *Expr) {
16316 if (const SCEV *S = Map.lookup(Expr))
16317 return S;
16318
16319 // Helper to check if S is a subtraction (A - B) where A != B, and if so,
16320 // return UMax(S, 1).
16321 auto RewriteSubtraction = [&](const SCEV *S) -> const SCEV * {
16322 SCEVUse LHS, RHS;
16323 if (MatchBinarySub(S, LHS, RHS)) {
16324 if (LHS > RHS)
16325 std::swap(LHS, RHS);
16326 if (NotEqual.contains({LHS, RHS})) {
16327 const SCEV *OneAlignedUp = getNextSCEVDivisibleByDivisor(
16328 SE.getOne(S->getType()), SE.getConstantMultiple(S), SE);
16329 return SE.getUMaxExpr(OneAlignedUp, S);
16330 }
16331 }
16332 return nullptr;
16333 };
16334
16335 // Check if Expr itself is a subtraction pattern with guard info.
16336 if (const SCEV *Rewritten = RewriteSubtraction(Expr))
16337 return Rewritten;
16338
16339 // Trip count expressions sometimes consist of adding 3 operands, i.e.
16340 // (Const + A + B). There may be guard info for A + B, and if so, apply
16341 // it.
16342 // TODO: Could more generally apply guards to Add sub-expressions.
16343 if (isa<SCEVConstant>(Expr->getOperand(0))) {
16344 if (Expr->getNumOperands() == 3) {
16345 const SCEV *Add =
16346 SE.getAddExpr(Expr->getOperand(1), Expr->getOperand(2));
16347 if (const SCEV *Rewritten = RewriteSubtraction(Add))
16348 return SE.getAddExpr(
16349 Expr->getOperand(0), Rewritten,
16350 ScalarEvolution::maskFlags(Expr->getNoWrapFlags(), FlagMask));
16351 if (const SCEV *S = Map.lookup(Add))
16352 return SE.getAddExpr(Expr->getOperand(0), S);
16353 }
16354
16355 // For expressions of the form (Const + A), check if we have guard info
16356 // for (Const + 1 + A), and rewrite to ((Const + 1 + A) - 1). This makes
16357 // sure we don't lose information when rewriting expressions based on
16358 // back-edge taken counts in some cases.
16359 if (Expr->getNumOperands() == 2) {
16360 const SCEV *S = nullptr;
16361 // Handle (-1 + 1 + A) without constructing SCEVs.
16362 if (match(Expr->getOperand(0), m_scev_AllOnes())) {
16363 S = Map.lookup(Expr->getOperand(1));
16364 } else {
16365 const SCEV *NewC =
16366 SE.getAddExpr(Expr->getOperand(0), SE.getOne(Expr->getType()));
16367 S = Map.lookup(SE.getAddExpr(NewC, Expr->getOperand(1)));
16368 }
16369 if (S)
16370 return SE.getAddExpr(S, SE.getMinusOne(Expr->getType()));
16371 }
16372 }
16374 bool Changed = false;
16375 for (SCEVUse Op : Expr->operands()) {
16376 Operands.push_back(
16378 Changed |= Op != Operands.back();
16379 }
16380 // We are only replacing operands with equivalent values, so transfer the
16381 // flags from the original expression.
16382 return !Changed ? Expr
16383 : SE.getAddExpr(Operands,
16385 Expr->getNoWrapFlags(), FlagMask));
16386 }
16387
16388 const SCEV *visitMulExpr(const SCEVMulExpr *Expr) {
16390 bool Changed = false;
16391 for (SCEVUse Op : Expr->operands()) {
16392 Operands.push_back(
16394 Changed |= Op != Operands.back();
16395 }
16396 // We are only replacing operands with equivalent values, so transfer the
16397 // flags from the original expression.
16398 return !Changed ? Expr
16399 : SE.getMulExpr(Operands,
16401 Expr->getNoWrapFlags(), FlagMask));
16402 }
16403 };
16404
16405 if (RewriteMap.empty() && NotEqual.empty())
16406 return Expr;
16407
16408 SCEVLoopGuardRewriter Rewriter(SE, *this);
16409 return Rewriter.visit(Expr);
16410}
16411
16412const SCEV *ScalarEvolution::applyLoopGuards(const SCEV *Expr, const Loop *L) {
16413 return applyLoopGuards(Expr, LoopGuards::collect(L, *this));
16414}
16415
16417 const LoopGuards &Guards) {
16418 return Guards.rewrite(Expr);
16419}
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:758
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:763
bool empty() const
Definition DenseMap.h:717
iterator find(const_arg_type_t< KeyT > Val)
Definition DenseMap.h:767
iterator end()
Definition DenseMap.h:687
DenseMapIterator< KeyT, ValueT, KeyInfoT, BucketT > iterator
Definition DenseMap.h:679
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:794
iterator find_as(const LookupKeyT &Val)
Alternate version of find() which allows a different, and possibly less expensive,...
Definition DenseMap.h:780
void swap(DenseMapBase &RHS)
Definition DenseMap.h:978
std::pair< iterator, bool > insert(const std::pair< KeyT, ValueT > &KV)
Definition DenseMap.h:828
std::pair< iterator, bool > try_emplace(KeyT &&Key, Ts &&...Args)
Definition DenseMap.h:857
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.
auto map_range(ContainerTy &&C, FuncTy F)
Return a range that applies F to the elements of C.
Definition STLExtras.h:366
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.