[llvm] [PGO] Add uniformity profile metadata and gate CHR (PR #221517)

via llvm-commits llvm-commits at lists.llvm.org
Sat Sep 5 21:09:56 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-pgo

Author: Yaxun (Sam) Liu (yxsamliu)

<details>
<summary>Changes</summary>

Device PGO branch weights count lanes, so a divergent branch can look
strongly biased even when each wave executes both paths. Applying CHR in
that case can duplicate control flow and increase register pressure
without reducing wave-level control flow.

Emit presence-only block and branch uniformity metadata during profile
use. Add a function marker to distinguish unprofiled from all-divergent
profiles. Use branch metadata to reject CHR scopes containing a profiled
divergent branch, while keeping the existing behavior for unprofiled
functions and select-only scopes.



---
Full diff: https://github.com/llvm/llvm-project/pull/221517.diff


6 Files Affected:

- (modified) llvm/include/llvm/IR/FixedMetadataKinds.def (+2) 
- (modified) llvm/lib/Transforms/Instrumentation/ControlHeightReduction.cpp (+47-10) 
- (modified) llvm/lib/Transforms/Instrumentation/PGOInstrumentation.cpp (+23-12) 
- (modified) llvm/test/Bitcode/block-uniformity-profile-metadata.ll (+8-10) 
- (added) llvm/test/Transforms/PGOProfile/chr-block-uniformity-profile.ll (+175) 
- (modified) llvm/unittests/Transforms/Instrumentation/PGOInstrumentationTest.cpp (+117) 


``````````diff
diff --git a/llvm/include/llvm/IR/FixedMetadataKinds.def b/llvm/include/llvm/IR/FixedMetadataKinds.def
index fee5e22cd61f4..76e862eae3aee 100644
--- a/llvm/include/llvm/IR/FixedMetadataKinds.def
+++ b/llvm/include/llvm/IR/FixedMetadataKinds.def
@@ -67,3 +67,5 @@ LLVM_FIXED_MD_KIND(MD_mem_cache_hint, "mem.cache_hint", 52)
 LLVM_FIXED_MD_KIND(MD_block_uniformity_profile, "block.uniformity.profile", 53)
 LLVM_FIXED_MD_KIND(MD_callgraph, "callgraph", 54)
 LLVM_FIXED_MD_KIND(MD_metadata_section_kind, "metadata_section_kind", 55)
+LLVM_FIXED_MD_KIND(MD_branch_uniformity_profile, "branch.uniformity.profile",
+                   56)
diff --git a/llvm/lib/Transforms/Instrumentation/ControlHeightReduction.cpp b/llvm/lib/Transforms/Instrumentation/ControlHeightReduction.cpp
index faf0e7debb3c1..1a4983f49bb49 100644
--- a/llvm/lib/Transforms/Instrumentation/ControlHeightReduction.cpp
+++ b/llvm/lib/Transforms/Instrumentation/ControlHeightReduction.cpp
@@ -285,16 +285,18 @@ class CHRScope {
 
 class CHR {
  public:
-  CHR(Function &Fin, BlockFrequencyInfo &BFIin, DominatorTree &DTin,
-      ProfileSummaryInfo &PSIin, RegionInfo &RIin,
-      OptimizationRemarkEmitter &OREin)
-      : F(Fin), BFI(BFIin), DT(DTin), PSI(PSIin), RI(RIin), ORE(OREin) {}
-
-  ~CHR() {
-    for (CHRScope *Scope : Scopes) {
-      delete Scope;
-    }
-  }
+   CHR(Function &Fin, BlockFrequencyInfo &BFIin, DominatorTree &DTin,
+       ProfileSummaryInfo &PSIin, RegionInfo &RIin,
+       OptimizationRemarkEmitter &OREin)
+       : F(Fin), BFI(BFIin), DT(DTin), PSI(PSIin), RI(RIin), ORE(OREin),
+         HasBlockUniformityProfile(
+             F.getMetadata(LLVMContext::MD_block_uniformity_profile)) {}
+
+   ~CHR() {
+     for (CHRScope *Scope : Scopes) {
+       delete Scope;
+     }
+   }
 
   bool run();
 
@@ -376,6 +378,11 @@ class CHR {
   OptimizationRemarkEmitter &ORE;
   CHRStats Stats;
 
+  // A function attachment indicates that block and branch uniformity profile
+  // metadata is available. This preserves CHR's existing behavior when no
+  // uniformity profile was collected.
+  bool HasBlockUniformityProfile;
+
   // All the true-biased regions in the function
   DenseSet<Region *> TrueBiasedRegionsGlobal;
   // All the false-biased regions in the function
@@ -1329,6 +1336,36 @@ static bool hasAtLeastTwoBiasedBranches(CHRScope *Scope) {
 void CHR::filterScopes(SmallVectorImpl<CHRScope *> &Input,
                        SmallVectorImpl<CHRScope *> &Output) {
   for (CHRScope *Scope : Input) {
+    // Branch weights count individual lanes. A divergent branch can therefore
+    // look strongly biased even when every wave executes both paths. Avoid
+    // applying CHR to such a scope because duplicating its control flow may
+    // increase live ranges and register pressure without eliminating control
+    // flow for the wave. Select-only scopes are unaffected because branch
+    // uniformity profile describes only conditional branches.
+    if (HasBlockUniformityProfile) {
+      Instruction *DivergentBranch = nullptr;
+      auto FindDivergentBranch = [&](const DenseSet<Region *> &Regions) {
+        for (Region *R : Regions) {
+          Instruction *Branch = R->getEntry()->getTerminator();
+          if (!Branch->getMetadata(LLVMContext::MD_branch_uniformity_profile)) {
+            DivergentBranch = Branch;
+            return;
+          }
+        }
+      };
+      FindDivergentBranch(Scope->TrueBiasedRegions);
+      if (!DivergentBranch)
+        FindDivergentBranch(Scope->FalseBiasedRegions);
+      if (DivergentBranch) {
+        ORE.emit([&]() {
+          return OptimizationRemarkMissed(DEBUG_TYPE, "DivergentBranch",
+                                          DivergentBranch)
+                 << "Drop scope containing a divergent branch";
+        });
+        continue;
+      }
+    }
+
     // Filter out the ones with only one region and no subs.
     if (!hasAtLeastTwoBiasedBranches(Scope)) {
       CHR_DEBUG(dbgs() << "Filtered out by biased branches truthy-regions "
diff --git a/llvm/lib/Transforms/Instrumentation/PGOInstrumentation.cpp b/llvm/lib/Transforms/Instrumentation/PGOInstrumentation.cpp
index 06ce7525a0fd8..8fe30b94c6a7e 100644
--- a/llvm/lib/Transforms/Instrumentation/PGOInstrumentation.cpp
+++ b/llvm/lib/Transforms/Instrumentation/PGOInstrumentation.cpp
@@ -1788,29 +1788,40 @@ void PGOUseFunc::setBlockUniformityAttribute() {
   if (ProfileRecord.UniformityBits.empty())
     return;
 
-  // Annotate uniformity on each instrumented IR basic block so later codegen
-  // passes (MachineFunction) can consume it without relying on fragile block
-  // numbering heuristics.
-  //
-  // Metadata kind: LLVMContext::MD_block_uniformity_profile
-  // Payload: i1 (true = uniform, false = divergent)
+  // Mark the function as having uniformity profile, then annotate each uniform
+  // instrumented IR basic block so later codegen passes (MachineFunction) can
+  // consume it without relying on fragile block numbering heuristics.
+  // Metadata presence on a terminator means uniform; divergent blocks have no
+  // terminator metadata.
 
   std::vector<BasicBlock *> InstrumentBBs;
   FuncInfo.getInstrumentBBs(InstrumentBBs);
 
   LLVMContext &Ctx = F.getContext();
-  Type *Int1Ty = Type::getInt1Ty(Ctx);
-
+  MDNode *UniformMD = MDNode::get(Ctx, {});
+  F.setMetadata(LLVMContext::MD_block_uniformity_profile, UniformMD);
+  DenseMap<CondBrInst *, bool> BranchUniformity;
   for (size_t I = 0, E = InstrumentBBs.size(); I < E; ++I) {
     BasicBlock *BB = InstrumentBBs[I];
     if (!BB || !BB->getTerminator())
       continue;
     bool IsUniform = ProfileRecord.isBlockUniform(I);
-    auto *MD = MDNode::get(
-        Ctx, ConstantAsMetadata::get(ConstantInt::get(Int1Ty, IsUniform)));
-    BB->getTerminator()->setMetadata(LLVMContext::MD_block_uniformity_profile,
-                                     MD);
+    // A counter placed in a block with a single conditional predecessor also
+    // measures the active lanes on that outgoing edge. Record the branch as
+    // uniform only when every instrumented outgoing edge is uniform.
+    if (BasicBlock *Pred = BB->getSinglePredecessor()) {
+      if (auto *Branch = dyn_cast<CondBrInst>(Pred->getTerminator())) {
+        auto It = BranchUniformity.try_emplace(Branch, true).first;
+        It->second &= IsUniform;
+      }
+    }
+    if (IsUniform)
+      BB->getTerminator()->setMetadata(LLVMContext::MD_block_uniformity_profile,
+                                       UniformMD);
   }
+  for (auto [Branch, IsUniform] : BranchUniformity)
+    if (IsUniform)
+      Branch->setMetadata(LLVMContext::MD_branch_uniformity_profile, UniformMD);
 
   LLVM_DEBUG({
     dbgs() << "PGO: Set block uniformity profile for " << F.getName() << ": ";
diff --git a/llvm/test/Bitcode/block-uniformity-profile-metadata.ll b/llvm/test/Bitcode/block-uniformity-profile-metadata.ll
index b69956fa26622..a98f1cebda69d 100644
--- a/llvm/test/Bitcode/block-uniformity-profile-metadata.ll
+++ b/llvm/test/Bitcode/block-uniformity-profile-metadata.ll
@@ -1,21 +1,19 @@
 ; RUN: llvm-as < %s | llvm-dis | FileCheck %s
 ; RUN: llvm-as < %s | llvm-dis | llvm-as | llvm-dis | FileCheck %s
 
-define void @branch_metadata(i1 %cond) {
-; CHECK-LABEL: define void @branch_metadata(
+define void @branch_metadata(i1 %cond) !block.uniformity.profile !0 {
+; CHECK-LABEL: define void @branch_metadata(i1 %cond) !block.uniformity.profile !0 {
 entry:
-  br i1 %cond, label %uniform, label %divergent, !block.uniformity.profile !0
-; CHECK: br i1 %cond, label %uniform, label %divergent, !block.uniformity.profile !0
+  br i1 %cond, label %uniform, label %divergent, !block.uniformity.profile !0, !branch.uniformity.profile !0
+; CHECK: br i1 %cond, label %uniform, label %divergent, !block.uniformity.profile !0, !branch.uniformity.profile !0
 
 uniform:
   ret void
 
 divergent:
-  br label %uniform, !block.uniformity.profile !1
-; CHECK: br label %uniform, !block.uniformity.profile !1
+  br label %uniform, !block.uniformity.profile !0
+; CHECK: br label %uniform, !block.uniformity.profile !0
 }
 
-; CHECK: !0 = !{i1 true}
-; CHECK: !1 = !{i1 false}
-!0 = !{i1 true}
-!1 = !{i1 false}
+; CHECK: !0 = !{}
+!0 = !{}
diff --git a/llvm/test/Transforms/PGOProfile/chr-block-uniformity-profile.ll b/llvm/test/Transforms/PGOProfile/chr-block-uniformity-profile.ll
new file mode 100644
index 0000000000000..f92c27a1173b7
--- /dev/null
+++ b/llvm/test/Transforms/PGOProfile/chr-block-uniformity-profile.ll
@@ -0,0 +1,175 @@
+; RUN: opt < %s -passes='require<profile-summary>,function(chr,instcombine,simplifycfg)' -S | FileCheck %s
+
+declare void @foo()
+
+; Preserve the existing behavior when a function has no block-uniformity
+; profile.
+define void @no_profile(ptr %ptr) !prof !14 {
+; CHECK-LABEL: define void @no_profile(
+; CHECK: entry.split.nonchr:
+entry:
+  %value = load i32, ptr %ptr
+  %bit0 = and i32 %value, 1
+  %cond0 = icmp eq i32 %bit0, 0
+  br i1 %cond0, label %bb1, label %bb0, !prof !15
+
+bb0:
+  call void @foo()
+  br label %bb1
+
+bb1:
+  %bit1 = and i32 %value, 2
+  %cond1 = icmp eq i32 %bit1, 0
+  br i1 %cond1, label %exit, label %bb2, !prof !15
+
+bb2:
+  call void @foo()
+  br label %exit
+
+exit:
+  ret void
+}
+
+; Uniform branches remain eligible for CHR when block-uniformity profile is
+; present.
+define void @uniform(ptr %ptr) !prof !14 !block.uniformity.profile !16 {
+; CHECK-LABEL: define void @uniform(
+; CHECK: entry.split.nonchr:
+entry:
+  %value = load i32, ptr %ptr
+  %bit0 = and i32 %value, 1
+  %cond0 = icmp eq i32 %bit0, 0
+  br i1 %cond0, label %bb1, label %bb0, !prof !15,
+      !branch.uniformity.profile !16
+
+bb0:
+  call void @foo()
+  br label %bb1
+
+bb1:
+  %bit1 = and i32 %value, 2
+  %cond1 = icmp eq i32 %bit1, 0
+  br i1 %cond1, label %exit, label %bb2, !prof !15,
+      !branch.uniformity.profile !16
+
+bb2:
+  call void @foo()
+  br label %exit
+
+exit:
+  ret void
+}
+
+; Once the function has block-uniformity profile, a branch without branch
+; uniformity metadata is divergent. Do not apply CHR to a scope containing one.
+define void @divergent(ptr %ptr) !prof !14 !block.uniformity.profile !16 {
+; CHECK-LABEL: define void @divergent(
+; CHECK-NOT: split
+; CHECK: ret void
+entry:
+  %value = load i32, ptr %ptr
+  %bit0 = and i32 %value, 1
+  %cond0 = icmp eq i32 %bit0, 0
+  br i1 %cond0, label %bb1, label %bb0, !prof !15,
+      !branch.uniformity.profile !16
+
+bb0:
+  call void @foo()
+  br label %bb1
+
+bb1:
+  %bit1 = and i32 %value, 2
+  %cond1 = icmp eq i32 %bit1, 0
+  br i1 %cond1, label %exit, label %bb2, !prof !15
+
+bb2:
+  call void @foo()
+  br label %exit
+
+exit:
+  ret void
+}
+
+; A divergent scope does not disable a separate uniform scope in the same
+; function.
+define void @per_scope(ptr %uniform_ptr, ptr %divergent_ptr) !prof !14 !block.uniformity.profile !16 {
+; CHECK-LABEL: define void @per_scope(
+; CHECK: entry.split.nonchr:
+; CHECK-NOT: after.uniform.split.nonchr:
+; CHECK: after.uniform:
+entry:
+  %uniform_value = load i32, ptr %uniform_ptr
+  %uniform_bit0 = and i32 %uniform_value, 1
+  %uniform_cond0 = icmp eq i32 %uniform_bit0, 0
+  br i1 %uniform_cond0, label %uniform.bb1, label %uniform.bb0, !prof !15,
+      !branch.uniformity.profile !16
+
+uniform.bb0:
+  call void @foo()
+  br label %uniform.bb1
+
+uniform.bb1:
+  %uniform_bit1 = and i32 %uniform_value, 2
+  %uniform_cond1 = icmp eq i32 %uniform_bit1, 0
+  br i1 %uniform_cond1, label %after.uniform, label %uniform.bb2, !prof !15,
+      !branch.uniformity.profile !16
+
+uniform.bb2:
+  call void @foo()
+  br label %after.uniform
+
+after.uniform:
+  %divergent_value = load i32, ptr %divergent_ptr
+  %divergent_bit0 = and i32 %divergent_value, 1
+  %divergent_cond0 = icmp eq i32 %divergent_bit0, 0
+  br i1 %divergent_cond0, label %divergent.bb1, label %divergent.bb0,
+      !prof !15
+
+divergent.bb0:
+  call void @foo()
+  br label %divergent.bb1
+
+divergent.bb1:
+  %divergent_bit1 = and i32 %divergent_value, 2
+  %divergent_cond1 = icmp eq i32 %divergent_bit1, 0
+  br i1 %divergent_cond1, label %exit, label %divergent.bb2, !prof !15
+
+divergent.bb2:
+  call void @foo()
+  br label %exit
+
+exit:
+  ret void
+}
+
+; Select-only scopes are not described by branch-uniformity profile and keep
+; their existing behavior.
+define i32 @select_only(i32 %value) !prof !14 !block.uniformity.profile !16 {
+; CHECK-LABEL: define i32 @select_only(
+; CHECK: entry.split.nonchr:
+entry:
+  %bit0 = and i32 %value, 1
+  %cond0 = icmp eq i32 %bit0, 0
+  %sum0 = select i1 %cond0, i32 %value, i32 42, !prof !15
+  %bit1 = and i32 %value, 2
+  %cond1 = icmp eq i32 %bit1, 0
+  %sum1 = select i1 %cond1, i32 %sum0, i32 43, !prof !15
+  ret i32 %sum1
+}
+
+!llvm.module.flags = !{!0}
+!0 = !{i32 1, !"ProfileSummary", !1}
+!1 = !{!2, !3, !4, !5, !6, !7, !8, !9}
+!2 = !{!"ProfileFormat", !"InstrProf"}
+!3 = !{!"TotalCount", i64 10000}
+!4 = !{!"MaxCount", i64 10}
+!5 = !{!"MaxInternalCount", i64 1}
+!6 = !{!"MaxFunctionCount", i64 1000}
+!7 = !{!"NumCounts", i64 1}
+!8 = !{!"NumFunctions", i64 1}
+!9 = !{!"DetailedSummary", !10}
+!10 = !{!11}
+!11 = !{i32 999999, i64 1, i32 1}
+!14 = !{!"function_entry_count", i64 100}
+!15 = !{!"branch_weights", i32 0, i32 1}
+!16 = !{}
diff --git a/llvm/unittests/Transforms/Instrumentation/PGOInstrumentationTest.cpp b/llvm/unittests/Transforms/Instrumentation/PGOInstrumentationTest.cpp
index 9ccb13934cbd3..b93efe7839bf6 100644
--- a/llvm/unittests/Transforms/Instrumentation/PGOInstrumentationTest.cpp
+++ b/llvm/unittests/Transforms/Instrumentation/PGOInstrumentationTest.cpp
@@ -8,9 +8,13 @@
 
 #include "llvm/Transforms/Instrumentation/PGOInstrumentation.h"
 #include "llvm/AsmParser/Parser.h"
+#include "llvm/IR/InstIterator.h"
+#include "llvm/IR/IntrinsicInst.h"
 #include "llvm/IR/Module.h"
 #include "llvm/Passes/PassBuilder.h"
 #include "llvm/ProfileData/InstrProf.h"
+#include "llvm/ProfileData/InstrProfWriter.h"
+#include "llvm/Testing/Support/Error.h"
 
 #include "gmock/gmock.h"
 #include "gtest/gtest.h"
@@ -187,4 +191,117 @@ TEST_P(PGOInstrumentationGenTest, Instrumented) {
   EXPECT_FALSE(IRInstrVar->isDeclaration());
 }
 
+TEST(PGOInstrumentationUseTest, BlockUniformityMetadataUsesPresence) {
+  static constexpr StringRef Code = R"(
+    define i32 @f(i1 %cond) {
+    entry:
+      br i1 %cond, label %then, label %else
+    then:
+      ret i32 1
+    else:
+      ret i32 0
+    })";
+
+  for (bool IsUniform : {false, true}) {
+    LLVMContext Context;
+    SMDiagnostic ParseError;
+    auto GenModule = parseAssemblyString(Code, ParseError, Context);
+    ASSERT_THAT(GenModule, NotNull());
+
+    PassBuilder PB;
+    LoopAnalysisManager LAM;
+    FunctionAnalysisManager FAM;
+    CGSCCAnalysisManager CGAM;
+    ModuleAnalysisManager MAM;
+    PB.registerModuleAnalyses(MAM);
+    PB.registerCGSCCAnalyses(CGAM);
+    PB.registerFunctionAnalyses(FAM);
+    PB.registerLoopAnalyses(LAM);
+    PB.crossRegisterProxies(LAM, FAM, CGAM, MAM);
+
+    ModulePassManager GenMPM;
+    GenMPM.addPass(PGOInstrumentationGen());
+    GenMPM.run(*GenModule, MAM);
+
+    Function *GenFunction = GenModule->getFunction("f");
+    ASSERT_THAT(GenFunction, NotNull());
+
+    uint64_t FunctionHash = 0;
+    unsigned NumCounters = 0;
+    std::vector<std::string> BlockNames;
+    for (Instruction &I : instructions(*GenFunction)) {
+      auto *Counter = dyn_cast<InstrProfCntrInstBase>(&I);
+      if (!Counter)
+        continue;
+
+      if (BlockNames.empty()) {
+        FunctionHash = Counter->getHash()->getZExtValue();
+        NumCounters = Counter->getNumCounters()->getZExtValue();
+        BlockNames.resize(NumCounters);
+      }
+
+      unsigned Index = Counter->getIndex()->getZExtValue();
+      ASSERT_LT(Index, NumCounters);
+      BlockNames[Index] = I.getParent()->getName().str();
+    }
+
+    ASSERT_GE(NumCounters, 2u);
+    for (const std::string &BlockName : BlockNames)
+      ASSERT_FALSE(BlockName.empty());
+
+    std::vector<uint8_t> UniformityBits((NumCounters + 7) / 8,
+                                        IsUniform ? 0xFF : 0);
+    std::string ProfileName = getIRPGOFuncName(*GenFunction);
+    NamedInstrProfRecord Record(ProfileName, FunctionHash,
+                                std::vector<uint64_t>(NumCounters, 10));
+    Record.UniformityBits = std::move(UniformityBits);
+
+    InstrProfWriter Writer;
+    ASSERT_THAT_ERROR(Writer.mergeProfileKind(InstrProfKind::IRInstrumentation),
+                      Succeeded());
+    Writer.addRecord(std::move(Record),
+                     [](Error E) { ADD_FAILURE() << toString(std::move(E)); });
+    auto Profile = Writer.writeBuffer();
+    ASSERT_THAT(Profile, NotNull());
+
+    auto FS = makeIntrusiveRefCnt<vfs::InMemoryFileSystem>();
+    ASSERT_TRUE(FS->addFile("/profile.profdata", 0, std::move(Profile)));
+
+    auto UseModule = parseAssemblyString(Code, ParseError, Context);
+    ASSERT_THAT(UseModule, NotNull());
+    ModulePassManager UseMPM;
+    UseMPM.addPass(PGOInstrumentationUse("/profile.profdata", "", false, FS));
+    UseMPM.run(*UseModule, MAM);
+
+    Function *UseFunction = UseModule->getFunction("f");
+    ASSERT_THAT(UseFunction, NotNull());
+    MDNode *FunctionMD =
+        UseFunction->getMetadata(LLVMContext::MD_block_uniformity_profile);
+    ASSERT_THAT(FunctionMD, NotNull());
+    EXPECT_EQ(FunctionMD->getNumOperands(), 0u);
+
+    auto *Branch =
+        cast<CondBrInst>(UseFunction->getEntryBlock().getTerminator());
+    MDNode *BranchMD =
+        Branch->getMetadata(LLVMContext::MD_branch_uniformity_profile);
+    EXPECT_EQ(BranchMD != nullptr, IsUniform);
+    if (BranchMD)
+      EXPECT_EQ(BranchMD->getNumOperands(), 0u);
+
+    for (unsigned I = 0; I < NumCounters; ++I) {
+      BasicBlock *BB = nullptr;
+      for (BasicBlock &Candidate : *UseFunction)
+        if (Candidate.getName() == BlockNames[I])
+          BB = &Candidate;
+      ASSERT_THAT(BB, NotNull());
+
+      MDNode *MD = BB->getTerminator()->getMetadata(
+          LLVMContext::MD_block_uniformity_profile);
+      EXPECT_EQ(MD != nullptr, IsUniform);
+      if (MD)
+        EXPECT_EQ(MD->getNumOperands(), 0u);
+    }
+  }
+}
+
 } // end anonymous namespace

``````````

</details>


https://github.com/llvm/llvm-project/pull/221517


More information about the llvm-commits mailing list