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/MDBuilder.h"
20 #include "llvm/ProfileData/InstrProfReader.h"
21 #include "llvm/Support/Endian.h"
22 #include "llvm/Support/FileSystem.h"
23 #include "llvm/Support/MD5.h"
24 
25 using namespace clang;
26 using namespace CodeGen;
27 
28 void CodeGenPGO::setFuncName(StringRef Name,
29                              llvm::GlobalValue::LinkageTypes Linkage) {
30   StringRef RawFuncName = Name;
31 
32   // Function names may be prefixed with a binary '1' to indicate
33   // that the backend should not modify the symbols due to any platform
34   // naming convention. Do not include that '1' in the PGO profile name.
35   if (RawFuncName[0] == '\1')
36     RawFuncName = RawFuncName.substr(1);
37 
38   FuncName = RawFuncName;
39   if (llvm::GlobalValue::isLocalLinkage(Linkage)) {
40     // For local symbols, prepend the main file name to distinguish them.
41     // Do not include the full path in the file name since there's no guarantee
42     // that it will stay the same, e.g., if the files are checked out from
43     // version control in different locations.
44     if (CGM.getCodeGenOpts().MainFileName.empty())
45       FuncName = FuncName.insert(0, "<unknown>:");
46     else
47       FuncName = FuncName.insert(0, CGM.getCodeGenOpts().MainFileName + ":");
48   }
49 }
50 
51 void CodeGenPGO::setFuncName(llvm::Function *Fn) {
52   setFuncName(Fn->getName(), Fn->getLinkage());
53 }
54 
55 void CodeGenPGO::setVarLinkage(llvm::GlobalValue::LinkageTypes Linkage) {
56   // Set the linkage for variables based on the function linkage.  Usually, we
57   // want to match it, but available_externally and extern_weak both have the
58   // wrong semantics.
59   VarLinkage = Linkage;
60   switch (VarLinkage) {
61   case llvm::GlobalValue::ExternalWeakLinkage:
62     VarLinkage = llvm::GlobalValue::LinkOnceAnyLinkage;
63     break;
64   case llvm::GlobalValue::AvailableExternallyLinkage:
65     VarLinkage = llvm::GlobalValue::LinkOnceODRLinkage;
66     break;
67   default:
68     break;
69   }
70 }
71 
72 static llvm::Function *getRegisterFunc(CodeGenModule &CGM) {
73   return CGM.getModule().getFunction("__llvm_profile_register_functions");
74 }
75 
76 static llvm::BasicBlock *getOrInsertRegisterBB(CodeGenModule &CGM) {
77   // Don't do this for Darwin.  compiler-rt uses linker magic.
78   if (CGM.getTarget().getTriple().isOSDarwin())
79     return nullptr;
80 
81   // Only need to insert this once per module.
82   if (llvm::Function *RegisterF = getRegisterFunc(CGM))
83     return &RegisterF->getEntryBlock();
84 
85   // Construct the function.
86   auto *VoidTy = llvm::Type::getVoidTy(CGM.getLLVMContext());
87   auto *RegisterFTy = llvm::FunctionType::get(VoidTy, false);
88   auto *RegisterF = llvm::Function::Create(RegisterFTy,
89                                            llvm::GlobalValue::InternalLinkage,
90                                            "__llvm_profile_register_functions",
91                                            &CGM.getModule());
92   RegisterF->setUnnamedAddr(true);
93   if (CGM.getCodeGenOpts().DisableRedZone)
94     RegisterF->addFnAttr(llvm::Attribute::NoRedZone);
95 
96   // Construct and return the entry block.
97   auto *BB = llvm::BasicBlock::Create(CGM.getLLVMContext(), "", RegisterF);
98   CGBuilderTy Builder(BB);
99   Builder.CreateRetVoid();
100   return BB;
101 }
102 
103 static llvm::Constant *getOrInsertRuntimeRegister(CodeGenModule &CGM) {
104   auto *VoidTy = llvm::Type::getVoidTy(CGM.getLLVMContext());
105   auto *VoidPtrTy = llvm::Type::getInt8PtrTy(CGM.getLLVMContext());
106   auto *RuntimeRegisterTy = llvm::FunctionType::get(VoidTy, VoidPtrTy, false);
107   return CGM.getModule().getOrInsertFunction("__llvm_profile_register_function",
108                                              RuntimeRegisterTy);
109 }
110 
111 static bool isMachO(const CodeGenModule &CGM) {
112   return CGM.getTarget().getTriple().isOSBinFormatMachO();
113 }
114 
115 static StringRef getCountersSection(const CodeGenModule &CGM) {
116   return isMachO(CGM) ? "__DATA,__llvm_prf_cnts" : "__llvm_prf_cnts";
117 }
118 
119 static StringRef getNameSection(const CodeGenModule &CGM) {
120   return isMachO(CGM) ? "__DATA,__llvm_prf_names" : "__llvm_prf_names";
121 }
122 
123 static StringRef getDataSection(const CodeGenModule &CGM) {
124   return isMachO(CGM) ? "__DATA,__llvm_prf_data" : "__llvm_prf_data";
125 }
126 
127 llvm::GlobalVariable *CodeGenPGO::buildDataVar() {
128   // Create name variable.
129   llvm::LLVMContext &Ctx = CGM.getLLVMContext();
130   auto *VarName = llvm::ConstantDataArray::getString(Ctx, getFuncName(),
131                                                      false);
132   auto *Name = new llvm::GlobalVariable(CGM.getModule(), VarName->getType(),
133                                         true, VarLinkage, VarName,
134                                         getFuncVarName("name"));
135   Name->setSection(getNameSection(CGM));
136   Name->setAlignment(1);
137 
138   // Create data variable.
139   auto *Int32Ty = llvm::Type::getInt32Ty(Ctx);
140   auto *Int64Ty = llvm::Type::getInt64Ty(Ctx);
141   auto *Int8PtrTy = llvm::Type::getInt8PtrTy(Ctx);
142   auto *Int64PtrTy = llvm::Type::getInt64PtrTy(Ctx);
143   llvm::GlobalVariable *Data = nullptr;
144   if (RegionCounters) {
145     llvm::Type *DataTypes[] = {
146       Int32Ty, Int32Ty, Int64Ty, Int8PtrTy, Int64PtrTy
147     };
148     auto *DataTy = llvm::StructType::get(Ctx, makeArrayRef(DataTypes));
149     llvm::Constant *DataVals[] = {
150       llvm::ConstantInt::get(Int32Ty, getFuncName().size()),
151       llvm::ConstantInt::get(Int32Ty, NumRegionCounters),
152       llvm::ConstantInt::get(Int64Ty, FunctionHash),
153       llvm::ConstantExpr::getBitCast(Name, Int8PtrTy),
154       llvm::ConstantExpr::getBitCast(RegionCounters, Int64PtrTy)
155     };
156     Data =
157       new llvm::GlobalVariable(CGM.getModule(), DataTy, true, VarLinkage,
158                                llvm::ConstantStruct::get(DataTy, DataVals),
159                                getFuncVarName("data"));
160 
161     // All the data should be packed into an array in its own section.
162     Data->setSection(getDataSection(CGM));
163     Data->setAlignment(8);
164   }
165 
166   // Create coverage mapping data variable.
167   if (!CoverageMapping.empty())
168     CGM.getCoverageMapping()->addFunctionMappingRecord(Name, getFuncName(),
169                                                        FunctionHash,
170                                                        CoverageMapping);
171 
172   // Hide all these symbols so that we correctly get a copy for each
173   // executable.  The profile format expects names and counters to be
174   // contiguous, so references into shared objects would be invalid.
175   if (!llvm::GlobalValue::isLocalLinkage(VarLinkage)) {
176     Name->setVisibility(llvm::GlobalValue::HiddenVisibility);
177     if (Data) {
178       Data->setVisibility(llvm::GlobalValue::HiddenVisibility);
179       RegionCounters->setVisibility(llvm::GlobalValue::HiddenVisibility);
180     }
181   }
182 
183   // Make sure the data doesn't get deleted.
184   if (Data) CGM.addUsedGlobal(Data);
185   return Data;
186 }
187 
188 void CodeGenPGO::emitInstrumentationData() {
189   if (!RegionCounters)
190     return;
191 
192   // Build the data.
193   auto *Data = buildDataVar();
194 
195   // Register the data.
196   auto *RegisterBB = getOrInsertRegisterBB(CGM);
197   if (!RegisterBB)
198     return;
199   CGBuilderTy Builder(RegisterBB->getTerminator());
200   auto *VoidPtrTy = llvm::Type::getInt8PtrTy(CGM.getLLVMContext());
201   Builder.CreateCall(getOrInsertRuntimeRegister(CGM),
202                      Builder.CreateBitCast(Data, VoidPtrTy));
203 }
204 
205 llvm::Function *CodeGenPGO::emitInitialization(CodeGenModule &CGM) {
206   if (!CGM.getCodeGenOpts().ProfileInstrGenerate)
207     return nullptr;
208 
209   assert(CGM.getModule().getFunction("__llvm_profile_init") == nullptr &&
210          "profile initialization already emitted");
211 
212   // Get the function to call at initialization.
213   llvm::Constant *RegisterF = getRegisterFunc(CGM);
214   if (!RegisterF)
215     return nullptr;
216 
217   // Create the initialization function.
218   auto *VoidTy = llvm::Type::getVoidTy(CGM.getLLVMContext());
219   auto *F = llvm::Function::Create(llvm::FunctionType::get(VoidTy, false),
220                                    llvm::GlobalValue::InternalLinkage,
221                                    "__llvm_profile_init", &CGM.getModule());
222   F->setUnnamedAddr(true);
223   F->addFnAttr(llvm::Attribute::NoInline);
224   if (CGM.getCodeGenOpts().DisableRedZone)
225     F->addFnAttr(llvm::Attribute::NoRedZone);
226 
227   // Add the basic block and the necessary calls.
228   CGBuilderTy Builder(llvm::BasicBlock::Create(CGM.getLLVMContext(), "", F));
229   Builder.CreateCall(RegisterF);
230   Builder.CreateRetVoid();
231 
232   return F;
233 }
234 
235 namespace {
236 /// \brief Stable hasher for PGO region counters.
237 ///
238 /// PGOHash produces a stable hash of a given function's control flow.
239 ///
240 /// Changing the output of this hash will invalidate all previously generated
241 /// profiles -- i.e., don't do it.
242 ///
243 /// \note  When this hash does eventually change (years?), we still need to
244 /// support old hashes.  We'll need to pull in the version number from the
245 /// profile data format and use the matching hash function.
246 class PGOHash {
247   uint64_t Working;
248   unsigned Count;
249   llvm::MD5 MD5;
250 
251   static const int NumBitsPerType = 6;
252   static const unsigned NumTypesPerWord = sizeof(uint64_t) * 8 / NumBitsPerType;
253   static const unsigned TooBig = 1u << NumBitsPerType;
254 
255 public:
256   /// \brief Hash values for AST nodes.
257   ///
258   /// Distinct values for AST nodes that have region counters attached.
259   ///
260   /// These values must be stable.  All new members must be added at the end,
261   /// and no members should be removed.  Changing the enumeration value for an
262   /// AST node will affect the hash of every function that contains that node.
263   enum HashType : unsigned char {
264     None = 0,
265     LabelStmt = 1,
266     WhileStmt,
267     DoStmt,
268     ForStmt,
269     CXXForRangeStmt,
270     ObjCForCollectionStmt,
271     SwitchStmt,
272     CaseStmt,
273     DefaultStmt,
274     IfStmt,
275     CXXTryStmt,
276     CXXCatchStmt,
277     ConditionalOperator,
278     BinaryOperatorLAnd,
279     BinaryOperatorLOr,
280     BinaryConditionalOperator,
281 
282     // Keep this last.  It's for the static assert that follows.
283     LastHashType
284   };
285   static_assert(LastHashType <= TooBig, "Too many types in HashType");
286 
287   // TODO: When this format changes, take in a version number here, and use the
288   // old hash calculation for file formats that used the old hash.
289   PGOHash() : Working(0), Count(0) {}
290   void combine(HashType Type);
291   uint64_t finalize();
292 };
293 const int PGOHash::NumBitsPerType;
294 const unsigned PGOHash::NumTypesPerWord;
295 const unsigned PGOHash::TooBig;
296 
297   /// A RecursiveASTVisitor that fills a map of statements to PGO counters.
298   struct MapRegionCounters : public RecursiveASTVisitor<MapRegionCounters> {
299     /// The next counter value to assign.
300     unsigned NextCounter;
301     /// The function hash.
302     PGOHash Hash;
303     /// The map of statements to counters.
304     llvm::DenseMap<const Stmt *, unsigned> &CounterMap;
305 
306     MapRegionCounters(llvm::DenseMap<const Stmt *, unsigned> &CounterMap)
307         : NextCounter(0), CounterMap(CounterMap) {}
308 
309     // Blocks and lambdas are handled as separate functions, so we need not
310     // traverse them in the parent context.
311     bool TraverseBlockExpr(BlockExpr *BE) { return true; }
312     bool TraverseLambdaBody(LambdaExpr *LE) { return true; }
313     bool TraverseCapturedStmt(CapturedStmt *CS) { return true; }
314 
315     bool VisitDecl(const Decl *D) {
316       switch (D->getKind()) {
317       default:
318         break;
319       case Decl::Function:
320       case Decl::CXXMethod:
321       case Decl::CXXConstructor:
322       case Decl::CXXDestructor:
323       case Decl::CXXConversion:
324       case Decl::ObjCMethod:
325       case Decl::Block:
326       case Decl::Captured:
327         CounterMap[D->getBody()] = NextCounter++;
328         break;
329       }
330       return true;
331     }
332 
333     bool VisitStmt(const Stmt *S) {
334       auto Type = getHashType(S);
335       if (Type == PGOHash::None)
336         return true;
337 
338       CounterMap[S] = NextCounter++;
339       Hash.combine(Type);
340       return true;
341     }
342     PGOHash::HashType getHashType(const Stmt *S) {
343       switch (S->getStmtClass()) {
344       default:
345         break;
346       case Stmt::LabelStmtClass:
347         return PGOHash::LabelStmt;
348       case Stmt::WhileStmtClass:
349         return PGOHash::WhileStmt;
350       case Stmt::DoStmtClass:
351         return PGOHash::DoStmt;
352       case Stmt::ForStmtClass:
353         return PGOHash::ForStmt;
354       case Stmt::CXXForRangeStmtClass:
355         return PGOHash::CXXForRangeStmt;
356       case Stmt::ObjCForCollectionStmtClass:
357         return PGOHash::ObjCForCollectionStmt;
358       case Stmt::SwitchStmtClass:
359         return PGOHash::SwitchStmt;
360       case Stmt::CaseStmtClass:
361         return PGOHash::CaseStmt;
362       case Stmt::DefaultStmtClass:
363         return PGOHash::DefaultStmt;
364       case Stmt::IfStmtClass:
365         return PGOHash::IfStmt;
366       case Stmt::CXXTryStmtClass:
367         return PGOHash::CXXTryStmt;
368       case Stmt::CXXCatchStmtClass:
369         return PGOHash::CXXCatchStmt;
370       case Stmt::ConditionalOperatorClass:
371         return PGOHash::ConditionalOperator;
372       case Stmt::BinaryConditionalOperatorClass:
373         return PGOHash::BinaryConditionalOperator;
374       case Stmt::BinaryOperatorClass: {
375         const BinaryOperator *BO = cast<BinaryOperator>(S);
376         if (BO->getOpcode() == BO_LAnd)
377           return PGOHash::BinaryOperatorLAnd;
378         if (BO->getOpcode() == BO_LOr)
379           return PGOHash::BinaryOperatorLOr;
380         break;
381       }
382       }
383       return PGOHash::None;
384     }
385   };
386 
387   /// A StmtVisitor that propagates the raw counts through the AST and
388   /// records the count at statements where the value may change.
389   struct ComputeRegionCounts : public ConstStmtVisitor<ComputeRegionCounts> {
390     /// PGO state.
391     CodeGenPGO &PGO;
392 
393     /// A flag that is set when the current count should be recorded on the
394     /// next statement, such as at the exit of a loop.
395     bool RecordNextStmtCount;
396 
397     /// The map of statements to count values.
398     llvm::DenseMap<const Stmt *, uint64_t> &CountMap;
399 
400     /// BreakContinueStack - Keep counts of breaks and continues inside loops.
401     struct BreakContinue {
402       uint64_t BreakCount;
403       uint64_t ContinueCount;
404       BreakContinue() : BreakCount(0), ContinueCount(0) {}
405     };
406     SmallVector<BreakContinue, 8> BreakContinueStack;
407 
408     ComputeRegionCounts(llvm::DenseMap<const Stmt *, uint64_t> &CountMap,
409                         CodeGenPGO &PGO)
410         : PGO(PGO), RecordNextStmtCount(false), CountMap(CountMap) {}
411 
412     void RecordStmtCount(const Stmt *S) {
413       if (RecordNextStmtCount) {
414         CountMap[S] = PGO.getCurrentRegionCount();
415         RecordNextStmtCount = false;
416       }
417     }
418 
419     void VisitStmt(const Stmt *S) {
420       RecordStmtCount(S);
421       for (Stmt::const_child_range I = S->children(); I; ++I) {
422         if (*I)
423          this->Visit(*I);
424       }
425     }
426 
427     void VisitFunctionDecl(const FunctionDecl *D) {
428       // Counter tracks entry to the function body.
429       RegionCounter Cnt(PGO, D->getBody());
430       Cnt.beginRegion();
431       CountMap[D->getBody()] = PGO.getCurrentRegionCount();
432       Visit(D->getBody());
433     }
434 
435     // Skip lambda expressions. We visit these as FunctionDecls when we're
436     // generating them and aren't interested in the body when generating a
437     // parent context.
438     void VisitLambdaExpr(const LambdaExpr *LE) {}
439 
440     void VisitCapturedDecl(const CapturedDecl *D) {
441       // Counter tracks entry to the capture body.
442       RegionCounter Cnt(PGO, D->getBody());
443       Cnt.beginRegion();
444       CountMap[D->getBody()] = PGO.getCurrentRegionCount();
445       Visit(D->getBody());
446     }
447 
448     void VisitObjCMethodDecl(const ObjCMethodDecl *D) {
449       // Counter tracks entry to the method body.
450       RegionCounter Cnt(PGO, D->getBody());
451       Cnt.beginRegion();
452       CountMap[D->getBody()] = PGO.getCurrentRegionCount();
453       Visit(D->getBody());
454     }
455 
456     void VisitBlockDecl(const BlockDecl *D) {
457       // Counter tracks entry to the block body.
458       RegionCounter Cnt(PGO, D->getBody());
459       Cnt.beginRegion();
460       CountMap[D->getBody()] = PGO.getCurrentRegionCount();
461       Visit(D->getBody());
462     }
463 
464     void VisitReturnStmt(const ReturnStmt *S) {
465       RecordStmtCount(S);
466       if (S->getRetValue())
467         Visit(S->getRetValue());
468       PGO.setCurrentRegionUnreachable();
469       RecordNextStmtCount = true;
470     }
471 
472     void VisitGotoStmt(const GotoStmt *S) {
473       RecordStmtCount(S);
474       PGO.setCurrentRegionUnreachable();
475       RecordNextStmtCount = true;
476     }
477 
478     void VisitLabelStmt(const LabelStmt *S) {
479       RecordNextStmtCount = false;
480       // Counter tracks the block following the label.
481       RegionCounter Cnt(PGO, S);
482       Cnt.beginRegion();
483       CountMap[S] = PGO.getCurrentRegionCount();
484       Visit(S->getSubStmt());
485     }
486 
487     void VisitBreakStmt(const BreakStmt *S) {
488       RecordStmtCount(S);
489       assert(!BreakContinueStack.empty() && "break not in a loop or switch!");
490       BreakContinueStack.back().BreakCount += PGO.getCurrentRegionCount();
491       PGO.setCurrentRegionUnreachable();
492       RecordNextStmtCount = true;
493     }
494 
495     void VisitContinueStmt(const ContinueStmt *S) {
496       RecordStmtCount(S);
497       assert(!BreakContinueStack.empty() && "continue stmt not in a loop!");
498       BreakContinueStack.back().ContinueCount += PGO.getCurrentRegionCount();
499       PGO.setCurrentRegionUnreachable();
500       RecordNextStmtCount = true;
501     }
502 
503     void VisitWhileStmt(const WhileStmt *S) {
504       RecordStmtCount(S);
505       // Counter tracks the body of the loop.
506       RegionCounter Cnt(PGO, S);
507       BreakContinueStack.push_back(BreakContinue());
508       // Visit the body region first so the break/continue adjustments can be
509       // included when visiting the condition.
510       Cnt.beginRegion();
511       CountMap[S->getBody()] = PGO.getCurrentRegionCount();
512       Visit(S->getBody());
513       Cnt.adjustForControlFlow();
514 
515       // ...then go back and propagate counts through the condition. The count
516       // at the start of the condition is the sum of the incoming edges,
517       // the backedge from the end of the loop body, and the edges from
518       // continue statements.
519       BreakContinue BC = BreakContinueStack.pop_back_val();
520       Cnt.setCurrentRegionCount(Cnt.getParentCount() +
521                                 Cnt.getAdjustedCount() + BC.ContinueCount);
522       CountMap[S->getCond()] = PGO.getCurrentRegionCount();
523       Visit(S->getCond());
524       Cnt.adjustForControlFlow();
525       Cnt.applyAdjustmentsToRegion(BC.BreakCount + BC.ContinueCount);
526       RecordNextStmtCount = true;
527     }
528 
529     void VisitDoStmt(const DoStmt *S) {
530       RecordStmtCount(S);
531       // Counter tracks the body of the loop.
532       RegionCounter Cnt(PGO, S);
533       BreakContinueStack.push_back(BreakContinue());
534       Cnt.beginRegion(/*AddIncomingFallThrough=*/true);
535       CountMap[S->getBody()] = PGO.getCurrentRegionCount();
536       Visit(S->getBody());
537       Cnt.adjustForControlFlow();
538 
539       BreakContinue BC = BreakContinueStack.pop_back_val();
540       // The count at the start of the condition is equal to the count at the
541       // end of the body. The adjusted count does not include either the
542       // fall-through count coming into the loop or the continue count, so add
543       // both of those separately. This is coincidentally the same equation as
544       // with while loops but for different reasons.
545       Cnt.setCurrentRegionCount(Cnt.getParentCount() +
546                                 Cnt.getAdjustedCount() + BC.ContinueCount);
547       CountMap[S->getCond()] = PGO.getCurrentRegionCount();
548       Visit(S->getCond());
549       Cnt.adjustForControlFlow();
550       Cnt.applyAdjustmentsToRegion(BC.BreakCount + BC.ContinueCount);
551       RecordNextStmtCount = true;
552     }
553 
554     void VisitForStmt(const ForStmt *S) {
555       RecordStmtCount(S);
556       if (S->getInit())
557         Visit(S->getInit());
558       // Counter tracks the body of the loop.
559       RegionCounter Cnt(PGO, S);
560       BreakContinueStack.push_back(BreakContinue());
561       // Visit the body region first. (This is basically the same as a while
562       // loop; see further comments in VisitWhileStmt.)
563       Cnt.beginRegion();
564       CountMap[S->getBody()] = PGO.getCurrentRegionCount();
565       Visit(S->getBody());
566       Cnt.adjustForControlFlow();
567 
568       // The increment is essentially part of the body but it needs to include
569       // the count for all the continue statements.
570       if (S->getInc()) {
571         Cnt.setCurrentRegionCount(PGO.getCurrentRegionCount() +
572                                   BreakContinueStack.back().ContinueCount);
573         CountMap[S->getInc()] = PGO.getCurrentRegionCount();
574         Visit(S->getInc());
575         Cnt.adjustForControlFlow();
576       }
577 
578       BreakContinue BC = BreakContinueStack.pop_back_val();
579 
580       // ...then go back and propagate counts through the condition.
581       if (S->getCond()) {
582         Cnt.setCurrentRegionCount(Cnt.getParentCount() +
583                                   Cnt.getAdjustedCount() +
584                                   BC.ContinueCount);
585         CountMap[S->getCond()] = PGO.getCurrentRegionCount();
586         Visit(S->getCond());
587         Cnt.adjustForControlFlow();
588       }
589       Cnt.applyAdjustmentsToRegion(BC.BreakCount + BC.ContinueCount);
590       RecordNextStmtCount = true;
591     }
592 
593     void VisitCXXForRangeStmt(const CXXForRangeStmt *S) {
594       RecordStmtCount(S);
595       Visit(S->getRangeStmt());
596       Visit(S->getBeginEndStmt());
597       // Counter tracks the body of the loop.
598       RegionCounter Cnt(PGO, S);
599       BreakContinueStack.push_back(BreakContinue());
600       // Visit the body region first. (This is basically the same as a while
601       // loop; see further comments in VisitWhileStmt.)
602       Cnt.beginRegion();
603       CountMap[S->getLoopVarStmt()] = PGO.getCurrentRegionCount();
604       Visit(S->getLoopVarStmt());
605       Visit(S->getBody());
606       Cnt.adjustForControlFlow();
607 
608       // The increment is essentially part of the body but it needs to include
609       // the count for all the continue statements.
610       Cnt.setCurrentRegionCount(PGO.getCurrentRegionCount() +
611                                 BreakContinueStack.back().ContinueCount);
612       CountMap[S->getInc()] = PGO.getCurrentRegionCount();
613       Visit(S->getInc());
614       Cnt.adjustForControlFlow();
615 
616       BreakContinue BC = BreakContinueStack.pop_back_val();
617 
618       // ...then go back and propagate counts through the condition.
619       Cnt.setCurrentRegionCount(Cnt.getParentCount() +
620                                 Cnt.getAdjustedCount() +
621                                 BC.ContinueCount);
622       CountMap[S->getCond()] = PGO.getCurrentRegionCount();
623       Visit(S->getCond());
624       Cnt.adjustForControlFlow();
625       Cnt.applyAdjustmentsToRegion(BC.BreakCount + BC.ContinueCount);
626       RecordNextStmtCount = true;
627     }
628 
629     void VisitObjCForCollectionStmt(const ObjCForCollectionStmt *S) {
630       RecordStmtCount(S);
631       Visit(S->getElement());
632       // Counter tracks the body of the loop.
633       RegionCounter Cnt(PGO, S);
634       BreakContinueStack.push_back(BreakContinue());
635       Cnt.beginRegion();
636       CountMap[S->getBody()] = PGO.getCurrentRegionCount();
637       Visit(S->getBody());
638       BreakContinue BC = BreakContinueStack.pop_back_val();
639       Cnt.adjustForControlFlow();
640       Cnt.applyAdjustmentsToRegion(BC.BreakCount + BC.ContinueCount);
641       RecordNextStmtCount = true;
642     }
643 
644     void VisitSwitchStmt(const SwitchStmt *S) {
645       RecordStmtCount(S);
646       Visit(S->getCond());
647       PGO.setCurrentRegionUnreachable();
648       BreakContinueStack.push_back(BreakContinue());
649       Visit(S->getBody());
650       // If the switch is inside a loop, add the continue counts.
651       BreakContinue BC = BreakContinueStack.pop_back_val();
652       if (!BreakContinueStack.empty())
653         BreakContinueStack.back().ContinueCount += BC.ContinueCount;
654       // Counter tracks the exit block of the switch.
655       RegionCounter ExitCnt(PGO, S);
656       ExitCnt.beginRegion();
657       RecordNextStmtCount = true;
658     }
659 
660     void VisitCaseStmt(const CaseStmt *S) {
661       RecordNextStmtCount = false;
662       // Counter for this particular case. This counts only jumps from the
663       // switch header and does not include fallthrough from the case before
664       // this one.
665       RegionCounter Cnt(PGO, S);
666       Cnt.beginRegion(/*AddIncomingFallThrough=*/true);
667       CountMap[S] = Cnt.getCount();
668       RecordNextStmtCount = true;
669       Visit(S->getSubStmt());
670     }
671 
672     void VisitDefaultStmt(const DefaultStmt *S) {
673       RecordNextStmtCount = false;
674       // Counter for this default case. This does not include fallthrough from
675       // the previous case.
676       RegionCounter Cnt(PGO, S);
677       Cnt.beginRegion(/*AddIncomingFallThrough=*/true);
678       CountMap[S] = Cnt.getCount();
679       RecordNextStmtCount = true;
680       Visit(S->getSubStmt());
681     }
682 
683     void VisitIfStmt(const IfStmt *S) {
684       RecordStmtCount(S);
685       // Counter tracks the "then" part of an if statement. The count for
686       // the "else" part, if it exists, will be calculated from this counter.
687       RegionCounter Cnt(PGO, S);
688       Visit(S->getCond());
689 
690       Cnt.beginRegion();
691       CountMap[S->getThen()] = PGO.getCurrentRegionCount();
692       Visit(S->getThen());
693       Cnt.adjustForControlFlow();
694 
695       if (S->getElse()) {
696         Cnt.beginElseRegion();
697         CountMap[S->getElse()] = PGO.getCurrentRegionCount();
698         Visit(S->getElse());
699         Cnt.adjustForControlFlow();
700       }
701       Cnt.applyAdjustmentsToRegion(0);
702       RecordNextStmtCount = true;
703     }
704 
705     void VisitCXXTryStmt(const CXXTryStmt *S) {
706       RecordStmtCount(S);
707       Visit(S->getTryBlock());
708       for (unsigned I = 0, E = S->getNumHandlers(); I < E; ++I)
709         Visit(S->getHandler(I));
710       // Counter tracks the continuation block of the try statement.
711       RegionCounter Cnt(PGO, S);
712       Cnt.beginRegion();
713       RecordNextStmtCount = true;
714     }
715 
716     void VisitCXXCatchStmt(const CXXCatchStmt *S) {
717       RecordNextStmtCount = false;
718       // Counter tracks the catch statement's handler block.
719       RegionCounter Cnt(PGO, S);
720       Cnt.beginRegion();
721       CountMap[S] = PGO.getCurrentRegionCount();
722       Visit(S->getHandlerBlock());
723     }
724 
725     void VisitAbstractConditionalOperator(
726         const AbstractConditionalOperator *E) {
727       RecordStmtCount(E);
728       // Counter tracks the "true" part of a conditional operator. The
729       // count in the "false" part will be calculated from this counter.
730       RegionCounter Cnt(PGO, E);
731       Visit(E->getCond());
732 
733       Cnt.beginRegion();
734       CountMap[E->getTrueExpr()] = PGO.getCurrentRegionCount();
735       Visit(E->getTrueExpr());
736       Cnt.adjustForControlFlow();
737 
738       Cnt.beginElseRegion();
739       CountMap[E->getFalseExpr()] = PGO.getCurrentRegionCount();
740       Visit(E->getFalseExpr());
741       Cnt.adjustForControlFlow();
742 
743       Cnt.applyAdjustmentsToRegion(0);
744       RecordNextStmtCount = true;
745     }
746 
747     void VisitBinLAnd(const BinaryOperator *E) {
748       RecordStmtCount(E);
749       // Counter tracks the right hand side of a logical and operator.
750       RegionCounter Cnt(PGO, E);
751       Visit(E->getLHS());
752       Cnt.beginRegion();
753       CountMap[E->getRHS()] = PGO.getCurrentRegionCount();
754       Visit(E->getRHS());
755       Cnt.adjustForControlFlow();
756       Cnt.applyAdjustmentsToRegion(0);
757       RecordNextStmtCount = true;
758     }
759 
760     void VisitBinLOr(const BinaryOperator *E) {
761       RecordStmtCount(E);
762       // Counter tracks the right hand side of a logical or operator.
763       RegionCounter Cnt(PGO, E);
764       Visit(E->getLHS());
765       Cnt.beginRegion();
766       CountMap[E->getRHS()] = PGO.getCurrentRegionCount();
767       Visit(E->getRHS());
768       Cnt.adjustForControlFlow();
769       Cnt.applyAdjustmentsToRegion(0);
770       RecordNextStmtCount = true;
771     }
772   };
773 }
774 
775 void PGOHash::combine(HashType Type) {
776   // Check that we never combine 0 and only have six bits.
777   assert(Type && "Hash is invalid: unexpected type 0");
778   assert(unsigned(Type) < TooBig && "Hash is invalid: too many types");
779 
780   // Pass through MD5 if enough work has built up.
781   if (Count && Count % NumTypesPerWord == 0) {
782     using namespace llvm::support;
783     uint64_t Swapped = endian::byte_swap<uint64_t, little>(Working);
784     MD5.update(llvm::makeArrayRef((uint8_t *)&Swapped, sizeof(Swapped)));
785     Working = 0;
786   }
787 
788   // Accumulate the current type.
789   ++Count;
790   Working = Working << NumBitsPerType | Type;
791 }
792 
793 uint64_t PGOHash::finalize() {
794   // Use Working as the hash directly if we never used MD5.
795   if (Count <= NumTypesPerWord)
796     // No need to byte swap here, since none of the math was endian-dependent.
797     // This number will be byte-swapped as required on endianness transitions,
798     // so we will see the same value on the other side.
799     return Working;
800 
801   // Check for remaining work in Working.
802   if (Working)
803     MD5.update(Working);
804 
805   // Finalize the MD5 and return the hash.
806   llvm::MD5::MD5Result Result;
807   MD5.final(Result);
808   using namespace llvm::support;
809   return endian::read<uint64_t, little, unaligned>(Result);
810 }
811 
812 static void emitRuntimeHook(CodeGenModule &CGM) {
813   const char *const RuntimeVarName = "__llvm_profile_runtime";
814   const char *const RuntimeUserName = "__llvm_profile_runtime_user";
815   if (CGM.getModule().getGlobalVariable(RuntimeVarName))
816     return;
817 
818   // Declare the runtime hook.
819   llvm::LLVMContext &Ctx = CGM.getLLVMContext();
820   auto *Int32Ty = llvm::Type::getInt32Ty(Ctx);
821   auto *Var = new llvm::GlobalVariable(CGM.getModule(), Int32Ty, false,
822                                        llvm::GlobalValue::ExternalLinkage,
823                                        nullptr, RuntimeVarName);
824 
825   // Make a function that uses it.
826   auto *User = llvm::Function::Create(llvm::FunctionType::get(Int32Ty, false),
827                                       llvm::GlobalValue::LinkOnceODRLinkage,
828                                       RuntimeUserName, &CGM.getModule());
829   User->addFnAttr(llvm::Attribute::NoInline);
830   if (CGM.getCodeGenOpts().DisableRedZone)
831     User->addFnAttr(llvm::Attribute::NoRedZone);
832   CGBuilderTy Builder(llvm::BasicBlock::Create(CGM.getLLVMContext(), "", User));
833   auto *Load = Builder.CreateLoad(Var);
834   Builder.CreateRet(Load);
835 
836   // Create a use of the function.  Now the definition of the runtime variable
837   // should get pulled in, along with any static initializears.
838   CGM.addUsedGlobal(User);
839 }
840 
841 void CodeGenPGO::checkGlobalDecl(GlobalDecl GD) {
842   // Make sure we only emit coverage mapping for one constructor/destructor.
843   // Clang emits several functions for the constructor and the destructor of
844   // a class. Every function is instrumented, but we only want to provide
845   // coverage for one of them. Because of that we only emit the coverage mapping
846   // for the base constructor/destructor.
847   if ((isa<CXXConstructorDecl>(GD.getDecl()) &&
848        GD.getCtorType() != Ctor_Base) ||
849       (isa<CXXDestructorDecl>(GD.getDecl()) &&
850        GD.getDtorType() != Dtor_Base)) {
851     SkipCoverageMapping = true;
852   }
853 }
854 
855 void CodeGenPGO::assignRegionCounters(const Decl *D, llvm::Function *Fn) {
856   bool InstrumentRegions = CGM.getCodeGenOpts().ProfileInstrGenerate;
857   llvm::IndexedInstrProfReader *PGOReader = CGM.getPGOReader();
858   if (!InstrumentRegions && !PGOReader)
859     return;
860   if (D->isImplicit())
861     return;
862   CGM.ClearUnusedCoverageMapping(D);
863   setFuncName(Fn);
864   setVarLinkage(Fn->getLinkage());
865 
866   mapRegionCounters(D);
867   if (InstrumentRegions) {
868     emitRuntimeHook(CGM);
869     emitCounterVariables();
870     if (CGM.getCodeGenOpts().CoverageMapping)
871       emitCounterRegionMapping(D);
872   }
873   if (PGOReader) {
874     SourceManager &SM = CGM.getContext().getSourceManager();
875     loadRegionCounts(PGOReader, SM.isInMainFile(D->getLocation()));
876     computeRegionCounts(D);
877     applyFunctionAttributes(PGOReader, Fn);
878   }
879 }
880 
881 void CodeGenPGO::mapRegionCounters(const Decl *D) {
882   RegionCounterMap.reset(new llvm::DenseMap<const Stmt *, unsigned>);
883   MapRegionCounters Walker(*RegionCounterMap);
884   if (const FunctionDecl *FD = dyn_cast_or_null<FunctionDecl>(D))
885     Walker.TraverseDecl(const_cast<FunctionDecl *>(FD));
886   else if (const ObjCMethodDecl *MD = dyn_cast_or_null<ObjCMethodDecl>(D))
887     Walker.TraverseDecl(const_cast<ObjCMethodDecl *>(MD));
888   else if (const BlockDecl *BD = dyn_cast_or_null<BlockDecl>(D))
889     Walker.TraverseDecl(const_cast<BlockDecl *>(BD));
890   else if (const CapturedDecl *CD = dyn_cast_or_null<CapturedDecl>(D))
891     Walker.TraverseDecl(const_cast<CapturedDecl *>(CD));
892   assert(Walker.NextCounter > 0 && "no entry counter mapped for decl");
893   NumRegionCounters = Walker.NextCounter;
894   FunctionHash = Walker.Hash.finalize();
895 }
896 
897 void CodeGenPGO::emitCounterRegionMapping(const Decl *D) {
898   if (SkipCoverageMapping)
899     return;
900   // Don't map the functions inside the system headers
901   auto Loc = D->getBody()->getLocStart();
902   if (CGM.getContext().getSourceManager().isInSystemHeader(Loc))
903     return;
904 
905   llvm::raw_string_ostream OS(CoverageMapping);
906   CoverageMappingGen MappingGen(*CGM.getCoverageMapping(),
907                                 CGM.getContext().getSourceManager(),
908                                 CGM.getLangOpts(), RegionCounterMap.get());
909   MappingGen.emitCounterMapping(D, OS);
910   OS.flush();
911 }
912 
913 void
914 CodeGenPGO::emitEmptyCounterMapping(const Decl *D, StringRef FuncName,
915                                     llvm::GlobalValue::LinkageTypes Linkage) {
916   if (SkipCoverageMapping)
917     return;
918   setFuncName(FuncName, Linkage);
919   setVarLinkage(Linkage);
920 
921   // Don't map the functions inside the system headers
922   auto Loc = D->getBody()->getLocStart();
923   if (CGM.getContext().getSourceManager().isInSystemHeader(Loc))
924     return;
925 
926   llvm::raw_string_ostream OS(CoverageMapping);
927   CoverageMappingGen MappingGen(*CGM.getCoverageMapping(),
928                                 CGM.getContext().getSourceManager(),
929                                 CGM.getLangOpts());
930   MappingGen.emitEmptyMapping(D, OS);
931   OS.flush();
932   buildDataVar();
933 }
934 
935 void CodeGenPGO::computeRegionCounts(const Decl *D) {
936   StmtCountMap.reset(new llvm::DenseMap<const Stmt *, uint64_t>);
937   ComputeRegionCounts Walker(*StmtCountMap, *this);
938   if (const FunctionDecl *FD = dyn_cast_or_null<FunctionDecl>(D))
939     Walker.VisitFunctionDecl(FD);
940   else if (const ObjCMethodDecl *MD = dyn_cast_or_null<ObjCMethodDecl>(D))
941     Walker.VisitObjCMethodDecl(MD);
942   else if (const BlockDecl *BD = dyn_cast_or_null<BlockDecl>(D))
943     Walker.VisitBlockDecl(BD);
944   else if (const CapturedDecl *CD = dyn_cast_or_null<CapturedDecl>(D))
945     Walker.VisitCapturedDecl(const_cast<CapturedDecl *>(CD));
946 }
947 
948 void
949 CodeGenPGO::applyFunctionAttributes(llvm::IndexedInstrProfReader *PGOReader,
950                                     llvm::Function *Fn) {
951   if (!haveRegionCounts())
952     return;
953 
954   uint64_t MaxFunctionCount = PGOReader->getMaximumFunctionCount();
955   uint64_t FunctionCount = getRegionCount(0);
956   if (FunctionCount >= (uint64_t)(0.3 * (double)MaxFunctionCount))
957     // Turn on InlineHint attribute for hot functions.
958     // FIXME: 30% is from preliminary tuning on SPEC, it may not be optimal.
959     Fn->addFnAttr(llvm::Attribute::InlineHint);
960   else if (FunctionCount <= (uint64_t)(0.01 * (double)MaxFunctionCount))
961     // Turn on Cold attribute for cold functions.
962     // FIXME: 1% is from preliminary tuning on SPEC, it may not be optimal.
963     Fn->addFnAttr(llvm::Attribute::Cold);
964 }
965 
966 void CodeGenPGO::emitCounterVariables() {
967   llvm::LLVMContext &Ctx = CGM.getLLVMContext();
968   llvm::ArrayType *CounterTy = llvm::ArrayType::get(llvm::Type::getInt64Ty(Ctx),
969                                                     NumRegionCounters);
970   RegionCounters =
971     new llvm::GlobalVariable(CGM.getModule(), CounterTy, false, VarLinkage,
972                              llvm::Constant::getNullValue(CounterTy),
973                              getFuncVarName("counters"));
974   RegionCounters->setAlignment(8);
975   RegionCounters->setSection(getCountersSection(CGM));
976 }
977 
978 void CodeGenPGO::emitCounterIncrement(CGBuilderTy &Builder, unsigned Counter) {
979   if (!RegionCounters)
980     return;
981   llvm::Value *Addr =
982     Builder.CreateConstInBoundsGEP2_64(RegionCounters, 0, Counter);
983   llvm::Value *Count = Builder.CreateLoad(Addr, "pgocount");
984   Count = Builder.CreateAdd(Count, Builder.getInt64(1));
985   Builder.CreateStore(Count, Addr);
986 }
987 
988 void CodeGenPGO::loadRegionCounts(llvm::IndexedInstrProfReader *PGOReader,
989                                   bool IsInMainFile) {
990   CGM.getPGOStats().addVisited(IsInMainFile);
991   RegionCounts.clear();
992   if (std::error_code EC = PGOReader->getFunctionCounts(
993           getFuncName(), FunctionHash, RegionCounts)) {
994     if (EC == llvm::instrprof_error::unknown_function)
995       CGM.getPGOStats().addMissing(IsInMainFile);
996     else if (EC == llvm::instrprof_error::hash_mismatch)
997       CGM.getPGOStats().addMismatched(IsInMainFile);
998     else if (EC == llvm::instrprof_error::malformed)
999       // TODO: Consider a more specific warning for this case.
1000       CGM.getPGOStats().addMismatched(IsInMainFile);
1001     RegionCounts.clear();
1002   }
1003 }
1004 
1005 void CodeGenPGO::destroyRegionCounters() {
1006   RegionCounterMap.reset();
1007   StmtCountMap.reset();
1008   RegionCounts.clear();
1009   RegionCounters = nullptr;
1010 }
1011 
1012 /// \brief Calculate what to divide by to scale weights.
1013 ///
1014 /// Given the maximum weight, calculate a divisor that will scale all the
1015 /// weights to strictly less than UINT32_MAX.
1016 static uint64_t calculateWeightScale(uint64_t MaxWeight) {
1017   return MaxWeight < UINT32_MAX ? 1 : MaxWeight / UINT32_MAX + 1;
1018 }
1019 
1020 /// \brief Scale an individual branch weight (and add 1).
1021 ///
1022 /// Scale a 64-bit weight down to 32-bits using \c Scale.
1023 ///
1024 /// According to Laplace's Rule of Succession, it is better to compute the
1025 /// weight based on the count plus 1, so universally add 1 to the value.
1026 ///
1027 /// \pre \c Scale was calculated by \a calculateWeightScale() with a weight no
1028 /// greater than \c Weight.
1029 static uint32_t scaleBranchWeight(uint64_t Weight, uint64_t Scale) {
1030   assert(Scale && "scale by 0?");
1031   uint64_t Scaled = Weight / Scale + 1;
1032   assert(Scaled <= UINT32_MAX && "overflow 32-bits");
1033   return Scaled;
1034 }
1035 
1036 llvm::MDNode *CodeGenPGO::createBranchWeights(uint64_t TrueCount,
1037                                               uint64_t FalseCount) {
1038   // Check for empty weights.
1039   if (!TrueCount && !FalseCount)
1040     return nullptr;
1041 
1042   // Calculate how to scale down to 32-bits.
1043   uint64_t Scale = calculateWeightScale(std::max(TrueCount, FalseCount));
1044 
1045   llvm::MDBuilder MDHelper(CGM.getLLVMContext());
1046   return MDHelper.createBranchWeights(scaleBranchWeight(TrueCount, Scale),
1047                                       scaleBranchWeight(FalseCount, Scale));
1048 }
1049 
1050 llvm::MDNode *CodeGenPGO::createBranchWeights(ArrayRef<uint64_t> Weights) {
1051   // We need at least two elements to create meaningful weights.
1052   if (Weights.size() < 2)
1053     return nullptr;
1054 
1055   // Check for empty weights.
1056   uint64_t MaxWeight = *std::max_element(Weights.begin(), Weights.end());
1057   if (MaxWeight == 0)
1058     return nullptr;
1059 
1060   // Calculate how to scale down to 32-bits.
1061   uint64_t Scale = calculateWeightScale(MaxWeight);
1062 
1063   SmallVector<uint32_t, 16> ScaledWeights;
1064   ScaledWeights.reserve(Weights.size());
1065   for (uint64_t W : Weights)
1066     ScaledWeights.push_back(scaleBranchWeight(W, Scale));
1067 
1068   llvm::MDBuilder MDHelper(CGM.getLLVMContext());
1069   return MDHelper.createBranchWeights(ScaledWeights);
1070 }
1071 
1072 llvm::MDNode *CodeGenPGO::createLoopWeights(const Stmt *Cond,
1073                                             RegionCounter &Cnt) {
1074   if (!haveRegionCounts())
1075     return nullptr;
1076   uint64_t LoopCount = Cnt.getCount();
1077   uint64_t CondCount = 0;
1078   bool Found = getStmtCount(Cond, CondCount);
1079   assert(Found && "missing expected loop condition count");
1080   (void)Found;
1081   if (CondCount == 0)
1082     return nullptr;
1083   return createBranchWeights(LoopCount,
1084                              std::max(CondCount, LoopCount) - LoopCount);
1085 }
1086