1//===- StraightLineStrengthReduce.cpp - -----------------------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file implements straight-line strength reduction (SLSR). Unlike loop
10// strength reduction, this algorithm is designed to reduce arithmetic
11// redundancy in straight-line code instead of loops. It has proven to be
12// effective in simplifying arithmetic statements derived from an unrolled loop.
13// It can also simplify the logic of SeparateConstOffsetFromGEP.
14//
15// There are many optimizations we can perform in the domain of SLSR.
16// We look for strength reduction candidates in the following forms:
17//
18// Form Add: B + i * S
19// Form Mul: (B + i) * S
20// Form GEP: &B[i * S]
21//
22// where S is an integer variable, and i is a constant integer. If we found two
23// candidates S1 and S2 in the same form and S1 dominates S2, we may rewrite S2
24// in a simpler way with respect to S1 (index delta). For example,
25//
26// S1: X = B + i * S
27// S2: Y = B + i' * S => X + (i' - i) * S
28//
29// S1: X = (B + i) * S
30// S2: Y = (B + i') * S => X + (i' - i) * S
31//
32// S1: X = &B[i * S]
33// S2: Y = &B[i' * S] => &X[(i' - i) * S]
34//
35// Note: (i' - i) * S is folded to the extent possible.
36//
37// For Add and GEP forms, we can also rewrite a candidate in a simpler way
38// with respect to other dominating candidates if their B or S are different
39// but other parts are the same. For example,
40//
41// Base Delta:
42// S1: X = B + i * S
43// S2: Y = B' + i * S => X + (B' - B)
44//
45// S1: X = &B [i * S]
46// S2: Y = &B'[i * S] => X + (B' - B)
47//
48// Stride Delta:
49// S1: X = B + i * S
50// S2: Y = B + i * S' => X + i * (S' - S)
51//
52// S1: X = &B[i * S]
53// S2: Y = &B[i * S'] => X + i * (S' - S)
54//
55// PS: Stride delta rewrite on Mul form is usually non-profitable, and Base
56// delta rewrite sometimes is profitable, so we do not support them on Mul.
57//
58// This rewriting is in general a good idea. The code patterns we focus on
59// usually come from loop unrolling, so the delta is likely the same
60// across iterations and can be reused. When that happens, the optimized form
61// takes only one add starting from the second iteration.
62//
63// When such rewriting is possible, we call S1 a "basis" of S2. When S2 has
64// multiple bases, we choose to rewrite S2 with respect to its "immediate"
65// basis, the basis that is the closest ancestor in the dominator tree.
66//
67// TODO:
68//
69// - Floating point arithmetics when fast math is enabled.
70
71#include "llvm/Transforms/Scalar/StraightLineStrengthReduce.h"
72#include "ScalarOptions.h"
73#include "llvm/ADT/APInt.h"
74#include "llvm/ADT/DepthFirstIterator.h"
75#include "llvm/ADT/SetVector.h"
76#include "llvm/ADT/SmallPtrSet.h"
77#include "llvm/ADT/SmallVector.h"
78#include "llvm/ADT/Statistic.h"
79#include "llvm/Analysis/ScalarEvolution.h"
80#include "llvm/Analysis/ScalarEvolutionExpressions.h"
81#include "llvm/Analysis/TargetTransformInfo.h"
82#include "llvm/Analysis/ValueTracking.h"
83#include "llvm/IR/Constants.h"
84#include "llvm/IR/DataLayout.h"
85#include "llvm/IR/DerivedTypes.h"
86#include "llvm/IR/Dominators.h"
87#include "llvm/IR/GetElementPtrTypeIterator.h"
88#include "llvm/IR/IRBuilder.h"
89#include "llvm/IR/Instruction.h"
90#include "llvm/IR/Instructions.h"
91#include "llvm/IR/Module.h"
92#include "llvm/IR/Operator.h"
93#include "llvm/IR/PatternMatch.h"
94#include "llvm/IR/Type.h"
95#include "llvm/IR/Value.h"
96#include "llvm/InitializePasses.h"
97#include "llvm/Pass.h"
98#include "llvm/Support/Casting.h"
99#include "llvm/Support/DebugCounter.h"
100#include "llvm/Support/ErrorHandling.h"
101#include "llvm/Transforms/Scalar.h"
102#include "llvm/Transforms/Utils/Local.h"
103#include <cassert>
104#include <cstdint>
105#include <limits>
106#include <list>
107#include <queue>
108#include <vector>
109
110using namespace llvm;
111using namespace PatternMatch;
112
113#define DEBUG_TYPE "slsr"
114
115static const unsigned UnknownAddressSpace =
116 std::numeric_limits<unsigned>::max();
117
118DEBUG_COUNTER(StraightLineStrengthReduceCounter, "slsr-counter",
119 "Controls whether rewriteCandidate is executed.");
120
121STATISTIC(NumSCEVCandidateBasisDifferences,
122 "Number of candidate-basis SCEV differences computed by SLSR");
123
124namespace {
125
126class StraightLineStrengthReduceLegacyPass : public FunctionPass {
127 const DataLayout *DL = nullptr;
128
129public:
130 static char ID;
131
132 StraightLineStrengthReduceLegacyPass() : FunctionPass(ID) {
133 initializeStraightLineStrengthReduceLegacyPassPass(
134 *PassRegistry::getPassRegistry());
135 }
136
137 void getAnalysisUsage(AnalysisUsage &AU) const override {
138 AU.addRequired<DominatorTreeWrapperPass>();
139 AU.addRequired<ScalarEvolutionWrapperPass>();
140 AU.addRequired<TargetTransformInfoWrapperPass>();
141 // We do not modify the shape of the CFG.
142 AU.setPreservesCFG();
143 }
144
145 bool doInitialization(Module &M) override {
146 DL = &M.getDataLayout();
147 return false;
148 }
149
150 bool runOnFunction(Function &F) override;
151};
152
153class StraightLineStrengthReduce {
154public:
155 StraightLineStrengthReduce(const DataLayout *DL, DominatorTree *DT,
156 ScalarEvolution *SE, TargetTransformInfo *TTI)
157 : DL(DL), DT(DT), SE(SE), TTI(TTI) {}
158
159 // SLSR candidate. Such a candidate must be in one of the forms described in
160 // the header comments.
161 struct Candidate {
162 enum Kind {
163 Invalid, // reserved for the default constructor
164 Add, // B + i * S
165 Mul, // (B + i) * S
166 GEP, // &B[..][i * S][..]
167 };
168
169 enum DKind {
170 InvalidDelta, // reserved for the default constructor
171 IndexDelta, // Delta is a constant from Index
172 BaseDelta, // Delta is a constant or variable from Base
173 StrideDelta, // Delta is a constant or variable from Stride
174 };
175
176 Candidate() = default;
177 Candidate(Kind CT, const SCEV *B, ConstantInt *Idx, Value *S,
178 Instruction *I, const SCEV *StrideSCEV)
179 : CandidateKind(CT), Base(B), Index(Idx), Stride(S), Ins(I),
180 StrideSCEV(StrideSCEV) {}
181
182 Kind CandidateKind = Invalid;
183
184 const SCEV *Base = nullptr;
185 // TODO: Swap Index and Stride's name.
186 // Note that Index and Stride of a GEP candidate do not necessarily have the
187 // same integer type. In that case, during rewriting, Stride will be
188 // sign-extended or truncated to Index's type.
189 ConstantInt *Index = nullptr;
190
191 Value *Stride = nullptr;
192
193 // The instruction this candidate corresponds to. It helps us to rewrite a
194 // candidate with respect to its immediate basis. Note that one instruction
195 // can correspond to multiple candidates depending on how you associate the
196 // expression. For instance,
197 //
198 // (a + 1) * (b + 2)
199 //
200 // can be treated as
201 //
202 // <Base: a, Index: 1, Stride: b + 2>
203 //
204 // or
205 //
206 // <Base: b, Index: 2, Stride: a + 1>
207 Instruction *Ins = nullptr;
208
209 // Points to the immediate basis of this candidate, or nullptr if we cannot
210 // find any basis for this candidate.
211 Candidate *Basis = nullptr;
212
213 DKind DeltaKind = InvalidDelta;
214
215 // Store SCEV of Stride to compute delta from different strides
216 const SCEV *StrideSCEV = nullptr;
217
218 // Points to (Y - X) that will be used to rewrite this candidate.
219 Value *Delta = nullptr;
220
221 // List of instructions whose poison-generating annotations must be dropped
222 // if this candidate is used as the basis of an executed rewrite.
223 SmallVector<Instruction *> DropList;
224
225 /// Cost model: Evaluate the computational efficiency of the candidate.
226 ///
227 /// Efficiency levels (higher is better):
228 /// ZeroInst (5) - [Variable] or [Const]
229 /// OneInstOneVar (4) - [Variable + Const] or [Variable * Const]
230 /// OneInstTwoVar (3) - [Variable + Variable] or [Variable * Variable]
231 /// TwoInstOneVar (2) - [Const + Const * Variable]
232 /// TwoInstTwoVar (1) - [Variable + Const * Variable]
233 enum EfficiencyLevel : unsigned {
234 Unknown = 0,
235 TwoInstTwoVar = 1,
236 TwoInstOneVar = 2,
237 OneInstTwoVar = 3,
238 OneInstOneVar = 4,
239 ZeroInst = 5
240 };
241
242 static EfficiencyLevel
243 getComputationEfficiency(Kind CandidateKind, const ConstantInt *Index,
244 const Value *Stride, const SCEV *Base = nullptr) {
245 bool IsConstantBase = false;
246 bool IsZeroBase = false;
247 // When evaluating the efficiency of a rewrite, if the Base's SCEV is
248 // not available, conservatively assume the base is not constant.
249 if (auto *ConstBase = dyn_cast_or_null<SCEVConstant>(Val: Base)) {
250 IsConstantBase = true;
251 IsZeroBase = ConstBase->getValue()->isZero();
252 }
253
254 bool IsConstantStride = isa<ConstantInt>(Val: Stride);
255 bool IsZeroStride =
256 IsConstantStride && cast<ConstantInt>(Val: Stride)->isZero();
257 // All constants
258 if (IsConstantBase && IsConstantStride)
259 return ZeroInst;
260
261 // (Base + Index) * Stride
262 if (CandidateKind == Mul) {
263 if (IsZeroStride)
264 return ZeroInst;
265 if (Index->isZero())
266 return (IsConstantStride || IsConstantBase) ? OneInstOneVar
267 : OneInstTwoVar;
268
269 if (IsConstantBase)
270 return IsZeroBase && (Index->isOne() || Index->isMinusOne())
271 ? ZeroInst
272 : OneInstOneVar;
273
274 if (IsConstantStride) {
275 auto *CI = cast<ConstantInt>(Val: Stride);
276 return (CI->isOne() || CI->isMinusOne()) ? OneInstOneVar
277 : TwoInstOneVar;
278 }
279 return TwoInstTwoVar;
280 }
281
282 // Base + Index * Stride
283 assert(CandidateKind == Add || CandidateKind == GEP);
284 if (Index->isZero() || IsZeroStride)
285 return ZeroInst;
286
287 bool IsSimpleIndex = Index->isOne() || Index->isMinusOne();
288
289 if (IsConstantBase)
290 return IsZeroBase ? (IsSimpleIndex ? ZeroInst : OneInstOneVar)
291 : (IsSimpleIndex ? OneInstOneVar : TwoInstOneVar);
292
293 if (IsConstantStride)
294 return IsZeroStride ? ZeroInst : OneInstOneVar;
295
296 if (IsSimpleIndex)
297 return OneInstTwoVar;
298
299 return TwoInstTwoVar;
300 }
301
302 // Evaluate if the given delta is profitable to rewrite this candidate.
303 bool isProfitableRewrite(const Value &Delta, const DKind DeltaKind) const {
304 // This function cannot accurately evaluate the profit of whole expression
305 // with context. A candidate (B + I * S) cannot express whether this
306 // instruction needs to compute on its own (I * S), which may be shared
307 // with other candidates or may need instructions to compute.
308 // If the rewritten form has the same strength, still rewrite to
309 // (X + Delta) since it may expose more CSE opportunities on Delta, as
310 // unrolled loops usually have identical Delta for each unrolled body.
311 //
312 // Note, this function should only be used on Index Delta rewrite.
313 // Base and Stride delta need context info to evaluate the register
314 // pressure impact from variable delta.
315 return getComputationEfficiency(CandidateKind, Index, Stride, Base) <=
316 getRewriteEfficiency(Delta, DeltaKind);
317 }
318
319 // Evaluate the rewrite efficiency of this candidate with its Basis
320 EfficiencyLevel getRewriteEfficiency() const {
321 return Basis ? getRewriteEfficiency(Delta: *Delta, DeltaKind) : Unknown;
322 }
323
324 // Evaluate the rewrite efficiency of this candidate with a given delta
325 EfficiencyLevel getRewriteEfficiency(const Value &Delta,
326 const DKind DeltaKind) const {
327 switch (DeltaKind) {
328 case BaseDelta: // [X + Delta]
329 return getComputationEfficiency(
330 CandidateKind,
331 Index: ConstantInt::get(Ty: cast<IntegerType>(Val: Delta.getType()), V: 1), Stride: &Delta);
332 case StrideDelta: // [X + Index * Delta]
333 return getComputationEfficiency(CandidateKind, Index, Stride: &Delta);
334 case IndexDelta: // [X + Delta * Stride]
335 return getComputationEfficiency(CandidateKind,
336 Index: cast<ConstantInt>(Val: &Delta), Stride);
337 default:
338 return Unknown;
339 }
340 }
341
342 bool isHighEfficiency() const {
343 return getComputationEfficiency(CandidateKind, Index, Stride, Base) >=
344 OneInstOneVar;
345 }
346
347 // Verify that this candidate has valid delta components relative to the
348 // basis
349 bool hasValidDelta(const Candidate &Basis) const {
350 switch (DeltaKind) {
351 case IndexDelta:
352 // Index differs, Base and Stride must match
353 return Base == Basis.Base && StrideSCEV == Basis.StrideSCEV;
354 case StrideDelta:
355 // Stride differs, Base and Index must match
356 return Base == Basis.Base && Index == Basis.Index;
357 case BaseDelta:
358 // Base differs, Stride and Index must match
359 return StrideSCEV == Basis.StrideSCEV && Index == Basis.Index;
360 default:
361 return false;
362 }
363 }
364 };
365
366 bool runOnFunction(Function &F);
367
368private:
369 // Fetch straight-line basis for rewriting C, update C.Basis to point to it,
370 // and store the delta between C and its Basis in C.Delta.
371 void setBasisAndDeltaFor(Candidate &C);
372 // Returns whether the candidate can be folded into an addressing mode.
373 bool isFoldable(const Candidate &C, TargetTransformInfo *TTI);
374
375 // Checks whether I is in a candidate form. If so, adds all the matching forms
376 // to Candidates, and tries to find the immediate basis for each of them.
377 void allocateCandidatesAndFindBasis(Instruction *I);
378
379 // Allocate candidates and find bases for Add instructions.
380 void allocateCandidatesAndFindBasisForAdd(Instruction *I);
381
382 // Given I = LHS + RHS, factors RHS into i * S and makes (LHS + i * S) a
383 // candidate.
384 void allocateCandidatesAndFindBasisForAdd(Value *LHS, Value *RHS,
385 Instruction *I);
386 // Allocate candidates and find bases for Mul instructions.
387 void allocateCandidatesAndFindBasisForMul(Instruction *I);
388
389 // Splits LHS into Base + Index and, if succeeds, calls
390 // allocateCandidatesAndFindBasis.
391 void allocateCandidatesAndFindBasisForMul(Value *LHS, Value *RHS,
392 Instruction *I);
393
394 // Allocate candidates and find bases for GetElementPtr instructions.
395 void allocateCandidatesAndFindBasisForGEP(GetElementPtrInst *GEP);
396
397 // Adds the given form <CT, B, Idx, S> to Candidates, and finds its immediate
398 // basis.
399 void allocateCandidatesAndFindBasis(Candidate::Kind CT, const SCEV *B,
400 ConstantInt *Idx, Value *S,
401 Instruction *I);
402
403 // Rewrites candidate C with respect to Basis.
404 void rewriteCandidate(const Candidate &C);
405
406 // Emit code that computes the "bump" from Basis to C.
407 static Value *emitBump(const Candidate &Basis, const Candidate &C,
408 IRBuilder<> &Builder, const DataLayout *DL);
409
410 const DataLayout *DL = nullptr;
411 DominatorTree *DT = nullptr;
412 ScalarEvolution *SE;
413 TargetTransformInfo *TTI = nullptr;
414 std::list<Candidate> Candidates;
415
416 // Map from SCEV to instructions that represent the value,
417 // instructions are sorted in depth-first order.
418 DenseMap<const SCEV *, SmallSetVector<Instruction *, 2>> SCEVToInsts;
419
420 using SCEVUnknownSet = SmallPtrSet<const SCEVUnknown *, 4>;
421 DenseMap<const SCEV *, SCEVUnknownSet> SCEVUnknownsCache;
422
423 // Record the dependency between instructions. If C.Basis == B, we would have
424 // {B.Ins -> {C.Ins, ...}}.
425 MapVector<Instruction *, std::vector<Instruction *>> DependencyGraph;
426
427 // Map between each instruction and its possible candidates.
428 DenseMap<Instruction *, SmallVector<Candidate *, 3>> RewriteCandidates;
429
430 // All instructions that have candidates sort in topological order based on
431 // dependency graph, from roots to leaves.
432 std::vector<Instruction *> SortedCandidateInsts;
433
434 // Record all instructions that are already rewritten and will be removed
435 // later.
436 std::vector<Instruction *> DeadInstructions;
437
438 // Classify candidates against Delta kind
439 class CandidateDictTy {
440 public:
441 using CandsTy = SmallVector<Candidate *, 8>;
442 using BBToCandsTy = DenseMap<const BasicBlock *, CandsTy>;
443
444 private:
445 // Index delta Basis must have the same (Base, StrideSCEV, Inst.Type)
446 using IndexDeltaKeyTy = std::tuple<const SCEV *, const SCEV *, Type *>;
447 DenseMap<IndexDeltaKeyTy, BBToCandsTy> IndexDeltaCandidates;
448
449 // Base delta Basis must have the same (StrideSCEV, Index, Inst.Type)
450 using BaseDeltaKeyTy = std::tuple<const SCEV *, ConstantInt *, Type *>;
451 DenseMap<BaseDeltaKeyTy, BBToCandsTy> BaseDeltaCandidates;
452
453 // Stride delta Basis must have the same (Base, Index, Inst.Type)
454 using StrideDeltaKeyTy = std::tuple<const SCEV *, ConstantInt *, Type *>;
455 DenseMap<StrideDeltaKeyTy, BBToCandsTy> StrideDeltaCandidates;
456
457 public:
458 // TODO: Disable index delta on GEP after we completely move
459 // from typed GEP to PtrAdd.
460 const BBToCandsTy *getCandidatesWithDeltaKind(const Candidate &C,
461 Candidate::DKind K) const {
462 assert(K != Candidate::InvalidDelta);
463 if (K == Candidate::IndexDelta) {
464 IndexDeltaKeyTy IndexDeltaKey(C.Base, C.StrideSCEV, C.Ins->getType());
465 auto It = IndexDeltaCandidates.find(Val: IndexDeltaKey);
466 if (It != IndexDeltaCandidates.end())
467 return &It->second;
468 } else if (K == Candidate::BaseDelta) {
469 BaseDeltaKeyTy BaseDeltaKey(C.StrideSCEV, C.Index, C.Ins->getType());
470 auto It = BaseDeltaCandidates.find(Val: BaseDeltaKey);
471 if (It != BaseDeltaCandidates.end())
472 return &It->second;
473 } else {
474 assert(K == Candidate::StrideDelta);
475 StrideDeltaKeyTy StrideDeltaKey(C.Base, C.Index, C.Ins->getType());
476 auto It = StrideDeltaCandidates.find(Val: StrideDeltaKey);
477 if (It != StrideDeltaCandidates.end())
478 return &It->second;
479 }
480 return nullptr;
481 }
482
483 // Pointers to C must remain valid until CandidateDict is cleared.
484 void add(Candidate &C) {
485 Type *ValueType = C.Ins->getType();
486 BasicBlock *BB = C.Ins->getParent();
487 IndexDeltaKeyTy IndexDeltaKey(C.Base, C.StrideSCEV, ValueType);
488 BaseDeltaKeyTy BaseDeltaKey(C.StrideSCEV, C.Index, ValueType);
489 StrideDeltaKeyTy StrideDeltaKey(C.Base, C.Index, ValueType);
490 IndexDeltaCandidates[IndexDeltaKey][BB].push_back(Elt: &C);
491 BaseDeltaCandidates[BaseDeltaKey][BB].push_back(Elt: &C);
492 StrideDeltaCandidates[StrideDeltaKey][BB].push_back(Elt: &C);
493 }
494 // Remove all mappings from set
495 void clear() {
496 IndexDeltaCandidates.clear();
497 BaseDeltaCandidates.clear();
498 StrideDeltaCandidates.clear();
499 }
500 } CandidateDict;
501
502 const SCEV *getAndRecordSCEV(Value *V) {
503 auto *S = SE->getSCEV(V);
504 if (isa<Instruction>(Val: V) && !(isa<SCEVCouldNotCompute>(Val: S) ||
505 isa<SCEVUnknown>(Val: S) || isa<SCEVConstant>(Val: S)))
506 SCEVToInsts[S].insert(X: cast<Instruction>(Val: V));
507
508 return S;
509 }
510
511 bool candidatePredicate(Candidate *Basis, Candidate &C, Candidate::DKind K);
512
513 bool hasSameSCEVUnknowns(const SCEV *A, const SCEV *B);
514
515 bool searchFrom(const CandidateDictTy::BBToCandsTy &BBToCands, Candidate &C,
516 Candidate::DKind K);
517
518 // Get the nearest instruction before CI that represents the value of S,
519 // return nullptr if no instruction is associated with S or S is not a
520 // reusable expression.
521 Value *getNearestValueOfSCEV(const SCEV *S, const Instruction *CI) const {
522 if (isa<SCEVCouldNotCompute>(Val: S))
523 return nullptr;
524
525 if (auto *SU = dyn_cast<SCEVUnknown>(Val: S))
526 return SU->getValue();
527 if (auto *SC = dyn_cast<SCEVConstant>(Val: S))
528 return SC->getValue();
529
530 auto It = SCEVToInsts.find(Val: S);
531 if (It == SCEVToInsts.end())
532 return nullptr;
533
534 // Instructions are sorted in depth-first order, so search for the nearest
535 // instruction by walking the list in reverse order.
536 for (Instruction *I : reverse(C: It->second))
537 if (DT->dominates(Def: I, User: CI))
538 return I;
539
540 return nullptr;
541 }
542
543 struct DeltaInfo {
544 Candidate *Cand;
545 Candidate::DKind DeltaKind;
546 Value *Delta;
547
548 DeltaInfo()
549 : Cand(nullptr), DeltaKind(Candidate::InvalidDelta), Delta(nullptr) {}
550 DeltaInfo(Candidate *Cand, Candidate::DKind DeltaKind, Value *Delta)
551 : Cand(Cand), DeltaKind(DeltaKind), Delta(Delta) {}
552 operator bool() const { return Cand != nullptr; }
553 };
554
555 friend raw_ostream &operator<<(raw_ostream &OS, const DeltaInfo &DI);
556
557 DeltaInfo compressPath(Candidate &C, Candidate *Basis) const;
558
559 Candidate *pickRewriteCandidate(Instruction *I) const;
560 void sortCandidateInstructions();
561 Value *getDelta(const Candidate &C, const Candidate &Basis,
562 Candidate::DKind K) const;
563 static bool isSimilar(Candidate &C, Candidate &Basis, Candidate::DKind K);
564
565 // Add Basis -> C in DependencyGraph and propagate
566 // C.Stride and C.Delta's dependency to C
567 void addDependency(Candidate &C, Candidate *Basis) {
568 if (Basis)
569 DependencyGraph[Basis->Ins].emplace_back(args&: C.Ins);
570
571 // If any candidate of Inst has a basis, then Inst will be rewritten,
572 // C must be rewritten after rewriting Inst, so we need to propagate
573 // the dependency to C
574 auto PropagateDependency = [&](Instruction *Inst) {
575 if (auto CandsIt = RewriteCandidates.find(Val: Inst);
576 CandsIt != RewriteCandidates.end() &&
577 llvm::any_of(Range&: CandsIt->second,
578 P: [](Candidate *Cand) { return Cand->Basis; }))
579 DependencyGraph[Inst].emplace_back(args&: C.Ins);
580 };
581
582 // If C has a variable delta and the delta is a candidate,
583 // propagate its dependency to C
584 if (auto *DeltaInst = dyn_cast_or_null<Instruction>(Val: C.Delta))
585 PropagateDependency(DeltaInst);
586
587 // If the stride is a candidate, propagate its dependency to C
588 if (auto *StrideInst = dyn_cast<Instruction>(Val: C.Stride))
589 PropagateDependency(StrideInst);
590 };
591};
592
593inline raw_ostream &operator<<(raw_ostream &OS,
594 const StraightLineStrengthReduce::Candidate &C) {
595 OS << "Ins: " << *C.Ins << "\n Base: " << *C.Base
596 << "\n Index: " << *C.Index << "\n Stride: " << *C.Stride
597 << "\n StrideSCEV: " << *C.StrideSCEV;
598 if (C.Basis)
599 OS << "\n Delta: " << *C.Delta << "\n Basis: \n [ " << *C.Basis << " ]";
600 return OS;
601}
602
603[[maybe_unused]] LLVM_DUMP_METHOD inline raw_ostream &
604operator<<(raw_ostream &OS, const StraightLineStrengthReduce::DeltaInfo &DI) {
605 OS << "Cand: " << *DI.Cand << "\n";
606 OS << "Delta Kind: ";
607 switch (DI.DeltaKind) {
608 case StraightLineStrengthReduce::Candidate::IndexDelta:
609 OS << "Index";
610 break;
611 case StraightLineStrengthReduce::Candidate::BaseDelta:
612 OS << "Base";
613 break;
614 case StraightLineStrengthReduce::Candidate::StrideDelta:
615 OS << "Stride";
616 break;
617 default:
618 break;
619 }
620 OS << "\nDelta: " << *DI.Delta;
621 return OS;
622}
623
624} // end anonymous namespace
625
626char StraightLineStrengthReduceLegacyPass::ID = 0;
627
628INITIALIZE_PASS_BEGIN(StraightLineStrengthReduceLegacyPass, "slsr",
629 "Straight line strength reduction", false, false)
630INITIALIZE_PASS_DEPENDENCY(DominatorTreeWrapperPass)
631INITIALIZE_PASS_DEPENDENCY(ScalarEvolutionWrapperPass)
632INITIALIZE_PASS_DEPENDENCY(TargetTransformInfoWrapperPass)
633INITIALIZE_PASS_END(StraightLineStrengthReduceLegacyPass, "slsr",
634 "Straight line strength reduction", false, false)
635
636FunctionPass *llvm::createStraightLineStrengthReducePass() {
637 return new StraightLineStrengthReduceLegacyPass();
638}
639
640// A helper function that unifies the bitwidth of A and B.
641static void unifyBitWidth(APInt &A, APInt &B) {
642 if (A.getBitWidth() < B.getBitWidth())
643 A = A.sext(width: B.getBitWidth());
644 else if (A.getBitWidth() > B.getBitWidth())
645 B = B.sext(width: A.getBitWidth());
646}
647
648// Whether sign-extending V to a wider type may not distribute over arithmetic,
649// i.e. the narrow value does not sign-extend linearly. Only an add/sub/mul/shl
650// carrying the `nsw` flag is known to sign-extend linearly; anything else is
651// treated conservatively as possibly wrapping. This notably covers
652// `xor X, signmask`, which merely flips the sign bit but ScalarEvolution models
653// as a non-nsw `add X, signmask` (so sext does not distribute over it).
654static bool mayHaveSignedWrap(const Value *V) {
655 // OverflowingBinaryOperator covers exactly add/sub/mul/shl.
656 const auto *OBO = dyn_cast<OverflowingBinaryOperator>(Val: V);
657 return !OBO || !OBO->hasNoSignedWrap();
658}
659
660// True when the GEP index is narrower than the index width, i.e. it is
661// implicitly sign-extended to the index width (not the pointer width) of the
662// address space before the address computation. A value already at or wider
663// than the index width is not sign-extended (it is used as-is or truncated), so
664// it cannot trigger the non-distributing-sext problem.
665static bool isSignExtendedGepIndex(const Value *Idx, GetElementPtrInst *GEP,
666 const DataLayout *DL) {
667 return Idx->getType()->getIntegerBitWidth() <
668 DL->getIndexSizeInBits(AS: GEP->getAddressSpace());
669}
670
671// A narrow GEP index is sign-extended to the index width before the address
672// computation. SLSR's Stride-delta rewrite turns two such GEPs into
673// Basis + Index * (Sc - Sb), so the stride difference Sc - Sb is reconstructed
674// in the sign-extended domain. This requires sext(Sc) == sext(Sb) +
675// sext(Delta).
676//
677// This screens the rewritten candidate's stride Sc = Sb + Delta: if Sc is
678// computed by a possibly-wrapping op, sext(Sc) does not equal sext(Sb) +
679// sext(Delta) and the rewrite would produce a wrong pointer.
680static bool isSafeToFactorGepIndex(const Value *Idx, GetElementPtrInst *GEP,
681 const DataLayout *DL) {
682 return !isSignExtendedGepIndex(Idx, GEP, DL) || !mayHaveSignedWrap(V: Idx);
683}
684
685Value *StraightLineStrengthReduce::getDelta(const Candidate &C,
686 const Candidate &Basis,
687 Candidate::DKind K) const {
688 if (K == Candidate::IndexDelta) {
689 APInt Idx = C.Index->getValue();
690 APInt BasisIdx = Basis.Index->getValue();
691 unifyBitWidth(A&: Idx, B&: BasisIdx);
692 APInt IndexDelta = Idx - BasisIdx;
693 IntegerType *DeltaType =
694 IntegerType::get(C&: C.Ins->getContext(), NumBits: IndexDelta.getBitWidth());
695 return ConstantInt::get(Ty: DeltaType, V: IndexDelta);
696 } else if (K == Candidate::BaseDelta || K == Candidate::StrideDelta) {
697 const SCEV *BasisPart =
698 (K == Candidate::BaseDelta) ? Basis.Base : Basis.StrideSCEV;
699 const SCEV *CandPart = (K == Candidate::BaseDelta) ? C.Base : C.StrideSCEV;
700 ++NumSCEVCandidateBasisDifferences;
701 const SCEV *Diff = SE->getMinusSCEV(LHS: CandPart, RHS: BasisPart);
702 return getNearestValueOfSCEV(S: Diff, CI: C.Ins);
703 }
704 return nullptr;
705}
706
707bool StraightLineStrengthReduce::isSimilar(Candidate &C, Candidate &Basis,
708 Candidate::DKind K) {
709 bool SameType = false;
710 switch (K) {
711 case Candidate::StrideDelta:
712 SameType = C.StrideSCEV->getType() == Basis.StrideSCEV->getType();
713 break;
714 case Candidate::BaseDelta:
715 SameType = C.Base->getType() == Basis.Base->getType();
716 break;
717 case Candidate::IndexDelta:
718 SameType = true;
719 break;
720 default:;
721 }
722 return SameType && Basis.Ins != C.Ins &&
723 Basis.CandidateKind == C.CandidateKind;
724}
725
726bool StraightLineStrengthReduce::hasSameSCEVUnknowns(const SCEV *A,
727 const SCEV *B) {
728 auto CacheUnknowns = [&](const SCEV *Root) {
729 auto [It, Inserted] = SCEVUnknownsCache.try_emplace(Key: Root);
730 if (!Inserted)
731 return;
732
733 struct Collector {
734 SCEVUnknownSet &Unknowns;
735
736 bool follow(const SCEV *S) {
737 if (auto *Unknown = dyn_cast<SCEVUnknown>(Val: S))
738 Unknowns.insert(Ptr: Unknown);
739 return true;
740 }
741 bool isDone() const { return false; }
742 } C{.Unknowns: It->second};
743 visitAll(Root, Visitor&: C);
744 };
745 CacheUnknowns(A);
746 CacheUnknowns(B);
747
748 return SCEVUnknownsCache.find(Val: A)->second == SCEVUnknownsCache.find(Val: B)->second;
749}
750
751// Try to find a Delta that C can reuse Basis to rewrite.
752// Set C.Delta, C.Basis, and C.DeltaKind if found.
753// Return true if found a constant delta.
754// Return false if not found or the delta is not a constant.
755bool StraightLineStrengthReduce::candidatePredicate(Candidate *Basis,
756 Candidate &C,
757 Candidate::DKind K) {
758 if (!isSimilar(C, Basis&: *Basis, K))
759 return false;
760
761 // Once a reusable delta is found, only a constant delta can improve it.
762 // Different symbolic leaves cannot cancel to a constant, so such a basis
763 // cannot improve C. Skip it and continue searching older candidates.
764 if (C.Delta && K != Candidate::IndexDelta) {
765 const SCEV *CandidateSCEV =
766 K == Candidate::BaseDelta ? C.Base : C.StrideSCEV;
767 const SCEV *BasisSCEV =
768 K == Candidate::BaseDelta ? Basis->Base : Basis->StrideSCEV;
769 if (!hasSameSCEVUnknowns(A: CandidateSCEV, B: BasisSCEV))
770 return false;
771 }
772
773 assert(DT->dominates(Basis->Ins, C.Ins));
774 Value *Delta = getDelta(C, Basis: *Basis, K);
775 if (!Delta)
776 return false;
777
778 // For a GEP Stride-delta rewrite g2 = g1 + Index * Delta, the addresses are
779 // computed from the sign-extended strides, so this requires
780 // sext(Sc) == sext(Sb) + sext(Delta).
781 //
782 // The rewritten candidate's stride Sc = Sb + Delta is already screened
783 // broadly at allocation time (allocateCandidatesAndFindBasis): a wrapping Sc
784 // breaks the identity for any Delta. The basis's stride Sb = Sc - Delta only
785 // needs screening when Delta folds to a *constant*: then sext(Sb) + C can
786 // differ from sext(Sc) if Sb wraps. For a *variable* Delta the basis may wrap
787 // and still be sound, because the candidate stride carries the no-wrap
788 // guarantee (e.g. Sc is an `add nsw`, as in stride_var); rejecting it would
789 // pessimize those.
790 if (K == Candidate::StrideDelta && C.CandidateKind == Candidate::GEP &&
791 isa<ConstantInt>(Val: Delta)) {
792 auto *BasisGEP = cast<GetElementPtrInst>(Val: Basis->Ins);
793 if (!isSafeToFactorGepIndex(Idx: Basis->Stride, GEP: BasisGEP, DL))
794 return false;
795 }
796
797 // IndexDelta rewrite is not always profitable, e.g.,
798 // X = B + 8 * S
799 // Y = B + S,
800 // rewriting Y to X - 7 * S is probably a bad idea.
801 // So, we need to check if the rewrite form's computation efficiency
802 // is better than the original form.
803 if (K == Candidate::IndexDelta &&
804 !C.isProfitableRewrite(Delta: *Delta, DeltaKind: Candidate::IndexDelta))
805 return false;
806
807 // Record delta if none has been found yet, or the new delta is
808 // a constant that is better than the existing delta.
809 if (!C.Delta || isa<ConstantInt>(Val: Delta)) {
810 C.Delta = Delta;
811 C.Basis = Basis;
812 C.DeltaKind = K;
813 }
814 return isa<ConstantInt>(Val: C.Delta);
815}
816
817// return true if find a Basis with constant delta and stop searching,
818// return false if did not find a Basis or the delta is not a constant
819// and continue searching for a Basis with constant delta
820bool StraightLineStrengthReduce::searchFrom(
821 const CandidateDictTy::BBToCandsTy &BBToCands, Candidate &C,
822 Candidate::DKind K) {
823
824 // Stride delta rewrite on Mul form is usually non-profitable, and Base
825 // delta rewrite sometimes is profitable, so we do not support them on Mul.
826 if (C.CandidateKind == Candidate::Mul && K != Candidate::IndexDelta)
827 return false;
828
829 // Search dominating candidates by walking the immediate-dominator chain
830 // from the candidate's defining block upward. Visiting blocks in this
831 // order ensures we prefer the closest dominating basis.
832 const BasicBlock *BB = C.Ins->getParent();
833 while (BB) {
834 auto It = BBToCands.find(Val: BB);
835 if (It != BBToCands.end())
836 for (Candidate *Basis : reverse(C: It->second))
837 if (candidatePredicate(Basis, C, K))
838 return true;
839
840 const DomTreeNode *Node = DT->getNode(BB);
841 if (!Node)
842 break;
843 Node = Node->getIDom();
844 BB = Node ? Node->getBlock() : nullptr;
845 }
846 return false;
847}
848
849void StraightLineStrengthReduce::setBasisAndDeltaFor(Candidate &C) {
850 if (const auto *BaseDeltaCandidates =
851 CandidateDict.getCandidatesWithDeltaKind(C, K: Candidate::BaseDelta))
852 if (searchFrom(BBToCands: *BaseDeltaCandidates, C, K: Candidate::BaseDelta)) {
853 LLVM_DEBUG(dbgs() << "Found delta from Base: " << *C.Delta << "\n");
854 return;
855 }
856
857 if (const auto *StrideDeltaCandidates =
858 CandidateDict.getCandidatesWithDeltaKind(C, K: Candidate::StrideDelta))
859 if (searchFrom(BBToCands: *StrideDeltaCandidates, C, K: Candidate::StrideDelta)) {
860 LLVM_DEBUG(dbgs() << "Found delta from Stride: " << *C.Delta << "\n");
861 return;
862 }
863
864 if (const auto *IndexDeltaCandidates =
865 CandidateDict.getCandidatesWithDeltaKind(C, K: Candidate::IndexDelta))
866 if (searchFrom(BBToCands: *IndexDeltaCandidates, C, K: Candidate::IndexDelta)) {
867 LLVM_DEBUG(dbgs() << "Found delta from Index: " << *C.Delta << "\n");
868 return;
869 }
870
871 // If we did not find a constant delta, we might have found a variable delta
872 if (C.Delta) {
873 LLVM_DEBUG({
874 dbgs() << "Found delta from ";
875 if (C.DeltaKind == Candidate::BaseDelta)
876 dbgs() << "Base: ";
877 else
878 dbgs() << "Stride: ";
879 dbgs() << *C.Delta << "\n";
880 });
881 assert(C.DeltaKind != Candidate::InvalidDelta && C.Basis);
882 }
883}
884
885// Compress the path from `Basis` to the deepest Basis in the Basis chain
886// to avoid non-profitable data dependency and improve ILP.
887// X = A + 1
888// Y = X + 1
889// Z = Y + 1
890// ->
891// X = A + 1
892// Y = A + 2
893// Z = A + 3
894// Return the delta info for C aginst the new Basis
895auto StraightLineStrengthReduce::compressPath(Candidate &C,
896 Candidate *Basis) const
897 -> DeltaInfo {
898 if (!Basis || !Basis->Basis || C.CandidateKind == Candidate::Mul)
899 return {};
900 Candidate *Root = Basis;
901 Value *NewDelta = nullptr;
902 auto NewKind = Candidate::InvalidDelta;
903
904 while (Root->Basis) {
905 Candidate *NextRoot = Root->Basis;
906 if (C.Base == NextRoot->Base && C.StrideSCEV == NextRoot->StrideSCEV &&
907 isSimilar(C, Basis&: *NextRoot, K: Candidate::IndexDelta)) {
908 ConstantInt *CI =
909 cast<ConstantInt>(Val: getDelta(C, Basis: *NextRoot, K: Candidate::IndexDelta));
910 if (CI->isZero() || CI->isOne() || isa<SCEVConstant>(Val: C.StrideSCEV)) {
911 Root = NextRoot;
912 NewKind = Candidate::IndexDelta;
913 NewDelta = CI;
914 continue;
915 }
916 }
917
918 const SCEV *CandPart = nullptr;
919 const SCEV *BasisPart = nullptr;
920 auto CurrKind = Candidate::InvalidDelta;
921 if (C.Base == NextRoot->Base && C.Index == NextRoot->Index) {
922 CandPart = C.StrideSCEV;
923 BasisPart = NextRoot->StrideSCEV;
924 CurrKind = Candidate::StrideDelta;
925 } else if (C.StrideSCEV == NextRoot->StrideSCEV &&
926 C.Index == NextRoot->Index) {
927 CandPart = C.Base;
928 BasisPart = NextRoot->Base;
929 CurrKind = Candidate::BaseDelta;
930 } else
931 break;
932
933 assert(CandPart && BasisPart);
934 if (!isSimilar(C, Basis&: *NextRoot, K: CurrKind))
935 break;
936
937 // Path compression folds a constant Stride-delta directly against the
938 // deeper basis NextRoot, bypassing candidatePredicate's wrap guard. With a
939 // constant delta sext(Sb) + C can differ from sext(Sc) if the deeper
940 // basis's stride wraps, so do not compress past such a basis (mirrors the
941 // check in candidatePredicate).
942 if (CurrKind == Candidate::StrideDelta &&
943 C.CandidateKind == Candidate::GEP &&
944 !isSafeToFactorGepIndex(Idx: NextRoot->Stride,
945 GEP: cast<GetElementPtrInst>(Val: NextRoot->Ins), DL))
946 break;
947
948 ++NumSCEVCandidateBasisDifferences;
949 if (auto DeltaVal =
950 dyn_cast<SCEVConstant>(Val: SE->getMinusSCEV(LHS: CandPart, RHS: BasisPart))) {
951 Root = NextRoot;
952 NewDelta = DeltaVal->getValue();
953 NewKind = CurrKind;
954 } else
955 break;
956 }
957
958 if (Root != Basis) {
959 assert(NewKind != Candidate::InvalidDelta && NewDelta);
960 LLVM_DEBUG(dbgs() << "Found new Basis with " << *NewDelta
961 << " from path compression.\n");
962 return {Root, NewKind, NewDelta};
963 }
964
965 return {};
966}
967
968// Topologically sort candidate instructions based on their relationship in
969// dependency graph.
970void StraightLineStrengthReduce::sortCandidateInstructions() {
971 SortedCandidateInsts.clear();
972 // An instruction may have multiple candidates that get different Basis
973 // instructions, and each candidate can get dependencies from Basis and
974 // Stride when Stride will also be rewritten by SLSR. Hence, an instruction
975 // may have multiple dependencies. Use InDegree to ensure all dependencies
976 // processed before processing itself.
977 DenseMap<Instruction *, int> InDegree;
978 for (auto &KV : DependencyGraph) {
979 InDegree.try_emplace(Key: KV.first, Args: 0);
980
981 for (auto *Child : KV.second) {
982 InDegree[Child]++;
983 }
984 }
985 std::queue<Instruction *> WorkList;
986 DenseSet<Instruction *> Visited;
987
988 for (auto &KV : DependencyGraph)
989 if (InDegree[KV.first] == 0)
990 WorkList.push(x: KV.first);
991
992 while (!WorkList.empty()) {
993 Instruction *I = WorkList.front();
994 WorkList.pop();
995 if (!Visited.insert(V: I).second)
996 continue;
997
998 SortedCandidateInsts.push_back(x: I);
999
1000 for (auto *Next : DependencyGraph[I]) {
1001 auto &Degree = InDegree[Next];
1002 if (--Degree == 0)
1003 WorkList.push(x: Next);
1004 }
1005 }
1006
1007 assert(SortedCandidateInsts.size() == DependencyGraph.size() &&
1008 "Dependency graph should not have cycles");
1009}
1010
1011auto StraightLineStrengthReduce::pickRewriteCandidate(Instruction *I) const
1012 -> Candidate * {
1013 // Return the candidate of instruction I that has the highest profit.
1014 auto It = RewriteCandidates.find(Val: I);
1015 if (It == RewriteCandidates.end())
1016 return nullptr;
1017
1018 Candidate *BestC = nullptr;
1019 auto BestEfficiency = Candidate::Unknown;
1020 for (Candidate *C : reverse(C: It->second))
1021 if (C->Basis) {
1022 auto Efficiency = C->getRewriteEfficiency();
1023 if (Efficiency > BestEfficiency) {
1024 BestEfficiency = Efficiency;
1025 BestC = C;
1026 }
1027 }
1028
1029 return BestC;
1030}
1031
1032static bool isGEPFoldable(GetElementPtrInst *GEP,
1033 const TargetTransformInfo *TTI) {
1034 SmallVector<const Value *, 4> Indices(GEP->indices());
1035 return TTI->getGEPCost(
1036 PointeeType: GEP->getSourceElementType(), Ptr: GEP->getPointerOperand(), Operands: Indices,
1037 /*CostKind*/ TTI::TargetCostKind::TCK_SizeAndLatency) ==
1038 TargetTransformInfo::TCC_Free;
1039}
1040
1041// Returns whether (Base + Index * Stride) can be folded to an addressing mode.
1042static bool isAddFoldable(const SCEV *Base, ConstantInt *Index, Value *Stride,
1043 TargetTransformInfo *TTI) {
1044 // Index->getSExtValue() may crash if Index is wider than 64-bit.
1045 return Index->getBitWidth() <= 64 &&
1046 TTI->isLegalAddressingMode(Ty: Base->getType(), BaseGV: nullptr, BaseOffset: 0, HasBaseReg: true,
1047 Scale: Index->getSExtValue(), AddrSpace: UnknownAddressSpace);
1048}
1049
1050bool StraightLineStrengthReduce::isFoldable(const Candidate &C,
1051 TargetTransformInfo *TTI) {
1052 if (C.CandidateKind == Candidate::Add)
1053 return isAddFoldable(Base: C.Base, Index: C.Index, Stride: C.Stride, TTI);
1054 if (C.CandidateKind == Candidate::GEP)
1055 return isGEPFoldable(GEP: cast<GetElementPtrInst>(Val: C.Ins), TTI);
1056 return false;
1057}
1058
1059void StraightLineStrengthReduce::allocateCandidatesAndFindBasis(
1060 Candidate::Kind CT, const SCEV *B, ConstantInt *Idx, Value *S,
1061 Instruction *I) {
1062 bool IsSafe = CT != Candidate::GEP ||
1063 isSafeToFactorGepIndex(Idx: S, GEP: cast<GetElementPtrInst>(Val: I), DL);
1064 // Record the SCEV of S that we may use it as a variable delta.
1065 // Ensure that we rewrite C with a existing IR that reproduces delta value.
1066
1067 Candidate C(CT, B, Idx, S, I, getAndRecordSCEV(V: S));
1068 // If we can fold I into an addressing mode, computing I is likely free or
1069 // takes only one instruction. So, we don't need to analyze or rewrite it.
1070 //
1071 // Currently, this algorithm can at best optimize complex computations into
1072 // a `variable +/* constant` form. However, some targets have stricter
1073 // constraints on the their addressing mode.
1074 // For example, a `variable + constant` can only be folded to an addressing
1075 // mode if the constant falls within a certain range.
1076 // So, we also check if the instruction is already high efficient enough
1077 // for the strength reduction algorithm.
1078 if (IsSafe && !isFoldable(C, TTI) && !C.isHighEfficiency()) {
1079 setBasisAndDeltaFor(C);
1080
1081 // Compress unnecessary rewrite to improve ILP
1082 if (auto Res = compressPath(C, Basis: C.Basis)) {
1083 C.Basis = Res.Cand;
1084 C.DeltaKind = Res.DeltaKind;
1085 C.Delta = Res.Delta;
1086 }
1087 }
1088 // Regardless of whether we find a basis for C, we need to push C to the
1089 // candidate list so that it can be the basis of other candidates.
1090 LLVM_DEBUG(dbgs() << "Allocated Candidate: " << C << "\n");
1091 Candidates.push_back(x: C);
1092 RewriteCandidates[C.Ins].push_back(Elt: &Candidates.back());
1093 // Only add to the dict if this instruction is safe to reuse as a basis. By
1094 // doing this early we avoid calling canReuseInstruction repeatedly for the
1095 // same instruction. The DropList is stored on the Candidate so the flags can
1096 // be dropped only if this candidate is used by an executed rewrite.
1097 if (!ScalarOptions::Global.enable_poison_reuse_guard ||
1098 SE->canReuseInstruction(S: SE->getSCEV(V: I), I, DropPoisonGeneratingInsts&: Candidates.back().DropList)) {
1099 CandidateDict.add(C&: Candidates.back());
1100 }
1101}
1102
1103void StraightLineStrengthReduce::allocateCandidatesAndFindBasis(
1104 Instruction *I) {
1105 switch (I->getOpcode()) {
1106 case Instruction::Add:
1107 allocateCandidatesAndFindBasisForAdd(I);
1108 break;
1109 case Instruction::Mul:
1110 allocateCandidatesAndFindBasisForMul(I);
1111 break;
1112 case Instruction::GetElementPtr:
1113 allocateCandidatesAndFindBasisForGEP(GEP: cast<GetElementPtrInst>(Val: I));
1114 break;
1115 }
1116}
1117
1118void StraightLineStrengthReduce::allocateCandidatesAndFindBasisForAdd(
1119 Instruction *I) {
1120 // Try matching B + i * S.
1121 if (!isa<IntegerType>(Val: I->getType()))
1122 return;
1123
1124 assert(I->getNumOperands() == 2 && "isn't I an add?");
1125 Value *LHS = I->getOperand(i: 0), *RHS = I->getOperand(i: 1);
1126 allocateCandidatesAndFindBasisForAdd(LHS, RHS, I);
1127 if (LHS != RHS)
1128 allocateCandidatesAndFindBasisForAdd(LHS: RHS, RHS: LHS, I);
1129}
1130
1131void StraightLineStrengthReduce::allocateCandidatesAndFindBasisForAdd(
1132 Value *LHS, Value *RHS, Instruction *I) {
1133 Value *S = nullptr;
1134 ConstantInt *Idx = nullptr;
1135 if (match(V: RHS, P: m_Mul(L: m_Value(V&: S), R: m_ConstantInt(CI&: Idx)))) {
1136 // I = LHS + RHS = LHS + Idx * S
1137 allocateCandidatesAndFindBasis(CT: Candidate::Add, B: SE->getSCEV(V: LHS), Idx, S, I);
1138 } else if (match(V: RHS, P: m_Shl(L: m_Value(V&: S), R: m_ConstantInt(CI&: Idx)))) {
1139 // I = LHS + RHS = LHS + (S << Idx) = LHS + S * (1 << Idx)
1140 APInt One(Idx->getBitWidth(), 1);
1141 Idx = ConstantInt::get(Context&: Idx->getContext(), V: One << Idx->getValue());
1142 allocateCandidatesAndFindBasis(CT: Candidate::Add, B: SE->getSCEV(V: LHS), Idx, S, I);
1143 } else {
1144 // At least, I = LHS + 1 * RHS
1145 ConstantInt *One = ConstantInt::get(Ty: cast<IntegerType>(Val: I->getType()), V: 1);
1146 allocateCandidatesAndFindBasis(CT: Candidate::Add, B: SE->getSCEV(V: LHS), Idx: One, S: RHS,
1147 I);
1148 }
1149}
1150
1151// Returns true if A matches B + C where C is constant.
1152static bool matchesAdd(Value *A, Value *&B, ConstantInt *&C) {
1153 return match(V: A, P: m_c_Add(L: m_Value(V&: B), R: m_ConstantInt(CI&: C)));
1154}
1155
1156// Returns true if A matches B | C where C is constant.
1157static bool matchesOr(Value *A, Value *&B, ConstantInt *&C) {
1158 return match(V: A, P: m_c_Or(L: m_Value(V&: B), R: m_ConstantInt(CI&: C)));
1159}
1160
1161void StraightLineStrengthReduce::allocateCandidatesAndFindBasisForMul(
1162 Value *LHS, Value *RHS, Instruction *I) {
1163 Value *B = nullptr;
1164 ConstantInt *Idx = nullptr;
1165 if (matchesAdd(A: LHS, B, C&: Idx)) {
1166 // If LHS is in the form of "Base + Index", then I is in the form of
1167 // "(Base + Index) * RHS".
1168 allocateCandidatesAndFindBasis(CT: Candidate::Mul, B: SE->getSCEV(V: B), Idx, S: RHS, I);
1169 } else if (matchesOr(A: LHS, B, C&: Idx) && haveNoCommonBitsSet(LHSCache: B, RHSCache: Idx, SQ: *DL)) {
1170 // If LHS is in the form of "Base | Index" and Base and Index have no common
1171 // bits set, then
1172 // Base | Index = Base + Index
1173 // and I is thus in the form of "(Base + Index) * RHS".
1174 allocateCandidatesAndFindBasis(CT: Candidate::Mul, B: SE->getSCEV(V: B), Idx, S: RHS, I);
1175 } else {
1176 // Otherwise, at least try the form (LHS + 0) * RHS.
1177 ConstantInt *Zero = ConstantInt::get(Ty: cast<IntegerType>(Val: I->getType()), V: 0);
1178 allocateCandidatesAndFindBasis(CT: Candidate::Mul, B: SE->getSCEV(V: LHS), Idx: Zero, S: RHS,
1179 I);
1180 }
1181}
1182
1183void StraightLineStrengthReduce::allocateCandidatesAndFindBasisForMul(
1184 Instruction *I) {
1185 // Try matching (B + i) * S.
1186 // TODO: we could extend SLSR to float and vector types.
1187 if (!isa<IntegerType>(Val: I->getType()))
1188 return;
1189
1190 assert(I->getNumOperands() == 2 && "isn't I a mul?");
1191 Value *LHS = I->getOperand(i: 0), *RHS = I->getOperand(i: 1);
1192 allocateCandidatesAndFindBasisForMul(LHS, RHS, I);
1193 if (LHS != RHS) {
1194 // Symmetrically, try to split RHS to Base + Index.
1195 allocateCandidatesAndFindBasisForMul(LHS: RHS, RHS: LHS, I);
1196 }
1197}
1198
1199void StraightLineStrengthReduce::allocateCandidatesAndFindBasisForGEP(
1200 GetElementPtrInst *GEP) {
1201 // TODO: handle vector GEPs
1202 if (GEP->getType()->isVectorTy())
1203 return;
1204
1205 SmallVector<SCEVUse, 4> IndexExprs;
1206 for (Use &Idx : GEP->indices())
1207 IndexExprs.push_back(Elt: SE->getSCEV(V: Idx));
1208
1209 gep_type_iterator GTI = gep_type_begin(GEP);
1210 for (unsigned I = 1, E = GEP->getNumOperands(); I != E; ++I, ++GTI) {
1211 if (GTI.isStruct())
1212 continue;
1213
1214 SCEVUse OrigIndexExpr = IndexExprs[I - 1];
1215 IndexExprs[I - 1] = SE->getZero(Ty: OrigIndexExpr.getPointer()->getType());
1216
1217 // The base of this candidate is GEP's base plus the offsets of all
1218 // indices except this current one.
1219 SCEVUse BaseExpr = SE->getGEPExpr(GEP: cast<GEPOperator>(Val: GEP), IndexExprs);
1220 Value *ArrayIdx = GEP->getOperand(i_nocapture: I);
1221 uint64_t ElementSize = GTI.getSequentialElementStride(DL: *DL);
1222 IntegerType *PtrIdxTy = cast<IntegerType>(Val: DL->getIndexType(PtrTy: GEP->getType()));
1223 // If the element size overflows the type, truncate.
1224 ConstantInt *ElementSizeIdx =
1225 ConstantInt::getSigned(Ty: PtrIdxTy, V: ElementSize, /*ImplicitTrunc=*/true);
1226 if (ArrayIdx->getType()->getIntegerBitWidth() <=
1227 DL->getIndexSizeInBits(AS: GEP->getAddressSpace())) {
1228 // Skip factoring if ArrayIdx is wider than the index size, because
1229 // ArrayIdx is implicitly truncated to the index size.
1230 allocateCandidatesAndFindBasis(CT: Candidate::GEP, B: BaseExpr, Idx: ElementSizeIdx,
1231 S: ArrayIdx, I: GEP);
1232 }
1233 // When ArrayIdx is the sext of a value, we try to factor that value as
1234 // well. Handling this case is important because array indices are
1235 // typically sign-extended to the pointer index size.
1236 Value *TruncatedArrayIdx = nullptr;
1237 if (match(V: ArrayIdx, P: m_SExt(Op: m_Value(V&: TruncatedArrayIdx))) &&
1238 TruncatedArrayIdx->getType()->getIntegerBitWidth() <=
1239 DL->getIndexSizeInBits(AS: GEP->getAddressSpace())) {
1240 // Skip factoring if TruncatedArrayIdx is wider than the pointer size,
1241 // because TruncatedArrayIdx is implicitly truncated to the pointer size.
1242 allocateCandidatesAndFindBasis(CT: Candidate::GEP, B: BaseExpr, Idx: ElementSizeIdx,
1243 S: TruncatedArrayIdx, I: GEP);
1244 }
1245
1246 IndexExprs[I - 1] = OrigIndexExpr;
1247 }
1248}
1249
1250Value *StraightLineStrengthReduce::emitBump(const Candidate &Basis,
1251 const Candidate &C,
1252 IRBuilder<> &Builder,
1253 const DataLayout *DL) {
1254 auto CreateMul = [&](Value *LHS, Value *RHS) {
1255 if (ConstantInt *CR = dyn_cast<ConstantInt>(Val: RHS)) {
1256 const APInt &ConstRHS = CR->getValue();
1257 IntegerType *DeltaType =
1258 IntegerType::get(C&: C.Ins->getContext(), NumBits: ConstRHS.getBitWidth());
1259 if (ConstRHS.isPowerOf2()) {
1260 ConstantInt *Exponent =
1261 ConstantInt::get(Ty: DeltaType, V: ConstRHS.logBase2());
1262 return Builder.CreateShl(LHS, RHS: Exponent);
1263 }
1264 if (ConstRHS.isNegatedPowerOf2()) {
1265 ConstantInt *Exponent =
1266 ConstantInt::get(Ty: DeltaType, V: (-ConstRHS).logBase2());
1267 return Builder.CreateNeg(V: Builder.CreateShl(LHS, RHS: Exponent));
1268 }
1269 }
1270
1271 return Builder.CreateMul(LHS, RHS);
1272 };
1273
1274 Value *Delta = C.Delta;
1275 // If Delta is 0, C is a fully redundant of C.Basis,
1276 // just replace C.Ins with Basis.Ins
1277 if (ConstantInt *CI = dyn_cast<ConstantInt>(Val: Delta);
1278 CI && CI->getValue().isZero())
1279 return nullptr;
1280
1281 if (C.DeltaKind == Candidate::IndexDelta) {
1282 APInt IndexDelta = cast<ConstantInt>(Val: C.Delta)->getValue();
1283 // IndexDelta
1284 // X = B + i * S
1285 // Y = B + i` * S
1286 // = B + (i + IndexDelta) * S
1287 // = B + i * S + IndexDelta * S
1288 // = X + IndexDelta * S
1289 // Bump = (i' - i) * S
1290
1291 // Common case 1: if (i' - i) is 1, Bump = S.
1292 if (IndexDelta == 1)
1293 return C.Stride;
1294 // Common case 2: if (i' - i) is -1, Bump = -S.
1295 if (IndexDelta.isAllOnes())
1296 return Builder.CreateNeg(V: C.Stride);
1297
1298 IntegerType *DeltaType =
1299 IntegerType::get(C&: Basis.Ins->getContext(), NumBits: IndexDelta.getBitWidth());
1300 Value *ExtendedStride = Builder.CreateSExtOrTrunc(V: C.Stride, DestTy: DeltaType);
1301
1302 return CreateMul(ExtendedStride, C.Delta);
1303 }
1304
1305 assert(C.DeltaKind == Candidate::StrideDelta ||
1306 C.DeltaKind == Candidate::BaseDelta);
1307 assert(C.CandidateKind != Candidate::Mul);
1308 // StrideDelta
1309 // X = B + i * S
1310 // Y = B + i * S'
1311 // = B + i * (S + StrideDelta)
1312 // = B + i * S + i * StrideDelta
1313 // = X + i * StrideDelta
1314 // Bump = i * (S' - S)
1315 //
1316 // BaseDelta
1317 // X = B + i * S
1318 // Y = B' + i * S
1319 // = (B + BaseDelta) + i * S
1320 // = X + BaseDelta
1321 // Bump = (B' - B).
1322 Value *Bump = C.Delta;
1323 if (C.DeltaKind == Candidate::StrideDelta) {
1324 // If this value is consumed by a GEP, promote StrideDelta before doing
1325 // StrideDelta * Index to ensure the same semantics as the original GEP.
1326 if (C.CandidateKind == Candidate::GEP) {
1327 auto *GEP = cast<GetElementPtrInst>(Val: C.Ins);
1328 Type *NewScalarIndexTy =
1329 DL->getIndexType(PtrTy: GEP->getPointerOperandType()->getScalarType());
1330 Bump = Builder.CreateSExtOrTrunc(V: Bump, DestTy: NewScalarIndexTy);
1331 }
1332 if (!C.Index->isOne()) {
1333 Value *ExtendedIndex =
1334 Builder.CreateSExtOrTrunc(V: C.Index, DestTy: Bump->getType());
1335 Bump = CreateMul(Bump, ExtendedIndex);
1336 }
1337 }
1338 return Bump;
1339}
1340
1341void StraightLineStrengthReduce::rewriteCandidate(const Candidate &C) {
1342 if (!DebugCounter::shouldExecute(Counter&: StraightLineStrengthReduceCounter))
1343 return;
1344
1345 const Candidate &Basis = *C.Basis;
1346 assert(C.Delta && C.CandidateKind == Basis.CandidateKind &&
1347 C.hasValidDelta(Basis));
1348
1349 for (Instruction *I : Basis.DropList)
1350 I->dropPoisonGeneratingAnnotations();
1351
1352 IRBuilder<> Builder(C.Ins);
1353 Value *Bump = emitBump(Basis, C, Builder, DL);
1354 Value *Reduced = nullptr; // equivalent to but weaker than C.Ins
1355 // If delta is 0, C is a fully redundant of Basis, and Bump is nullptr,
1356 // just replace C.Ins with Basis.Ins
1357 if (!Bump)
1358 Reduced = Basis.Ins;
1359 else {
1360 switch (C.CandidateKind) {
1361 case Candidate::Add:
1362 case Candidate::Mul: {
1363 // C = Basis + Bump
1364 Value *NegBump;
1365 if (match(V: Bump, P: m_Neg(V: m_Value(V&: NegBump)))) {
1366 // If Bump is a neg instruction, emit C = Basis - (-Bump).
1367 Reduced = Builder.CreateSub(LHS: Basis.Ins, RHS: NegBump);
1368 // We only use the negative argument of Bump, and Bump itself may be
1369 // trivially dead.
1370 RecursivelyDeleteTriviallyDeadInstructions(V: Bump);
1371 } else {
1372 // It's tempting to preserve nsw on Bump and/or Reduced. However, it's
1373 // usually unsound, e.g.,
1374 //
1375 // X = (-2 +nsw 1) *nsw INT_MAX
1376 // Y = (-2 +nsw 3) *nsw INT_MAX
1377 // =>
1378 // Y = X + 2 * INT_MAX
1379 //
1380 // Neither + and * in the resultant expression are nsw.
1381 Reduced = Builder.CreateAdd(LHS: Basis.Ins, RHS: Bump);
1382 }
1383 break;
1384 }
1385 case Candidate::GEP: {
1386 bool InBounds = cast<GetElementPtrInst>(Val: C.Ins)->isInBounds();
1387 // C = (char *)Basis + Bump
1388 Reduced = Builder.CreatePtrAdd(Ptr: Basis.Ins, Offset: Bump, Name: "", NW: InBounds);
1389 break;
1390 }
1391 default:
1392 llvm_unreachable("C.CandidateKind is invalid");
1393 };
1394 Reduced->takeName(V: C.Ins);
1395 }
1396 C.Ins->replaceAllUsesWith(V: Reduced);
1397 DeadInstructions.push_back(x: C.Ins);
1398}
1399
1400bool StraightLineStrengthReduceLegacyPass::runOnFunction(Function &F) {
1401 if (skipFunction(F))
1402 return false;
1403
1404 auto *TTI = &getAnalysis<TargetTransformInfoWrapperPass>().getTTI(F);
1405 auto *DT = &getAnalysis<DominatorTreeWrapperPass>().getDomTree();
1406 auto *SE = &getAnalysis<ScalarEvolutionWrapperPass>().getSE();
1407 return StraightLineStrengthReduce(DL, DT, SE, TTI).runOnFunction(F);
1408}
1409
1410bool StraightLineStrengthReduce::runOnFunction(Function &F) {
1411 LLVM_DEBUG(dbgs() << "SLSR on Function: " << F.getName() << "\n");
1412 // Traverse the dominator tree in the depth-first order. This order makes sure
1413 // all bases of a candidate are in Candidates when we process it.
1414 for (const auto Node : depth_first(G: DT))
1415 for (auto &I : *(Node->getBlock()))
1416 allocateCandidatesAndFindBasis(I: &I);
1417
1418 // Build the dependency graph and sort candidate instructions from dependency
1419 // roots to leaves
1420 for (auto &C : Candidates) {
1421 DependencyGraph.try_emplace(Key: C.Ins);
1422 addDependency(C, Basis: C.Basis);
1423 }
1424 sortCandidateInstructions();
1425
1426 // Rewrite candidates in the topological order that rewrites a Candidate
1427 // always before rewriting its Basis
1428 for (Instruction *I : reverse(C&: SortedCandidateInsts))
1429 if (Candidate *C = pickRewriteCandidate(I))
1430 rewriteCandidate(C: *C);
1431
1432 for (auto *DeadIns : DeadInstructions)
1433 // A dead instruction may be another dead instruction's op,
1434 // don't delete an instruction twice
1435 if (DeadIns->getParent())
1436 RecursivelyDeleteTriviallyDeadInstructions(V: DeadIns);
1437
1438 bool Ret = !DeadInstructions.empty();
1439 DeadInstructions.clear();
1440 DependencyGraph.clear();
1441 RewriteCandidates.clear();
1442 SortedCandidateInsts.clear();
1443 // First clear all references to candidates in the list
1444 CandidateDict.clear();
1445 // Then destroy the list
1446 Candidates.clear();
1447 return Ret;
1448}
1449
1450PreservedAnalyses
1451StraightLineStrengthReducePass::run(Function &F, FunctionAnalysisManager &AM) {
1452 const DataLayout *DL = &F.getDataLayout();
1453 auto *DT = &AM.getResult<DominatorTreeAnalysis>(IR&: F);
1454 auto *SE = &AM.getResult<ScalarEvolutionAnalysis>(IR&: F);
1455 auto *TTI = &AM.getResult<TargetIRAnalysis>(IR&: F);
1456
1457 if (!StraightLineStrengthReduce(DL, DT, SE, TTI).runOnFunction(F))
1458 return PreservedAnalyses::all();
1459
1460 PreservedAnalyses PA;
1461 PA.preserveSet<CFGAnalyses>();
1462 PA.preserve<ScalarEvolutionAnalysis>();
1463 PA.preserve<TargetIRAnalysis>();
1464 return PA;
1465}
1466