[llvm] [AArch64] Add pattern for reassociating 3-way FADD to form FADDP (PR #216687)
via llvm-commits
llvm-commits at lists.llvm.org
Mon Aug 17 03:30:32 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-backend-aarch64
Author: Hari Limaye (hazzlim)
<details>
<summary>Changes</summary>
Combine `fadd(fadd(a[0], b), a[1]) -> fadd(fadd(a[0], a[1]), b)` to match FADDP patterns.
---
Full diff: https://github.com/llvm/llvm-project/pull/216687.diff
2 Files Affected:
- (modified) llvm/lib/Target/AArch64/AArch64ISelLowering.cpp (+68-14)
- (modified) llvm/test/CodeGen/AArch64/arm64-addp.ll (+145)
``````````diff
diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index e9fdaaab967af..cfc39c3bdaab8 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -22290,8 +22290,71 @@ static SDValue performANDCombine(SDNode *N,
return SDValue();
}
+static bool hasPairwiseAdd(unsigned Opcode, EVT VT, bool FullFP16) {
+ switch (Opcode) {
+ case ISD::STRICT_FADD:
+ case ISD::FADD:
+ return (FullFP16 && VT == MVT::f16) || VT == MVT::f32 || VT == MVT::f64;
+ case ISD::ADD:
+ return VT == MVT::i64;
+ default:
+ return false;
+ }
+}
+
+static SDValue performPairwiseFADDCombine(SDNode *N,
+ TargetLowering::DAGCombinerInfo &DCI,
+ const AArch64Subtarget &Subtarget) {
+ assert(N->getFlags().hasAllowReassociation() &&
+ "Expected AllowReassociation!");
+
+ if (!DCI.isAfterLegalizeDAG())
+ return SDValue();
+
+ EVT VT = N->getValueType(0);
+ if (!Subtarget.hasNEON() ||
+ !hasPairwiseAdd(N->getOpcode(), VT, Subtarget.hasFullFP16()))
+ return SDValue();
+
+ using namespace llvm::SDPatternMatch;
+
+ SDValue OuterExtract;
+ SDValue InnerAdd;
+ uint64_t OuterLane;
+ if (!sd_match(
+ N, m_FAdd(m_Value(OuterExtract,
+ m_ExtractElt(m_Value(), m_ConstInt(OuterLane))),
+ m_Value(InnerAdd, m_OneUse(m_FAdd(m_Value(), m_Value()))))))
+ return SDValue();
+
+ if (!InnerAdd->getFlags().hasAllowReassociation())
+ return SDValue();
+
+ SDValue OuterVec = OuterExtract.getOperand(0);
+ SDValue InnerExtract;
+ SDValue OtherVal;
+ uint64_t InnerLane;
+ if (!sd_match(InnerAdd, m_FAdd(m_Value(InnerExtract,
+ m_ExtractElt(m_Specific(OuterVec),
+ m_ConstInt(InnerLane))),
+ m_Value(OtherVal))))
+ return SDValue();
+
+ if (InnerLane > 1 || OuterLane > 1 || InnerLane == OuterLane)
+ return SDValue();
+
+ SDNodeFlags Flags = N->getFlags();
+ Flags &= InnerAdd->getFlags();
+ SDLoc DL(N);
+ SelectionDAG &DAG = DCI.DAG;
+ SDValue Pair =
+ DAG.getNode(ISD::FADD, DL, VT, InnerExtract, OuterExtract, Flags);
+ return DAG.getNode(ISD::FADD, DL, VT, Pair, OtherVal, Flags);
+}
+
static SDValue performFADDCombine(SDNode *N,
- TargetLowering::DAGCombinerInfo &DCI) {
+ TargetLowering::DAGCombinerInfo &DCI,
+ const AArch64Subtarget &Subtarget) {
SelectionDAG &DAG = DCI.DAG;
SDValue LHS = N->getOperand(0);
SDValue RHS = N->getOperand(1);
@@ -22323,19 +22386,10 @@ static SDValue performFADDCombine(SDNode *N,
if (SDValue R = ReassocComplex(RHS, LHS))
return R;
- return SDValue();
-}
+ if (SDValue R = performPairwiseFADDCombine(N, DCI, Subtarget))
+ return R;
-static bool hasPairwiseAdd(unsigned Opcode, EVT VT, bool FullFP16) {
- switch (Opcode) {
- case ISD::STRICT_FADD:
- case ISD::FADD:
- return (FullFP16 && VT == MVT::f16) || VT == MVT::f32 || VT == MVT::f64;
- case ISD::ADD:
- return VT == MVT::i64;
- default:
- return false;
- }
+ return SDValue();
}
static SDValue getPTest(SelectionDAG &DAG, EVT VT, SDValue Pg, SDValue Op,
@@ -31378,7 +31432,7 @@ SDValue AArch64TargetLowering::PerformDAGCombine(SDNode *N,
case ISD::AND:
return performANDCombine(N, DCI);
case ISD::FADD:
- return performFADDCombine(N, DCI);
+ return performFADDCombine(N, DCI, *Subtarget);
case ISD::INTRINSIC_WO_CHAIN:
return performIntrinsicCombine(N, DCI, Subtarget);
case ISD::ANY_EXTEND:
diff --git a/llvm/test/CodeGen/AArch64/arm64-addp.ll b/llvm/test/CodeGen/AArch64/arm64-addp.ll
index 64617bd0f359d..3aa5c0554c09e 100644
--- a/llvm/test/CodeGen/AArch64/arm64-addp.ll
+++ b/llvm/test/CodeGen/AArch64/arm64-addp.ll
@@ -164,3 +164,148 @@ entry:
%b = add <32 x i8> %s, %a
ret <32 x i8> %b
}
+
+define double @reassociate_three_term_fadd_lhs_lhs(<2 x double> %a, double %b) nounwind {
+; CHECK-LABEL: reassociate_three_term_fadd_lhs_lhs:
+; CHECK: // %bb.0:
+; CHECK-NEXT: faddp.2d d0, v0
+; CHECK-NEXT: fadd d0, d0, d1
+; CHECK-NEXT: ret
+ %lane0 = extractelement <2 x double> %a, i32 0
+ %lane1 = extractelement <2 x double> %a, i32 1
+ %inner = fadd fast double %lane1, %b
+ %result = fadd fast double %inner, %lane0
+ ret double %result
+}
+
+define double @reassociate_three_term_fadd_lhs_rhs(<2 x double> %a, double %b) nounwind {
+; CHECK-LABEL: reassociate_three_term_fadd_lhs_rhs:
+; CHECK: // %bb.0:
+; CHECK-NEXT: faddp.2d d0, v0
+; CHECK-NEXT: fadd d0, d0, d1
+; CHECK-NEXT: ret
+ %lane0 = extractelement <2 x double> %a, i32 0
+ %lane1 = extractelement <2 x double> %a, i32 1
+ %inner = fadd fast double %b, %lane1
+ %result = fadd fast double %inner, %lane0
+ ret double %result
+}
+
+define double @reassociate_three_term_fadd_rhs_lhs(<2 x double> %a, double %b) nounwind {
+; CHECK-LABEL: reassociate_three_term_fadd_rhs_lhs:
+; CHECK: // %bb.0:
+; CHECK-NEXT: faddp.2d d0, v0
+; CHECK-NEXT: fadd d0, d0, d1
+; CHECK-NEXT: ret
+ %lane0 = extractelement <2 x double> %a, i32 0
+ %lane1 = extractelement <2 x double> %a, i32 1
+ %inner = fadd fast double %lane1, %b
+ %result = fadd fast double %lane0, %inner
+ ret double %result
+}
+
+define double @reassociate_three_term_fadd_rhs_rhs(<2 x double> %a, double %b) nounwind {
+; CHECK-LABEL: reassociate_three_term_fadd_rhs_rhs:
+; CHECK: // %bb.0:
+; CHECK-NEXT: faddp.2d d0, v0
+; CHECK-NEXT: fadd d0, d0, d1
+; CHECK-NEXT: ret
+ %lane0 = extractelement <2 x double> %a, i32 0
+ %lane1 = extractelement <2 x double> %a, i32 1
+ %inner = fadd fast double %b, %lane1
+ %result = fadd fast double %lane0, %inner
+ ret double %result
+}
+
+define float @reassociate_three_term_fadd_f32(<2 x float> %a, float %b) nounwind {
+; CHECK-LABEL: reassociate_three_term_fadd_f32:
+; CHECK: // %bb.0:
+; CHECK-NEXT: // kill: def $d0 killed $d0 def $q0
+; CHECK-NEXT: faddp.2s s0, v0
+; CHECK-NEXT: fadd s0, s0, s1
+; CHECK-NEXT: ret
+ %lane0 = extractelement <2 x float> %a, i32 0
+ %lane1 = extractelement <2 x float> %a, i32 1
+ %inner = fadd reassoc float %lane1, %b
+ %result = fadd reassoc float %inner, %lane0
+ ret float %result
+}
+
+define half @reassociate_three_term_fadd_f16(<2 x half> %a, half %b) nounwind "target-features"="+fullfp16" {
+; CHECK-LABEL: reassociate_three_term_fadd_f16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: // kill: def $d0 killed $d0 def $q0
+; CHECK-NEXT: faddp.2h h0, v0
+; CHECK-NEXT: fadd h0, h0, h1
+; CHECK-NEXT: ret
+ %lane0 = extractelement <2 x half> %a, i32 0
+ %lane1 = extractelement <2 x half> %a, i32 1
+ %inner = fadd reassoc half %lane1, %b
+ %result = fadd reassoc half %inner, %lane0
+ ret half %result
+}
+
+define half @do_not_reassociate_f16_without_fullfp16(<2 x half> %a, half %b) nounwind {
+; CHECK-LABEL: do_not_reassociate_f16_without_fullfp16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: // kill: def $d0 killed $d0 def $q0
+; CHECK-NEXT: mov h2, v0[1]
+; CHECK-NEXT: fcvt s2, h2
+; CHECK-NEXT: fcvt s1, h1
+; CHECK-NEXT: fadd s1, s2, s1
+; CHECK-NEXT: fcvt h1, s1
+; CHECK-NEXT: fcvt s0, h0
+; CHECK-NEXT: fcvt s1, h1
+; CHECK-NEXT: fadd s0, s1, s0
+; CHECK-NEXT: fcvt h0, s0
+; CHECK-NEXT: ret
+ %lane0 = extractelement <2 x half> %a, i32 0
+ %lane1 = extractelement <2 x half> %a, i32 1
+ %inner = fadd reassoc half %lane1, %b
+ %result = fadd reassoc half %inner, %lane0
+ ret half %result
+}
+
+define double @do_not_reassociate_when_inner_lacks_reassoc(<2 x double> %a, double %b) nounwind {
+; CHECK-LABEL: do_not_reassociate_when_inner_lacks_reassoc:
+; CHECK: // %bb.0:
+; CHECK-NEXT: mov d2, v0[1]
+; CHECK-NEXT: fadd d1, d2, d1
+; CHECK-NEXT: fadd d0, d1, d0
+; CHECK-NEXT: ret
+ %lane0 = extractelement <2 x double> %a, i32 0
+ %lane1 = extractelement <2 x double> %a, i32 1
+ %inner = fadd double %lane1, %b
+ %result = fadd reassoc double %inner, %lane0
+ ret double %result
+}
+
+define double @do_not_reassociate_multiuse_inner(<2 x double> %a, double %b, ptr %p) nounwind {
+; CHECK-LABEL: do_not_reassociate_multiuse_inner:
+; CHECK: // %bb.0:
+; CHECK-NEXT: mov d2, v0[1]
+; CHECK-NEXT: fadd d1, d2, d1
+; CHECK-NEXT: str d1, [x0]
+; CHECK-NEXT: fadd d0, d1, d0
+; CHECK-NEXT: ret
+ %lane0 = extractelement <2 x double> %a, i32 0
+ %lane1 = extractelement <2 x double> %a, i32 1
+ %inner = fadd reassoc double %lane1, %b
+ store double %inner, ptr %p
+ %result = fadd reassoc double %inner, %lane0
+ ret double %result
+}
+
+define double @do_not_reassociate_strict_three_term_fadd(<2 x double> %a, double %b) nounwind {
+; CHECK-LABEL: do_not_reassociate_strict_three_term_fadd:
+; CHECK: // %bb.0:
+; CHECK-NEXT: mov d2, v0[1]
+; CHECK-NEXT: fadd d1, d2, d1
+; CHECK-NEXT: fadd d0, d1, d0
+; CHECK-NEXT: ret
+ %lane0 = extractelement <2 x double> %a, i32 0
+ %lane1 = extractelement <2 x double> %a, i32 1
+ %inner = fadd double %lane1, %b
+ %result = fadd double %inner, %lane0
+ ret double %result
+}
``````````
</details>
https://github.com/llvm/llvm-project/pull/216687
More information about the llvm-commits
mailing list