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