[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