1 //===--- CodeGenPGO.cpp - PGO Instrumentation for LLVM CodeGen --*- C++ -*-===//
2 //
3 //                     The LLVM Compiler Infrastructure
4 //
5 // This file is distributed under the University of Illinois Open Source
6 // License. See LICENSE.TXT for details.
7 //
8 //===----------------------------------------------------------------------===//
9 //
10 // Instrumentation-based profile-guided optimization
11 //
12 //===----------------------------------------------------------------------===//
13 
14 #include "CodeGenPGO.h"
15 #include "CodeGenFunction.h"
16 #include "CoverageMappingGen.h"
17 #include "clang/AST/RecursiveASTVisitor.h"
18 #include "clang/AST/StmtVisitor.h"
19 #include "llvm/IR/Intrinsics.h"
20 #include "llvm/IR/MDBuilder.h"
21 #include "llvm/ProfileData/InstrProfReader.h"
22 #include "llvm/Support/Endian.h"
23 #include "llvm/Support/FileSystem.h"
24 #include "llvm/Support/MD5.h"
25 
26 using namespace clang;
27 using namespace CodeGen;
28 
29 void CodeGenPGO::setFuncName(StringRef Name,
30                              llvm::GlobalValue::LinkageTypes Linkage) {
31   StringRef RawFuncName = Name;
32 
33   // Function names may be prefixed with a binary '1' to indicate
34   // that the backend should not modify the symbols due to any platform
35   // naming convention. Do not include that '1' in the PGO profile name.
36   if (RawFuncName[0] == '\1')
37     RawFuncName = RawFuncName.substr(1);
38 
39   FuncName = RawFuncName;
40   if (llvm::GlobalValue::isLocalLinkage(Linkage)) {
41     // For local symbols, prepend the main file name to distinguish them.
42     // Do not include the full path in the file name since there's no guarantee
43     // that it will stay the same, e.g., if the files are checked out from
44     // version control in different locations.
45     if (CGM.getCodeGenOpts().MainFileName.empty())
46       FuncName = FuncName.insert(0, "<unknown>:");
47     else
48       FuncName = FuncName.insert(0, CGM.getCodeGenOpts().MainFileName + ":");
49   }
50 
51   // If we're generating a profile, create a variable for the name.
52   if (CGM.getCodeGenOpts().ProfileInstrGenerate)
53     createFuncNameVar(Linkage);
54 }
55 
56 void CodeGenPGO::setFuncName(llvm::Function *Fn) {
57   setFuncName(Fn->getName(), Fn->getLinkage());
58 }
59 
60 void CodeGenPGO::createFuncNameVar(llvm::GlobalValue::LinkageTypes Linkage) {
61   // We generally want to match the function's linkage, but available_externally
62   // and extern_weak both have the wrong semantics, and anything that doesn't
63   // need to link across compilation units doesn't need to be visible at all.
64   if (Linkage == llvm::GlobalValue::ExternalWeakLinkage)
65     Linkage = llvm::GlobalValue::LinkOnceAnyLinkage;
66   else if (Linkage == llvm::GlobalValue::AvailableExternallyLinkage)
67     Linkage = llvm::GlobalValue::LinkOnceODRLinkage;
68   else if (Linkage == llvm::GlobalValue::InternalLinkage ||
69            Linkage == llvm::GlobalValue::ExternalLinkage)
70     Linkage = llvm::GlobalValue::PrivateLinkage;
71 
72   auto *Value =
73       llvm::ConstantDataArray::getString(CGM.getLLVMContext(), FuncName, false);
74   FuncNameVar =
75       new llvm::GlobalVariable(CGM.getModule(), Value->getType(), true, Linkage,
76                                Value, "__llvm_profile_name_" + FuncName);
77 
78   // Hide the symbol so that we correctly get a copy for each executable.
79   if (!llvm::GlobalValue::isLocalLinkage(FuncNameVar->getLinkage()))
80     FuncNameVar->setVisibility(llvm::GlobalValue::HiddenVisibility);
81 }
82 
83 namespace {
84 /// \brief Stable hasher for PGO region counters.
85 ///
86 /// PGOHash produces a stable hash of a given function's control flow.
87 ///
88 /// Changing the output of this hash will invalidate all previously generated
89 /// profiles -- i.e., don't do it.
90 ///
91 /// \note  When this hash does eventually change (years?), we still need to
92 /// support old hashes.  We'll need to pull in the version number from the
93 /// profile data format and use the matching hash function.
94 class PGOHash {
95   uint64_t Working;
96   unsigned Count;
97   llvm::MD5 MD5;
98 
99   static const int NumBitsPerType = 6;
100   static const unsigned NumTypesPerWord = sizeof(uint64_t) * 8 / NumBitsPerType;
101   static const unsigned TooBig = 1u << NumBitsPerType;
102 
103 public:
104   /// \brief Hash values for AST nodes.
105   ///
106   /// Distinct values for AST nodes that have region counters attached.
107   ///
108   /// These values must be stable.  All new members must be added at the end,
109   /// and no members should be removed.  Changing the enumeration value for an
110   /// AST node will affect the hash of every function that contains that node.
111   enum HashType : unsigned char {
112     None = 0,
113     LabelStmt = 1,
114     WhileStmt,
115     DoStmt,
116     ForStmt,
117     CXXForRangeStmt,
118     ObjCForCollectionStmt,
119     SwitchStmt,
120     CaseStmt,
121     DefaultStmt,
122     IfStmt,
123     CXXTryStmt,
124     CXXCatchStmt,
125     ConditionalOperator,
126     BinaryOperatorLAnd,
127     BinaryOperatorLOr,
128     BinaryConditionalOperator,
129 
130     // Keep this last.  It's for the static assert that follows.
131     LastHashType
132   };
133   static_assert(LastHashType <= TooBig, "Too many types in HashType");
134 
135   // TODO: When this format changes, take in a version number here, and use the
136   // old hash calculation for file formats that used the old hash.
137   PGOHash() : Working(0), Count(0) {}
138   void combine(HashType Type);
139   uint64_t finalize();
140 };
141 const int PGOHash::NumBitsPerType;
142 const unsigned PGOHash::NumTypesPerWord;
143 const unsigned PGOHash::TooBig;
144 
145 /// A RecursiveASTVisitor that fills a map of statements to PGO counters.
146 struct MapRegionCounters : public RecursiveASTVisitor<MapRegionCounters> {
147   /// The next counter value to assign.
148   unsigned NextCounter;
149   /// The function hash.
150   PGOHash Hash;
151   /// The map of statements to counters.
152   llvm::DenseMap<const Stmt *, unsigned> &CounterMap;
153 
154   MapRegionCounters(llvm::DenseMap<const Stmt *, unsigned> &CounterMap)
155       : NextCounter(0), CounterMap(CounterMap) {}
156 
157   // Blocks and lambdas are handled as separate functions, so we need not
158   // traverse them in the parent context.
159   bool TraverseBlockExpr(BlockExpr *BE) { return true; }
160   bool TraverseLambdaBody(LambdaExpr *LE) { return true; }
161   bool TraverseCapturedStmt(CapturedStmt *CS) { return true; }
162 
163   bool VisitDecl(const Decl *D) {
164     switch (D->getKind()) {
165     default:
166       break;
167     case Decl::Function:
168     case Decl::CXXMethod:
169     case Decl::CXXConstructor:
170     case Decl::CXXDestructor:
171     case Decl::CXXConversion:
172     case Decl::ObjCMethod:
173     case Decl::Block:
174     case Decl::Captured:
175       CounterMap[D->getBody()] = NextCounter++;
176       break;
177     }
178     return true;
179   }
180 
181   bool VisitStmt(const Stmt *S) {
182     auto Type = getHashType(S);
183     if (Type == PGOHash::None)
184       return true;
185 
186     CounterMap[S] = NextCounter++;
187     Hash.combine(Type);
188     return true;
189   }
190   PGOHash::HashType getHashType(const Stmt *S) {
191     switch (S->getStmtClass()) {
192     default:
193       break;
194     case Stmt::LabelStmtClass:
195       return PGOHash::LabelStmt;
196     case Stmt::WhileStmtClass:
197       return PGOHash::WhileStmt;
198     case Stmt::DoStmtClass:
199       return PGOHash::DoStmt;
200     case Stmt::ForStmtClass:
201       return PGOHash::ForStmt;
202     case Stmt::CXXForRangeStmtClass:
203       return PGOHash::CXXForRangeStmt;
204     case Stmt::ObjCForCollectionStmtClass:
205       return PGOHash::ObjCForCollectionStmt;
206     case Stmt::SwitchStmtClass:
207       return PGOHash::SwitchStmt;
208     case Stmt::CaseStmtClass:
209       return PGOHash::CaseStmt;
210     case Stmt::DefaultStmtClass:
211       return PGOHash::DefaultStmt;
212     case Stmt::IfStmtClass:
213       return PGOHash::IfStmt;
214     case Stmt::CXXTryStmtClass:
215       return PGOHash::CXXTryStmt;
216     case Stmt::CXXCatchStmtClass:
217       return PGOHash::CXXCatchStmt;
218     case Stmt::ConditionalOperatorClass:
219       return PGOHash::ConditionalOperator;
220     case Stmt::BinaryConditionalOperatorClass:
221       return PGOHash::BinaryConditionalOperator;
222     case Stmt::BinaryOperatorClass: {
223       const BinaryOperator *BO = cast<BinaryOperator>(S);
224       if (BO->getOpcode() == BO_LAnd)
225         return PGOHash::BinaryOperatorLAnd;
226       if (BO->getOpcode() == BO_LOr)
227         return PGOHash::BinaryOperatorLOr;
228       break;
229     }
230     }
231     return PGOHash::None;
232   }
233 };
234 
235 /// A StmtVisitor that propagates the raw counts through the AST and
236 /// records the count at statements where the value may change.
237 struct ComputeRegionCounts : public ConstStmtVisitor<ComputeRegionCounts> {
238   /// PGO state.
239   CodeGenPGO &PGO;
240 
241   /// A flag that is set when the current count should be recorded on the
242   /// next statement, such as at the exit of a loop.
243   bool RecordNextStmtCount;
244 
245   /// The map of statements to count values.
246   llvm::DenseMap<const Stmt *, uint64_t> &CountMap;
247 
248   /// BreakContinueStack - Keep counts of breaks and continues inside loops.
249   struct BreakContinue {
250     uint64_t BreakCount;
251     uint64_t ContinueCount;
252     BreakContinue() : BreakCount(0), ContinueCount(0) {}
253   };
254   SmallVector<BreakContinue, 8> BreakContinueStack;
255 
256   ComputeRegionCounts(llvm::DenseMap<const Stmt *, uint64_t> &CountMap,
257                       CodeGenPGO &PGO)
258       : PGO(PGO), RecordNextStmtCount(false), CountMap(CountMap) {}
259 
260   void RecordStmtCount(const Stmt *S) {
261     if (RecordNextStmtCount) {
262       CountMap[S] = PGO.getCurrentRegionCount();
263       RecordNextStmtCount = false;
264     }
265   }
266 
267   /// Set and return the current count.
268   uint64_t setCount(uint64_t Count) {
269     PGO.setCurrentRegionCount(Count);
270     return Count;
271   }
272 
273   void VisitStmt(const Stmt *S) {
274     RecordStmtCount(S);
275     for (Stmt::const_child_range I = S->children(); I; ++I) {
276       if (*I)
277         this->Visit(*I);
278     }
279   }
280 
281   void VisitFunctionDecl(const FunctionDecl *D) {
282     // Counter tracks entry to the function body.
283     uint64_t BodyCount = setCount(PGO.getRegionCount(D->getBody()));
284     CountMap[D->getBody()] = BodyCount;
285     Visit(D->getBody());
286   }
287 
288   // Skip lambda expressions. We visit these as FunctionDecls when we're
289   // generating them and aren't interested in the body when generating a
290   // parent context.
291   void VisitLambdaExpr(const LambdaExpr *LE) {}
292 
293   void VisitCapturedDecl(const CapturedDecl *D) {
294     // Counter tracks entry to the capture body.
295     uint64_t BodyCount = setCount(PGO.getRegionCount(D->getBody()));
296     CountMap[D->getBody()] = BodyCount;
297     Visit(D->getBody());
298   }
299 
300   void VisitObjCMethodDecl(const ObjCMethodDecl *D) {
301     // Counter tracks entry to the method body.
302     uint64_t BodyCount = setCount(PGO.getRegionCount(D->getBody()));
303     CountMap[D->getBody()] = BodyCount;
304     Visit(D->getBody());
305   }
306 
307   void VisitBlockDecl(const BlockDecl *D) {
308     // Counter tracks entry to the block body.
309     uint64_t BodyCount = setCount(PGO.getRegionCount(D->getBody()));
310     CountMap[D->getBody()] = BodyCount;
311     Visit(D->getBody());
312   }
313 
314   void VisitReturnStmt(const ReturnStmt *S) {
315     RecordStmtCount(S);
316     if (S->getRetValue())
317       Visit(S->getRetValue());
318     PGO.setCurrentRegionUnreachable();
319     RecordNextStmtCount = true;
320   }
321 
322   void VisitCXXThrowExpr(const CXXThrowExpr *E) {
323     RecordStmtCount(E);
324     if (E->getSubExpr())
325       Visit(E->getSubExpr());
326     PGO.setCurrentRegionUnreachable();
327     RecordNextStmtCount = true;
328   }
329 
330   void VisitGotoStmt(const GotoStmt *S) {
331     RecordStmtCount(S);
332     PGO.setCurrentRegionUnreachable();
333     RecordNextStmtCount = true;
334   }
335 
336   void VisitLabelStmt(const LabelStmt *S) {
337     RecordNextStmtCount = false;
338     // Counter tracks the block following the label.
339     uint64_t BlockCount = setCount(PGO.getRegionCount(S));
340     CountMap[S] = BlockCount;
341     Visit(S->getSubStmt());
342   }
343 
344   void VisitBreakStmt(const BreakStmt *S) {
345     RecordStmtCount(S);
346     assert(!BreakContinueStack.empty() && "break not in a loop or switch!");
347     BreakContinueStack.back().BreakCount += PGO.getCurrentRegionCount();
348     PGO.setCurrentRegionUnreachable();
349     RecordNextStmtCount = true;
350   }
351 
352   void VisitContinueStmt(const ContinueStmt *S) {
353     RecordStmtCount(S);
354     assert(!BreakContinueStack.empty() && "continue stmt not in a loop!");
355     BreakContinueStack.back().ContinueCount += PGO.getCurrentRegionCount();
356     PGO.setCurrentRegionUnreachable();
357     RecordNextStmtCount = true;
358   }
359 
360   void VisitWhileStmt(const WhileStmt *S) {
361     RecordStmtCount(S);
362     uint64_t ParentCount = PGO.getCurrentRegionCount();
363 
364     BreakContinueStack.push_back(BreakContinue());
365     // Visit the body region first so the break/continue adjustments can be
366     // included when visiting the condition.
367     uint64_t BodyCount = setCount(PGO.getRegionCount(S));
368     CountMap[S->getBody()] = PGO.getCurrentRegionCount();
369     Visit(S->getBody());
370     uint64_t BackedgeCount = PGO.getCurrentRegionCount();
371 
372     // ...then go back and propagate counts through the condition. The count
373     // at the start of the condition is the sum of the incoming edges,
374     // the backedge from the end of the loop body, and the edges from
375     // continue statements.
376     BreakContinue BC = BreakContinueStack.pop_back_val();
377     uint64_t CondCount =
378         setCount(ParentCount + BackedgeCount + BC.ContinueCount);
379     CountMap[S->getCond()] = CondCount;
380     Visit(S->getCond());
381     setCount(BC.BreakCount + CondCount - BodyCount);
382     RecordNextStmtCount = true;
383   }
384 
385   void VisitDoStmt(const DoStmt *S) {
386     RecordStmtCount(S);
387     uint64_t LoopCount = PGO.getRegionCount(S);
388 
389     BreakContinueStack.push_back(BreakContinue());
390     // The count doesn't include the fallthrough from the parent scope. Add it.
391     uint64_t BodyCount = setCount(LoopCount + PGO.getCurrentRegionCount());
392     CountMap[S->getBody()] = BodyCount;
393     Visit(S->getBody());
394     uint64_t BackedgeCount = PGO.getCurrentRegionCount();
395 
396     BreakContinue BC = BreakContinueStack.pop_back_val();
397     // The count at the start of the condition is equal to the count at the
398     // end of the body, plus any continues.
399     uint64_t CondCount = setCount(BackedgeCount + BC.ContinueCount);
400     CountMap[S->getCond()] = CondCount;
401     Visit(S->getCond());
402     setCount(BC.BreakCount + CondCount - LoopCount);
403     RecordNextStmtCount = true;
404   }
405 
406   void VisitForStmt(const ForStmt *S) {
407     RecordStmtCount(S);
408     if (S->getInit())
409       Visit(S->getInit());
410 
411     uint64_t ParentCount = PGO.getCurrentRegionCount();
412 
413     BreakContinueStack.push_back(BreakContinue());
414     // Visit the body region first. (This is basically the same as a while
415     // loop; see further comments in VisitWhileStmt.)
416     uint64_t BodyCount = setCount(PGO.getRegionCount(S));
417     CountMap[S->getBody()] = BodyCount;
418     Visit(S->getBody());
419     uint64_t BackedgeCount = PGO.getCurrentRegionCount();
420     BreakContinue BC = BreakContinueStack.pop_back_val();
421 
422     // The increment is essentially part of the body but it needs to include
423     // the count for all the continue statements.
424     if (S->getInc()) {
425       uint64_t IncCount = setCount(BackedgeCount + BC.ContinueCount);
426       CountMap[S->getInc()] = IncCount;
427       Visit(S->getInc());
428     }
429 
430     // ...then go back and propagate counts through the condition.
431     uint64_t CondCount =
432         setCount(ParentCount + BackedgeCount + BC.ContinueCount);
433     if (S->getCond()) {
434       CountMap[S->getCond()] = CondCount;
435       Visit(S->getCond());
436     }
437     setCount(BC.BreakCount + CondCount - BodyCount);
438     RecordNextStmtCount = true;
439   }
440 
441   void VisitCXXForRangeStmt(const CXXForRangeStmt *S) {
442     RecordStmtCount(S);
443     Visit(S->getLoopVarStmt());
444     Visit(S->getRangeStmt());
445     Visit(S->getBeginEndStmt());
446 
447     uint64_t ParentCount = PGO.getCurrentRegionCount();
448 
449     BreakContinueStack.push_back(BreakContinue());
450     // Visit the body region first. (This is basically the same as a while
451     // loop; see further comments in VisitWhileStmt.)
452     uint64_t BodyCount = setCount(PGO.getRegionCount(S));
453     CountMap[S->getBody()] = BodyCount;
454     Visit(S->getBody());
455     uint64_t BackedgeCount = PGO.getCurrentRegionCount();
456     BreakContinue BC = BreakContinueStack.pop_back_val();
457 
458     // The increment is essentially part of the body but it needs to include
459     // the count for all the continue statements.
460     uint64_t IncCount = setCount(BackedgeCount + BC.ContinueCount);
461     CountMap[S->getInc()] = IncCount;
462     Visit(S->getInc());
463 
464     // ...then go back and propagate counts through the condition.
465     uint64_t CondCount =
466         setCount(ParentCount + BackedgeCount + BC.ContinueCount);
467     CountMap[S->getCond()] = CondCount;
468     Visit(S->getCond());
469     setCount(BC.BreakCount + CondCount - BodyCount);
470     RecordNextStmtCount = true;
471   }
472 
473   void VisitObjCForCollectionStmt(const ObjCForCollectionStmt *S) {
474     RecordStmtCount(S);
475     Visit(S->getElement());
476     uint64_t ParentCount = PGO.getCurrentRegionCount();
477     BreakContinueStack.push_back(BreakContinue());
478     // Counter tracks the body of the loop.
479     uint64_t BodyCount = setCount(PGO.getRegionCount(S));
480     CountMap[S->getBody()] = BodyCount;
481     Visit(S->getBody());
482     uint64_t BackedgeCount = PGO.getCurrentRegionCount();
483     BreakContinue BC = BreakContinueStack.pop_back_val();
484 
485     setCount(BC.BreakCount + ParentCount + BackedgeCount + BC.ContinueCount -
486              BodyCount);
487     RecordNextStmtCount = true;
488   }
489 
490   void VisitSwitchStmt(const SwitchStmt *S) {
491     RecordStmtCount(S);
492     Visit(S->getCond());
493     PGO.setCurrentRegionUnreachable();
494     BreakContinueStack.push_back(BreakContinue());
495     Visit(S->getBody());
496     // If the switch is inside a loop, add the continue counts.
497     BreakContinue BC = BreakContinueStack.pop_back_val();
498     if (!BreakContinueStack.empty())
499       BreakContinueStack.back().ContinueCount += BC.ContinueCount;
500     // Counter tracks the exit block of the switch.
501     setCount(PGO.getRegionCount(S));
502     RecordNextStmtCount = true;
503   }
504 
505   void VisitSwitchCase(const SwitchCase *S) {
506     RecordNextStmtCount = false;
507     // Counter for this particular case. This counts only jumps from the
508     // switch header and does not include fallthrough from the case before
509     // this one.
510     uint64_t CaseCount = PGO.getRegionCount(S);
511     setCount(PGO.getCurrentRegionCount() + CaseCount);
512     // We need the count without fallthrough in the mapping, so it's more useful
513     // for branch probabilities.
514     CountMap[S] = CaseCount;
515     RecordNextStmtCount = true;
516     Visit(S->getSubStmt());
517   }
518 
519   void VisitIfStmt(const IfStmt *S) {
520     RecordStmtCount(S);
521     uint64_t ParentCount = PGO.getCurrentRegionCount();
522     Visit(S->getCond());
523 
524     // Counter tracks the "then" part of an if statement. The count for
525     // the "else" part, if it exists, will be calculated from this counter.
526     uint64_t ThenCount = setCount(PGO.getRegionCount(S));
527     CountMap[S->getThen()] = ThenCount;
528     Visit(S->getThen());
529     uint64_t OutCount = PGO.getCurrentRegionCount();
530 
531     uint64_t ElseCount = ParentCount - ThenCount;
532     if (S->getElse()) {
533       setCount(ElseCount);
534       CountMap[S->getElse()] = ElseCount;
535       Visit(S->getElse());
536       OutCount += PGO.getCurrentRegionCount();
537     } else
538       OutCount += ElseCount;
539     setCount(OutCount);
540     RecordNextStmtCount = true;
541   }
542 
543   void VisitCXXTryStmt(const CXXTryStmt *S) {
544     RecordStmtCount(S);
545     Visit(S->getTryBlock());
546     for (unsigned I = 0, E = S->getNumHandlers(); I < E; ++I)
547       Visit(S->getHandler(I));
548     // Counter tracks the continuation block of the try statement.
549     setCount(PGO.getRegionCount(S));
550     RecordNextStmtCount = true;
551   }
552 
553   void VisitCXXCatchStmt(const CXXCatchStmt *S) {
554     RecordNextStmtCount = false;
555     // Counter tracks the catch statement's handler block.
556     uint64_t CatchCount = setCount(PGO.getRegionCount(S));
557     CountMap[S] = CatchCount;
558     Visit(S->getHandlerBlock());
559   }
560 
561   void VisitAbstractConditionalOperator(const AbstractConditionalOperator *E) {
562     RecordStmtCount(E);
563     uint64_t ParentCount = PGO.getCurrentRegionCount();
564     Visit(E->getCond());
565 
566     // Counter tracks the "true" part of a conditional operator. The
567     // count in the "false" part will be calculated from this counter.
568     uint64_t TrueCount = setCount(PGO.getRegionCount(E));
569     CountMap[E->getTrueExpr()] = TrueCount;
570     Visit(E->getTrueExpr());
571     uint64_t OutCount = PGO.getCurrentRegionCount();
572 
573     uint64_t FalseCount = setCount(ParentCount - TrueCount);
574     CountMap[E->getFalseExpr()] = FalseCount;
575     Visit(E->getFalseExpr());
576     OutCount += PGO.getCurrentRegionCount();
577 
578     setCount(OutCount);
579     RecordNextStmtCount = true;
580   }
581 
582   void VisitBinLAnd(const BinaryOperator *E) {
583     RecordStmtCount(E);
584     uint64_t ParentCount = PGO.getCurrentRegionCount();
585     Visit(E->getLHS());
586     // Counter tracks the right hand side of a logical and operator.
587     uint64_t RHSCount = setCount(PGO.getRegionCount(E));
588     CountMap[E->getRHS()] = RHSCount;
589     Visit(E->getRHS());
590     setCount(ParentCount + RHSCount - PGO.getCurrentRegionCount());
591     RecordNextStmtCount = true;
592   }
593 
594   void VisitBinLOr(const BinaryOperator *E) {
595     RecordStmtCount(E);
596     uint64_t ParentCount = PGO.getCurrentRegionCount();
597     Visit(E->getLHS());
598     // Counter tracks the right hand side of a logical or operator.
599     uint64_t RHSCount = setCount(PGO.getRegionCount(E));
600     CountMap[E->getRHS()] = RHSCount;
601     Visit(E->getRHS());
602     setCount(ParentCount + RHSCount - PGO.getCurrentRegionCount());
603     RecordNextStmtCount = true;
604   }
605 };
606 }
607 
608 void PGOHash::combine(HashType Type) {
609   // Check that we never combine 0 and only have six bits.
610   assert(Type && "Hash is invalid: unexpected type 0");
611   assert(unsigned(Type) < TooBig && "Hash is invalid: too many types");
612 
613   // Pass through MD5 if enough work has built up.
614   if (Count && Count % NumTypesPerWord == 0) {
615     using namespace llvm::support;
616     uint64_t Swapped = endian::byte_swap<uint64_t, little>(Working);
617     MD5.update(llvm::makeArrayRef((uint8_t *)&Swapped, sizeof(Swapped)));
618     Working = 0;
619   }
620 
621   // Accumulate the current type.
622   ++Count;
623   Working = Working << NumBitsPerType | Type;
624 }
625 
626 uint64_t PGOHash::finalize() {
627   // Use Working as the hash directly if we never used MD5.
628   if (Count <= NumTypesPerWord)
629     // No need to byte swap here, since none of the math was endian-dependent.
630     // This number will be byte-swapped as required on endianness transitions,
631     // so we will see the same value on the other side.
632     return Working;
633 
634   // Check for remaining work in Working.
635   if (Working)
636     MD5.update(Working);
637 
638   // Finalize the MD5 and return the hash.
639   llvm::MD5::MD5Result Result;
640   MD5.final(Result);
641   using namespace llvm::support;
642   return endian::read<uint64_t, little, unaligned>(Result);
643 }
644 
645 void CodeGenPGO::checkGlobalDecl(GlobalDecl GD) {
646   // Make sure we only emit coverage mapping for one constructor/destructor.
647   // Clang emits several functions for the constructor and the destructor of
648   // a class. Every function is instrumented, but we only want to provide
649   // coverage for one of them. Because of that we only emit the coverage mapping
650   // for the base constructor/destructor.
651   if ((isa<CXXConstructorDecl>(GD.getDecl()) &&
652        GD.getCtorType() != Ctor_Base) ||
653       (isa<CXXDestructorDecl>(GD.getDecl()) &&
654        GD.getDtorType() != Dtor_Base)) {
655     SkipCoverageMapping = true;
656   }
657 }
658 
659 void CodeGenPGO::assignRegionCounters(const Decl *D, llvm::Function *Fn) {
660   bool InstrumentRegions = CGM.getCodeGenOpts().ProfileInstrGenerate;
661   llvm::IndexedInstrProfReader *PGOReader = CGM.getPGOReader();
662   if (!InstrumentRegions && !PGOReader)
663     return;
664   if (D->isImplicit())
665     return;
666   CGM.ClearUnusedCoverageMapping(D);
667   setFuncName(Fn);
668 
669   mapRegionCounters(D);
670   if (CGM.getCodeGenOpts().CoverageMapping)
671     emitCounterRegionMapping(D);
672   if (PGOReader) {
673     SourceManager &SM = CGM.getContext().getSourceManager();
674     loadRegionCounts(PGOReader, SM.isInMainFile(D->getLocation()));
675     computeRegionCounts(D);
676     applyFunctionAttributes(PGOReader, Fn);
677   }
678 }
679 
680 void CodeGenPGO::mapRegionCounters(const Decl *D) {
681   RegionCounterMap.reset(new llvm::DenseMap<const Stmt *, unsigned>);
682   MapRegionCounters Walker(*RegionCounterMap);
683   if (const FunctionDecl *FD = dyn_cast_or_null<FunctionDecl>(D))
684     Walker.TraverseDecl(const_cast<FunctionDecl *>(FD));
685   else if (const ObjCMethodDecl *MD = dyn_cast_or_null<ObjCMethodDecl>(D))
686     Walker.TraverseDecl(const_cast<ObjCMethodDecl *>(MD));
687   else if (const BlockDecl *BD = dyn_cast_or_null<BlockDecl>(D))
688     Walker.TraverseDecl(const_cast<BlockDecl *>(BD));
689   else if (const CapturedDecl *CD = dyn_cast_or_null<CapturedDecl>(D))
690     Walker.TraverseDecl(const_cast<CapturedDecl *>(CD));
691   assert(Walker.NextCounter > 0 && "no entry counter mapped for decl");
692   NumRegionCounters = Walker.NextCounter;
693   FunctionHash = Walker.Hash.finalize();
694 }
695 
696 void CodeGenPGO::emitCounterRegionMapping(const Decl *D) {
697   if (SkipCoverageMapping)
698     return;
699   // Don't map the functions inside the system headers
700   auto Loc = D->getBody()->getLocStart();
701   if (CGM.getContext().getSourceManager().isInSystemHeader(Loc))
702     return;
703 
704   std::string CoverageMapping;
705   llvm::raw_string_ostream OS(CoverageMapping);
706   CoverageMappingGen MappingGen(*CGM.getCoverageMapping(),
707                                 CGM.getContext().getSourceManager(),
708                                 CGM.getLangOpts(), RegionCounterMap.get());
709   MappingGen.emitCounterMapping(D, OS);
710   OS.flush();
711 
712   if (CoverageMapping.empty())
713     return;
714 
715   CGM.getCoverageMapping()->addFunctionMappingRecord(
716       FuncNameVar, FuncName, FunctionHash, CoverageMapping);
717 }
718 
719 void
720 CodeGenPGO::emitEmptyCounterMapping(const Decl *D, StringRef Name,
721                                     llvm::GlobalValue::LinkageTypes Linkage) {
722   if (SkipCoverageMapping)
723     return;
724   // Don't map the functions inside the system headers
725   auto Loc = D->getBody()->getLocStart();
726   if (CGM.getContext().getSourceManager().isInSystemHeader(Loc))
727     return;
728 
729   std::string CoverageMapping;
730   llvm::raw_string_ostream OS(CoverageMapping);
731   CoverageMappingGen MappingGen(*CGM.getCoverageMapping(),
732                                 CGM.getContext().getSourceManager(),
733                                 CGM.getLangOpts());
734   MappingGen.emitEmptyMapping(D, OS);
735   OS.flush();
736 
737   if (CoverageMapping.empty())
738     return;
739 
740   setFuncName(Name, Linkage);
741   CGM.getCoverageMapping()->addFunctionMappingRecord(
742       FuncNameVar, FuncName, FunctionHash, CoverageMapping);
743 }
744 
745 void CodeGenPGO::computeRegionCounts(const Decl *D) {
746   StmtCountMap.reset(new llvm::DenseMap<const Stmt *, uint64_t>);
747   ComputeRegionCounts Walker(*StmtCountMap, *this);
748   if (const FunctionDecl *FD = dyn_cast_or_null<FunctionDecl>(D))
749     Walker.VisitFunctionDecl(FD);
750   else if (const ObjCMethodDecl *MD = dyn_cast_or_null<ObjCMethodDecl>(D))
751     Walker.VisitObjCMethodDecl(MD);
752   else if (const BlockDecl *BD = dyn_cast_or_null<BlockDecl>(D))
753     Walker.VisitBlockDecl(BD);
754   else if (const CapturedDecl *CD = dyn_cast_or_null<CapturedDecl>(D))
755     Walker.VisitCapturedDecl(const_cast<CapturedDecl *>(CD));
756 }
757 
758 void
759 CodeGenPGO::applyFunctionAttributes(llvm::IndexedInstrProfReader *PGOReader,
760                                     llvm::Function *Fn) {
761   if (!haveRegionCounts())
762     return;
763 
764   uint64_t MaxFunctionCount = PGOReader->getMaximumFunctionCount();
765   uint64_t FunctionCount = getRegionCount(0);
766   if (FunctionCount >= (uint64_t)(0.3 * (double)MaxFunctionCount))
767     // Turn on InlineHint attribute for hot functions.
768     // FIXME: 30% is from preliminary tuning on SPEC, it may not be optimal.
769     Fn->addFnAttr(llvm::Attribute::InlineHint);
770   else if (FunctionCount <= (uint64_t)(0.01 * (double)MaxFunctionCount))
771     // Turn on Cold attribute for cold functions.
772     // FIXME: 1% is from preliminary tuning on SPEC, it may not be optimal.
773     Fn->addFnAttr(llvm::Attribute::Cold);
774 }
775 
776 void CodeGenPGO::emitCounterIncrement(CGBuilderTy &Builder, const Stmt *S) {
777   if (!CGM.getCodeGenOpts().ProfileInstrGenerate || !RegionCounterMap)
778     return;
779   if (!Builder.GetInsertPoint())
780     return;
781 
782   unsigned Counter = (*RegionCounterMap)[S];
783   auto *I8PtrTy = llvm::Type::getInt8PtrTy(CGM.getLLVMContext());
784   Builder.CreateCall4(CGM.getIntrinsic(llvm::Intrinsic::instrprof_increment),
785                       llvm::ConstantExpr::getBitCast(FuncNameVar, I8PtrTy),
786                       Builder.getInt64(FunctionHash),
787                       Builder.getInt32(NumRegionCounters),
788                       Builder.getInt32(Counter));
789 }
790 
791 void CodeGenPGO::loadRegionCounts(llvm::IndexedInstrProfReader *PGOReader,
792                                   bool IsInMainFile) {
793   CGM.getPGOStats().addVisited(IsInMainFile);
794   RegionCounts.clear();
795   if (std::error_code EC =
796           PGOReader->getFunctionCounts(FuncName, FunctionHash, RegionCounts)) {
797     if (EC == llvm::instrprof_error::unknown_function)
798       CGM.getPGOStats().addMissing(IsInMainFile);
799     else if (EC == llvm::instrprof_error::hash_mismatch)
800       CGM.getPGOStats().addMismatched(IsInMainFile);
801     else if (EC == llvm::instrprof_error::malformed)
802       // TODO: Consider a more specific warning for this case.
803       CGM.getPGOStats().addMismatched(IsInMainFile);
804     RegionCounts.clear();
805   }
806 }
807 
808 /// \brief Calculate what to divide by to scale weights.
809 ///
810 /// Given the maximum weight, calculate a divisor that will scale all the
811 /// weights to strictly less than UINT32_MAX.
812 static uint64_t calculateWeightScale(uint64_t MaxWeight) {
813   return MaxWeight < UINT32_MAX ? 1 : MaxWeight / UINT32_MAX + 1;
814 }
815 
816 /// \brief Scale an individual branch weight (and add 1).
817 ///
818 /// Scale a 64-bit weight down to 32-bits using \c Scale.
819 ///
820 /// According to Laplace's Rule of Succession, it is better to compute the
821 /// weight based on the count plus 1, so universally add 1 to the value.
822 ///
823 /// \pre \c Scale was calculated by \a calculateWeightScale() with a weight no
824 /// greater than \c Weight.
825 static uint32_t scaleBranchWeight(uint64_t Weight, uint64_t Scale) {
826   assert(Scale && "scale by 0?");
827   uint64_t Scaled = Weight / Scale + 1;
828   assert(Scaled <= UINT32_MAX && "overflow 32-bits");
829   return Scaled;
830 }
831 
832 llvm::MDNode *CodeGenPGO::createBranchWeights(uint64_t TrueCount,
833                                               uint64_t FalseCount) {
834   // Check for empty weights.
835   if (!TrueCount && !FalseCount)
836     return nullptr;
837 
838   // Calculate how to scale down to 32-bits.
839   uint64_t Scale = calculateWeightScale(std::max(TrueCount, FalseCount));
840 
841   llvm::MDBuilder MDHelper(CGM.getLLVMContext());
842   return MDHelper.createBranchWeights(scaleBranchWeight(TrueCount, Scale),
843                                       scaleBranchWeight(FalseCount, Scale));
844 }
845 
846 llvm::MDNode *CodeGenPGO::createBranchWeights(ArrayRef<uint64_t> Weights) {
847   // We need at least two elements to create meaningful weights.
848   if (Weights.size() < 2)
849     return nullptr;
850 
851   // Check for empty weights.
852   uint64_t MaxWeight = *std::max_element(Weights.begin(), Weights.end());
853   if (MaxWeight == 0)
854     return nullptr;
855 
856   // Calculate how to scale down to 32-bits.
857   uint64_t Scale = calculateWeightScale(MaxWeight);
858 
859   SmallVector<uint32_t, 16> ScaledWeights;
860   ScaledWeights.reserve(Weights.size());
861   for (uint64_t W : Weights)
862     ScaledWeights.push_back(scaleBranchWeight(W, Scale));
863 
864   llvm::MDBuilder MDHelper(CGM.getLLVMContext());
865   return MDHelper.createBranchWeights(ScaledWeights);
866 }
867 
868 llvm::MDNode *CodeGenPGO::createLoopWeights(const Stmt *Cond,
869                                             uint64_t LoopCount) {
870   if (!haveRegionCounts())
871     return nullptr;
872   Optional<uint64_t> CondCount = getStmtCount(Cond);
873   assert(CondCount.hasValue() && "missing expected loop condition count");
874   if (*CondCount == 0)
875     return nullptr;
876   return createBranchWeights(LoopCount,
877                              std::max(*CondCount, LoopCount) - LoopCount);
878 }
879