[llvm-branch-commits] [llvm] [IR] Define GPU wave-profile metadata (PR #225837)
Yaxun Liu via llvm-branch-commits
llvm-branch-commits at lists.llvm.org
Mon Sep 28 18:02:51 PDT 2026
https://github.com/yxsamliu updated https://github.com/llvm/llvm-project/pull/225837
>From 37656e78fb9b5ee5ac3c942fb2f8082cdc263829 Mon Sep 17 00:00:00 2001
From: "Yaxun (Sam) Liu" <yaxun.liu at amd.com>
Date: Wed, 23 Sep 2026 11:06:42 -0400
Subject: [PATCH] [IR] Define GPU wave-profile metadata
Existing offload GPU profile counters measure lane executions, while
GPU instructions execute at wave granularity under an active-lane
mask. A block visited by every wave can therefore look cold when only
a few lanes are active. Scalar branch weights also cannot represent a
divergent wave visiting both successors before reconverging. These
profiles are a poor fit for optimizations that estimate work performed
by a wave.
Lane counters remain useful for measuring per-lane branch selectivity
and estimating work that scales with the number of active lanes.
Wave counters cannot replace them: a visit with one active lane and a
visit with all lanes active both count as one. The two profiles provide
complementary information about instruction execution and lane activity.
Introduce wave.profile and wave.profile.block metadata to represent
measured wave visits to IR blocks. Each dynamic visit with at least one
active lane contributes one to the count. A measured zero is distinct
from an unavailable count. Wave counts do not obey scalar flow
conservation: one wave may visit both sides of a branch and reconverge
once. The original entry count provides a stable normalization point.
These counts are profitability hints, not proofs of uniform execution
or transformation legality.
Spill placement and register-allocation heuristics can use wave counts
to estimate how often spill and reload instructions execute.
Code-motion and if-conversion cost models can also use them to compare
work across divergent paths, while retaining lane counts for costs
that depend on lane activity.
Store counts in a function-level table, with block identities and
recorded successors attached to terminators. Add verification and
shared helpers to read, write and preserve the metadata, rejecting
stale or ambiguous mappings. Optimizer integration is left to
subsequent changes.
---
llvm/docs/LangRef.md | 47 +
llvm/include/llvm/IR/FixedMetadataKinds.def | 2 +
llvm/include/llvm/IR/ProfDataUtils.h | 74 ++
llvm/lib/IR/Instructions.cpp | 1 +
llvm/lib/IR/ProfDataUtils.cpp | 396 +++++++++
llvm/lib/IR/Verifier.cpp | 30 +
llvm/test/Bitcode/wave-profile.ll | 40 +
.../Transforms/InstCombine/wave-profile.ll | 26 +
llvm/test/Verifier/wave-profile.ll | 91 ++
llvm/unittests/IR/CMakeLists.txt | 1 +
llvm/unittests/IR/ProfDataUtilsTest.cpp | 817 ++++++++++++++++++
11 files changed, 1525 insertions(+)
create mode 100644 llvm/test/Bitcode/wave-profile.ll
create mode 100644 llvm/test/Transforms/InstCombine/wave-profile.ll
create mode 100644 llvm/test/Verifier/wave-profile.ll
create mode 100644 llvm/unittests/IR/ProfDataUtilsTest.cpp
diff --git a/llvm/docs/LangRef.md b/llvm/docs/LangRef.md
index 12e0942aeb4b2..2d5ed2f363917 100644
--- a/llvm/docs/LangRef.md
+++ b/llvm/docs/LangRef.md
@@ -7897,6 +7897,53 @@ section is not marked as readable or writable and it uses the section flag
!0 = !{}
```
+(md_wave_profile)=
+
+#### '`wave.profile`' Metadata
+
+`wave.profile` records measured GPU wave visits on a function definition.
+A divergent wave can visit both successors and reconverge once, so these
+counts do not obey scalar flow conservation. They are profiling hints for
+profitability, not guarantees that justify correctness transformations.
+
+The function node contains `i64` operands: format version 2, a function
+identity given by the XXH3 64-bit hash of its complete name (including any
+promotion or specialization suffix), and a table of unsigned wave counts
+indexed by block identity. Renaming invalidates the profile. Identity zero
+holds the measured original entry count.
+A producer must omit the profile if that entry count is unavailable; it must
+not infer wave counts from lane counts or branch weights.
+
+Each profiled block's terminator has `wave.profile.block` metadata with
+`i64` operands: version 2, the function identity, the block identity, a
+zero-or-one measured flag, and the identities of its successors in order.
+A measured zero is an observation. A zero measured flag marks a missing or
+invalidated count; consumers must ignore the table entry in that case. Blocks
+without measurements use a zero placeholder, including blocks without
+instrumentation and blocks created by later transformations.
+
+Block identities refer to the original count table, independent of block
+layout. Extraction checks the function identity, unique block identities,
+and recorded edges. Exact extraction requires a complete mapping. Partial
+extraction retains unambiguous blocks and invalidates the source and old and
+new targets of a changed edge. Missing or duplicated identities also
+invalidate affected blocks. Unsupported versions, conflicting function
+records, and stale mappings remain valid IR but must not supply usable counts.
+
+A transform may preserve counts only for the same execution events. It may
+refresh the edge snapshot after validating the incoming mapping, keeping
+original identities and assigning unmeasured identities to new blocks. It
+must not restore already invalid counts or transfer a measured count to a
+block that gains executions, such as a new loop header. Transforms may drop
+the metadata. Cloning or inlining does not establish a valid count mapping.
+
+The original entry count remains the normalization anchor if that block is
+removed. It is separate from the mapped count of the current entry block.
+Consumers must handle missing or unmeasured data and a zero normalization
+count, and must not assume an IR count applies to every machine block
+generated from that IR block. Wave visits do not by themselves measure
+memory traffic, occupancy, or the complete cost of a spill.
+
(md_uniformity_profile)=
#### '`uniformity.profile`' Metadata
diff --git a/llvm/include/llvm/IR/FixedMetadataKinds.def b/llvm/include/llvm/IR/FixedMetadataKinds.def
index 2f2a4d5926876..c666a3c07419a 100644
--- a/llvm/include/llvm/IR/FixedMetadataKinds.def
+++ b/llvm/include/llvm/IR/FixedMetadataKinds.def
@@ -71,3 +71,5 @@ LLVM_FIXED_MD_KIND(MD_atomic_ignore_denormal_mode, "atomic.ignore.denormal.mode"
LLVM_FIXED_MD_KIND(MD_branch_uniformity_profile, "branch.uniformity.profile",
57)
LLVM_FIXED_MD_KIND(MD_uniformity_profile, "uniformity.profile", 58)
+LLVM_FIXED_MD_KIND(MD_wave_profile, "wave.profile", 59)
+LLVM_FIXED_MD_KIND(MD_wave_profile_block, "wave.profile.block", 60)
diff --git a/llvm/include/llvm/IR/ProfDataUtils.h b/llvm/include/llvm/IR/ProfDataUtils.h
index 2d77ef33e7056..c87467ae3fff2 100644
--- a/llvm/include/llvm/IR/ProfDataUtils.h
+++ b/llvm/include/llvm/IR/ProfDataUtils.h
@@ -15,15 +15,19 @@
#ifndef LLVM_IR_PROFDATAUTILS_H
#define LLVM_IR_PROFDATAUTILS_H
+#include "llvm/ADT/BitVector.h"
#include "llvm/ADT/STLFunctionalExtras.h"
#include "llvm/ADT/SmallVector.h"
#include "llvm/IR/Metadata.h"
+#include "llvm/IR/TrackingMDRef.h"
+#include "llvm/IR/ValueHandle.h"
#include "llvm/Support/CommandLine.h"
#include "llvm/Support/Compiler.h"
#include <cstddef>
#include <type_traits>
namespace llvm {
+class CondBrInst;
struct MDProfLabels {
LLVM_ABI static const char *BranchWeights;
LLVM_ABI static const char *ValueProfile;
@@ -35,6 +39,76 @@ struct MDProfLabels {
extern LLVM_ABI cl::opt<bool> ProfcheckDisableMetadataFixes;
+/// Attach direct wave visits to stable block identities. These counts do not
+/// obey scalar flow conservation. The original entry count must be measured.
+LLVM_ABI void setBlockWaveCounts(Function &F, ArrayRef<uint64_t> Counts);
+
+/// Attach direct wave visits and identify which blocks have measured counts.
+/// Unmeasured blocks use a zero placeholder in \p Counts and must be ignored by
+/// consumers.
+LLVM_ABI void setBlockWaveCounts(Function &F, ArrayRef<uint64_t> Counts,
+ const BitVector &HasCounts);
+
+/// Remove direct wave visits and their block identities.
+LLVM_ABI void clearBlockWaveCounts(Function &F);
+
+/// Extract wave visits in current function block order when every recorded
+/// block identity and control-flow edge still matches. If \p HasCounts is
+/// provided, it identifies measured blocks; otherwise profiles containing an
+/// unmeasured block are rejected. Clear outputs and return false for missing,
+/// stale, or unsupported metadata.
+LLVM_ABI bool extractBlockWaveCounts(const Function &F,
+ SmallVectorImpl<uint64_t> &Counts,
+ BitVector *HasCounts = nullptr);
+
+/// Extract wave visits for the subset of current blocks whose recorded
+/// identity and local control flow remain unambiguous. Missing, duplicated, or
+/// redirected blocks are returned as unmeasured instead of rejecting the
+/// complete function profile. \p EntryCount remains valid as the function
+/// invocation count even if the original entry block no longer exists. It is
+/// separate from the current entry block's mapped count and is required to
+/// normalize the other block counts.
+LLVM_ABI bool extractMappedBlockWaveCounts(const Function &F,
+ SmallVectorImpl<uint64_t> &Counts,
+ BitVector &HasCounts,
+ uint64_t &EntryCount);
+
+/// Update the recorded successor order after swapping a conditional branch's
+/// successors. Ignore malformed or unsupported block metadata.
+LLVM_ABI void swapBlockWaveCountSuccessors(CondBrInst &BI);
+
+/// Preserve validated wave counts across a CFG rewrite that retains the
+/// execution events of existing blocks. Call invalidate() for a surviving
+/// block whose executions change, then restore() after completing the rewrite.
+/// New blocks acquire unmeasured identities; previously invalid counts are
+/// never made valid. The original normalization count and identity space remain
+/// intact, including when the original entry has disappeared. A snapshot is
+/// consumed by restore(). Later helper updates and direct metadata edits take
+/// precedence; if restoration conflicts with pending invalidations, the whole
+/// profile is cleared. Missing metadata on replacement terminators is restored.
+/// Call invalidate() before discarding an independently edited block record;
+/// restoration cannot observe an attachment that has already been removed.
+class LLVM_ABI BlockWaveCountPreserver {
+ struct BlockProfile {
+ WeakVH Block;
+ WeakVH Terminator;
+ TrackingMDNodeRef Metadata;
+ SmallVector<llvm::Metadata *, 6> Operands;
+ unsigned Id;
+ bool HasCount;
+ };
+ Function &F;
+ TrackingMDNodeRef Profile;
+ SmallVector<Metadata *, 8> ProfileOperands;
+ SmallVector<BlockProfile> Blocks;
+ bool Invalidated = false;
+
+public:
+ explicit BlockWaveCountPreserver(Function &F);
+ void invalidate(const BasicBlock &BB);
+ void restore();
+};
+
/// Profile-based loop metadata that should be accessed only by using
/// \c llvm::getLoopEstimatedTripCount and \c llvm::setLoopEstimatedTripCount.
LLVM_ABI extern const char *LLVMLoopEstimatedTripCount;
diff --git a/llvm/lib/IR/Instructions.cpp b/llvm/lib/IR/Instructions.cpp
index 325850b0a880d..1fe095e9bdae6 100644
--- a/llvm/lib/IR/Instructions.cpp
+++ b/llvm/lib/IR/Instructions.cpp
@@ -1263,6 +1263,7 @@ void CondBrInst::swapSuccessors() {
// Update profile metadata if present and it matches our structural
// expectations.
swapProfMetadata();
+ swapBlockWaveCountSuccessors(*this);
}
//===----------------------------------------------------------------------===//
diff --git a/llvm/lib/IR/ProfDataUtils.cpp b/llvm/lib/IR/ProfDataUtils.cpp
index 34d46cb062bc3..c86fae903cde4 100644
--- a/llvm/lib/IR/ProfDataUtils.cpp
+++ b/llvm/lib/IR/ProfDataUtils.cpp
@@ -12,9 +12,11 @@
#include "llvm/IR/ProfDataUtils.h"
+#include "llvm/ADT/DenseMap.h"
#include "llvm/ADT/STLExtras.h"
#include "llvm/ADT/STLFunctionalExtras.h"
#include "llvm/ADT/SmallVector.h"
+#include "llvm/IR/CFG.h"
#include "llvm/IR/Constants.h"
#include "llvm/IR/Function.h"
#include "llvm/IR/Instructions.h"
@@ -22,6 +24,8 @@
#include "llvm/IR/MDBuilder.h"
#include "llvm/IR/Metadata.h"
#include "llvm/Support/CommandLine.h"
+#include "llvm/Support/xxhash.h"
+#include <optional>
using namespace llvm;
@@ -29,6 +33,398 @@ namespace llvm {
extern cl::opt<bool> ProfcheckDisableMetadataFixes;
}
+static uint64_t getWaveProfileFunctionId(const Function &F) {
+ return xxh3_64bits(F.getName());
+}
+
+static const MDNode *getWaveProfileTable(const Function &F) {
+ SmallVector<MDNode *> Profiles;
+ F.getMetadata(LLVMContext::MD_wave_profile, Profiles);
+ if (F.isDeclaration() || Profiles.size() != 1)
+ return nullptr;
+ const MDNode *MD = Profiles.front();
+ if (MD->getNumOperands() < 3)
+ return nullptr;
+ for (const MDOperand &Op : MD->operands()) {
+ const auto *CI = mdconst::dyn_extract_or_null<ConstantInt>(Op);
+ if (!CI || !CI->getType()->isIntegerTy(64))
+ return nullptr;
+ }
+ if (mdconst::extract<ConstantInt>(MD->getOperand(0))->getZExtValue() != 2 ||
+ mdconst::extract<ConstantInt>(MD->getOperand(1))->getZExtValue() !=
+ getWaveProfileFunctionId(F))
+ return nullptr;
+ return MD;
+}
+
+namespace {
+struct BlockWaveProfile {
+ uint64_t FunctionId;
+ uint64_t Id;
+ bool HasCount;
+ SmallVector<uint64_t, 2> Successors;
+};
+} // namespace
+
+static std::optional<BlockWaveProfile> getBlockWaveProfile(const MDNode *MD) {
+ if (!MD || MD->getNumOperands() < 4)
+ return std::nullopt;
+ for (const MDOperand &Op : MD->operands()) {
+ const auto *CI = mdconst::dyn_extract_or_null<ConstantInt>(Op);
+ if (!CI || !CI->getType()->isIntegerTy(64))
+ return std::nullopt;
+ }
+ auto GetInt = [&](unsigned I) {
+ return mdconst::extract<ConstantInt>(MD->getOperand(I))->getZExtValue();
+ };
+ if (GetInt(0) != 2 || GetInt(3) > 1)
+ return std::nullopt;
+ BlockWaveProfile Profile{GetInt(1), GetInt(2), GetInt(3) != 0, {}};
+ for (unsigned I = 4, E = MD->getNumOperands(); I != E; ++I)
+ Profile.Successors.push_back(GetInt(I));
+ return Profile;
+}
+
+void llvm::swapBlockWaveCountSuccessors(CondBrInst &BI) {
+ MDNode *MD = BI.getMetadata(LLVMContext::MD_wave_profile_block);
+ std::optional<BlockWaveProfile> Profile = getBlockWaveProfile(MD);
+ if (!Profile || Profile->Successors.size() != 2)
+ return;
+ SmallVector<Metadata *> Ops(MD->op_begin(), MD->op_end());
+ std::swap(Ops[4], Ops[5]);
+ BI.setMetadata(LLVMContext::MD_wave_profile_block,
+ MDNode::get(BI.getContext(), Ops));
+}
+
+void llvm::clearBlockWaveCounts(Function &F) {
+ F.setMetadata(LLVMContext::MD_wave_profile, nullptr);
+ for (BasicBlock &BB : F)
+ BB.getTerminator()->setMetadata(LLVMContext::MD_wave_profile_block,
+ nullptr);
+}
+
+static Metadata *getWaveProfileInt(LLVMContext &Context, uint64_t Value) {
+ return ConstantAsMetadata::get(
+ ConstantInt::get(Type::getInt64Ty(Context), Value));
+}
+
+static void writeWaveProfile(Function &F, ArrayRef<Metadata *> Table,
+ const DenseMap<const BasicBlock *, unsigned> &Ids,
+ const BitVector &HasCounts) {
+ LLVMContext &Context = F.getContext();
+ for (auto [Index, BB] : enumerate(F)) {
+ SmallVector<Metadata *> Ops{Table[0], Table[1],
+ getWaveProfileInt(Context, Ids.lookup(&BB)),
+ getWaveProfileInt(Context, HasCounts[Index])};
+ for (const BasicBlock *Succ : successors(&BB))
+ Ops.push_back(getWaveProfileInt(Context, Ids.lookup(Succ)));
+ BB.getTerminator()->setMetadata(LLVMContext::MD_wave_profile_block,
+ MDNode::get(Context, Ops));
+ }
+ // Every update invalidates older snapshots, even if only block records
+ // changed and their terminators are subsequently removed.
+ F.setMetadata(LLVMContext::MD_wave_profile,
+ MDNode::getDistinct(Context, Table));
+}
+
+void llvm::setBlockWaveCounts(Function &F, ArrayRef<uint64_t> Counts) {
+ BitVector HasCounts(Counts.size(), true);
+ setBlockWaveCounts(F, Counts, HasCounts);
+}
+
+void llvm::setBlockWaveCounts(Function &F, ArrayRef<uint64_t> Counts,
+ const BitVector &HasCounts) {
+ assert(Counts.size() == F.size() && "one wave counter per IR block");
+ assert(HasCounts.size() == F.size() &&
+ "one wave-count validity bit per IR block");
+ assert(!Counts.empty() && HasCounts.test(0) &&
+ "entry wave count must be measured");
+ clearBlockWaveCounts(F);
+
+ LLVMContext &Context = F.getContext();
+ SmallVector<Metadata *> Table{
+ getWaveProfileInt(Context, 2),
+ getWaveProfileInt(Context, getWaveProfileFunctionId(F))};
+ for (uint64_t Count : Counts)
+ Table.push_back(getWaveProfileInt(Context, Count));
+
+ DenseMap<const BasicBlock *, unsigned> Ids;
+ for (const BasicBlock &BB : F)
+ Ids.try_emplace(&BB, Ids.size());
+ writeWaveProfile(F, Table, Ids, HasCounts);
+}
+
+bool llvm::extractBlockWaveCounts(const Function &F,
+ SmallVectorImpl<uint64_t> &Counts,
+ BitVector *HasCounts) {
+ Counts.clear();
+ if (HasCounts)
+ HasCounts->clear();
+ auto Fail = [&]() {
+ Counts.clear();
+ if (HasCounts)
+ HasCounts->clear();
+ return false;
+ };
+ const MDNode *MD = getWaveProfileTable(F);
+ if (!MD || MD->getNumOperands() != F.size() + 2)
+ return Fail();
+ const uint64_t FunctionId =
+ mdconst::extract<ConstantInt>(MD->getOperand(1))->getZExtValue();
+
+ SmallVector<const BasicBlock *> BlocksById(F.size());
+ SmallVector<BlockWaveProfile> BlockProfiles;
+ for (const BasicBlock &BB : F) {
+ auto Block = getBlockWaveProfile(
+ BB.getTerminator()->getMetadata(LLVMContext::MD_wave_profile_block));
+ if (!Block || Block->FunctionId != FunctionId ||
+ Block->Id >= BlocksById.size() || BlocksById[Block->Id])
+ return Fail();
+ BlocksById[Block->Id] = &BB;
+ BlockProfiles.push_back(std::move(*Block));
+ }
+ if (BlocksById[0] != &F.getEntryBlock())
+ return Fail();
+
+ BitVector ExtractedHasCounts;
+ for (auto [BB, Block] : zip(F, BlockProfiles)) {
+ if (Block.Successors.size() != BB.getTerminator()->getNumSuccessors())
+ return Fail();
+ for (auto [Succ, ExpectedId] : zip(successors(&BB), Block.Successors))
+ if (ExpectedId >= BlocksById.size() || BlocksById[ExpectedId] != Succ)
+ return Fail();
+ Counts.push_back(mdconst::extract<ConstantInt>(MD->getOperand(Block.Id + 2))
+ ->getZExtValue());
+ ExtractedHasCounts.push_back(Block.HasCount);
+ }
+ if (!HasCounts && ExtractedHasCounts.count() != ExtractedHasCounts.size())
+ return Fail();
+ if (!ExtractedHasCounts.test(0))
+ return Fail();
+ if (HasCounts)
+ *HasCounts = std::move(ExtractedHasCounts);
+ return true;
+}
+
+namespace {
+struct WaveProfileMapping {
+ const MDNode *Table;
+ SmallVector<unsigned> BlockIds;
+ BitVector HasCounts;
+ BitVector DuplicateIds;
+
+ WaveProfileMapping(const MDNode *Table, unsigned NumBlocks)
+ : Table(Table), BlockIds(NumBlocks, Table->getNumOperands() - 2),
+ HasCounts(NumBlocks), DuplicateIds(Table->getNumOperands() - 2) {}
+};
+} // namespace
+
+static std::optional<WaveProfileMapping> mapBlockWaveCounts(const Function &F) {
+ const MDNode *MD = getWaveProfileTable(F);
+ if (!MD)
+ return std::nullopt;
+ const uint64_t FunctionId =
+ mdconst::extract<ConstantInt>(MD->getOperand(1))->getZExtValue();
+
+ const unsigned NumProfileBlocks = MD->getNumOperands() - 2;
+ WaveProfileMapping Result(MD, F.size());
+ SmallVector<std::optional<BlockWaveProfile>> BlockProfiles(F.size());
+ SmallVector<const BasicBlock *> BlocksById(NumProfileBlocks);
+ DenseMap<const BasicBlock *, unsigned> CurrentIndices;
+ for (auto [Index, BB] : enumerate(F)) {
+ CurrentIndices.try_emplace(&BB, Index);
+ const MDNode *BlockMD =
+ BB.getTerminator()->getMetadata(LLVMContext::MD_wave_profile_block);
+ if (!BlockMD)
+ continue;
+ auto Block = getBlockWaveProfile(BlockMD);
+ if (!Block || Block->FunctionId != FunctionId ||
+ Block->Id >= NumProfileBlocks)
+ return std::nullopt;
+ unsigned BlockId = Block->Id;
+ BlockProfiles[Index] = std::move(Block);
+ Result.BlockIds[Index] = BlockId;
+ if (BlocksById[BlockId])
+ Result.DuplicateIds.set(BlockId);
+ else
+ BlocksById[BlockId] = &BB;
+ }
+
+ for (auto [Index, BB] : enumerate(F)) {
+ const auto &Block = BlockProfiles[Index];
+ if (!Block || Result.DuplicateIds.test(Result.BlockIds[Index]))
+ continue;
+ Result.HasCounts[Index] = Block->HasCount;
+ }
+
+ // A changed edge makes the source count and both the old and new target
+ // counts ambiguous. Invalidate that local neighborhood while retaining
+ // independent blocks whose recorded execution event still matches.
+ auto Invalidate = [&](const BasicBlock *BB) {
+ auto It = CurrentIndices.find(BB);
+ if (It != CurrentIndices.end())
+ Result.HasCounts.reset(It->second);
+ };
+ for (auto [Index, BB] : enumerate(F)) {
+ const auto &Block = BlockProfiles[Index];
+ if (!Block || Result.DuplicateIds.test(Result.BlockIds[Index])) {
+ for (const BasicBlock *Succ : successors(&BB))
+ Invalidate(Succ);
+ continue;
+ }
+
+ if (Block->Successors.size() != BB.getTerminator()->getNumSuccessors()) {
+ Result.HasCounts.reset(Index);
+ for (const BasicBlock *Succ : successors(&BB))
+ Invalidate(Succ);
+ for (uint64_t OldSuccId : Block->Successors) {
+ if (OldSuccId < BlocksById.size() &&
+ !Result.DuplicateIds.test(OldSuccId))
+ Invalidate(BlocksById[OldSuccId]);
+ }
+ continue;
+ }
+
+ for (auto [Succ, OldSuccId] : zip(successors(&BB), Block->Successors)) {
+ auto Current = CurrentIndices.find(Succ);
+ unsigned SuccId = Current == CurrentIndices.end()
+ ? NumProfileBlocks
+ : Result.BlockIds[Current->second];
+ if (SuccId < NumProfileBlocks && !Result.DuplicateIds.test(SuccId) &&
+ OldSuccId == SuccId)
+ continue;
+
+ Result.HasCounts.reset(Index);
+ Invalidate(Succ);
+ if (OldSuccId < BlocksById.size() && !Result.DuplicateIds.test(OldSuccId))
+ Invalidate(BlocksById[OldSuccId]);
+ }
+ }
+
+ if (Result.DuplicateIds.test(0))
+ return std::nullopt;
+ if (const BasicBlock *OriginalEntry = BlocksById[0]) {
+ unsigned EntryIndex = CurrentIndices.lookup(OriginalEntry);
+ if (!BlockProfiles[EntryIndex]->HasCount)
+ return std::nullopt;
+ }
+ return Result;
+}
+
+bool llvm::extractMappedBlockWaveCounts(const Function &F,
+ SmallVectorImpl<uint64_t> &Counts,
+ BitVector &HasCounts,
+ uint64_t &EntryCount) {
+ Counts.clear();
+ HasCounts.clear();
+ EntryCount = 0;
+ auto Mapping = mapBlockWaveCounts(F);
+ if (!Mapping)
+ return false;
+
+ unsigned NumProfileBlocks = Mapping->Table->getNumOperands() - 2;
+ for (unsigned Id : Mapping->BlockIds)
+ Counts.push_back(
+ Id < NumProfileBlocks && !Mapping->DuplicateIds.test(Id)
+ ? mdconst::extract<ConstantInt>(Mapping->Table->getOperand(Id + 2))
+ ->getZExtValue()
+ : 0);
+ HasCounts = std::move(Mapping->HasCounts);
+ EntryCount = mdconst::extract<ConstantInt>(Mapping->Table->getOperand(2))
+ ->getZExtValue();
+ return true;
+}
+
+// Wave records contain only constants. Save their operands by value because a
+// tracking reference follows in-place edits rather than preserving old values.
+static bool hasWaveProfileOperands(const MDNode *MD,
+ ArrayRef<Metadata *> Operands) {
+ return MD && equal(MD->operands(), Operands,
+ [](const MDOperand &Op, Metadata *Saved) {
+ return Op.get() == Saved;
+ });
+}
+
+BlockWaveCountPreserver::BlockWaveCountPreserver(Function &F) : F(F) {
+ auto Mapping = mapBlockWaveCounts(F);
+ if (!Mapping)
+ return;
+ Profile.reset(F.getMetadata(LLVMContext::MD_wave_profile));
+ ProfileOperands.assign(Profile->op_begin(), Profile->op_end());
+ for (auto [Index, BB] : enumerate(F)) {
+ Instruction *Term = BB.getTerminator();
+ MDNode *MD = Term->getMetadata(LLVMContext::MD_wave_profile_block);
+ SmallVector<Metadata *, 6> Operands;
+ if (MD)
+ Operands.assign(MD->op_begin(), MD->op_end());
+ Blocks.push_back({&BB, Term, TrackingMDNodeRef(MD), std::move(Operands),
+ Mapping->BlockIds[Index], Mapping->HasCounts[Index]});
+ }
+}
+
+void BlockWaveCountPreserver::invalidate(const BasicBlock &BB) {
+ for (BlockProfile &Block : Blocks)
+ if (Block.Block == &BB) {
+ Block.HasCount = false;
+ Invalidated = true;
+ }
+}
+
+void BlockWaveCountPreserver::restore() {
+ TrackingMDNodeRef SavedProfile(std::move(Profile));
+ if (!SavedProfile)
+ return;
+ auto Abandon = [&]() {
+ // Skipping reconstruction must not leave invalidated counts readable.
+ if (Invalidated)
+ clearBlockWaveCounts(F);
+ };
+ if (getWaveProfileTable(F) != SavedProfile ||
+ !hasWaveProfileOperands(SavedProfile, ProfileOperands))
+ return Abandon();
+
+ DenseMap<const BasicBlock *, const BlockProfile *> Saved;
+ for (const BlockProfile &Block : Blocks) {
+ Value *V = Block.Block;
+ auto *BB = dyn_cast_or_null<BasicBlock>(V);
+ if (!BB || BB->getParent() != &F)
+ continue;
+ if (Block.Metadata &&
+ !hasWaveProfileOperands(Block.Metadata, Block.Operands))
+ return Abandon();
+ Instruction *Term = BB->getTerminator();
+ MDNode *MD = Term->getMetadata(LLVMContext::MD_wave_profile_block);
+ // Missing metadata is expected on a replacement terminator, but clearing
+ // metadata on a surviving terminator is an independent invalidation.
+ if (MD != Block.Metadata && (MD || Term == Block.Terminator))
+ return Abandon();
+ Saved.try_emplace(BB, &Block);
+ }
+
+ SmallVector<Metadata *> Table(SavedProfile->op_begin(),
+ SavedProfile->op_end());
+ unsigned NumOriginalIds = Table.size() - 2;
+ BitVector UsedIds(NumOriginalIds);
+ DenseMap<const BasicBlock *, unsigned> Ids;
+ BitVector HasCounts;
+ for (const BasicBlock &BB : F) {
+ const BlockProfile *Block = Saved.lookup(&BB);
+ unsigned Id = Block ? Block->Id : NumOriginalIds;
+ // ID zero owns the normalization count. An invalid current block must not
+ // turn that anchor into an unmeasured count.
+ if (Id >= NumOriginalIds || UsedIds[Id] || (Id == 0 && !Block->HasCount)) {
+ Id = Table.size() - 2;
+ Table.push_back(getWaveProfileInt(F.getContext(), 0));
+ } else {
+ UsedIds.set(Id);
+ }
+ Ids[&BB] = Id;
+ HasCounts.push_back(Block && Block->HasCount);
+ }
+
+ writeWaveProfile(F, Table, Ids, HasCounts);
+}
+
// MD_prof nodes have the following layout
//
// In general:
diff --git a/llvm/lib/IR/Verifier.cpp b/llvm/lib/IR/Verifier.cpp
index 7d3ee1de61061..534da1eefcdeb 100644
--- a/llvm/lib/IR/Verifier.cpp
+++ b/llvm/lib/IR/Verifier.cpp
@@ -591,6 +591,11 @@ void Verifier::visitGlobalValue(const GlobalValue &GV) {
if (GO->hasMetadata(LLVMContext::MD_uniformity_profile))
Check(isa<Function>(GO) && !GO->isDeclaration(),
"uniformity.profile is only valid on function definitions", GO);
+ if (GO->hasMetadata(LLVMContext::MD_wave_profile))
+ Check(isa<Function>(GO) && !GO->isDeclaration(),
+ "wave.profile is only valid on function definitions", GO);
+ Check(!GO->hasMetadata(LLVMContext::MD_wave_profile_block),
+ "wave.profile.block is only valid on terminators", GO);
if (const MDNode *Associated =
GO->getMetadata(LLVMContext::MD_associated)) {
@@ -2818,6 +2823,16 @@ void Verifier::verifyUnknownProfileMetadata(MDNode *MD) {
void Verifier::verifyFunctionMetadata(
ArrayRef<std::pair<unsigned, MDNode *>> MDs) {
for (const auto &Pair : MDs) {
+ if (Pair.first == LLVMContext::MD_wave_profile) {
+ Check(Pair.second->getNumOperands() >= 3,
+ "wave.profile requires a version, function ID, and wave counts",
+ Pair.second);
+ for (const MDOperand &Op : Pair.second->operands()) {
+ const auto *CI = mdconst::dyn_extract_or_null<ConstantInt>(Op);
+ Check(CI && CI->getType()->isIntegerTy(64),
+ "wave.profile operands must be i64", Pair.second);
+ }
+ }
if (Pair.first == LLVMContext::MD_uniformity_profile)
Check(Pair.second->getNumOperands() == 0,
"uniformity.profile must be an empty node", Pair.second);
@@ -6112,6 +6127,21 @@ void Verifier::visitInstruction(Instruction &I) {
Check(!I.getMetadata(LLVMContext::MD_uniformity_profile),
"uniformity.profile is only valid on function definitions", &I);
+ Check(!I.getMetadata(LLVMContext::MD_wave_profile),
+ "wave.profile is only valid on function definitions", &I);
+ if (MDNode *MD = I.getMetadata(LLVMContext::MD_wave_profile_block)) {
+ Check(I.isTerminator(), "wave.profile.block is only valid on terminators",
+ &I);
+ Check(MD->getNumOperands() >= 4,
+ "wave.profile.block requires a version, function ID, block ID, and "
+ "count-valid flag",
+ MD);
+ for (const MDOperand &Op : MD->operands()) {
+ const auto *CI = mdconst::dyn_extract_or_null<ConstantInt>(Op);
+ Check(CI && CI->getType()->isIntegerTy(64),
+ "wave.profile.block operands must be i64", MD);
+ }
+ }
if (MDNode *MD = I.getMetadata(LLVMContext::MD_branch_uniformity_profile)) {
Check(isa<CondBrInst>(I),
diff --git a/llvm/test/Bitcode/wave-profile.ll b/llvm/test/Bitcode/wave-profile.ll
new file mode 100644
index 0000000000000..b909165c046db
--- /dev/null
+++ b/llvm/test/Bitcode/wave-profile.ll
@@ -0,0 +1,40 @@
+; RUN: llvm-as %s -o %t.bc
+; RUN: llvm-dis %t.bc -o - | FileCheck %s
+; RUN: opt -passes=verify -disable-output %t.bc
+; RUN: verify-uselistorder %s
+
+; Measured zero and an uninstrumented block have different validity flags.
+; The serializer must retain both channels without treating wave counts as
+; ordinary function-entry or branch-flow counts.
+
+; CHECK-LABEL: define void @sparse(
+; CHECK-SAME: !wave.profile [[PROFILE:![0-9]+]]
+define void @sparse(i1 %condition) !wave.profile !0 {
+entry:
+ br i1 %condition, label %observed_zero, label %unmeasured, !wave.profile.block !1
+
+observed_zero:
+; CHECK: ret void, !wave.profile.block [[ZERO:![0-9]+]]
+ ret void, !wave.profile.block !2
+
+unmeasured:
+; CHECK: ret void, !wave.profile.block [[UNMEASURED:![0-9]+]]
+ ret void, !wave.profile.block !3
+}
+
+; Unknown versions and stale identities remain representable after transforms.
+; CHECK-LABEL: define void @future(
+; CHECK-SAME: !wave.profile [[FUTURE:![0-9]+]]
+define void @future() !wave.profile !4 {
+ ret void
+}
+
+; CHECK: [[PROFILE]] = !{i64 2, i64 123, i64 32, i64 0, i64 0}
+; CHECK: [[ZERO]] = !{i64 2, i64 123, i64 1, i64 1}
+; CHECK: [[UNMEASURED]] = !{i64 2, i64 123, i64 2, i64 0}
+; CHECK: [[FUTURE]] = !{i64 3, i64 456, i64 0}
+!0 = !{i64 2, i64 123, i64 32, i64 0, i64 0}
+!1 = !{i64 2, i64 123, i64 0, i64 1, i64 1, i64 2}
+!2 = !{i64 2, i64 123, i64 1, i64 1}
+!3 = !{i64 2, i64 123, i64 2, i64 0}
+!4 = !{i64 3, i64 456, i64 0}
diff --git a/llvm/test/Transforms/InstCombine/wave-profile.ll b/llvm/test/Transforms/InstCombine/wave-profile.ll
new file mode 100644
index 0000000000000..c51255e5b742d
--- /dev/null
+++ b/llvm/test/Transforms/InstCombine/wave-profile.ll
@@ -0,0 +1,26 @@
+; RUN: opt -S -passes='instcombine,verify' %s | FileCheck %s
+
+; InstCombine canonicalizes a negated branch condition by swapping successors.
+; The stable successor identities must follow that equivalent edge swap.
+
+define void @_Z14divergent_loopPVi(i1 %condition) !wave.profile !0 {
+; CHECK-LABEL: define void @_Z14divergent_loopPVi(
+; CHECK: entry:
+; CHECK-NEXT: br i1 %condition, label %right, label %left, !wave.profile.block [[ENTRY:![0-9]+]]
+entry:
+ %not = xor i1 %condition, true
+ br i1 %not, label %left, label %right, !wave.profile.block !1
+
+left:
+ ret void, !wave.profile.block !2
+
+right:
+ ret void, !wave.profile.block !3
+}
+
+; CHECK: [[ENTRY]] = !{i64 2, i64 2480672276464841217, i64 0, i64 1, i64 2, i64 1}
+
+!0 = !{i64 2, i64 2480672276464841217, i64 100, i64 10, i64 90}
+!1 = !{i64 2, i64 2480672276464841217, i64 0, i64 1, i64 1, i64 2}
+!2 = !{i64 2, i64 2480672276464841217, i64 1, i64 1}
+!3 = !{i64 2, i64 2480672276464841217, i64 2, i64 1}
diff --git a/llvm/test/Verifier/wave-profile.ll b/llvm/test/Verifier/wave-profile.ll
new file mode 100644
index 0000000000000..eea0df43f9507
--- /dev/null
+++ b/llvm/test/Verifier/wave-profile.ll
@@ -0,0 +1,91 @@
+; RUN: split-file %s %t
+; RUN: opt -passes=verify -disable-output %t/valid.ll
+; RUN: not opt -passes=verify -disable-output %t/instruction.ll 2>&1 | FileCheck %s --check-prefix=LOCATION
+; RUN: not opt -passes=verify -disable-output %t/global.ll 2>&1 | FileCheck %s --check-prefix=LOCATION
+; RUN: not opt -passes=verify -disable-output %t/short.ll 2>&1 | FileCheck %s --check-prefix=SHORT
+; RUN: not opt -passes=verify -disable-output %t/type.ll 2>&1 | FileCheck %s --check-prefix=TYPE
+; RUN: not opt -passes=verify -disable-output %t/duplicate.ll 2>&1 | FileCheck %s --check-prefix=TYPE
+; RUN: sed 's/!0 !wave.profile !1/!1 !wave.profile !0/' %t/duplicate.ll | not opt -passes=verify -disable-output 2>&1 | FileCheck %s --check-prefix=TYPE
+; RUN: not opt -passes=verify -disable-output %t/block-location.ll 2>&1 | FileCheck %s --check-prefix=BLOCK-LOCATION
+; RUN: not opt -passes=verify -disable-output %t/block-short.ll 2>&1 | FileCheck %s --check-prefix=BLOCK-SHORT
+; RUN: not opt -passes=verify -disable-output %t/block-type.ll 2>&1 | FileCheck %s --check-prefix=BLOCK-TYPE
+; RUN: not opt -passes=verify -disable-output %t/declaration.ll 2>&1 | FileCheck %s --check-prefix=LOCATION
+; RUN: not opt -passes=verify -disable-output %t/block-function.ll 2>&1 | FileCheck %s --check-prefix=BLOCK-LOCATION
+; RUN: not opt -passes=verify -disable-output %t/block-global.ll 2>&1 | FileCheck %s --check-prefix=BLOCK-LOCATION
+
+; LOCATION: wave.profile is only valid on function definitions
+; SHORT: wave.profile requires a version, function ID, and wave counts
+; TYPE: wave.profile operands must be i64
+; BLOCK-LOCATION: wave.profile.block is only valid on terminators
+; BLOCK-SHORT: wave.profile.block requires a version, function ID, block ID, and count-valid flag
+; BLOCK-TYPE: wave.profile.block operands must be i64
+
+;--- valid.ll
+; Stale fingerprints, unknown versions, and stale block counts are ignored by
+; consumers, not rejected by the verifier after an optimization changes IR.
+define void @stale() !wave.profile !0 {
+ ret void
+}
+!0 = !{i64 1, i64 0, i64 100, i64 200}
+
+;--- instruction.ll
+define void @bad() {
+ ret void, !wave.profile !0
+}
+!0 = !{i64 1, i64 0, i64 100}
+
+;--- global.ll
+ at bad = global i32 0, !wave.profile !0
+!0 = !{i64 1, i64 0, i64 100}
+
+;--- short.ll
+define void @bad() !wave.profile !0 {
+ ret void
+}
+!0 = !{i64 1, i64 0}
+
+;--- type.ll
+define void @bad() !wave.profile !0 {
+ ret void
+}
+!0 = !{i64 1, i64 0, i32 100}
+
+;--- duplicate.ll
+define void @bad() !wave.profile !0 !wave.profile !1 {
+ ret void
+}
+!0 = !{i64 1, i64 0, i32 100}
+!1 = !{i64 1, i64 0, i64 100}
+
+;--- block-location.ll
+define void @bad() {
+ %value = add i32 1, 2, !wave.profile.block !0
+ ret void
+}
+!0 = !{i64 2, i64 0, i64 0, i64 1}
+
+;--- block-short.ll
+define void @bad() {
+ ret void, !wave.profile.block !0
+}
+!0 = !{i64 2, i64 0, i64 0}
+
+;--- block-type.ll
+define void @bad() {
+ ret void, !wave.profile.block !0
+}
+!0 = !{i64 2, i64 0, i32 0, i64 1}
+
+;--- declaration.ll
+declare !wave.profile !0 void @bad()
+!0 = !{i64 2, i64 0, i64 100}
+
+;--- block-function.ll
+define void @bad() !wave.profile.block !0 {
+ ret void
+}
+!0 = !{i64 2, i64 0, i64 0, i64 1}
+
+;--- block-global.ll
+ at bad = global i32 0, !wave.profile.block !0
+!0 = !{i64 2, i64 0, i64 0, i64 1}
diff --git a/llvm/unittests/IR/CMakeLists.txt b/llvm/unittests/IR/CMakeLists.txt
index df83992fadd1e..82e72b4739bff 100644
--- a/llvm/unittests/IR/CMakeLists.txt
+++ b/llvm/unittests/IR/CMakeLists.txt
@@ -43,6 +43,7 @@ add_llvm_unittest(IRTests
ModuleSummaryIndexTest.cpp
PassManagerTest.cpp
PatternMatch.cpp
+ ProfDataUtilsTest.cpp
ShuffleVectorInstTest.cpp
StructuralHashTest.cpp
RuntimeLibcallsTest.cpp
diff --git a/llvm/unittests/IR/ProfDataUtilsTest.cpp b/llvm/unittests/IR/ProfDataUtilsTest.cpp
new file mode 100644
index 0000000000000..38bbd2761f072
--- /dev/null
+++ b/llvm/unittests/IR/ProfDataUtilsTest.cpp
@@ -0,0 +1,817 @@
+//===- ProfDataUtilsTest.cpp - Profiling metadata tests ------------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+
+#include "llvm/IR/ProfDataUtils.h"
+#include "llvm/AsmParser/Parser.h"
+#include "llvm/IR/Constants.h"
+#include "llvm/IR/IRBuilder.h"
+#include "llvm/IR/Instructions.h"
+#include "llvm/IR/LLVMContext.h"
+#include "llvm/IR/Module.h"
+#include "llvm/IR/Verifier.h"
+#include "llvm/Support/SourceMgr.h"
+#include "llvm/Transforms/Utils/Cloning.h"
+#include "gtest/gtest.h"
+#include <initializer_list>
+
+using namespace llvm;
+
+namespace {
+static BitVector makeBitVector(unsigned Size,
+ std::initializer_list<unsigned> SetBits) {
+ BitVector Result(Size);
+ for (unsigned Index : SetBits)
+ Result.set(Index);
+ return Result;
+}
+
+class WaveProfileTest : public testing::Test {
+protected:
+ LLVMContext Context;
+ std::unique_ptr<Module> M;
+
+ void expectUnmeasured(Function &F, unsigned Index) {
+ SmallVector<uint64_t> Counts;
+ BitVector Valid;
+ uint64_t Entry;
+ if (extractMappedBlockWaveCounts(F, Counts, Valid, Entry))
+ EXPECT_FALSE(Valid[Index]);
+ if (extractBlockWaveCounts(F, Counts, &Valid))
+ EXPECT_FALSE(Valid[Index]);
+ }
+
+ void SetUp() override {
+ SMDiagnostic Error;
+ M = parseAssemblyString(R"(
+ define void @diamond(i1 %condition) {
+ entry:
+ br i1 %condition, label %left, label %right
+ left:
+ br label %exit
+ right:
+ br label %exit
+ exit:
+ ret void
+ })",
+ Error, Context);
+ ASSERT_TRUE(M);
+ }
+};
+
+TEST_F(WaveProfileTest, RoundTripAndReplacement) {
+ Function &F = *M->getFunction("diamond");
+ F.setEntryCount(6400);
+ SmallVector<uint64_t> Counts{99};
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts));
+ EXPECT_TRUE(Counts.empty());
+ setBlockWaveCounts(F, {100, 100, 100, 100});
+ EXPECT_TRUE(extractBlockWaveCounts(F, Counts));
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 100, 100, 100}));
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+
+ std::string Text;
+ raw_string_ostream OS(Text);
+ M->print(OS, nullptr);
+ SMDiagnostic Error;
+ std::unique_ptr<Module> Reloaded = parseAssemblyString(Text, Error, Context);
+ ASSERT_TRUE(Reloaded);
+ EXPECT_TRUE(
+ extractBlockWaveCounts(*Reloaded->getFunction("diamond"), Counts));
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 100, 100, 100}));
+
+ setBlockWaveCounts(F, {200, 0, 200, 200});
+ EXPECT_TRUE(extractBlockWaveCounts(F, Counts));
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{200, 0, 200, 200}));
+}
+
+TEST_F(WaveProfileTest, PreserveAcrossInstructionChanges) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 100, 100, 100});
+ cast<CondBrInst>(F.getEntryBlock().getTerminator())
+ ->setCondition(ConstantInt::getTrue(Context));
+ SmallVector<uint64_t> Counts{99};
+ EXPECT_TRUE(extractBlockWaveCounts(F, Counts));
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 100, 100, 100}));
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, PreserveAcrossConditionalSuccessorSwap) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ cast<CondBrInst>(F.getEntryBlock().getTerminator())->swapSuccessors();
+ SmallVector<uint64_t> Counts;
+ EXPECT_TRUE(extractBlockWaveCounts(F, Counts));
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 10, 90, 100}));
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, PreserveIdentityAcrossBlockReordering) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BasicBlock *Left = F.getEntryBlock().getNextNode();
+ Left->moveAfter(Left->getNextNode());
+ SmallVector<uint64_t> Counts;
+ EXPECT_TRUE(extractBlockWaveCounts(F, Counts));
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 90, 10, 100}));
+}
+
+TEST_F(WaveProfileTest, PreserveAcrossEntryCountChangeButRejectRename) {
+ Function &F = *M->getFunction("diamond");
+ F.setEntryCount(6400);
+ setBlockWaveCounts(F, {100, 100, 100, 100});
+ F.setEntryCount(3200);
+ SmallVector<uint64_t> Counts;
+ EXPECT_TRUE(extractBlockWaveCounts(F, Counts));
+ F.setName("specialized");
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts));
+}
+
+TEST_F(WaveProfileTest, RejectRedirectedEdge) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 100, 100, 100});
+ BasicBlock *Left = F.getEntryBlock().getNextNode();
+ BasicBlock *Right = Left->getNextNode();
+ cast<UncondBrInst>(Left->getTerminator())->setSuccessor(Right);
+ SmallVector<uint64_t> Counts;
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts));
+}
+
+TEST_F(WaveProfileTest, RejectDuplicateBlockIdentity) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 100, 100, 100});
+ BasicBlock *Left = F.getEntryBlock().getNextNode();
+ BasicBlock *Right = Left->getNextNode();
+ Right->getTerminator()->setMetadata(
+ LLVMContext::MD_wave_profile_block,
+ Left->getTerminator()->getMetadata(LLVMContext::MD_wave_profile_block));
+ SmallVector<uint64_t> Counts;
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts));
+}
+
+TEST_F(WaveProfileTest, RejectSplitBlock) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 100, 100, 100});
+ BasicBlock &Entry = F.getEntryBlock();
+ Entry.splitBasicBlock(Entry.begin(), "split");
+ SmallVector<uint64_t> Counts;
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts));
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, MapCountsAfterRemovingOriginalBlock) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BasicBlock *Left = F.getEntryBlock().getNextNode();
+ BasicBlock *Exit = &F.back();
+ Left->replaceAllUsesWith(Exit);
+ Left->eraseFromParent();
+
+ SmallVector<uint64_t> Counts;
+ BitVector HasCounts;
+ uint64_t EntryCount = 0;
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts, &HasCounts));
+ EXPECT_TRUE(extractMappedBlockWaveCounts(F, Counts, HasCounts, EntryCount));
+ EXPECT_EQ(EntryCount, 100u);
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 90, 100}));
+ EXPECT_EQ(HasCounts, makeBitVector(3, {1}));
+}
+
+TEST_F(WaveProfileTest, MapCountsAfterRemovingOriginalEntry) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ F.getEntryBlock().eraseFromParent();
+
+ SmallVector<uint64_t> Counts;
+ BitVector HasCounts;
+ uint64_t EntryCount = 0;
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts, &HasCounts));
+ EXPECT_TRUE(extractMappedBlockWaveCounts(F, Counts, HasCounts, EntryCount));
+ EXPECT_EQ(EntryCount, 100u);
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{10, 90, 100}));
+ EXPECT_EQ(HasCounts, makeBitVector(3, {0, 1, 2}));
+}
+
+TEST_F(WaveProfileTest, MapCountsAroundNewUnmeasuredBlock) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BasicBlock *Left = F.getEntryBlock().getNextNode();
+ BasicBlock *Exit = &F.back();
+ BasicBlock *Inserted = BasicBlock::Create(Context, "inserted", &F, Exit);
+ UncondBrInst::Create(Exit, Inserted);
+ cast<UncondBrInst>(Left->getTerminator())->setSuccessor(Inserted);
+
+ SmallVector<uint64_t> Counts;
+ BitVector HasCounts;
+ uint64_t EntryCount = 0;
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts, &HasCounts));
+ EXPECT_TRUE(extractMappedBlockWaveCounts(F, Counts, HasCounts, EntryCount));
+ EXPECT_EQ(EntryCount, 100u);
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 10, 90, 0, 100}));
+ EXPECT_EQ(HasCounts, makeBitVector(5, {0, 2}));
+}
+
+TEST_F(WaveProfileTest, MapCountsAroundDuplicateBlockIdentity) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BasicBlock *Left = F.getEntryBlock().getNextNode();
+ BasicBlock *Right = Left->getNextNode();
+ Right->getTerminator()->setMetadata(
+ LLVMContext::MD_wave_profile_block,
+ Left->getTerminator()->getMetadata(LLVMContext::MD_wave_profile_block));
+
+ SmallVector<uint64_t> Counts;
+ BitVector HasCounts;
+ uint64_t EntryCount = 0;
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts, &HasCounts));
+ EXPECT_TRUE(extractMappedBlockWaveCounts(F, Counts, HasCounts, EntryCount));
+ EXPECT_EQ(EntryCount, 100u);
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 0, 0, 100}));
+ EXPECT_EQ(HasCounts, makeBitVector(4, {}));
+}
+
+TEST_F(WaveProfileTest, MapCountsAroundRedirectedEdge) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BasicBlock *Left = F.getEntryBlock().getNextNode();
+ BasicBlock *Right = Left->getNextNode();
+ cast<UncondBrInst>(Left->getTerminator())->setSuccessor(Right);
+
+ SmallVector<uint64_t> Counts;
+ BitVector HasCounts;
+ uint64_t EntryCount = 0;
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts, &HasCounts));
+ EXPECT_TRUE(extractMappedBlockWaveCounts(F, Counts, HasCounts, EntryCount));
+ EXPECT_EQ(EntryCount, 100u);
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 10, 90, 100}));
+ EXPECT_EQ(HasCounts, makeBitVector(4, {0}));
+}
+
+TEST_F(WaveProfileTest, RepresentExplicitlyUnmeasuredBlocks) {
+ Function &F = *M->getFunction("diamond");
+ BasicBlock &Entry = F.getEntryBlock();
+ Entry.splitBasicBlock(Entry.begin(), "synthetic");
+
+ SmallVector<uint64_t> ExpectedCounts{100, 0, 100, 100, 100};
+ BitVector ExpectedHasCounts(ExpectedCounts.size(), true);
+ ExpectedHasCounts.reset(1);
+ setBlockWaveCounts(F, ExpectedCounts, ExpectedHasCounts);
+
+ SmallVector<uint64_t> Counts;
+ BitVector HasCounts;
+ EXPECT_TRUE(extractBlockWaveCounts(F, Counts, &HasCounts));
+ EXPECT_EQ(Counts, ExpectedCounts);
+ EXPECT_EQ(HasCounts, ExpectedHasCounts);
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts));
+ EXPECT_TRUE(Counts.empty());
+}
+
+TEST_F(WaveProfileTest, RejectUnmeasuredOriginalEntry) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ Instruction *EntryTerminator = F.getEntryBlock().getTerminator();
+ MDNode *EntryMD =
+ EntryTerminator->getMetadata(LLVMContext::MD_wave_profile_block);
+ SmallVector<Metadata *> Ops(EntryMD->op_begin(), EntryMD->op_end());
+ Ops[3] =
+ ConstantAsMetadata::get(ConstantInt::get(Type::getInt64Ty(Context), 0));
+ EntryTerminator->setMetadata(LLVMContext::MD_wave_profile_block,
+ MDNode::get(Context, Ops));
+
+ SmallVector<uint64_t> Counts;
+ BitVector HasCounts;
+ uint64_t EntryCount = 0;
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts, &HasCounts));
+ EXPECT_FALSE(extractMappedBlockWaveCounts(F, Counts, HasCounts, EntryCount));
+ EXPECT_EQ(EntryCount, 0u);
+ EXPECT_TRUE(Counts.empty());
+ EXPECT_TRUE(HasCounts.empty());
+}
+
+TEST_F(WaveProfileTest, RejectUnsupportedOrMalformedMetadata) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 100, 100, 100});
+ MDNode *Valid = F.getMetadata(LLVMContext::MD_wave_profile);
+ SmallVector<Metadata *> Ops(Valid->op_begin(), Valid->op_end());
+ SmallVector<uint64_t> Counts{99};
+
+ Ops[0] =
+ ConstantAsMetadata::get(ConstantInt::get(Type::getInt64Ty(Context), 3));
+ F.setMetadata(LLVMContext::MD_wave_profile, MDNode::get(Context, Ops));
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts));
+ EXPECT_TRUE(Counts.empty());
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+
+ Ops[0] = Valid->getOperand(0);
+ Ops[2] =
+ ConstantAsMetadata::get(ConstantInt::get(Type::getInt32Ty(Context), 100));
+ F.setMetadata(LLVMContext::MD_wave_profile, MDNode::get(Context, Ops));
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts));
+ EXPECT_TRUE(Counts.empty());
+
+ Ops[2] = nullptr;
+ F.setMetadata(LLVMContext::MD_wave_profile, MDNode::get(Context, Ops));
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts));
+ EXPECT_TRUE(Counts.empty());
+
+ Ops[2] = Valid->getOperand(2);
+ Ops.pop_back();
+ F.setMetadata(LLVMContext::MD_wave_profile, MDNode::get(Context, Ops));
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts));
+ EXPECT_TRUE(Counts.empty());
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+TEST_F(WaveProfileTest, TransferAcrossEdgeSplit) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 80, 60, 100});
+ BlockWaveCountPreserver Profile(F);
+ auto *Br = cast<CondBrInst>(F.getEntryBlock().getTerminator());
+ BasicBlock *Left = Br->getSuccessor(0);
+ BasicBlock *Edge = BasicBlock::Create(Context, "edge", &F);
+ UncondBrInst::Create(Left, Edge);
+ Br->setSuccessor(0, Edge);
+ Profile.restore();
+
+ SmallVector<uint64_t> Counts;
+ BitVector Valid;
+ uint64_t Entry;
+ ASSERT_TRUE(extractMappedBlockWaveCounts(F, Counts, Valid, Entry));
+ EXPECT_EQ(Entry, 100u);
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 80, 60, 100, 0}));
+ EXPECT_EQ(Valid, makeBitVector(5, {0, 1, 2, 3}));
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+
+ BlockWaveCountPreserver Again(F);
+ Again.restore();
+ ASSERT_TRUE(extractMappedBlockWaveCounts(F, Counts, Valid, Entry));
+ EXPECT_EQ(Entry, 100u);
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 80, 60, 100, 0}));
+ EXPECT_EQ(Valid, makeBitVector(5, {0, 1, 2, 3}));
+}
+
+TEST_F(WaveProfileTest, TransferDoesNotResurrectInvalidCounts) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 80, 60, 100});
+ auto *Br = cast<CondBrInst>(F.getEntryBlock().getTerminator());
+ BasicBlock *Left = Br->getSuccessor(0);
+ BasicBlock *Edge = BasicBlock::Create(Context, "edge", &F);
+ UncondBrInst::Create(Left, Edge);
+ Br->setSuccessor(0, Edge);
+
+ // Capture after an unsupported rewrite has invalidated the entry and left.
+ BlockWaveCountPreserver Profile(F);
+ Profile.restore();
+ SmallVector<uint64_t> Counts;
+ BitVector Valid;
+ uint64_t Entry;
+ ASSERT_TRUE(extractMappedBlockWaveCounts(F, Counts, Valid, Entry));
+ EXPECT_EQ(Entry, 100u);
+ EXPECT_EQ(Valid, makeBitVector(5, {2, 3}));
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, TransferInvalidatesChangedExecutionEvent) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 80, 60, 100});
+ BlockWaveCountPreserver Profile(F);
+ Profile.invalidate(F.getEntryBlock());
+ Profile.restore();
+ SmallVector<uint64_t> Counts;
+ BitVector Valid;
+ uint64_t Entry;
+ ASSERT_TRUE(extractMappedBlockWaveCounts(F, Counts, Valid, Entry));
+ EXPECT_EQ(Entry, 100u);
+ EXPECT_EQ(Valid, makeBitVector(4, {1, 2, 3}));
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, TransferDoesNotResurrectDuplicatedCounts) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 80, 60, 100});
+ auto *Br = cast<CondBrInst>(F.getEntryBlock().getTerminator());
+ BasicBlock *Left = Br->getSuccessor(0);
+ BasicBlock *Copy = BasicBlock::Create(Context, "copy", &F);
+ UncondBrInst *CopyBr = UncondBrInst::Create(Left->getSingleSuccessor(), Copy);
+ CopyBr->setMetadata(
+ LLVMContext::MD_wave_profile_block,
+ Left->getTerminator()->getMetadata(LLVMContext::MD_wave_profile_block));
+ BlockWaveCountPreserver Profile(F);
+ Profile.restore();
+ SmallVector<uint64_t> Counts;
+ BitVector Valid;
+ uint64_t Entry;
+ ASSERT_TRUE(extractMappedBlockWaveCounts(F, Counts, Valid, Entry));
+ EXPECT_EQ(Entry, 100u);
+ EXPECT_EQ(Valid, makeBitVector(5, {2}));
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, TransferAfterRemovingOriginalEntry) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 80, 60, 100});
+ BlockWaveCountPreserver Profile(F);
+ F.getEntryBlock().eraseFromParent();
+ Profile.restore();
+ BlockWaveCountPreserver Again(F);
+ Again.restore();
+
+ SmallVector<uint64_t> Counts;
+ BitVector Valid;
+ uint64_t Entry;
+ ASSERT_TRUE(extractMappedBlockWaveCounts(F, Counts, Valid, Entry));
+ EXPECT_EQ(Entry, 100u);
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{80, 60, 100}));
+ EXPECT_EQ(Valid, makeBitVector(3, {0, 1, 2}));
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, TransferDoesNotFollowReplacedBlock) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 80, 60, 100});
+ BlockWaveCountPreserver Profile(F);
+ BasicBlock *Left = F.getEntryBlock().getNextNode();
+ BasicBlock *Replacement = BasicBlock::Create(Context, "replacement", &F);
+ UncondBrInst::Create(Left->getSingleSuccessor(), Replacement);
+ Left->replaceAllUsesWith(Replacement);
+ Left->eraseFromParent();
+ Profile.restore();
+
+ SmallVector<uint64_t> Counts;
+ BitVector Valid;
+ uint64_t Entry;
+ ASSERT_TRUE(extractMappedBlockWaveCounts(F, Counts, Valid, Entry));
+ EXPECT_EQ(Entry, 100u);
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 60, 100, 0}));
+ EXPECT_EQ(Valid, makeBitVector(4, {0, 1, 2}));
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, SparseMeasuredZeroAndReplacement) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {32, 0, 32, 0}, makeBitVector(4, {0, 1, 2}));
+ SmallVector<uint64_t> Counts;
+ BitVector HasCounts;
+ uint64_t EntryCount;
+ ASSERT_TRUE(extractMappedBlockWaveCounts(F, Counts, HasCounts, EntryCount));
+ EXPECT_EQ(EntryCount, 32u);
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{32, 0, 32, 0}));
+ EXPECT_EQ(HasCounts, makeBitVector(4, {0, 1, 2}));
+
+ // Replacing a measured block with an unmeasured one clears its old state.
+ setBlockWaveCounts(F, {32, 0, 32, 0}, makeBitVector(4, {0, 2}));
+ ASSERT_TRUE(extractBlockWaveCounts(F, Counts, &HasCounts));
+ EXPECT_EQ(HasCounts, makeBitVector(4, {0, 2}));
+ clearBlockWaveCounts(F);
+ EXPECT_FALSE(F.hasMetadata(LLVMContext::MD_wave_profile));
+ for (BasicBlock &BB : F)
+ EXPECT_FALSE(
+ BB.getTerminator()->hasMetadata(LLVMContext::MD_wave_profile_block));
+ EXPECT_FALSE(extractMappedBlockWaveCounts(F, Counts, HasCounts, EntryCount));
+ EXPECT_TRUE(Counts.empty());
+ EXPECT_TRUE(HasCounts.empty());
+ EXPECT_EQ(EntryCount, 0u);
+}
+
+TEST_F(WaveProfileTest, RejectConflictingFunctionRecords) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ MDNode *First = F.getMetadata(LLVMContext::MD_wave_profile);
+ setBlockWaveCounts(F, {200, 20, 180, 200});
+ MDNode *Second = F.getMetadata(LLVMContext::MD_wave_profile);
+ for (bool Reverse : {false, true}) {
+ F.setMetadata(LLVMContext::MD_wave_profile, Reverse ? Second : First);
+ F.addMetadata(LLVMContext::MD_wave_profile, Reverse ? *First : *Second);
+ SmallVector<uint64_t> Counts{99};
+ BitVector HasCounts;
+ uint64_t EntryCount = 99;
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts, &HasCounts));
+ EXPECT_TRUE(Counts.empty());
+ EXPECT_FALSE(
+ extractMappedBlockWaveCounts(F, Counts, HasCounts, EntryCount));
+ EXPECT_TRUE(Counts.empty());
+ EXPECT_TRUE(HasCounts.empty());
+ EXPECT_EQ(EntryCount, 0u);
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+ }
+}
+
+TEST_F(WaveProfileTest, RejectSpecializedFunctionNames) {
+ Function &F = *M->getFunction("diamond");
+ for (StringRef Name : {"compute", "compute.llvm.123", "compute.__uniq.123",
+ "compute.content.123"}) {
+ SCOPED_TRACE(Name.str());
+ F.setName(Name);
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ SmallVector<uint64_t> Counts;
+ BitVector Valid;
+ uint64_t Entry;
+ ASSERT_TRUE(extractBlockWaveCounts(F, Counts));
+ ASSERT_TRUE(extractMappedBlockWaveCounts(F, Counts, Valid, Entry));
+ // Function specialization appends this suffix after any promotion suffix.
+ ValueToValueMapTy VMap;
+ Function *Clone = CloneFunction(&F, VMap);
+ Clone->setName(Name + ".specialized.1");
+ EXPECT_FALSE(extractBlockWaveCounts(*Clone, Counts));
+ EXPECT_TRUE(Counts.empty());
+ EXPECT_FALSE(extractMappedBlockWaveCounts(*Clone, Counts, Valid, Entry));
+ EXPECT_TRUE(Counts.empty());
+ EXPECT_TRUE(Valid.empty());
+ EXPECT_EQ(Entry, 0u);
+ EXPECT_FALSE(verifyFunction(*Clone, &errs()));
+ Clone->eraseFromParent();
+ }
+}
+
+TEST_F(WaveProfileTest, TransferRespectsNestedInvalidation) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BlockWaveCountPreserver Outer(F);
+ BlockWaveCountPreserver Inner(F);
+ Inner.invalidate(*F.getEntryBlock().getNextNode());
+ Inner.restore();
+ Outer.restore();
+
+ SmallVector<uint64_t> Counts;
+ BitVector Valid;
+ uint64_t Entry;
+ ASSERT_TRUE(extractMappedBlockWaveCounts(F, Counts, Valid, Entry));
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 10, 90, 100}));
+ EXPECT_EQ(Valid, makeBitVector(4, {0, 2, 3}));
+}
+
+TEST_F(WaveProfileTest, TransferRespectsProfileRemoval) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BlockWaveCountPreserver Profile(F);
+ clearBlockWaveCounts(F);
+ Profile.restore();
+ EXPECT_FALSE(F.hasMetadata(LLVMContext::MD_wave_profile));
+ for (BasicBlock &BB : F)
+ EXPECT_FALSE(
+ BB.getTerminator()->hasMetadata(LLVMContext::MD_wave_profile_block));
+}
+
+TEST_F(WaveProfileTest, TransferRespectsProfileReplacement) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BlockWaveCountPreserver Profile(F);
+ setBlockWaveCounts(F, {200, 20, 180, 200});
+ Profile.restore();
+ SmallVector<uint64_t> Counts;
+ ASSERT_TRUE(extractBlockWaveCounts(F, Counts));
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{200, 20, 180, 200}));
+}
+
+TEST_F(WaveProfileTest, TransferRespectsBlockMetadataRemoval) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BlockWaveCountPreserver Profile(F);
+ Instruction *LeftTerm = F.getEntryBlock().getNextNode()->getTerminator();
+ LeftTerm->setMetadata(LLVMContext::MD_wave_profile_block, nullptr);
+ Profile.restore();
+ EXPECT_FALSE(LeftTerm->hasMetadata(LLVMContext::MD_wave_profile_block));
+ SmallVector<uint64_t> Counts;
+ BitVector Valid;
+ uint64_t Entry;
+ ASSERT_TRUE(extractMappedBlockWaveCounts(F, Counts, Valid, Entry));
+ EXPECT_FALSE(Valid[1]);
+}
+
+TEST_F(WaveProfileTest, TransferAcrossTerminatorReplacement) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BlockWaveCountPreserver Profile(F);
+ BasicBlock *Left = F.getEntryBlock().getNextNode();
+ Instruction *OldTerm = Left->getTerminator();
+ UncondBrInst::Create(&F.back(), OldTerm->getIterator());
+ OldTerm->eraseFromParent();
+ Profile.restore();
+ SmallVector<uint64_t> Counts;
+ ASSERT_TRUE(extractBlockWaveCounts(F, Counts));
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{100, 10, 90, 100}));
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, RejectMalformedBlockRecordsAndIgnoreTheirSwap) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ auto *Br = cast<CondBrInst>(F.getEntryBlock().getTerminator());
+ MDNode *Original = Br->getMetadata(LLVMContext::MD_wave_profile_block);
+ // All record fields are i64, including IDs and the measured flag.
+ for (unsigned I = 0; I != Original->getNumOperands(); ++I) {
+ for (Metadata *Bad :
+ {static_cast<Metadata *>(nullptr),
+ static_cast<Metadata *>(ConstantAsMetadata::get(
+ ConstantInt::get(Type::getInt32Ty(Context), 0)))}) {
+ SCOPED_TRACE(I);
+ SmallVector<Metadata *> Ops(Original->op_begin(), Original->op_end());
+ Ops[I] = Bad;
+ MDNode *Malformed = MDNode::get(Context, Ops);
+ Br->setMetadata(LLVMContext::MD_wave_profile_block, Malformed);
+ SmallVector<uint64_t> Counts{99};
+ BitVector Valid;
+ uint64_t Entry = 99;
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts, &Valid));
+ EXPECT_TRUE(Counts.empty());
+ EXPECT_FALSE(extractMappedBlockWaveCounts(F, Counts, Valid, Entry));
+ EXPECT_TRUE(Counts.empty());
+ EXPECT_TRUE(Valid.empty());
+ EXPECT_EQ(Entry, 0u);
+ Br->swapSuccessors();
+ EXPECT_EQ(Br->getMetadata(LLVMContext::MD_wave_profile_block), Malformed);
+ Br->swapSuccessors();
+ }
+ }
+}
+
+TEST_F(WaveProfileTest, RejectUnsupportedBlockVersionAndInvalidFlag) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ auto *Br = cast<CondBrInst>(F.getEntryBlock().getTerminator());
+ MDNode *Original = Br->getMetadata(LLVMContext::MD_wave_profile_block);
+ for (unsigned I : {0u, 3u}) {
+ SmallVector<Metadata *> Ops(Original->op_begin(), Original->op_end());
+ Ops[I] =
+ ConstantAsMetadata::get(ConstantInt::get(Type::getInt64Ty(Context), 3));
+ MDNode *Unsupported = MDNode::get(Context, Ops);
+ Br->setMetadata(LLVMContext::MD_wave_profile_block, Unsupported);
+ SmallVector<uint64_t> Counts;
+ BitVector Valid;
+ uint64_t Entry;
+ EXPECT_FALSE(extractBlockWaveCounts(F, Counts, &Valid));
+ EXPECT_FALSE(extractMappedBlockWaveCounts(F, Counts, Valid, Entry));
+ Br->swapSuccessors();
+ EXPECT_EQ(Br->getMetadata(LLVMContext::MD_wave_profile_block), Unsupported);
+ Br->swapSuccessors();
+ if (I == 0)
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+ }
+}
+
+TEST_F(WaveProfileTest, TransferRespectsMetadataOnReplacementTerminator) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BlockWaveCountPreserver Profile(F);
+ BasicBlock *Left = F.getEntryBlock().getNextNode();
+ Instruction *OldTerm = Left->getTerminator();
+ MDNode *MD = OldTerm->getMetadata(LLVMContext::MD_wave_profile_block);
+ SmallVector<Metadata *> Ops(MD->op_begin(), MD->op_end());
+ Ops[3] =
+ ConstantAsMetadata::get(ConstantInt::get(Type::getInt64Ty(Context), 0));
+ auto *NewTerm = UncondBrInst::Create(&F.back(), OldTerm->getIterator());
+ NewTerm->setMetadata(LLVMContext::MD_wave_profile_block,
+ MDNode::get(Context, Ops));
+ OldTerm->eraseFromParent();
+ Profile.restore();
+ SmallVector<uint64_t> Counts;
+ BitVector Valid;
+ ASSERT_TRUE(extractBlockWaveCounts(F, Counts, &Valid));
+ EXPECT_EQ(Valid, makeBitVector(4, {0, 2, 3}));
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, TransferKeepsInvalidationsAfterSuccessorSwap) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BlockWaveCountPreserver Profile(F);
+ Profile.invalidate(*F.getEntryBlock().getNextNode());
+ auto *Br = cast<CondBrInst>(F.getEntryBlock().getTerminator());
+ IRBuilder<> Builder(Br);
+ Br->setCondition(Builder.CreateNot(Br->getCondition()));
+ Br->swapSuccessors();
+ Profile.restore();
+
+ expectUnmeasured(F, 1);
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, TransferKeepsInvalidationsAfterNestedRestore) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BasicBlock *Left = F.getEntryBlock().getNextNode();
+ BlockWaveCountPreserver Outer(F);
+ Outer.invalidate(*Left->getNextNode());
+ BlockWaveCountPreserver Inner(F);
+ Inner.invalidate(*Left);
+ Inner.restore();
+ Outer.restore();
+
+ expectUnmeasured(F, 1);
+ expectUnmeasured(F, 2);
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest,
+ TransferRespectsInvalidationBeforeTerminatorReplacement) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BasicBlock *Left = F.getEntryBlock().getNextNode();
+ BlockWaveCountPreserver Outer(F);
+ BlockWaveCountPreserver Inner(F);
+ Inner.invalidate(*Left);
+ Inner.restore();
+ Instruction *OldTerm = Left->getTerminator();
+ UncondBrInst::Create(&F.back(), OldTerm->getIterator());
+ OldTerm->eraseFromParent();
+ Outer.restore();
+
+ expectUnmeasured(F, 1);
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, TransferRespectsInPlaceMetadataChanges) {
+ Function &F = *M->getFunction("diamond");
+ for (bool ReplaceTerminator : {false, true}) {
+ SCOPED_TRACE(ReplaceTerminator);
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BlockWaveCountPreserver Profile(F);
+ Instruction *OldTerm = F.getEntryBlock().getNextNode()->getTerminator();
+ MDNode *MD = OldTerm->getMetadata(LLVMContext::MD_wave_profile_block);
+ MD->replaceOperandWith(3, ConstantAsMetadata::get(ConstantInt::get(
+ Type::getInt64Ty(Context), 0)));
+ if (ReplaceTerminator) {
+ UncondBrInst::Create(&F.back(), OldTerm->getIterator());
+ OldTerm->eraseFromParent();
+ }
+ Profile.restore();
+ expectUnmeasured(F, 1);
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+ }
+}
+
+TEST_F(WaveProfileTest, TransferKeepsInvalidationsAfterTableChange) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BlockWaveCountPreserver Profile(F);
+ Profile.invalidate(*F.getEntryBlock().getNextNode());
+ F.getMetadata(LLVMContext::MD_wave_profile)
+ ->replaceOperandWith(2, ConstantAsMetadata::get(ConstantInt::get(
+ Type::getInt64Ty(Context), 200)));
+ Profile.restore();
+ expectUnmeasured(F, 1);
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, TransferKeepsInvalidationsAfterProfileReplacement) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BlockWaveCountPreserver Profile(F);
+ Profile.invalidate(*F.getEntryBlock().getNextNode());
+ setBlockWaveCounts(F, {200, 20, 180, 200});
+ Profile.restore();
+ expectUnmeasured(F, 1);
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, TransferConsumesSnapshot) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BlockWaveCountPreserver Profile(F);
+ Profile.invalidate(*F.getEntryBlock().getNextNode());
+ Profile.restore();
+ Profile.restore();
+ SmallVector<uint64_t> Counts;
+ BitVector Valid;
+ uint64_t Entry;
+ ASSERT_TRUE(extractMappedBlockWaveCounts(F, Counts, Valid, Entry));
+ EXPECT_EQ(Valid, makeBitVector(4, {0, 2, 3}));
+ setBlockWaveCounts(F, {200, 20, 180, 200});
+ Profile.restore();
+ ASSERT_TRUE(extractBlockWaveCounts(F, Counts));
+ EXPECT_EQ(Counts, (SmallVector<uint64_t>{200, 20, 180, 200}));
+}
+
+TEST_F(WaveProfileTest,
+ TransferRespectsReattachedProfileAfterTerminatorReplacement) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 0, 90, 100});
+ BlockWaveCountPreserver Profile(F);
+ setBlockWaveCounts(F, {100, 0, 90, 100}, makeBitVector(4, {0, 2, 3}));
+ Instruction *OldTerm = F.getEntryBlock().getNextNode()->getTerminator();
+ UncondBrInst::Create(&F.back(), OldTerm->getIterator());
+ OldTerm->eraseFromParent();
+ Profile.restore();
+ expectUnmeasured(F, 1);
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+TEST_F(WaveProfileTest, TransferInvalidatesDiscardedMetadataEdit) {
+ Function &F = *M->getFunction("diamond");
+ setBlockWaveCounts(F, {100, 10, 90, 100});
+ BlockWaveCountPreserver Profile(F);
+ BasicBlock *Left = F.getEntryBlock().getNextNode();
+ Instruction *OldTerm = Left->getTerminator();
+ OldTerm->setMetadata(LLVMContext::MD_wave_profile_block, nullptr);
+ Profile.invalidate(*Left);
+ UncondBrInst::Create(&F.back(), OldTerm->getIterator());
+ OldTerm->eraseFromParent();
+ Profile.restore();
+ expectUnmeasured(F, 1);
+ EXPECT_FALSE(verifyModule(*M, &errs()));
+}
+
+} // namespace
More information about the llvm-branch-commits
mailing list