1//===-------- LoopDataPrefetch.cpp - Loop Data Prefetching Pass -----------===//
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 a Loop Data Prefetching Pass.
10//
11//===----------------------------------------------------------------------===//
12
13#include "llvm/Transforms/Scalar/LoopDataPrefetch.h"
14#include "ScalarOptions.h"
15#include "llvm/InitializePasses.h"
16
17#include "llvm/ADT/DepthFirstIterator.h"
18#include "llvm/ADT/Statistic.h"
19#include "llvm/Analysis/AssumptionCache.h"
20#include "llvm/Analysis/CodeMetrics.h"
21#include "llvm/Analysis/LoopInfo.h"
22#include "llvm/Analysis/OptimizationRemarkEmitter.h"
23#include "llvm/Analysis/ScalarEvolution.h"
24#include "llvm/Analysis/ScalarEvolutionExpressions.h"
25#include "llvm/Analysis/TargetTransformInfo.h"
26#include "llvm/IR/Dominators.h"
27#include "llvm/IR/Function.h"
28#include "llvm/Support/Debug.h"
29#include "llvm/Transforms/Scalar.h"
30#include "llvm/Transforms/Utils.h"
31#include "llvm/Transforms/Utils/ScalarEvolutionExpander.h"
32
33#define DEBUG_TYPE "loop-data-prefetch"
34
35using namespace llvm;
36
37STATISTIC(NumPrefetches, "Number of prefetches inserted");
38
39namespace {
40
41/// Loop prefetch implementation class.
42class LoopDataPrefetch {
43public:
44 LoopDataPrefetch(AssumptionCache *AC, DominatorTree *DT, LoopInfo *LI,
45 ScalarEvolution *SE, const TargetTransformInfo *TTI,
46 OptimizationRemarkEmitter *ORE)
47 : Opts(ScalarOptions::Global), AC(AC), DT(DT), LI(LI), SE(SE), TTI(TTI),
48 ORE(ORE) {}
49
50 bool run();
51
52private:
53 bool runOnLoop(Loop *L);
54
55 /// Check if the stride of the accesses is large enough to
56 /// warrant a prefetch.
57 bool isStrideLargeEnough(const SCEVAddRecExpr *AR, unsigned TargetMinStride);
58
59 unsigned getMinPrefetchStride(unsigned NumMemAccesses,
60 unsigned NumStridedMemAccesses,
61 unsigned NumPrefetches,
62 bool HasCall) {
63 if (Opts.min_prefetch_stride)
64 return *Opts.min_prefetch_stride;
65 return TTI->getMinPrefetchStride(NumMemAccesses, NumStridedMemAccesses,
66 NumPrefetches, HasCall);
67 }
68
69 unsigned getPrefetchDistance() {
70 if (Opts.prefetch_distance)
71 return *Opts.prefetch_distance;
72 return TTI->getPrefetchDistance();
73 }
74
75 unsigned getMaxPrefetchIterationsAhead() {
76 if (Opts.max_prefetch_iters_ahead)
77 return *Opts.max_prefetch_iters_ahead;
78 return TTI->getMaxPrefetchIterationsAhead();
79 }
80
81 bool doPrefetchWrites() {
82 return valueOr(X: Opts.loop_prefetch_writes, Default: TTI->enableWritePrefetching());
83 }
84
85 const ScalarOptions &Opts;
86 AssumptionCache *AC;
87 DominatorTree *DT;
88 LoopInfo *LI;
89 ScalarEvolution *SE;
90 const TargetTransformInfo *TTI;
91 OptimizationRemarkEmitter *ORE;
92};
93
94/// Legacy class for inserting loop data prefetches.
95class LoopDataPrefetchLegacyPass : public FunctionPass {
96public:
97 static char ID; // Pass ID, replacement for typeid
98 LoopDataPrefetchLegacyPass() : FunctionPass(ID) {
99 initializeLoopDataPrefetchLegacyPassPass(*PassRegistry::getPassRegistry());
100 }
101
102 void getAnalysisUsage(AnalysisUsage &AU) const override {
103 AU.addRequired<AssumptionCacheTracker>();
104 AU.addRequired<DominatorTreeWrapperPass>();
105 AU.addPreserved<DominatorTreeWrapperPass>();
106 AU.addRequired<LoopInfoWrapperPass>();
107 AU.addPreserved<LoopInfoWrapperPass>();
108 AU.addRequiredID(ID&: LoopSimplifyID);
109 AU.addPreservedID(ID&: LoopSimplifyID);
110 AU.addRequired<OptimizationRemarkEmitterWrapperPass>();
111 AU.addRequired<ScalarEvolutionWrapperPass>();
112 AU.addPreserved<ScalarEvolutionWrapperPass>();
113 AU.addRequired<TargetTransformInfoWrapperPass>();
114 }
115
116 bool runOnFunction(Function &F) override;
117 };
118}
119
120char LoopDataPrefetchLegacyPass::ID = 0;
121INITIALIZE_PASS_BEGIN(LoopDataPrefetchLegacyPass, "loop-data-prefetch",
122 "Loop Data Prefetch", false, false)
123INITIALIZE_PASS_DEPENDENCY(AssumptionCacheTracker)
124INITIALIZE_PASS_DEPENDENCY(TargetTransformInfoWrapperPass)
125INITIALIZE_PASS_DEPENDENCY(LoopInfoWrapperPass)
126INITIALIZE_PASS_DEPENDENCY(LoopSimplify)
127INITIALIZE_PASS_DEPENDENCY(OptimizationRemarkEmitterWrapperPass)
128INITIALIZE_PASS_DEPENDENCY(ScalarEvolutionWrapperPass)
129INITIALIZE_PASS_END(LoopDataPrefetchLegacyPass, "loop-data-prefetch",
130 "Loop Data Prefetch", false, false)
131
132FunctionPass *llvm::createLoopDataPrefetchPass() {
133 return new LoopDataPrefetchLegacyPass();
134}
135
136bool LoopDataPrefetch::isStrideLargeEnough(const SCEVAddRecExpr *AR,
137 unsigned TargetMinStride) {
138 // No need to check if any stride goes.
139 if (TargetMinStride <= 1)
140 return true;
141
142 const auto *ConstStride = dyn_cast<SCEVConstant>(Val: AR->getStepRecurrence(SE&: *SE));
143 // If MinStride is set, don't prefetch unless we can ensure that stride is
144 // larger.
145 if (!ConstStride)
146 return false;
147
148 unsigned AbsStride = std::abs(i: ConstStride->getAPInt().getSExtValue());
149 return TargetMinStride <= AbsStride;
150}
151
152PreservedAnalyses LoopDataPrefetchPass::run(Function &F,
153 FunctionAnalysisManager &AM) {
154 DominatorTree *DT = &AM.getResult<DominatorTreeAnalysis>(IR&: F);
155 LoopInfo *LI = &AM.getResult<LoopAnalysis>(IR&: F);
156 ScalarEvolution *SE = &AM.getResult<ScalarEvolutionAnalysis>(IR&: F);
157 AssumptionCache *AC = &AM.getResult<AssumptionAnalysis>(IR&: F);
158 OptimizationRemarkEmitter *ORE =
159 &AM.getResult<OptimizationRemarkEmitterAnalysis>(IR&: F);
160 const TargetTransformInfo *TTI = &AM.getResult<TargetIRAnalysis>(IR&: F);
161
162 LoopDataPrefetch LDP(AC, DT, LI, SE, TTI, ORE);
163 bool Changed = LDP.run();
164
165 if (Changed) {
166 PreservedAnalyses PA;
167 PA.preserve<DominatorTreeAnalysis>();
168 PA.preserve<LoopAnalysis>();
169 return PA;
170 }
171
172 return PreservedAnalyses::all();
173}
174
175bool LoopDataPrefetchLegacyPass::runOnFunction(Function &F) {
176 if (skipFunction(F))
177 return false;
178
179 DominatorTree *DT = &getAnalysis<DominatorTreeWrapperPass>().getDomTree();
180 LoopInfo *LI = &getAnalysis<LoopInfoWrapperPass>().getLoopInfo();
181 ScalarEvolution *SE = &getAnalysis<ScalarEvolutionWrapperPass>().getSE();
182 AssumptionCache *AC =
183 &getAnalysis<AssumptionCacheTracker>().getAssumptionCache(F);
184 OptimizationRemarkEmitter *ORE =
185 &getAnalysis<OptimizationRemarkEmitterWrapperPass>().getORE();
186 const TargetTransformInfo *TTI =
187 &getAnalysis<TargetTransformInfoWrapperPass>().getTTI(F);
188
189 LoopDataPrefetch LDP(AC, DT, LI, SE, TTI, ORE);
190 return LDP.run();
191}
192
193bool LoopDataPrefetch::run() {
194 // If PrefetchDistance is not set, don't run the pass. This gives an
195 // opportunity for targets to run this pass for selected subtargets only
196 // (whose TTI sets PrefetchDistance and CacheLineSize).
197 if (getPrefetchDistance() == 0 || TTI->getCacheLineSize() == 0) {
198 LLVM_DEBUG(dbgs() << "Please set both PrefetchDistance and CacheLineSize "
199 "for loop data prefetch.\n");
200 return false;
201 }
202
203 bool MadeChange = false;
204
205 for (Loop *I : *LI)
206 for (Loop *L : depth_first(G: I))
207 MadeChange |= runOnLoop(L);
208
209 return MadeChange;
210}
211
212/// A record for a potential prefetch made during the initial scan of the
213/// loop. This is used to let a single prefetch target multiple memory accesses.
214struct Prefetch {
215 /// The address formula for this prefetch as returned by ScalarEvolution.
216 const SCEVAddRecExpr *LSCEVAddRec;
217 /// The point of insertion for the prefetch instruction.
218 Instruction *InsertPt = nullptr;
219 /// True if targeting a write memory access.
220 bool Writes = false;
221 /// The (first seen) prefetched instruction.
222 Instruction *MemI = nullptr;
223
224 /// Constructor to create a new Prefetch for \p I.
225 Prefetch(const SCEVAddRecExpr *L, Instruction *I) : LSCEVAddRec(L) {
226 addInstruction(I);
227 };
228
229 /// Add the instruction \param I to this prefetch. If it's not the first
230 /// one, 'InsertPt' and 'Writes' will be updated as required.
231 /// \param PtrDiff the known constant address difference to the first added
232 /// instruction.
233 void addInstruction(Instruction *I, DominatorTree *DT = nullptr,
234 int64_t PtrDiff = 0) {
235 if (!InsertPt) {
236 MemI = I;
237 InsertPt = I;
238 Writes = isa<StoreInst>(Val: I);
239 } else {
240 BasicBlock *PrefBB = InsertPt->getParent();
241 BasicBlock *InsBB = I->getParent();
242 if (PrefBB != InsBB) {
243 BasicBlock *DomBB = DT->findNearestCommonDominator(A: PrefBB, B: InsBB);
244 if (DomBB != PrefBB)
245 InsertPt = DomBB->getTerminator();
246 }
247
248 if (isa<StoreInst>(Val: I) && PtrDiff == 0)
249 Writes = true;
250 }
251 }
252};
253
254bool LoopDataPrefetch::runOnLoop(Loop *L) {
255 bool MadeChange = false;
256
257 // Only prefetch in the inner-most loop
258 if (!L->isInnermost())
259 return MadeChange;
260
261 SmallPtrSet<const Value *, 32> EphValues;
262 CodeMetrics::collectEphemeralValues(L, AC, EphValues);
263
264 // Calculate the number of iterations ahead to prefetch
265 CodeMetrics Metrics;
266 bool HasCall = false;
267 for (const auto BB : L->blocks()) {
268 // If the loop already has prefetches, then assume that the user knows
269 // what they are doing and don't add any more.
270 for (auto &I : *BB) {
271 if (isa<CallInst>(Val: &I) || isa<InvokeInst>(Val: &I)) {
272 if (const Function *F = cast<CallBase>(Val&: I).getCalledFunction()) {
273 if (F->getIntrinsicID() == Intrinsic::prefetch)
274 return MadeChange;
275 if (TTI->isLoweredToCall(F))
276 HasCall = true;
277 } else { // indirect call.
278 HasCall = true;
279 }
280 }
281 }
282 Metrics.analyzeBasicBlock(BB, TTI: *TTI, EphValues);
283 }
284
285 if (!Metrics.NumInsts.isValid())
286 return MadeChange;
287
288 unsigned LoopSize = Metrics.NumInsts.getValue();
289 if (!LoopSize)
290 LoopSize = 1;
291
292 unsigned ItersAhead = getPrefetchDistance() / LoopSize;
293 if (!ItersAhead)
294 ItersAhead = 1;
295
296 if (ItersAhead > getMaxPrefetchIterationsAhead())
297 return MadeChange;
298
299 unsigned ConstantMaxTripCount = SE->getSmallConstantMaxTripCount(L);
300 if (ConstantMaxTripCount && ConstantMaxTripCount < ItersAhead + 1)
301 return MadeChange;
302
303 unsigned NumMemAccesses = 0;
304 unsigned NumStridedMemAccesses = 0;
305 SmallVector<Prefetch, 16> Prefetches;
306 for (const auto BB : L->blocks())
307 for (auto &I : *BB) {
308 Value *PtrValue;
309 Instruction *MemI;
310
311 if (LoadInst *LMemI = dyn_cast<LoadInst>(Val: &I)) {
312 MemI = LMemI;
313 PtrValue = LMemI->getPointerOperand();
314 } else if (StoreInst *SMemI = dyn_cast<StoreInst>(Val: &I)) {
315 if (!doPrefetchWrites()) continue;
316 MemI = SMemI;
317 PtrValue = SMemI->getPointerOperand();
318 } else continue;
319
320 unsigned PtrAddrSpace = PtrValue->getType()->getPointerAddressSpace();
321 if (!TTI->shouldPrefetchAddressSpace(AS: PtrAddrSpace))
322 continue;
323 NumMemAccesses++;
324 if (L->isLoopInvariant(V: PtrValue))
325 continue;
326
327 const SCEV *LSCEV = SE->getSCEV(V: PtrValue);
328 const SCEVAddRecExpr *LSCEVAddRec = dyn_cast<SCEVAddRecExpr>(Val: LSCEV);
329 if (!LSCEVAddRec)
330 continue;
331 NumStridedMemAccesses++;
332
333 // We don't want to double prefetch individual cache lines. If this
334 // access is known to be within one cache line of some other one that
335 // has already been prefetched, then don't prefetch this one as well.
336 bool DupPref = false;
337 for (auto &Pref : Prefetches) {
338 const SCEV *PtrDiff = SE->getMinusSCEV(LHS: LSCEVAddRec, RHS: Pref.LSCEVAddRec);
339 if (const SCEVConstant *ConstPtrDiff =
340 dyn_cast<SCEVConstant>(Val: PtrDiff)) {
341 int64_t PD = std::abs(i: ConstPtrDiff->getValue()->getSExtValue());
342 if (PD < (int64_t) TTI->getCacheLineSize()) {
343 Pref.addInstruction(I: MemI, DT, PtrDiff: PD);
344 DupPref = true;
345 break;
346 }
347 }
348 }
349 if (!DupPref)
350 Prefetches.push_back(Elt: Prefetch(LSCEVAddRec, MemI));
351 }
352
353 unsigned TargetMinStride =
354 getMinPrefetchStride(NumMemAccesses, NumStridedMemAccesses,
355 NumPrefetches: Prefetches.size(), HasCall);
356
357 LLVM_DEBUG(dbgs() << "Prefetching " << ItersAhead
358 << " iterations ahead (loop size: " << LoopSize << ") in "
359 << L->getHeader()->getParent()->getName() << ": " << *L);
360 LLVM_DEBUG(dbgs() << "Loop has: "
361 << NumMemAccesses << " memory accesses, "
362 << NumStridedMemAccesses << " strided memory accesses, "
363 << Prefetches.size() << " potential prefetch(es), "
364 << "a minimum stride of " << TargetMinStride << ", "
365 << (HasCall ? "calls" : "no calls") << ".\n");
366
367 for (auto &P : Prefetches) {
368 // Check if the stride of the accesses is large enough to warrant a
369 // prefetch.
370 if (!isStrideLargeEnough(AR: P.LSCEVAddRec, TargetMinStride))
371 continue;
372
373 BasicBlock *BB = P.InsertPt->getParent();
374 SCEVExpander SCEVE(*SE, "prefaddr");
375 const SCEV *NextLSCEV = SE->getAddExpr(
376 LHS: P.LSCEVAddRec,
377 RHS: SE->getMulExpr(LHS: SE->getConstant(Ty: P.LSCEVAddRec->getType(), V: ItersAhead),
378 RHS: P.LSCEVAddRec->getStepRecurrence(SE&: *SE)));
379 if (!SCEVE.isSafeToExpand(S: NextLSCEV))
380 continue;
381
382 unsigned PtrAddrSpace = NextLSCEV->getType()->getPointerAddressSpace();
383 Type *I8Ptr = PointerType::get(C&: BB->getContext(), AddressSpace: PtrAddrSpace);
384 Value *PrefPtrValue = SCEVE.expandCodeFor(SH: NextLSCEV, Ty: I8Ptr, I: P.InsertPt);
385
386 IRBuilder<> Builder(P.InsertPt);
387 Type *I32 = Type::getInt32Ty(C&: BB->getContext());
388 Builder.CreateIntrinsic(ID: Intrinsic::prefetch, OverloadTypes: PrefPtrValue->getType(),
389 Args: {PrefPtrValue, ConstantInt::get(Ty: I32, V: P.Writes),
390 ConstantInt::get(Ty: I32, V: 3),
391 ConstantInt::get(Ty: I32, V: 1)});
392 ++NumPrefetches;
393 LLVM_DEBUG(dbgs() << " Access: "
394 << *P.MemI->getOperand(isa<LoadInst>(P.MemI) ? 0 : 1)
395 << ", SCEV: " << *P.LSCEVAddRec << "\n");
396 ORE->emit(RemarkBuilder: [&]() {
397 return OptimizationRemark(DEBUG_TYPE, "Prefetched", P.MemI)
398 << "prefetched memory access";
399 });
400
401 MadeChange = true;
402 }
403
404 return MadeChange;
405}
406