[llvm] [NVPTX] Scalarize contract FMUL v2f32 to enable FMA fusion (PR #192815)
Princeton Ferro via llvm-commits
llvm-commits at lists.llvm.org
Sat Apr 18 16:32:18 PDT 2026
https://github.com/Prince781 updated https://github.com/llvm/llvm-project/pull/192815
>From c10c0440a6e9a49ee8aab73d41b0a5ac1732555e Mon Sep 17 00:00:00 2001
From: Princeton Ferro <pferro at nvidia.com>
Date: Sat, 18 Apr 2026 16:23:36 -0700
Subject: [PATCH] [NVPTX] Scalarize contract FMUL v2f32 to enable FMA fusion
By default, SM100+ legalizes FMUL for v2f32, which prevents FMA fusion
with scalar FADD on individual lanes. Scalarize an FMUL v2f32 when the
contract flag is set or allowFMA() is enabled, all lanes are consumed via
EXTRACT_VECTOR_ELT, and each extract has a single FADD user carrying the
contract flag.
Also convert PerformFADDCombine and PerformFADDCombineWithOperands from
static free functions to member functions of NVPTXTargetLowering so they
can call PerformScalarizeV2F32Op directly.
---
llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp | 62 +++++++++++++-------
llvm/lib/Target/NVPTX/NVPTXISelLowering.h | 10 ++++
llvm/test/CodeGen/NVPTX/fp-contract-f32x2.ll | 29 +++++++++
3 files changed, 80 insertions(+), 21 deletions(-)
diff --git a/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp b/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
index a5fd0a8724762..b403524daed82 100644
--- a/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
@@ -5748,15 +5748,12 @@ PerformADDCombineWithOperands(SDNode *N, SDValue N0, SDValue N1,
return SDValue();
}
-static SDValue
-PerformFADDCombineWithOperands(SDNode *N, SDValue N0, SDValue N1,
- TargetLowering::DAGCombinerInfo &DCI,
- CodeGenOptLevel OptLevel) {
+SDValue NVPTXTargetLowering::PerformFADDCombineWithOperands(
+ SDNode *N, SDValue N0, SDValue N1, TargetLowering::DAGCombinerInfo &DCI,
+ CodeGenOptLevel OptLevel) const {
EVT VT = N0.getValueType();
if (N0.getOpcode() == ISD::FMUL) {
- const auto *TLI = static_cast<const NVPTXTargetLowering *>(
- &DCI.DAG.getTargetLoweringInfo());
- if (!(TLI->allowFMA(DCI.DAG.getMachineFunction(), OptLevel) ||
+ if (!(allowFMA(DCI.DAG.getMachineFunction(), OptLevel) ||
(N->getFlags().hasAllowContract() &&
N0->getFlags().hasAllowContract())))
return SDValue();
@@ -6133,6 +6130,24 @@ static bool isNonCoalescableBuildVector(const SDValue &BV) {
return std::abs(Idx0->getSExtValue() - Idx1->getSExtValue()) != 1;
}
+/// Return true if FMUL v2f32 node \p N may be scalarized to fold each lane's
+/// product into a scalar FMA.
+bool NVPTXTargetLowering::MayFoldFMULIntoFMA(SDNode *N, MachineFunction &MF,
+ CodeGenOptLevel OptLevel) const {
+ if (N->getOpcode() != ISD::FMUL || N->getValueType(0) != MVT::v2f32)
+ return false;
+ const bool GlobalFMA = allowFMA(MF, OptLevel);
+ if (!N->getFlags().hasAllowContract() && !GlobalFMA)
+ return false;
+ return all_of(N->users(), [GlobalFMA](const SDNode *U) {
+ if (U->getOpcode() != ISD::EXTRACT_VECTOR_ELT || !U->hasOneUse())
+ return false;
+ const SDNode *FAdd = *U->users().begin();
+ return FAdd->getOpcode() == ISD::FADD &&
+ (GlobalFMA || FAdd->getFlags().hasAllowContract());
+ });
+}
+
/// Scalarize a v2f32 arithmetic node (FADD, FMUL, FSUB, FMA) when at least
/// one operand is a BUILD_VECTOR that repacks values from non-adjacent register
/// pairs. Without this combine the BUILD_VECTOR forces allocation of a
@@ -6156,15 +6171,20 @@ static bool isNonCoalescableBuildVector(const SDValue &BV) {
/// r0: f32 = fma a0, t1, c0
/// r1: f32 = fma a1, t2, c1
/// t4: v2f32 = BUILD_VECTOR r0, r1
-static SDValue PerformScalarizeV2F32Op(SDNode *N,
- TargetLowering::DAGCombinerInfo &DCI) {
+///
+/// Also scalarizes an FMUL v2f32 when MayFoldFMULIntoFMA is true. On sm100+,
+/// a wider float vector fmul legalizes to v2f32 operations, blocking the scalar
+/// FADD->FMA combiner. Scalarizing exposes each product to the FADD combiner,
+/// which can then fuse fadd(fmul(a, b), acc) -> fma.rn.f32.
+SDValue NVPTXTargetLowering::PerformScalarizeV2F32Op(
+ SDNode *N, TargetLowering::DAGCombinerInfo &DCI,
+ CodeGenOptLevel OptLevel) const {
EVT VT = N->getValueType(0);
if (VT != MVT::v2f32)
return SDValue();
- // Only scalarize when at least one operand is a BUILD_VECTOR whose elements
- // are guaranteed to reside in different register pairs.
- if (none_of(N->ops(), isNonCoalescableBuildVector))
+ if (none_of(N->ops(), isNonCoalescableBuildVector) &&
+ !MayFoldFMULIntoFMA(N, DCI.DAG.getMachineFunction(), OptLevel))
return SDValue();
SelectionDAG &DAG = DCI.DAG;
@@ -6195,17 +6215,17 @@ static SDValue PerformScalarizeV2F32Op(SDNode *N,
return DAG.getNode(ISD::BUILD_VECTOR, DL, VT, Res0, Res1);
}
-/// PerformFADDCombine - Target-specific dag combine xforms for ISD::FADD.
-///
-static SDValue PerformFADDCombine(SDNode *N,
- TargetLowering::DAGCombinerInfo &DCI,
- CodeGenOptLevel OptLevel) {
+/// Target-specific dag combine xforms for ISD::FADD.
+SDValue
+NVPTXTargetLowering::PerformFADDCombine(SDNode *N,
+ TargetLowering::DAGCombinerInfo &DCI,
+ CodeGenOptLevel OptLevel) const {
+ if (SDValue Result = PerformScalarizeV2F32Op(N, DCI, OptLevel))
+ return Result;
+
SDValue N0 = N->getOperand(0);
SDValue N1 = N->getOperand(1);
- if (SDValue Result = PerformScalarizeV2F32Op(N, DCI))
- return Result;
-
EVT VT = N0.getValueType();
if (VT.isVector() || !(VT == MVT::f32 || VT == MVT::f64))
return SDValue();
@@ -7059,7 +7079,7 @@ SDValue NVPTXTargetLowering::PerformDAGCombine(SDNode *N,
case ISD::FMA:
case ISD::FMUL:
case ISD::FSUB:
- return PerformScalarizeV2F32Op(N, DCI);
+ return PerformScalarizeV2F32Op(N, DCI, OptLevel);
case ISD::FMAXNUM:
case ISD::FMINNUM:
case ISD::FMAXIMUM:
diff --git a/llvm/lib/Target/NVPTX/NVPTXISelLowering.h b/llvm/lib/Target/NVPTX/NVPTXISelLowering.h
index d667662615e9a..3d22e199cd29f 100644
--- a/llvm/lib/Target/NVPTX/NVPTXISelLowering.h
+++ b/llvm/lib/Target/NVPTX/NVPTXISelLowering.h
@@ -238,6 +238,16 @@ class NVPTXTargetLowering : public TargetLowering {
void ReplaceNodeResults(SDNode *N, SmallVectorImpl<SDValue> &Results,
SelectionDAG &DAG) const override;
SDValue PerformDAGCombine(SDNode *N, DAGCombinerInfo &DCI) const override;
+
+ bool MayFoldFMULIntoFMA(SDNode *N, MachineFunction &MF,
+ CodeGenOptLevel OptLevel) const;
+ SDValue PerformScalarizeV2F32Op(SDNode *N, DAGCombinerInfo &DCI,
+ CodeGenOptLevel OptLevel) const;
+ SDValue PerformFADDCombineWithOperands(SDNode *N, SDValue N0, SDValue N1,
+ DAGCombinerInfo &DCI,
+ CodeGenOptLevel OptLevel) const;
+ SDValue PerformFADDCombine(SDNode *N, DAGCombinerInfo &DCI,
+ CodeGenOptLevel OptLevel) const;
};
} // namespace llvm
diff --git a/llvm/test/CodeGen/NVPTX/fp-contract-f32x2.ll b/llvm/test/CodeGen/NVPTX/fp-contract-f32x2.ll
index c4d4dfcc618d8..67ac240be3318 100644
--- a/llvm/test/CodeGen/NVPTX/fp-contract-f32x2.ll
+++ b/llvm/test/CodeGen/NVPTX/fp-contract-f32x2.ll
@@ -110,3 +110,32 @@ define <2 x float> @t3(<2 x float> %a, <2 x float> %b, <2 x float> %c) {
%v1 = fadd contract <2 x float> %v0, %c
ret <2 x float> %v1
}
+
+;; Dot product represented as `fmul+reduce.fadd`. With "contract" FMF, we
+;; should scalarize the fmul.f32x2 to allow for FMA fusion.
+define float @dot_reduce_contract(<8 x float> %a, <8 x float> %b) {
+; CHECK-LABEL: dot_reduce_contract(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<25>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param.v4.b32 {%r1, %r2, %r3, %r4}, [dot_reduce_contract_param_0+16];
+; CHECK-NEXT: ld.param.v4.b32 {%r5, %r6, %r7, %r8}, [dot_reduce_contract_param_0];
+; CHECK-NEXT: ld.param.v4.b32 {%r9, %r10, %r11, %r12}, [dot_reduce_contract_param_1+16];
+; CHECK-NEXT: ld.param.v4.b32 {%r13, %r14, %r15, %r16}, [dot_reduce_contract_param_1];
+; CHECK-NEXT: fma.rn.f32 %r17, %r5, %r13, 0f00000000;
+; CHECK-NEXT: fma.rn.f32 %r18, %r6, %r14, %r17;
+; CHECK-NEXT: fma.rn.f32 %r19, %r7, %r15, %r18;
+; CHECK-NEXT: fma.rn.f32 %r20, %r8, %r16, %r19;
+; CHECK-NEXT: fma.rn.f32 %r21, %r1, %r9, %r20;
+; CHECK-NEXT: fma.rn.f32 %r22, %r2, %r10, %r21;
+; CHECK-NEXT: fma.rn.f32 %r23, %r3, %r11, %r22;
+; CHECK-NEXT: fma.rn.f32 %r24, %r4, %r12, %r23;
+; CHECK-NEXT: st.param.b32 [func_retval0], %r24;
+; CHECK-NEXT: ret;
+ %mul = fmul contract <8 x float> %a, %b
+ %red = call contract float @llvm.vector.reduce.fadd.v8f32(float 0.0, <8 x float> %mul)
+ ret float %red
+}
+
+declare float @llvm.vector.reduce.fadd.v8f32(float, <8 x float>)
More information about the llvm-commits
mailing list