[llvm] b0a9b66 - [AMDGPU] Widen fminnum/fmaxnum bf16 to v2bf16 (#215042)
via llvm-commits
llvm-commits at lists.llvm.org
Tue Aug 11 11:55:43 PDT 2026
Author: LU-JOHN
Date: 2026-08-11T13:55:38-05:00
New Revision: b0a9b665638b0fc9dc97135238f977ee22fce15c
URL: https://github.com/llvm/llvm-project/commit/b0a9b665638b0fc9dc97135238f977ee22fce15c
DIFF: https://github.com/llvm/llvm-project/commit/b0a9b665638b0fc9dc97135238f977ee22fce15c.diff
LOG: [AMDGPU] Widen fminnum/fmaxnum bf16 to v2bf16 (#215042)
Perform scalar FMINNUM/FMAXNUM with bf16 more efficiently. Utilize
v2bf16 patterns.
Converting from a scalar to vector operation is implemented as a Promote
action.
---------
Signed-off-by: John Lu <John.Lu at amd.com>
Added:
Modified:
llvm/lib/CodeGen/SelectionDAG/LegalizeDAG.cpp
llvm/lib/Target/AMDGPU/SIISelLowering.cpp
llvm/lib/Target/AMDGPU/SIISelLowering.h
llvm/test/CodeGen/AMDGPU/bf16.ll
llvm/test/CodeGen/AMDGPU/minmax3-tree-reduction.ll
Removed:
################################################################################
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeDAG.cpp b/llvm/lib/CodeGen/SelectionDAG/LegalizeDAG.cpp
index 8fbd44918173f..e86d83a72ba75 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeDAG.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeDAG.cpp
@@ -5871,6 +5871,17 @@ void SelectionDAGLegalize::PromoteNode(SDNode *Node) {
case ISD::FMAXIMUMNUM:
case ISD::FPOW:
case ISD::FATAN2:
+ // Promote scalar operations to vector using SCALAR_TO_VECTOR
+ if (!OVT.isVector() && NVT.isVector() &&
+ NVT.getVectorElementType() == OVT) {
+ Tmp1 = DAG.getNode(ISD::SCALAR_TO_VECTOR, dl, NVT, Node->getOperand(0));
+ Tmp2 = DAG.getNode(ISD::SCALAR_TO_VECTOR, dl, NVT, Node->getOperand(1));
+ Tmp3 =
+ DAG.getNode(Node->getOpcode(), dl, NVT, Tmp1, Tmp2, Node->getFlags());
+ Results.push_back(DAG.getNode(ISD::EXTRACT_VECTOR_ELT, dl, OVT, Tmp3,
+ DAG.getConstant(0, dl, MVT::i32)));
+ break;
+ }
Tmp1 = DAG.getNode(ISD::FP_EXTEND, dl, NVT, Node->getOperand(0));
Tmp2 = DAG.getNode(ISD::FP_EXTEND, dl, NVT, Node->getOperand(1));
Tmp3 = DAG.getNode(Node->getOpcode(), dl, NVT, Tmp1, Tmp2);
@@ -6036,6 +6047,15 @@ void SelectionDAGLegalize::PromoteNode(SDNode *Node) {
case ISD::FEXP2:
case ISD::FEXP10:
case ISD::FCANONICALIZE:
+ // Promote scalar operations to vector using SCALAR_TO_VECTOR
+ if (!OVT.isVector() && NVT.isVector() &&
+ NVT.getVectorElementType() == OVT) {
+ Tmp1 = DAG.getNode(ISD::SCALAR_TO_VECTOR, dl, NVT, Node->getOperand(0));
+ Tmp2 = DAG.getNode(Node->getOpcode(), dl, NVT, Tmp1, Node->getFlags());
+ Results.push_back(DAG.getNode(ISD::EXTRACT_VECTOR_ELT, dl, OVT, Tmp2,
+ DAG.getConstant(0, dl, MVT::i32)));
+ break;
+ }
Tmp1 = DAG.getNode(ISD::FP_EXTEND, dl, NVT, Node->getOperand(0));
Tmp2 = DAG.getNode(Node->getOpcode(), dl, NVT, Tmp1);
Results.push_back(
diff --git a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
index c04bcfb0388f2..939a90b7d2461 100644
--- a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
+++ b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
@@ -244,16 +244,15 @@ SITargetLowering::SITargetLowering(const TargetMachine &TM,
// Only targets with packed bf16 instructions, e.g. gfx13.
if (Subtarget->hasBF16PackedInsts()) {
- // Turn fsub into fadd(x, fneg y) so it reuses the packed v_pk_add_bf16
- // path instead of promoting to f32.
- setOperationAction(ISD::FSUB, MVT::bf16, Expand);
- // Widen scalar fadd to a v2bf16 operation with an unused high lane.
- setOperationAction(ISD::FADD, MVT::bf16, Custom);
- // Widen scalar fcanonicalize to a v2bf16 operation with an unused high
+ // Don't use Expand for fsub - the DAG combiner will undo fadd+fneg back
+ // to fsub, causing a libcall (which doesn't exist for bf16). Instead,
+ // directly expand to widened v2bf16 operations.
+ setOperationAction(ISD::FSUB, MVT::bf16, Custom);
+ // Promote scalar operations to a v2bf16 operation with an unused high
// lane.
- setOperationAction(ISD::FCANONICALIZE, MVT::bf16, Custom);
- // Widen scalar fmul to a v2bf16 operation with an unused high lane.
- setOperationAction(ISD::FMUL, MVT::bf16, Custom);
+ for (unsigned Opc : {ISD::FADD, ISD::FMUL, ISD::FMAXNUM, ISD::FMINNUM,
+ ISD::FCANONICALIZE})
+ AddPromotedToType(Opc, MVT::bf16, MVT::v2bf16);
}
setOperationAction(ISD::FP_ROUND, MVT::bf16, Expand);
@@ -7716,6 +7715,7 @@ SDValue SITargetLowering::LowerOperation(SDValue Op, SelectionDAG &DAG) const {
case ISD::ABS:
case ISD::FABS:
case ISD::FNEG:
+ case ISD::FCANONICALIZE:
case ISD::BSWAP:
return splitUnaryVectorOp(Op, DAG);
case ISD::FP_TO_SINT_SAT:
@@ -7724,6 +7724,35 @@ SDValue SITargetLowering::LowerOperation(SDValue Op, SelectionDAG &DAG) const {
Op.getOperand(0).getValueType().getScalarType() == MVT::f32)
return splitUnaryVectorOp(Op, DAG);
return LowerFP_TO_INT_SAT(Op, DAG);
+ case ISD::FSUB:
+ if (Op.getValueType() == MVT::bf16) {
+ // Custom expansion:
+ // fsub bf16 %a, %b -> fadd v2bf16(widen %a), fneg v2bf16(widen %b)
+ // Then extract back to bf16.
+ //
+ // We create fneg on v2bf16 (not bf16) so the instruction selector can
+ // fold the negation into the packed add's neg_lo/neg_hi modifiers,
+ // generating a single v_pk_add_bf16 instruction. If we negate bf16 first,
+ // it becomes a separate v_xor instruction before widening.
+ SDLoc DL(Op);
+ SDValue Op0 = Op.getOperand(0);
+ SDValue Op1 = Op.getOperand(1);
+
+ // Widen both operands to v2bf16
+ SDValue Vec0 = DAG.getNode(ISD::SCALAR_TO_VECTOR, DL, MVT::v2bf16, Op0);
+ SDValue Vec1 = DAG.getNode(ISD::SCALAR_TO_VECTOR, DL, MVT::v2bf16, Op1);
+
+ // Create FNEG v2bf16 for the second operand
+ SDValue NegVec1 = DAG.getNode(ISD::FNEG, DL, MVT::v2bf16, Vec1);
+
+ // Perform FADD v2bf16
+ SDValue Result = DAG.getNode(ISD::FADD, DL, MVT::v2bf16, Vec0, NegVec1);
+
+ // Extract element 0 back to bf16
+ return DAG.getNode(ISD::EXTRACT_VECTOR_ELT, DL, MVT::bf16, Result,
+ DAG.getConstant(0, DL, MVT::i32));
+ }
+ return SDValue();
case ISD::FMINNUM:
case ISD::FMAXNUM:
return lowerFMINNUM_FMAXNUM(Op, DAG);
@@ -7761,16 +7790,9 @@ SDValue SITargetLowering::LowerOperation(SDValue Op, SelectionDAG &DAG) const {
case ISD::USUBSAT:
case ISD::SADDSAT:
case ISD::SSUBSAT:
- return splitBinaryVectorOp(Op, DAG);
case ISD::FADD:
case ISD::FMUL:
- if (Op.getValueType() == MVT::bf16)
- return lowerScalarBF16BinaryOp(Op, DAG);
return splitBinaryVectorOp(Op, DAG);
- case ISD::FCANONICALIZE:
- if (Op.getValueType() == MVT::bf16)
- return lowerScalarBF16FCanonicalize(Op, DAG);
- return splitUnaryVectorOp(Op, DAG);
case ISD::FCOPYSIGN:
return lowerFCOPYSIGN(Op, DAG);
case ISD::MUL:
@@ -8839,49 +8861,6 @@ SDValue SITargetLowering::lowerFP_ROUND(SDValue Op, SelectionDAG &DAG) const {
DAG.getTargetConstant(0, DL, MVT::i32));
}
-SDValue SITargetLowering::lowerScalarBF16BinaryOp(SDValue Op,
- SelectionDAG &DAG) const {
- assert(Subtarget->hasBF16PackedInsts());
-
- SDLoc DL(Op);
-
- auto WidenOperand = [&](SDValue Src) {
- if (Src.getOpcode() == ISD::FNEG) {
- SDValue WideSrc = DAG.getNode(ISD::SCALAR_TO_VECTOR, DL, MVT::v2bf16,
- Src.getOperand(0));
- return DAG.getNode(ISD::FNEG, DL, MVT::v2bf16, WideSrc, Src->getFlags());
- }
-
- return DAG.getNode(ISD::SCALAR_TO_VECTOR, DL, MVT::v2bf16, Src);
- };
-
- SDValue LHS = WidenOperand(Op.getOperand(0));
- SDValue RHS = WidenOperand(Op.getOperand(1));
- SDValue Result =
- DAG.getNode(Op.getOpcode(), DL, MVT::v2bf16, LHS, RHS, Op->getFlags());
-
- return DAG.getNode(ISD::EXTRACT_VECTOR_ELT, DL, MVT::bf16, Result,
- DAG.getConstant(0, DL, MVT::i32));
-}
-
-SDValue
-SITargetLowering::lowerScalarBF16FCanonicalize(SDValue Op,
- SelectionDAG &DAG) const {
- assert(Subtarget->hasBF16PackedInsts());
-
- SDLoc DL(Op);
- SDValue Src = Op.getOperand(0);
-
- // Widen to v2bf16, canonicalize with v_pk_mul_bf16, then extract.
- SDValue WideSrc = DAG.getNode(ISD::SCALAR_TO_VECTOR, DL, MVT::v2bf16, Src);
-
- SDValue Canonicalized =
- DAG.getNode(ISD::FCANONICALIZE, DL, MVT::v2bf16, WideSrc, Op->getFlags());
-
- return DAG.getNode(ISD::EXTRACT_VECTOR_ELT, DL, MVT::bf16, Canonicalized,
- DAG.getConstant(0, DL, MVT::i32));
-}
-
SDValue SITargetLowering::lowerFMINNUM_FMAXNUM(SDValue Op,
SelectionDAG &DAG) const {
EVT VT = Op.getValueType();
diff --git a/llvm/lib/Target/AMDGPU/SIISelLowering.h b/llvm/lib/Target/AMDGPU/SIISelLowering.h
index f3733ada688a5..4e667b70fdc91 100644
--- a/llvm/lib/Target/AMDGPU/SIISelLowering.h
+++ b/llvm/lib/Target/AMDGPU/SIISelLowering.h
@@ -169,8 +169,6 @@ class SITargetLowering final : public AMDGPUTargetLowering {
/// Custom lowering for ISD::FP_ROUND for MVT::f16.
SDValue lowerFP_ROUND(SDValue Op, SelectionDAG &DAG) const;
SDValue splitFP_ROUNDVectorOp(SDValue Op, SelectionDAG &DAG) const;
- SDValue lowerScalarBF16BinaryOp(SDValue Op, SelectionDAG &DAG) const;
- SDValue lowerScalarBF16FCanonicalize(SDValue Op, SelectionDAG &DAG) const;
SDValue lowerFMINNUM_FMAXNUM(SDValue Op, SelectionDAG &DAG) const;
SDValue lowerFMINIMUMNUM_FMAXIMUMNUM(SDValue Op, SelectionDAG &DAG) const;
SDValue lowerFMINIMUM_FMAXIMUM(SDValue Op, SelectionDAG &DAG) const;
diff --git a/llvm/test/CodeGen/AMDGPU/bf16.ll b/llvm/test/CodeGen/AMDGPU/bf16.ll
index 425ffdd5cd549..c95eeb1fbacd6 100644
--- a/llvm/test/CodeGen/AMDGPU/bf16.ll
+++ b/llvm/test/CodeGen/AMDGPU/bf16.ll
@@ -19220,10 +19220,7 @@ define bfloat @v_minnum_bf16(bfloat %a, bfloat %b) #0 {
; GFX1250: ; %bb.0:
; GFX1250-NEXT: s_wait_loadcnt_dscnt 0x0
; GFX1250-NEXT: s_wait_kmcnt 0x0
-; GFX1250-NEXT: v_dual_lshlrev_b32 v1, 16, v1 :: v_dual_lshlrev_b32 v0, 16, v0
-; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_1) | instskip(NEXT) | instid1(VALU_DEP_1)
-; GFX1250-NEXT: v_min_num_f32_e32 v0, v0, v1
-; GFX1250-NEXT: v_cvt_pk_bf16_f32 v0, v0, s0
+; GFX1250-NEXT: v_pk_min_num_bf16 v0, v0, v1
; GFX1250-NEXT: s_set_pc_i64 s[30:31]
%op = call bfloat @llvm.minnum.bf16(bfloat %a, bfloat %b)
ret bfloat %op
@@ -23718,10 +23715,7 @@ define bfloat @v_maxnum_bf16(bfloat %a, bfloat %b) #0 {
; GFX1250: ; %bb.0:
; GFX1250-NEXT: s_wait_loadcnt_dscnt 0x0
; GFX1250-NEXT: s_wait_kmcnt 0x0
-; GFX1250-NEXT: v_dual_lshlrev_b32 v1, 16, v1 :: v_dual_lshlrev_b32 v0, 16, v0
-; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_1) | instskip(NEXT) | instid1(VALU_DEP_1)
-; GFX1250-NEXT: v_max_num_f32_e32 v0, v0, v1
-; GFX1250-NEXT: v_cvt_pk_bf16_f32 v0, v0, s0
+; GFX1250-NEXT: v_pk_max_num_bf16 v0, v0, v1
; GFX1250-NEXT: s_set_pc_i64 s[30:31]
%op = call bfloat @llvm.maxnum.bf16(bfloat %a, bfloat %b)
ret bfloat %op
diff --git a/llvm/test/CodeGen/AMDGPU/minmax3-tree-reduction.ll b/llvm/test/CodeGen/AMDGPU/minmax3-tree-reduction.ll
index ed5989352c77e..2658847d1ca1b 100644
--- a/llvm/test/CodeGen/AMDGPU/minmax3-tree-reduction.ll
+++ b/llvm/test/CodeGen/AMDGPU/minmax3-tree-reduction.ll
@@ -485,7 +485,7 @@ define double @v_no_max3_maxnum_tree4_f64(double %a, double %b, double %c, doubl
ret double %result
}
-; Negative test: bf16 is promoted to f32 with conversions, tree combine cannot apply
+; Negative test: bf16 has no max3 on any target yet, tree combine must not fire
define bfloat @v_no_max3_maxnum_tree4_bf16(bfloat %a, bfloat %b, bfloat %c, bfloat %d) {
; GFX9-LABEL: v_no_max3_maxnum_tree4_bf16:
; GFX9: ; %bb.0:
@@ -522,17 +522,10 @@ define bfloat @v_no_max3_maxnum_tree4_bf16(bfloat %a, bfloat %b, bfloat %c, bflo
; GFX1250: ; %bb.0:
; GFX1250-NEXT: s_wait_loadcnt_dscnt 0x0
; GFX1250-NEXT: s_wait_kmcnt 0x0
-; GFX1250-NEXT: v_dual_lshlrev_b32 v1, 16, v1 :: v_dual_lshlrev_b32 v3, 16, v3
-; GFX1250-NEXT: v_dual_lshlrev_b32 v2, 16, v2 :: v_dual_lshlrev_b32 v0, 16, v0
-; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_1) | instskip(NEXT) | instid1(VALU_DEP_1)
-; GFX1250-NEXT: v_dual_max_num_f32 v2, v2, v3 :: v_dual_max_num_f32 v0, v0, v1
-; GFX1250-NEXT: v_cvt_pk_bf16_f32 v1, v2, s0
-; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_2) | instskip(NEXT) | instid1(VALU_DEP_1)
-; GFX1250-NEXT: v_cvt_pk_bf16_f32 v0, v0, s0
-; GFX1250-NEXT: v_dual_lshlrev_b32 v1, 16, v1 :: v_dual_lshlrev_b32 v0, 16, v0
-; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_1) | instskip(NEXT) | instid1(VALU_DEP_1)
-; GFX1250-NEXT: v_max_num_f32_e32 v0, v0, v1
-; GFX1250-NEXT: v_cvt_pk_bf16_f32 v0, v0, s0
+; GFX1250-NEXT: v_pk_max_num_bf16 v0, v0, v1
+; GFX1250-NEXT: v_pk_max_num_bf16 v1, v2, v3
+; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_1)
+; GFX1250-NEXT: v_pk_max_num_bf16 v0, v0, v1
; GFX1250-NEXT: s_set_pc_i64 s[30:31]
%max.ab = call bfloat @llvm.maxnum.bf16(bfloat %a, bfloat %b)
%max.cd = call bfloat @llvm.maxnum.bf16(bfloat %c, bfloat %d)
@@ -540,6 +533,54 @@ define bfloat @v_no_max3_maxnum_tree4_bf16(bfloat %a, bfloat %b, bfloat %c, bflo
ret bfloat %result
}
+; Negative test: bf16 has no min3 on any target yet, tree combine must not fire
+define bfloat @v_no_min3_minnum_tree4_bf16(bfloat %a, bfloat %b, bfloat %c, bfloat %d) {
+; GFX9-LABEL: v_no_min3_minnum_tree4_bf16:
+; GFX9: ; %bb.0:
+; GFX9-NEXT: s_waitcnt vmcnt(0) expcnt(0) lgkmcnt(0)
+; GFX9-NEXT: v_lshlrev_b32_e32 v1, 16, v1
+; GFX9-NEXT: v_lshlrev_b32_e32 v0, 16, v0
+; GFX9-NEXT: v_min_f32_e32 v0, v0, v1
+; GFX9-NEXT: v_bfe_u32 v1, v0, 16, 1
+; GFX9-NEXT: s_movk_i32 s4, 0x7fff
+; GFX9-NEXT: v_add3_u32 v1, v1, v0, s4
+; GFX9-NEXT: v_or_b32_e32 v4, 0x400000, v0
+; GFX9-NEXT: v_cmp_u_f32_e32 vcc, v0, v0
+; GFX9-NEXT: v_cndmask_b32_e32 v0, v1, v4, vcc
+; GFX9-NEXT: v_lshlrev_b32_e32 v1, 16, v3
+; GFX9-NEXT: v_lshlrev_b32_e32 v2, 16, v2
+; GFX9-NEXT: v_min_f32_e32 v1, v2, v1
+; GFX9-NEXT: v_bfe_u32 v2, v1, 16, 1
+; GFX9-NEXT: v_add3_u32 v2, v2, v1, s4
+; GFX9-NEXT: v_or_b32_e32 v3, 0x400000, v1
+; GFX9-NEXT: v_cmp_u_f32_e32 vcc, v1, v1
+; GFX9-NEXT: v_cndmask_b32_e32 v1, v2, v3, vcc
+; GFX9-NEXT: v_and_b32_e32 v1, 0xffff0000, v1
+; GFX9-NEXT: v_and_b32_e32 v0, 0xffff0000, v0
+; GFX9-NEXT: v_min_f32_e32 v0, v0, v1
+; GFX9-NEXT: v_bfe_u32 v1, v0, 16, 1
+; GFX9-NEXT: v_add3_u32 v1, v1, v0, s4
+; GFX9-NEXT: v_or_b32_e32 v2, 0x400000, v0
+; GFX9-NEXT: v_cmp_u_f32_e32 vcc, v0, v0
+; GFX9-NEXT: v_cndmask_b32_e32 v0, v1, v2, vcc
+; GFX9-NEXT: v_lshrrev_b32_e32 v0, 16, v0
+; GFX9-NEXT: s_setpc_b64 s[30:31]
+;
+; GFX1250-LABEL: v_no_min3_minnum_tree4_bf16:
+; GFX1250: ; %bb.0:
+; GFX1250-NEXT: s_wait_loadcnt_dscnt 0x0
+; GFX1250-NEXT: s_wait_kmcnt 0x0
+; GFX1250-NEXT: v_pk_min_num_bf16 v0, v0, v1
+; GFX1250-NEXT: v_pk_min_num_bf16 v1, v2, v3
+; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_1)
+; GFX1250-NEXT: v_pk_min_num_bf16 v0, v0, v1
+; GFX1250-NEXT: s_set_pc_i64 s[30:31]
+ %min.ab = call bfloat @llvm.minnum.bf16(bfloat %a, bfloat %b)
+ %min.cd = call bfloat @llvm.minnum.bf16(bfloat %c, bfloat %d)
+ %result = call bfloat @llvm.minnum.bf16(bfloat %min.ab, bfloat %min.cd)
+ ret bfloat %result
+}
+
; Two-level ternary tree
define float @v_max3_maxnum_ternary_2level_f32(
; GFX9-LABEL: v_max3_maxnum_ternary_2level_f32:
@@ -755,13 +796,3 @@ define <2 x half> @v_max3_maxnum_ternary_2level_v2f16(
%R = call <2 x half> @llvm.maxnum.v2f16(<2 x half> %AB, <2 x half> %C)
ret <2 x half> %R
}
-
-declare float @llvm.maxnum.f32(float, float)
-declare float @llvm.minnum.f32(float, float)
-declare float @llvm.maximum.f32(float, float)
-declare float @llvm.minimum.f32(float, float)
-declare half @llvm.maxnum.f16(half, half)
-declare double @llvm.maxnum.f64(double, double)
-declare bfloat @llvm.maxnum.bf16(bfloat, bfloat)
-declare <2 x float> @llvm.maxnum.v2f32(<2 x float>, <2 x float>)
-declare <2 x half> @llvm.maxnum.v2f16(<2 x half>, <2 x half>)
More information about the llvm-commits
mailing list