[llvm] [DAGCombiner] Reassociate chains of vector reductions (PR #206471)
Luke Lau via llvm-commits
llvm-commits at lists.llvm.org
Thu Jul 2 01:57:03 PDT 2026
================
@@ -1401,6 +1401,42 @@ SDValue DAGCombiner::reassociateReduction(unsigned RedOpc, unsigned Opc,
SDValue Op2 = DAG.getNode(Opc, DL, VT, B, D);
return DAG.getNode(Opc, DL, VT, Red, Op2);
}
+
+ // Reassociate a reduction chain so two reductions become adjacent and the
+ // folds above can merge them:
+ // op(vecreduce(X), op(vecreduce(Y), Z))
+ // -> op(vecreduce(op(X, Y)), Z)
+ // Applied to fixpoint by the combiner worklist, this collapses an
+ // arbitrarily long chain of reductions (such as the left-leaning chain SLP
+ // emits) into a single reduction.
+ auto FoldReductionChain = [&](SDValue Red0, SDValue Chain) -> SDValue {
+ SDValue X, Y, Z, RedY;
+ if (!sd_match(Red0, m_OneUse(m_UnaryOp(RedOpc, m_Value(X)))) ||
+ !sd_match(
+ Chain,
+ m_OneUse(m_c_BinOp(
+ Opc, m_Value(RedY, m_OneUse(m_UnaryOp(RedOpc, m_Value(Y)))),
+ m_Value(Z, m_Unless(m_UnaryOp(RedOpc, m_Value())))))) ||
+ X.getValueType() != Y.getValueType() ||
+ !hasOperation(Opc, X.getValueType()) ||
+ !TLI.shouldReassociateReduction(RedOpc, VT))
+ return SDValue();
+ if ((Opc == ISD::FADD || Opc == ISD::FMUL) &&
+ (!Chain->getFlags().hasAllowReassociation() ||
+ !Red0->getFlags().hasAllowReassociation() ||
+ !RedY->getFlags().hasAllowReassociation()))
+ return SDValue();
+ SelectionDAG::FlagInserter FlagsInserter(
+ DAG, Flags & Chain->getFlags() & Red0->getFlags() & RedY->getFlags());
+ SDValue Sum = DAG.getNode(Opc, DL, X.getValueType(), X, Y);
+ SDValue Red = DAG.getNode(RedOpc, DL, VT, Sum);
----------------
lukel97 wrote:
Nit, name this Op to match the combines above?
```suggestion
SDValue Op = DAG.getNode(Opc, DL, X.getValueType(), X, Y);
SDValue Red = DAG.getNode(RedOpc, DL, VT, Op);
```
https://github.com/llvm/llvm-project/pull/206471
More information about the llvm-commits
mailing list