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