[llvm] [VPlan] Cost scalar IV steps in replicate regions. (PR #221995)

via llvm-commits llvm-commits at lists.llvm.org
Tue Sep 8 06:16:26 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-llvm-transforms

Author: Florian Hahn (fhahn)

<details>
<summary>Changes</summary>

Replace the return 0 bail-out in VPScalarIVStepsRecipe::computeCost in
replicate regions by properly scaling by the execution probability of
the region.

VPlan-based estimates of execution probabilities depend on
https://github.com/llvm/llvm-project/pull/216172 (include in PR)

Relevant commit is https://github.com/llvm/llvm-project/commit/b13864475308bb69ca50604d8056d3045c1f2f0e

---

Patch is 52.32 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/221995.diff


19 Files Affected:

- (modified) llvm/include/llvm/Transforms/Vectorize/LoopVectorize.h (+2) 
- (modified) llvm/lib/Transforms/Vectorize/LoopVectorizationPlanner.h (+6-1) 
- (modified) llvm/lib/Transforms/Vectorize/LoopVectorize.cpp (+37-14) 
- (modified) llvm/lib/Transforms/Vectorize/VPlan.h (+40-3) 
- (modified) llvm/lib/Transforms/Vectorize/VPlanConstruction.cpp (+41-9) 
- (modified) llvm/lib/Transforms/Vectorize/VPlanHelpers.h (+3-3) 
- (modified) llvm/lib/Transforms/Vectorize/VPlanPredicator.cpp (+3-3) 
- (modified) llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp (+42-32) 
- (modified) llvm/lib/Transforms/Vectorize/VPlanTransforms.cpp (+7-3) 
- (modified) llvm/lib/Transforms/Vectorize/VPlanTransforms.h (+3-1) 
- (modified) llvm/lib/Transforms/Vectorize/VPlanUnroll.cpp (+5-4) 
- (modified) llvm/lib/Transforms/Vectorize/VPlanUtils.cpp (+15-11) 
- (modified) llvm/lib/Transforms/Vectorize/VPlanUtils.h (+3-2) 
- (modified) llvm/test/Transforms/LoopVectorize/AArch64/induction-costs-sve.ll (+2-46) 
- (modified) llvm/test/Transforms/LoopVectorize/AArch64/scalar-steps-cost.ll (+2-2) 
- (modified) llvm/test/Transforms/LoopVectorize/ARM/optsize_minsize.ll (+15-128) 
- (modified) llvm/test/Transforms/LoopVectorize/VPlan/execution-frequencies-match-bfi.ll (+6-5) 
- (modified) llvm/test/Transforms/LoopVectorize/X86/CostModel/store-scalarization-cost.ll (+2-2) 
- (modified) llvm/test/Transforms/LoopVectorize/replicate-region-cost-probability.ll (+5-4) 


``````````diff
diff --git a/llvm/include/llvm/Transforms/Vectorize/LoopVectorize.h b/llvm/include/llvm/Transforms/Vectorize/LoopVectorize.h
index 5eba1a0f8527c..73c607ad1cf5a 100644
--- a/llvm/include/llvm/Transforms/Vectorize/LoopVectorize.h
+++ b/llvm/include/llvm/Transforms/Vectorize/LoopVectorize.h
@@ -66,6 +66,7 @@ namespace llvm {
 
 class AssumptionCache;
 class BlockFrequencyInfo;
+class BranchProbabilityInfo;
 class DemandedBits;
 class DominatorTree;
 class Function;
@@ -146,6 +147,7 @@ struct LoopVectorizePass : public OptionalPassInfoMixin<LoopVectorizePass> {
   TargetTransformInfo *TTI;
   DominatorTree *DT;
   std::function<BlockFrequencyInfo &()> GetBFI;
+  std::function<const BranchProbabilityInfo &()> GetBPI;
   TargetLibraryInfo *TLI;
   DemandedBits *DB;
   AssumptionCache *AC;
diff --git a/llvm/lib/Transforms/Vectorize/LoopVectorizationPlanner.h b/llvm/lib/Transforms/Vectorize/LoopVectorizationPlanner.h
index 972f8f8e642e2..6b9a2effd6422 100644
--- a/llvm/lib/Transforms/Vectorize/LoopVectorizationPlanner.h
+++ b/llvm/lib/Transforms/Vectorize/LoopVectorizationPlanner.h
@@ -36,6 +36,7 @@ class GeneratedRTChecks;
 
 namespace llvm {
 
+class BranchProbabilityInfo;
 class LoopInfo;
 class DominatorTree;
 class LoopVectorizationLegality;
@@ -891,6 +892,9 @@ class LoopVectorizationPlanner {
 
   OptimizationRemarkEmitter *ORE;
 
+  /// Lazily fetch BranchProbabilityInfo, independent of BlockFrequencyInfo.
+  std::function<const BranchProbabilityInfo &()> GetBPI;
+
   SmallVector<VPlanPtr, 4> VPlans;
 
   /// Profitable vector factors.
@@ -921,7 +925,8 @@ class LoopVectorizationPlanner {
       const TargetTransformInfo &TTI, LoopVectorizationLegality *Legal,
       std::unique_ptr<LoopVectorizationCostModel> CM,
       VFSelectionContext &Config, InterleavedAccessInfo &IAI,
-      PredicatedScalarEvolution &PSE, OptimizationRemarkEmitter *ORE);
+      PredicatedScalarEvolution &PSE, OptimizationRemarkEmitter *ORE,
+      std::function<const BranchProbabilityInfo &()> GetBPI);
 
   ~LoopVectorizationPlanner();
 
diff --git a/llvm/lib/Transforms/Vectorize/LoopVectorize.cpp b/llvm/lib/Transforms/Vectorize/LoopVectorize.cpp
index e027dfa834759..661e1f0957170 100644
--- a/llvm/lib/Transforms/Vectorize/LoopVectorize.cpp
+++ b/llvm/lib/Transforms/Vectorize/LoopVectorize.cpp
@@ -80,6 +80,7 @@
 #include "llvm/Analysis/AssumptionCache.h"
 #include "llvm/Analysis/BasicAliasAnalysis.h"
 #include "llvm/Analysis/BlockFrequencyInfo.h"
+#include "llvm/Analysis/BranchProbabilityInfo.h"
 #include "llvm/Analysis/CFG.h"
 #include "llvm/Analysis/CodeMetrics.h"
 #include "llvm/Analysis/CycleAnalysis.h"
@@ -5347,8 +5348,19 @@ void VPCostContext::invalidateWideningDecision(Instruction *I,
                          LoopVectorizationCostModel::CM_InvalidatedDecision, 0);
 }
 
-uint64_t VPCostContext::getPredBlockCostDivisor(BasicBlock *BB) const {
-  return CM.getPredBlockCostDivisor(CostKind, BB);
+uint64_t
+VPCostContext::getPredBlockCostDivisor(const VPRegionBlock *Region) const {
+  if (CostKind == TTI::TCK_CodeSize)
+    return 1;
+  std::optional<VPExecutionFrequency> Freq =
+      Region->getEntryBranchOnMask()->getExecutionFrequency();
+  if (!Freq)
+    return 1;
+  // A recorded frequency is neither zero nor always-executing, so the
+  // probability is non-zero and the division below is safe.
+  return divideNearest(
+      BranchProbability::getDenominator(),
+      vputils::getExecutionProbability(Freq->Freq).getNumerator());
 }
 
 bool VPCostContext::willBeScalarized(Instruction *I, ElementCount VF) const {
@@ -5669,9 +5681,11 @@ LoopVectorizationPlanner::LoopVectorizationPlanner(
     const TargetTransformInfo &TTI, LoopVectorizationLegality *Legal,
     std::unique_ptr<LoopVectorizationCostModel> CM, VFSelectionContext &Config,
     InterleavedAccessInfo &IAI, PredicatedScalarEvolution &PSE,
-    OptimizationRemarkEmitter *ORE)
+    OptimizationRemarkEmitter *ORE,
+    std::function<const BranchProbabilityInfo &()> GetBPI)
     : OrigLoop(L), LI(LI), DT(DT), TLI(TLI), TTI(TTI), Legal(Legal),
-      CM(std::move(CM)), Config(Config), IAI(IAI), PSE(PSE), ORE(ORE) {}
+      CM(std::move(CM)), Config(Config), IAI(IAI), PSE(PSE), ORE(ORE),
+      GetBPI(GetBPI) {}
 
 LoopVectorizationPlanner::~LoopVectorizationPlanner() = default;
 
@@ -6331,16 +6345,16 @@ static bool verifyExecutionFrequenciesMatchBFI(VPlan &Plan, Loop *OrigLoop,
     Edges += VPBB->getNumSuccessors();
   uint64_t Tolerance = Edges + BranchProbability::getDenominator() / HeaderFreq;
 
-  DenseMap<const VPBasicBlock *, std::optional<BlockFrequency>> Frequencies =
-      vputils::computeExecutionFrequencies(Blocks);
+  DenseMap<const VPBasicBlock *, std::optional<VPExecutionFrequency>>
+      Frequencies = vputils::computeExecutionFrequencies(Blocks);
   for (const auto &[VPBB, BB] :
        zip_equal(drop_begin(Blocks), drop_begin(OrigRPO))) {
     // Compare at BranchProbability's coarser resolution, which is as precise as
     // BFI's frequencies get.
-    std::optional<BlockFrequency> Freq = Frequencies.lookup(VPBB);
+    std::optional<VPExecutionFrequency> Freq = Frequencies.lookup(VPBB);
     if (!Freq)
       continue;
-    BranchProbability Computed = vputils::getExecutionProbability(*Freq);
+    BranchProbability Computed = vputils::getExecutionProbability(Freq->Freq);
 
     // Clamp to the header's frequency, which BFI's rounding may exceed.
     uint64_t BBFreq = BFI.getBlockFreq(BB).getFrequency();
@@ -6379,9 +6393,9 @@ VPlanPtr LoopVectorizationPlanner::tryToBuildVPlan1() {
 
   // Create initial base VPlan0, to serve as common starting point for all
   // candidates built later for specific VF ranges.
-  auto VPlan0 = VPlanTransforms::buildVPlan0(OrigLoop, *LI,
-                                             Legal->getWidestInductionType(),
-                                             PSE, LVer ? &*LVer : nullptr);
+  auto VPlan0 = VPlanTransforms::buildVPlan0(
+      OrigLoop, *LI, Legal->getWidestInductionType(), PSE,
+      LVer ? &*LVer : nullptr, GetBPI);
 
   VPDominatorTree VPDT(*VPlan0);
   if (const LoopAccessInfo *LAI = Legal->getLAI())
@@ -7930,7 +7944,7 @@ bool LoopVectorizePass::processLoop(Loop *L) {
       L, LI, DT, TLI, *TTI, &LVL,
       std::make_unique<LoopVectorizationCostModel>(
           SEL, L, PSE, LI, &LVL, *TTI, TLI, AC, ORE, GetBFI, F, IAI, Config),
-      Config, IAI, PSE, ORE);
+      Config, IAI, PSE, ORE, GetBPI);
 
   EpilogueLowering EpilogueTailLoweringStatus =
       getEpilogueTailLowering(LVP.getCostModel(), L, ORE, LVL, Hints);
@@ -8340,12 +8354,21 @@ PreservedAnalyses LoopVectorizePass::run(Function &F,
 
   auto &MAMProxy = AM.getResult<ModuleAnalysisManagerFunctionProxy>(F);
   PSI = MAMProxy.getCachedResult<ProfileSummaryAnalysis>(*F.getParent());
-  GetBFI = [this, &AM, &F]() -> BlockFrequencyInfo & {
-    // CycleInfo cached by an earlier pass is invalidated when the CFG changes.
+  // CycleInfo cached by an earlier pass is invalidated when the CFG changes.
+  // Both BlockFrequencyAnalysis and BranchProbabilityAnalysis depend on it, so
+  // drop the stale result before either is (re-)computed.
+  auto ClearStaleCycleInfo = [this, &AM, &F] {
     if (CFGChanged && AM.getCachedResult<CycleAnalysis>(F))
       AM.clearAnalysis<CycleAnalysis>(F);
+  };
+  GetBFI = [&AM, &F, ClearStaleCycleInfo]() -> BlockFrequencyInfo & {
+    ClearStaleCycleInfo();
     return AM.getResult<BlockFrequencyAnalysis>(F);
   };
+  GetBPI = [&AM, &F, ClearStaleCycleInfo]() -> const BranchProbabilityInfo & {
+    ClearStaleCycleInfo();
+    return AM.getResult<BranchProbabilityAnalysis>(F);
+  };
   LoopVectorizeResult Result = runImpl(F);
   if (!Result.MadeAnyChange)
     return PreservedAnalyses::all();
diff --git a/llvm/lib/Transforms/Vectorize/VPlan.h b/llvm/lib/Transforms/Vectorize/VPlan.h
index e9f18506af7c2..af7ebccb713ef 100644
--- a/llvm/lib/Transforms/Vectorize/VPlan.h
+++ b/llvm/lib/Transforms/Vectorize/VPlan.h
@@ -1177,6 +1177,14 @@ struct VPRecipeWithIRFlags : public VPSingleDefRecipe, public VPIRFlags {
                                              VPCostContext &Ctx) const;
 };
 
+/// The frequency with which a recipe executes, relative to the entry of the
+/// loop region. IsEstimated is set if any branch weight it was composed from
+/// was estimated from static heuristics.
+struct VPExecutionFrequency {
+  BlockFrequency Freq;
+  bool IsEstimated;
+};
+
 /// Helper to manage IR metadata for recipes. It filters out metadata that
 /// cannot be propagated.
 class LLVM_ABI_FOR_TEST VPIRMetadata {
@@ -1186,6 +1194,10 @@ class LLVM_ABI_FOR_TEST VPIRMetadata {
   static constexpr StringLiteral ExecutionFrequencyMDName =
       "vplan.execution.frequency";
 
+  /// Name of the VPlan-internal metadata kind holding estimated branch weights.
+  static constexpr StringLiteral EstimatedProfileMDName =
+      "vplan.prof.estimated";
+
   /// Returns the ID of the metadata kind named \p Kind, taking the context from
   /// any attached node; all belong to the context of the VPlan's function.
   unsigned getMDKindID(StringRef Kind) const {
@@ -1193,6 +1205,12 @@ class LLVM_ABI_FOR_TEST VPIRMetadata {
     return Metadata.front().second->getContext().getMDKindID(Kind);
   }
 
+  /// Returns the node attached under the VPlan-internal metadata kind named
+  /// \p Kind, or nullptr if there is none.
+  MDNode *getInternalMetadata(StringRef Kind) const {
+    return Metadata.empty() ? nullptr : getMetadata(getMDKindID(Kind));
+  }
+
 public:
   VPIRMetadata() = default;
 
@@ -1240,16 +1258,35 @@ class LLVM_ABI_FOR_TEST VPIRMetadata {
   }
 
   /// Record that the recipe executes with frequency \p Freq, relative to the
-  /// entry of the loop region; see vputils::AlwaysExecutesFreq.
-  void setExecutionFrequency(std::optional<BlockFrequency> Freq,
+  /// entry of the loop region.
+  void setExecutionFrequency(std::optional<VPExecutionFrequency> Freq,
                              LLVMContext &Ctx);
 
   /// Returns the frequency recorded by setExecutionFrequency, if any.
-  std::optional<BlockFrequency> getExecutionFrequency() const;
+  std::optional<VPExecutionFrequency> getExecutionFrequency() const;
 
   /// Drop the frequency recorded by setExecutionFrequency, if any.
   void clearExecutionFrequency();
 
+  /// Returns the branch weights recorded for this terminator, preferring real
+  /// profile data over an estimate, or nullptr if there are none.
+  MDNode *getBranchWeights() const {
+    MDNode *Node = getMetadata(LLVMContext::MD_prof);
+    return Node ? Node : getInternalMetadata(EstimatedProfileMDName);
+  }
+
+  /// Returns true if the weights returned by getBranchWeights are estimated.
+  bool hasEstimatedBranchWeights() const {
+    return getInternalMetadata(EstimatedProfileMDName);
+  }
+
+  /// Set estimated branch weights to \p Node.
+  void setEstimatedBranchWeights(MDNode *Node) {
+    assert(!getMetadata(LLVMContext::MD_prof) &&
+           "real profile data takes precedence over an estimate");
+    setMetadata(Node->getContext().getMDKindID(EstimatedProfileMDName), Node);
+  }
+
 #if !defined(NDEBUG) || defined(LLVM_ENABLE_DUMP)
   /// Print metadata with node IDs.
   void print(raw_ostream &O, VPSlotTracker &SlotTracker) const;
diff --git a/llvm/lib/Transforms/Vectorize/VPlanConstruction.cpp b/llvm/lib/Transforms/Vectorize/VPlanConstruction.cpp
index a2006ebf4bb4f..384dc5bc85d4a 100644
--- a/llvm/lib/Transforms/Vectorize/VPlanConstruction.cpp
+++ b/llvm/lib/Transforms/Vectorize/VPlanConstruction.cpp
@@ -20,7 +20,9 @@
 #include "VPlanPatternMatch.h"
 #include "VPlanTransforms.h"
 #include "VPlanUtils.h"
+#include "llvm/ADT/Sequence.h"
 #include "llvm/ADT/SmallVectorExtras.h"
+#include "llvm/Analysis/BranchProbabilityInfo.h"
 #include "llvm/Analysis/Loads.h"
 #include "llvm/Analysis/LoopInfo.h"
 #include "llvm/Analysis/LoopIterator.h"
@@ -53,6 +55,12 @@ class PlainCFGBuilder {
   // Loop versioning for alias metadata.
   LoopVersioning *LVer;
 
+  // Lazily provides branch probabilities for the incoming IR.
+  function_ref<const BranchProbabilityInfo &()> GetBPI;
+
+  // The BranchProbabilityInfo returned by GetBPI, cached on first use.
+  const BranchProbabilityInfo *BPI = nullptr;
+
   // Vectorization plan that we are working on.
   std::unique_ptr<VPlan> Plan;
 
@@ -79,10 +87,12 @@ class PlainCFGBuilder {
 #endif
   VPValue *getOrCreateVPOperand(Value *IRVal);
   void createVPInstructionsForVPBB(VPBasicBlock *VPBB, BasicBlock *BB);
+  VPIRMetadata getTerminatorMetadata(Instruction &Term);
 
 public:
-  PlainCFGBuilder(Loop *Lp, LoopInfo *LI, LoopVersioning *LVer, Type *IdxTy)
-      : TheLoop(Lp), LI(LI), LVer(LVer),
+  PlainCFGBuilder(Loop *Lp, LoopInfo *LI, LoopVersioning *LVer, Type *IdxTy,
+                  function_ref<const BranchProbabilityInfo &()> GetBPI)
+      : TheLoop(Lp), LI(LI), LVer(LVer), GetBPI(GetBPI),
         Plan(std::make_unique<VPlan>(Lp, IdxTy)) {}
 
   /// Build plain CFG for TheLoop and connect it to Plan's entry.
@@ -183,6 +193,26 @@ VPValue *PlainCFGBuilder::getOrCreateVPOperand(Value *IRVal) {
   return NewVPVal;
 }
 
+// Returns the metadata to preserve for terminator \p Term.
+VPIRMetadata PlainCFGBuilder::getTerminatorMetadata(Instruction &Term) {
+  VPIRMetadata MD(Term);
+  if (!GetBPI || MD.getMetadata(LLVMContext::MD_prof))
+    return MD;
+  // Estimates are only read for edges inside the loop region.
+  if (!TheLoop->isInnermost() || Term.getParent() == TheLoop->getLoopLatch())
+    return MD;
+  // The weights describe the edges leaving Term in the order of its successors,
+  // matching the successor order of the VPBasicBlock created for Term's parent.
+  if (!BPI)
+    BPI = &GetBPI();
+  auto Weights = map_to_vector(seq(Term.getNumSuccessors()), [&](unsigned I) {
+    return BPI->getEdgeProbability(Term.getParent(), I).getNumerator();
+  });
+  MD.setEstimatedBranchWeights(
+      MDBuilder(Plan->getContext()).createBranchWeights(Weights));
+  return MD;
+}
+
 // Create new VPInstructions in a VPBasicBlock, given its BasicBlock
 // counterpart. This function must be invoked in RPO so that the operands of a
 // VPInstruction in \p BB have been visited before (except for Phi nodes).
@@ -207,7 +237,8 @@ void PlainCFGBuilder::createVPInstructionsForVPBB(VPBasicBlock *VPBB,
       // recipes.
       VPValue *Cond = getOrCreateVPOperand(Br->getCondition());
       VPIRBuilder.createNaryOp(VPInstruction::BranchOnCond, {Cond}, Inst, {},
-                               VPIRMetadata(*Inst), Inst->getDebugLoc());
+                               getTerminatorMetadata(*Inst),
+                               Inst->getDebugLoc());
       continue;
     }
 
@@ -219,7 +250,8 @@ void PlainCFGBuilder::createVPInstructionsForVPBB(VPBasicBlock *VPBB,
       for (auto Case : SI->cases())
         Ops.push_back(getOrCreateVPOperand(Case.getCaseValue()));
       VPIRBuilder.createNaryOp(Instruction::Switch, Ops, Inst, {},
-                               VPIRMetadata(*Inst), Inst->getDebugLoc());
+                               getTerminatorMetadata(*Inst),
+                               Inst->getDebugLoc());
       continue;
     }
 
@@ -598,11 +630,11 @@ static void addInitialSkeleton(VPlan &Plan, Type *InductionTy,
 /// To make RUN_VPLAN_PASS print initial VPlan.
 static void printAfterInitialConstruction(VPlan &) {}
 
-std::unique_ptr<VPlan>
-VPlanTransforms::buildVPlan0(Loop *TheLoop, LoopInfo &LI, Type *InductionTy,
-                             PredicatedScalarEvolution &PSE,
-                             LoopVersioning *LVer) {
-  PlainCFGBuilder Builder(TheLoop, &LI, LVer, InductionTy);
+std::unique_ptr<VPlan> VPlanTransforms::buildVPlan0(
+    Loop *TheLoop, LoopInfo &LI, Type *InductionTy,
+    PredicatedScalarEvolution &PSE, LoopVersioning *LVer,
+    function_ref<const BranchProbabilityInfo &()> GetBPI) {
+  PlainCFGBuilder Builder(TheLoop, &LI, LVer, InductionTy, GetBPI);
   std::unique_ptr<VPlan> VPlan0 = Builder.buildPlainCFG();
   addInitialSkeleton(*VPlan0, InductionTy, PSE, TheLoop);
   simplifyLiveInsWithSCEV(*VPlan0, PSE);
diff --git a/llvm/lib/Transforms/Vectorize/VPlanHelpers.h b/llvm/lib/Transforms/Vectorize/VPlanHelpers.h
index 2838a9e4895a6..a4ae2d57548a5 100644
--- a/llvm/lib/Transforms/Vectorize/VPlanHelpers.h
+++ b/llvm/lib/Transforms/Vectorize/VPlanHelpers.h
@@ -356,9 +356,9 @@ struct VPCostContext {
   /// transform replaced the original recipe.
   void invalidateWideningDecision(Instruction *I, ElementCount VF);
 
-  /// \returns how much the cost of a predicated block should be divided by.
-  /// Forwards to LoopVectorizationCostModel::getPredBlockCostDivisor.
-  uint64_t getPredBlockCostDivisor(BasicBlock *BB) const;
+  /// \returns how much the cost of the block predicated by replicate region
+  /// \p Region should be divided by.
+  uint64_t getPredBlockCostDivisor(const VPRegionBlock *Region) const;
 
   /// Returns true if \p I is known to be scalarized at \p VF.
   bool willBeScalarized(Instruction *I, ElementCount VF) const;
diff --git a/llvm/lib/Transforms/Vectorize/VPlanPredicator.cpp b/llvm/lib/Transforms/Vectorize/VPlanPredicator.cpp
index cbc41e82fa87c..e056b2b613b33 100644
--- a/llvm/lib/Transforms/Vectorize/VPlanPredicator.cpp
+++ b/llvm/lib/Transforms/Vectorize/VPlanPredicator.cpp
@@ -406,8 +406,8 @@ void VPlanTransforms::introduceMasksAndLinearize(VPlan &Plan) {
       Header);
   // Non-outer regions with VPBBs only are supported at the moment.
   auto Blocks = to_vector(VPBlockUtils::blocksAs<VPBasicBlock>(RPOT));
-  DenseMap<const VPBasicBlock *, std::optional<BlockFrequency>> Frequencies =
-      vputils::computeExecutionFrequencies(Blocks);
+  DenseMap<const VPBasicBlock *, std::optional<VPExecutionFrequency>>
+      Frequencies = vputils::computeExecutionFrequencies(Blocks);
 
   VPPredicator Predicator(Plan);
   for (VPBasicBlock *VPBB : Blocks) {
@@ -423,7 +423,7 @@ void VPlanTransforms::introduceMasksAndLinearize(VPlan &Plan) {
 
     // Mask all VPInstructions in the block and record the frequency with
     // which the masked recipes execute.
-    std::optional<BlockFrequency> Freq = Frequencies.lookup(VPBB);
+    std::optional<VPExecutionFrequency> Freq = Frequencies.lookup(VPBB);
     for (VPRecipeBase &R : *VPBB) {
       auto *VPI = dyn_cast<VPInstruction>(&R);
       if (!VPI)
diff --git a/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp b/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp
index ff212663a3fac..9a0b7dddfb963 100644
--- a/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp
+++ b/llvm/lib/Transforms/Vectorize/VPlanRecipes.cpp
@@ -2128,41 +2128,43 @@ void VPIRPhi::printRecipe(raw_ostream &O, const Twine &Indent,
 void VPIRMetadata::applyMetadata(Instruction &I) const {
   if (Metadata.empty())
     return;
-  // The execution frequency is VPlan-internal and must not reach IR.
+  // Frequencies and estimated branch weights are VPlan-internal and must not
+  // reach IR.
   unsigned ExecFreqKind = getMDKindID(ExecutionFrequencyMDName);
+  unsigned EstProfKind = getMDKindID(EstimatedProfileMDName);
   for (const auto &[Kind, Node] : Metadata)
-    if (Kind != ExecFreqKind)
+    if (Kind != ExecFreqKind && Kind != EstProfKind)
       I.setMetadata(Kind, Node);
 }
 
 /// Returns the execution frequency recorded in \p Node.
-static BlockFrequency getExecutionFrequencyFromMD(const MDNode *Node) {
+static VPExecutionFrequency getExecutionFrequencyFromMD(const MDNode *Node) {
+  assert(Node->getNumOperands() <= 2 && "unexpected frequency node shape");
   uint64_t Freq =
       mdconst::extract<ConstantInt>(Node->getOperand(0))->getZExtValue();
   assert(Freq <= vputils::AlwaysExecutesFreq &&
          "frequency cannot exceed the one of an always executing block");
-  return BlockFrequency(Freq);
+  return {BlockFrequency(Freq), Node->getNumOperands() == 2};
 }
 
-void VPIRMetadata::setExecutionFrequency(std::optional<BlockFrequency> Freq,
-                                         LLVMContext &Ctx) {
+void VPIRMetadata::setExecutionFrequency(
+    std::optional<VPExecutionFrequency> Freq, LLVMContext &Ctx) {
   // A recipe that never or always executes needs no annotation.
-  if (!Freq || Freq->getFrequency() == 0 ||
-      Freq->getFrequency() == vputils::AlwaysExecutesFreq)
+  if (!Freq || Freq->Freq.getFrequency() == 0 ||
+      Freq->Freq.getFrequency() == vputils::AlwaysExecutesFreq)
     return;
-  Constant *Frequency =
-      ConstantInt::get(Type::getInt64Ty(Ctx), Freq->getFrequency());
-...
[truncated]

``````````

</details>


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


More information about the llvm-commits mailing list