[llvm] [LV] NFC: Move PHI/result update out of transformToPartialReduction. (PR #228436)

via llvm-commits llvm-commits at lists.llvm.org
Fri Oct 2 06:20:16 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-llvm-transforms

Author: Sander de Smalen (sdesmalen-arm)

<details>
<summary>Changes</summary>

That lets `transformToPartialReduction` focus on just creating a partial reduction expression for each link in the chain, whereas `createPartialReductions` updates the PHI and ReductionResult to complete the work for the chain. It removes the need to know that the partial reduction is part of a 'chain' in `transformToPartialReduction`.

I'll rebase this with the better terminology after #<!-- -->222377 gets merged.

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


1 Files Affected:

- (modified) llvm/lib/Transforms/Vectorize/VPlanTransforms.cpp (+49-40) 


``````````diff
diff --git a/llvm/lib/Transforms/Vectorize/VPlanTransforms.cpp b/llvm/lib/Transforms/Vectorize/VPlanTransforms.cpp
index cc1c9a683e404..8490810852b45 100644
--- a/llvm/lib/Transforms/Vectorize/VPlanTransforms.cpp
+++ b/llvm/lib/Transforms/Vectorize/VPlanTransforms.cpp
@@ -5232,43 +5232,6 @@ static void transformToPartialReduction(const VPPartialReductionChain &Chain,
   VPExpressionRecipe *E = createPartialReductionExpression(PartialRed);
   E->insertBefore(WidenRecipe);
   PartialRed->replaceAllUsesWith(E);
-
-  // We only need to update the PHI node once, which is when we find the
-  // last reduction in the chain.
-  if (!IsLastInChain)
-    return;
-
-  // Scale the PHI and ReductionStartVector by the VFScaleFactor
-  assert(RdxPhi->getVFScaleFactor() == 1 && "scale factor must not be set");
-  RdxPhi->setVFScaleFactor(Chain.ScaleFactor);
-
-  auto *StartInst = cast<VPInstruction>(RdxPhi->getStartValue());
-  assert(StartInst->getOpcode() == VPInstruction::ReductionStartVector);
-  auto *NewScaleFactor = Plan.getConstantInt(32, Chain.ScaleFactor);
-  StartInst->setOperand(2, NewScaleFactor);
-
-  // If this is the last value in a sub-reduction chain, then update the PHI
-  // node to start at `0` and update the reduction-result to subtract from
-  // the PHI's start value.
-  if (Chain.RK != RecurKind::Sub && Chain.RK != RecurKind::FSub)
-    return;
-
-  VPValue *OldStartValue = StartInst->getOperand(0);
-  StartInst->setOperand(0, StartInst->getOperand(1));
-
-  // Replace reduction_result by 'sub (startval, reductionresult)'.
-  VPInstruction *RdxResult = vputils::findComputeReductionResult(RdxPhi);
-  assert(RdxResult && "Could not find reduction result");
-
-  VPBuilder Builder = VPBuilder::getToInsertAfter(RdxResult);
-  unsigned SubOpc = Chain.RK == RecurKind::FSub ? Instruction::BinaryOps::FSub
-                                                : Instruction::BinaryOps::Sub;
-  VPInstruction *NewResult = Builder.createNaryOp(
-      SubOpc, {OldStartValue, RdxResult}, VPIRFlags::getDefaultFlags(SubOpc),
-      RdxPhi->getDebugLoc());
-  RdxResult->replaceUsesWithIf(
-      NewResult,
-      [&NewResult](VPUser &U, unsigned Idx) { return &U != NewResult; });
 }
 
 /// Returns the cost of a link in a partial-reduction chain for a given VF.
@@ -5504,6 +5467,42 @@ getScaledReductions(VPReductionPHIRecipe *RedPhiR) {
 }
 } // namespace
 
+// Scale the PHI and ReductionStartVector by \p Factor and if the recurrence is
+// a sub-recurrence, negate the reduction result.
+static void updatePartialReductionPhiAndResult(VPlan &Plan,
+                                               VPReductionPHIRecipe *Phi,
+                                               unsigned Factor, RecurKind RK) {
+  assert(Phi->getVFScaleFactor() == 1 && "scale factor must not be set");
+  Phi->setVFScaleFactor(Factor);
+
+  auto *StartInst = cast<VPInstruction>(Phi->getStartValue());
+  assert(StartInst->getOpcode() == VPInstruction::ReductionStartVector);
+  auto *NewScaleFactor = Plan.getConstantInt(32, Factor);
+  StartInst->setOperand(2, NewScaleFactor);
+
+  if (RK != RecurKind::Sub && RK != RecurKind::FSub)
+    return;
+
+  // Update the PHI node to start at `0` and update the reduction-result
+  // to subtract from the PHI's start value.
+  VPValue *OldStartValue = StartInst->getOperand(0);
+  StartInst->setOperand(0, StartInst->getOperand(1));
+
+  // Replace reduction_result by 'sub (startval, reductionresult)'.
+  VPInstruction *RdxResult = vputils::findComputeReductionResult(Phi);
+  assert(RdxResult && "Could not find reduction result");
+
+  VPBuilder Builder = VPBuilder::getToInsertAfter(RdxResult);
+  unsigned SubOpc = RK == RecurKind::FSub ? Instruction::BinaryOps::FSub
+                                          : Instruction::BinaryOps::Sub;
+  VPInstruction *NewResult = Builder.createNaryOp(
+      SubOpc, {OldStartValue, RdxResult}, VPIRFlags::getDefaultFlags(SubOpc),
+      Phi->getDebugLoc());
+  RdxResult->replaceUsesWithIf(
+      NewResult,
+      [&NewResult](VPUser &U, unsigned Idx) { return &U != NewResult; });
+}
+
 void VPlanTransforms::createPartialReductions(VPlan &Plan,
                                               VPCostContext &CostCtx,
                                               VFRange &Range) {
@@ -5668,9 +5667,19 @@ void VPlanTransforms::createPartialReductions(VPlan &Plan,
       Chains.clear();
   }
 
-  for (auto &[Phi, Chains] : ChainsByPhi)
-    for (const VPPartialReductionChain &Chain : Chains)
-      transformToPartialReduction(Chain, Plan, Phi);
+  for (auto &[Phi, Chain] : ChainsByPhi) {
+    if (Chain.empty())
+      continue;
+
+    for (const VPPartialReductionChain &Link : Chain)
+      transformToPartialReduction(Link, Plan, Phi);
+
+    // After transforming all links in the chain, the PHI node and result need
+    // updating. Note that we can pick any link in the chain for this, as the
+    // ScaleFactor and RecurKind must match for all links in the chain.
+    const VPPartialReductionChain &Link = Chain[0];
+    updatePartialReductionPhiAndResult(Plan, Phi, Link.ScaleFactor, Link.RK);
+  }
 }
 
 void VPlanTransforms::makeMemOpWideningDecisions(VPlan &Plan, VFRange &Range,

``````````

</details>


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


More information about the llvm-commits mailing list