[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