[llvm] [AMDGPU] Fix crash on strict fptrunc to bf16 (PR #208037)
Arseniy Obolenskiy via llvm-commits
llvm-commits at lists.llvm.org
Thu Jul 9 06:30:23 PDT 2026
https://github.com/aobolensk updated https://github.com/llvm/llvm-project/pull/208037
>From 37dcbbdce73ef5ecdd1814df592803a5f600b166 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Tue, 7 Jul 2026 17:50:57 +0200
Subject: [PATCH 1/3] [AMDGPU] Fix crash on strict fptrunc to bf16
STRICT_FP_ROUND to bf16 was not custom lowered, so it reached the default legalizer and crashed
Handle it in lowerFP_ROUND like the non-strict case (since it is a valid case in AMDGPU)
---
llvm/lib/Target/AMDGPU/SIISelLowering.cpp | 22 ++++--
.../CodeGen/AMDGPU/strict_fptrunc_bf16.ll | 69 +++++++++++++++++++
2 files changed, 86 insertions(+), 5 deletions(-)
create mode 100644 llvm/test/CodeGen/AMDGPU/strict_fptrunc_bf16.ll
diff --git a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
index cb38d9081f16f..0cec187f19975 100644
--- a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
+++ b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
@@ -1023,7 +1023,8 @@ SITargetLowering::SITargetLowering(const TargetMachine &TM,
setOperationAction(ISD::MUL, MVT::i1, Promote);
if (Subtarget->hasBF16ConversionInsts()) {
- setOperationAction(ISD::FP_ROUND, {MVT::bf16, MVT::v2bf16}, Custom);
+ setOperationAction({ISD::FP_ROUND, ISD::STRICT_FP_ROUND},
+ {MVT::bf16, MVT::v2bf16}, Custom);
setOperationAction(ISD::BUILD_VECTOR, MVT::v2bf16, Legal);
}
@@ -8651,7 +8652,8 @@ SDValue SITargetLowering::splitFP_ROUNDVectorOp(SDValue Op,
}
SDValue SITargetLowering::lowerFP_ROUND(SDValue Op, SelectionDAG &DAG) const {
- SDValue Src = Op.getOperand(0);
+ bool IsStrict = Op->isStrictFPOpcode();
+ SDValue Src = Op.getOperand(IsStrict ? 1 : 0);
EVT SrcVT = Src.getValueType();
EVT DstVT = Op.getValueType();
@@ -8662,8 +8664,15 @@ SDValue SITargetLowering::lowerFP_ROUND(SDValue Op, SelectionDAG &DAG) const {
return SrcVT == MVT::v2f32 ? Op : splitFP_ROUNDVectorOp(Op, DAG);
}
- if (SrcVT.getScalarType() != MVT::f64)
+ if (SrcVT.getScalarType() != MVT::f64) {
+ if (IsStrict && DstVT.getScalarType() == MVT::bf16) {
+ SDLoc DL(Op);
+ SDValue Result = DAG.getNode(ISD::FP_ROUND, DL, DstVT, Src,
+ DAG.getTargetConstant(0, DL, MVT::i32));
+ return DAG.getMergeValues({Result, Op.getOperand(0)}, DL);
+ }
return Op;
+ }
SDLoc DL(Op);
if (DstVT == MVT::f16) {
@@ -8694,8 +8703,11 @@ SDValue SITargetLowering::lowerFP_ROUND(SDValue Op, SelectionDAG &DAG) const {
// hardware f32 -> bf16 instruction.
EVT F32VT = SrcVT.changeElementType(*DAG.getContext(), MVT::f32);
SDValue Rod = expandRoundInexactToOdd(F32VT, Src, DL, DAG);
- return DAG.getNode(ISD::FP_ROUND, DL, DstVT, Rod,
- DAG.getTargetConstant(0, DL, MVT::i32));
+ SDValue Result = DAG.getNode(ISD::FP_ROUND, DL, DstVT, Rod,
+ DAG.getTargetConstant(0, DL, MVT::i32));
+ if (IsStrict)
+ return DAG.getMergeValues({Result, Op.getOperand(0)}, DL);
+ return Result;
}
SDValue SITargetLowering::lowerFMINNUM_FMAXNUM(SDValue Op,
diff --git a/llvm/test/CodeGen/AMDGPU/strict_fptrunc_bf16.ll b/llvm/test/CodeGen/AMDGPU/strict_fptrunc_bf16.ll
new file mode 100644
index 0000000000000..912985ce93ca6
--- /dev/null
+++ b/llvm/test/CodeGen/AMDGPU/strict_fptrunc_bf16.ll
@@ -0,0 +1,69 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 4
+; RUN: llc -mtriple=amdgcn -mcpu=gfx950 < %s | FileCheck --check-prefixes=GFX950 %s
+; RUN: llc -mtriple=amdgcn -mcpu=gfx1250 -mattr=-real-true16 < %s | FileCheck --check-prefixes=GFX1250 %s
+
+define amdgpu_ps void @strict_fptrunc_f32_to_bf16(float %a, ptr %out) #0 {
+; GFX950-LABEL: strict_fptrunc_f32_to_bf16:
+; GFX950: ; %bb.0:
+; GFX950-NEXT: v_mov_b32_e32 v3, v2
+; GFX950-NEXT: v_mov_b32_e32 v2, v1
+; GFX950-NEXT: v_cvt_pk_bf16_f32 v0, v0, s0
+; GFX950-NEXT: flat_store_short v[2:3], v0
+; GFX950-NEXT: s_endpgm
+;
+; GFX1250-LABEL: strict_fptrunc_f32_to_bf16:
+; GFX1250: ; %bb.0:
+; GFX1250-NEXT: s_setreg_imm32_b32 hwreg(HW_REG_WAVE_MODE, 25, 1), 1 ; msbs: dst=0 src0=0 src1=0 src2=0
+; GFX1250-NEXT: v_dual_mov_b32 v3, v2 :: v_dual_mov_b32 v2, v1
+; GFX1250-NEXT: v_cvt_pk_bf16_f32 v0, v0, s0
+; GFX1250-NEXT: flat_store_b16 v[2:3], v0
+; GFX1250-NEXT: s_endpgm
+ %cvt = call bfloat @llvm.experimental.constrained.fptrunc.bf16.f32(float %a, metadata !"round.tonearest", metadata !"fpexcept.strict")
+ store bfloat %cvt, ptr %out
+ ret void
+}
+
+define amdgpu_ps void @strict_fptrunc_f64_to_bf16(double %a, ptr %out) #0 {
+; GFX950-LABEL: strict_fptrunc_f64_to_bf16:
+; GFX950: ; %bb.0:
+; GFX950-NEXT: v_cvt_f32_f64_e32 v6, v[0:1]
+; GFX950-NEXT: v_cvt_f64_f32_e32 v[4:5], v6
+; GFX950-NEXT: v_and_b32_e32 v7, 1, v6
+; GFX950-NEXT: v_cmp_gt_f64_e64 s[2:3], |v[0:1]|, |v[4:5]|
+; GFX950-NEXT: v_cmp_nlg_f64_e32 vcc, v[0:1], v[4:5]
+; GFX950-NEXT: v_cmp_eq_u32_e64 s[0:1], 1, v7
+; GFX950-NEXT: v_cndmask_b32_e64 v0, -1, 1, s[2:3]
+; GFX950-NEXT: v_add_u32_e32 v0, v6, v0
+; GFX950-NEXT: s_or_b64 vcc, vcc, s[0:1]
+; GFX950-NEXT: v_cndmask_b32_e32 v0, v0, v6, vcc
+; GFX950-NEXT: v_cvt_pk_bf16_f32 v0, v0, s0
+; GFX950-NEXT: flat_store_short v[2:3], v0
+; GFX950-NEXT: s_endpgm
+;
+; GFX1250-LABEL: strict_fptrunc_f64_to_bf16:
+; GFX1250: ; %bb.0:
+; GFX1250-NEXT: s_setreg_imm32_b32 hwreg(HW_REG_WAVE_MODE, 25, 1), 1 ; msbs: dst=0 src0=0 src1=0 src2=0
+; GFX1250-NEXT: v_cvt_f32_f64_e32 v6, v[0:1]
+; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_1) | instskip(NEXT) | instid1(VALU_DEP_1)
+; GFX1250-NEXT: v_cvt_f64_f32_e32 v[4:5], v6
+; GFX1250-NEXT: v_cmp_gt_f64_e64 s0, |v[0:1]|, |v[4:5]|
+; GFX1250-NEXT: v_cmp_nlg_f64_e32 vcc_lo, v[0:1], v[4:5]
+; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_2) | instskip(NEXT) | instid1(VALU_DEP_1)
+; GFX1250-NEXT: v_cndmask_b32_e64 v0, -1, 1, s0
+; GFX1250-NEXT: v_dual_add_nc_u32 v0, v6, v0 :: v_dual_bitop2_b32 v7, 1, v6 bitop3:0x40
+; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_1) | instskip(SKIP_2) | instid1(VALU_DEP_1)
+; GFX1250-NEXT: v_cmp_eq_u32_e64 s0, 1, v7
+; GFX1250-NEXT: s_or_b32 vcc_lo, vcc_lo, s0
+; GFX1250-NEXT: v_cndmask_b32_e32 v0, v0, v6, vcc_lo
+; GFX1250-NEXT: v_cvt_pk_bf16_f32 v0, v0, s0
+; GFX1250-NEXT: flat_store_b16 v[2:3], v0
+; GFX1250-NEXT: s_endpgm
+ %cvt = call bfloat @llvm.experimental.constrained.fptrunc.bf16.f64(double %a, metadata !"round.tonearest", metadata !"fpexcept.strict")
+ store bfloat %cvt, ptr %out
+ ret void
+}
+
+declare bfloat @llvm.experimental.constrained.fptrunc.bf16.f32(float, metadata, metadata)
+declare bfloat @llvm.experimental.constrained.fptrunc.bf16.f64(double, metadata, metadata)
+
+attributes #0 = { strictfp }
>From 9a44df74aebf572685850a121e2c9c76dafc89a0 Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Wed, 8 Jul 2026 09:11:22 +0200
Subject: [PATCH 2/3] Address comments
---
llvm/lib/Target/AMDGPU/SIISelLowering.cpp | 18 ++++++------------
llvm/lib/Target/AMDGPU/VOP3Instructions.td | 2 +-
2 files changed, 7 insertions(+), 13 deletions(-)
diff --git a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
index 0cec187f19975..378560f824eca 100644
--- a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
+++ b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
@@ -8664,15 +8664,8 @@ SDValue SITargetLowering::lowerFP_ROUND(SDValue Op, SelectionDAG &DAG) const {
return SrcVT == MVT::v2f32 ? Op : splitFP_ROUNDVectorOp(Op, DAG);
}
- if (SrcVT.getScalarType() != MVT::f64) {
- if (IsStrict && DstVT.getScalarType() == MVT::bf16) {
- SDLoc DL(Op);
- SDValue Result = DAG.getNode(ISD::FP_ROUND, DL, DstVT, Src,
- DAG.getTargetConstant(0, DL, MVT::i32));
- return DAG.getMergeValues({Result, Op.getOperand(0)}, DL);
- }
+ if (SrcVT.getScalarType() != MVT::f64)
return Op;
- }
SDLoc DL(Op);
if (DstVT == MVT::f16) {
@@ -8703,11 +8696,12 @@ SDValue SITargetLowering::lowerFP_ROUND(SDValue Op, SelectionDAG &DAG) const {
// hardware f32 -> bf16 instruction.
EVT F32VT = SrcVT.changeElementType(*DAG.getContext(), MVT::f32);
SDValue Rod = expandRoundInexactToOdd(F32VT, Src, DL, DAG);
- SDValue Result = DAG.getNode(ISD::FP_ROUND, DL, DstVT, Rod,
- DAG.getTargetConstant(0, DL, MVT::i32));
if (IsStrict)
- return DAG.getMergeValues({Result, Op.getOperand(0)}, DL);
- return Result;
+ return DAG.getNode(
+ ISD::STRICT_FP_ROUND, DL, {DstVT, MVT::Other},
+ {Op.getOperand(0), Rod, DAG.getTargetConstant(0, DL, MVT::i32)});
+ return DAG.getNode(ISD::FP_ROUND, DL, DstVT, Rod,
+ DAG.getTargetConstant(0, DL, MVT::i32));
}
SDValue SITargetLowering::lowerFMINNUM_FMAXNUM(SDValue Op,
diff --git a/llvm/lib/Target/AMDGPU/VOP3Instructions.td b/llvm/lib/Target/AMDGPU/VOP3Instructions.td
index f2a2f2b3a2499..ebbe89981e4ef 100644
--- a/llvm/lib/Target/AMDGPU/VOP3Instructions.td
+++ b/llvm/lib/Target/AMDGPU/VOP3Instructions.td
@@ -1794,7 +1794,7 @@ class VOP3_CVT_SR_FP16_TiedInput_Profile<VOPProfile P> : VOP3_CVT_SCALE_F1632_FP
// FIXME: GlobalISel cannot distinguish f16 and bf16 and may start using bf16 patterns
// instead of less complex f16. Disable GlobalISel for these for now.
-def bf16_fpround : PatFrag <(ops node:$src0), (fpround $src0), [{ return true; }]> {
+def bf16_fpround : PatFrag <(ops node:$src0), (any_fpround $src0), [{ return true; }]> {
let GISelPredicateCode = [{return false;}];
}
>From 8aa3675b83f61d0621914ee742f738d85f889c1c Mon Sep 17 00:00:00 2001
From: Arseniy Obolenskiy <arseniy.obolenskiy at amd.com>
Date: Thu, 9 Jul 2026 15:29:43 +0200
Subject: [PATCH 3/3] Address comments
---
llvm/lib/Target/AMDGPU/SIISelLowering.cpp | 3 +-
.../CodeGen/AMDGPU/strict_fptrunc_bf16.ll | 82 +++++++++++++++++++
2 files changed, 84 insertions(+), 1 deletion(-)
diff --git a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
index 378560f824eca..69c6795a2ce87 100644
--- a/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
+++ b/llvm/lib/Target/AMDGPU/SIISelLowering.cpp
@@ -8696,10 +8696,11 @@ SDValue SITargetLowering::lowerFP_ROUND(SDValue Op, SelectionDAG &DAG) const {
// hardware f32 -> bf16 instruction.
EVT F32VT = SrcVT.changeElementType(*DAG.getContext(), MVT::f32);
SDValue Rod = expandRoundInexactToOdd(F32VT, Src, DL, DAG);
- if (IsStrict)
+ if (IsStrict) {
return DAG.getNode(
ISD::STRICT_FP_ROUND, DL, {DstVT, MVT::Other},
{Op.getOperand(0), Rod, DAG.getTargetConstant(0, DL, MVT::i32)});
+ }
return DAG.getNode(ISD::FP_ROUND, DL, DstVT, Rod,
DAG.getTargetConstant(0, DL, MVT::i32));
}
diff --git a/llvm/test/CodeGen/AMDGPU/strict_fptrunc_bf16.ll b/llvm/test/CodeGen/AMDGPU/strict_fptrunc_bf16.ll
index 912985ce93ca6..1aa8815e63868 100644
--- a/llvm/test/CodeGen/AMDGPU/strict_fptrunc_bf16.ll
+++ b/llvm/test/CodeGen/AMDGPU/strict_fptrunc_bf16.ll
@@ -63,7 +63,89 @@ define amdgpu_ps void @strict_fptrunc_f64_to_bf16(double %a, ptr %out) #0 {
ret void
}
+define amdgpu_ps void @strict_fptrunc_v2f32_to_v2bf16(<2 x float> %a, ptr %out) #0 {
+; GFX950-LABEL: strict_fptrunc_v2f32_to_v2bf16:
+; GFX950: ; %bb.0:
+; GFX950-NEXT: v_cvt_pk_bf16_f32 v0, v0, v1
+; GFX950-NEXT: flat_store_dword v[2:3], v0
+; GFX950-NEXT: s_endpgm
+;
+; GFX1250-LABEL: strict_fptrunc_v2f32_to_v2bf16:
+; GFX1250: ; %bb.0:
+; GFX1250-NEXT: s_setreg_imm32_b32 hwreg(HW_REG_WAVE_MODE, 25, 1), 1 ; msbs: dst=0 src0=0 src1=0 src2=0
+; GFX1250-NEXT: v_cvt_pk_bf16_f32 v0, v0, v1
+; GFX1250-NEXT: flat_store_b32 v[2:3], v0
+; GFX1250-NEXT: s_endpgm
+ %cvt = call <2 x bfloat> @llvm.experimental.constrained.fptrunc.v2bf16.v2f32(<2 x float> %a, metadata !"round.tonearest", metadata !"fpexcept.strict")
+ store <2 x bfloat> %cvt, ptr %out
+ ret void
+}
+
+define amdgpu_ps void @strict_fptrunc_v2f64_to_v2bf16(<2 x double> %a, ptr %out) #0 {
+; GFX950-LABEL: strict_fptrunc_v2f64_to_v2bf16:
+; GFX950: ; %bb.0:
+; GFX950-NEXT: v_cvt_f32_f64_e32 v8, v[2:3]
+; GFX950-NEXT: v_and_b32_e32 v6, 1, v8
+; GFX950-NEXT: v_cmp_ne_u32_e32 vcc, 0, v6
+; GFX950-NEXT: v_cvt_f64_f32_e32 v[6:7], v8
+; GFX950-NEXT: v_cmp_gt_f64_e64 s[2:3], |v[2:3]|, |v[6:7]|
+; GFX950-NEXT: v_cmp_nlg_f64_e64 s[0:1], v[2:3], v[6:7]
+; GFX950-NEXT: v_cvt_f32_f64_e32 v9, v[0:1]
+; GFX950-NEXT: v_cndmask_b32_e64 v2, -1, 1, s[2:3]
+; GFX950-NEXT: v_add_u32_e32 v2, v8, v2
+; GFX950-NEXT: s_or_b64 vcc, vcc, s[0:1]
+; GFX950-NEXT: v_cndmask_b32_e32 v6, v2, v8, vcc
+; GFX950-NEXT: v_cvt_f64_f32_e32 v[2:3], v9
+; GFX950-NEXT: v_and_b32_e32 v10, 1, v9
+; GFX950-NEXT: v_cmp_gt_f64_e64 s[2:3], |v[0:1]|, |v[2:3]|
+; GFX950-NEXT: v_cmp_ne_u32_e32 vcc, 0, v10
+; GFX950-NEXT: v_cmp_nlg_f64_e64 s[0:1], v[0:1], v[2:3]
+; GFX950-NEXT: v_cndmask_b32_e64 v0, -1, 1, s[2:3]
+; GFX950-NEXT: v_add_u32_e32 v0, v9, v0
+; GFX950-NEXT: s_or_b64 vcc, vcc, s[0:1]
+; GFX950-NEXT: v_cndmask_b32_e32 v0, v0, v9, vcc
+; GFX950-NEXT: v_cvt_pk_bf16_f32 v0, v0, v6
+; GFX950-NEXT: flat_store_dword v[4:5], v0
+; GFX950-NEXT: s_endpgm
+;
+; GFX1250-LABEL: strict_fptrunc_v2f64_to_v2bf16:
+; GFX1250: ; %bb.0:
+; GFX1250-NEXT: s_setreg_imm32_b32 hwreg(HW_REG_WAVE_MODE, 25, 1), 1 ; msbs: dst=0 src0=0 src1=0 src2=0
+; GFX1250-NEXT: v_cvt_f32_f64_e32 v10, v[2:3]
+; GFX1250-NEXT: v_cvt_f32_f64_e32 v11, v[0:1]
+; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_2) | instskip(NEXT) | instid1(VALU_DEP_2)
+; GFX1250-NEXT: v_cvt_f64_f32_e32 v[6:7], v10
+; GFX1250-NEXT: v_cvt_f64_f32_e32 v[8:9], v11
+; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_2) | instskip(SKIP_1) | instid1(VALU_DEP_3)
+; GFX1250-NEXT: v_cmp_gt_f64_e64 s1, |v[2:3]|, |v[6:7]|
+; GFX1250-NEXT: v_cmp_nlg_f64_e32 vcc_lo, v[2:3], v[6:7]
+; GFX1250-NEXT: v_cmp_nlg_f64_e64 s0, v[0:1], v[8:9]
+; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_3) | instskip(SKIP_1) | instid1(VALU_DEP_2)
+; GFX1250-NEXT: v_cndmask_b32_e64 v2, -1, 1, s1
+; GFX1250-NEXT: v_cmp_gt_f64_e64 s1, |v[0:1]|, |v[8:9]|
+; GFX1250-NEXT: v_dual_add_nc_u32 v1, v10, v2 :: v_dual_bitop2_b32 v13, 1, v11 bitop3:0x40
+; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_1) | instskip(NEXT) | instid1(VALU_DEP_3)
+; GFX1250-NEXT: v_cmp_ne_u32_e64 s2, 0, v13
+; GFX1250-NEXT: v_cndmask_b32_e64 v0, -1, 1, s1
+; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_1) | instskip(NEXT) | instid1(VALU_DEP_1)
+; GFX1250-NEXT: v_dual_add_nc_u32 v0, v11, v0 :: v_dual_bitop2_b32 v12, 1, v10 bitop3:0x40
+; GFX1250-NEXT: v_cmp_ne_u32_e64 s1, 0, v12
+; GFX1250-NEXT: s_or_b32 vcc_lo, s1, vcc_lo
+; GFX1250-NEXT: v_cndmask_b32_e32 v1, v1, v10, vcc_lo
+; GFX1250-NEXT: s_or_b32 vcc_lo, s2, s0
+; GFX1250-NEXT: v_cndmask_b32_e32 v0, v0, v11, vcc_lo
+; GFX1250-NEXT: s_delay_alu instid0(VALU_DEP_1)
+; GFX1250-NEXT: v_cvt_pk_bf16_f32 v0, v0, v1
+; GFX1250-NEXT: flat_store_b32 v[4:5], v0
+; GFX1250-NEXT: s_endpgm
+ %cvt = call <2 x bfloat> @llvm.experimental.constrained.fptrunc.v2bf16.v2f64(<2 x double> %a, metadata !"round.tonearest", metadata !"fpexcept.strict")
+ store <2 x bfloat> %cvt, ptr %out
+ ret void
+}
+
declare bfloat @llvm.experimental.constrained.fptrunc.bf16.f32(float, metadata, metadata)
declare bfloat @llvm.experimental.constrained.fptrunc.bf16.f64(double, metadata, metadata)
+declare <2 x bfloat> @llvm.experimental.constrained.fptrunc.v2bf16.v2f32(<2 x float>, metadata, metadata)
+declare <2 x bfloat> @llvm.experimental.constrained.fptrunc.v2bf16.v2f64(<2 x double>, metadata, metadata)
attributes #0 = { strictfp }
More information about the llvm-commits
mailing list