pub(crate) use super::f16_encode::{f32_to_f16_bits, f32_to_finite_f16_bits};
#[inline]
pub(crate) fn f16_bits_to_f32(bits: u16) -> f32 {
let sign = ((bits >> 15) & 0x1) as u32;
let exp = ((bits >> 10) & 0x1f) as u32;
let frac = (bits & 0x03ff) as u32;
let f32_bits = match (exp, frac) {
(0, 0) => sign << 31,
(0, _) => {
let mut mant = frac;
let mut e = -14i32;
while (mant & 0x0400) == 0 {
mant <<= 1;
e -= 1;
}
mant &= 0x03ff;
(sign << 31) | (((e + 127) as u32) << 23) | (mant << 13)
}
(0x1f, 0) => (sign << 31) | 0x7f80_0000,
(0x1f, _) => (sign << 31) | 0x7f80_0000 | (frac << 13),
_ => (sign << 31) | (((exp as i32 - 15 + 127) as u32) << 23) | (frac << 13),
};
f32::from_bits(f32_bits)
}
#[inline]
pub(crate) fn bf16_bits_to_f32(bits: u16) -> f32 {
f32::from_bits((bits as u32) << 16)
}
#[cfg(test)]
mod tests {
use super::*;
fn f32_to_bf16_bits_reference(v: f32) -> u16 {
let bits = v.to_bits();
if v.is_nan() {
return ((bits >> 16) as u16) | 0x0040;
}
let round_bit = (bits >> 16) & 1;
let half = 0x7fff + round_bit;
(bits.wrapping_add(half) >> 16) as u16
}
fn is_signaling_nan_bits(bits: u16) -> bool {
let exp = (bits >> 10) & 0x1f;
let frac = bits & 0x03ff;
exp == 0x1f && frac != 0 && (frac & 0x0200) == 0
}
#[test]
fn f16_widen_narrow_composition_round_trips_self_consistently() {
for bits in 0u32..=0xffff {
let bits = bits as u16;
let widened = f16_bits_to_f32(bits);
let exp = (bits >> 10) & 0x1f;
let frac = bits & 0x03ff;
if exp == 0x1f && frac != 0 {
assert!(
widened.is_nan(),
"f16 NaN bits {bits:#06x} must widen to NaN"
);
continue;
}
let narrowed = f32_to_f16_bits(widened);
assert_eq!(
narrowed, bits,
"roundtrip mismatch: bits={bits:#06x} widened={widened} narrowed={narrowed:#06x}"
);
}
}
#[test]
fn f16_bits_to_f32_matches_independent_half_crate_oracle() {
let mut checked = 0u32;
for bits in 0u32..=0xffff {
let bits = bits as u16;
if is_signaling_nan_bits(bits) {
continue;
}
let ours = f16_bits_to_f32(bits).to_bits();
let oracle = half::f16::from_bits(bits).to_f32().to_bits();
assert_eq!(
ours, oracle,
"f16_bits_to_f32({bits:#06x}) diverges from the `half` crate oracle: \
ours={ours:#010x} oracle={oracle:#010x}"
);
checked += 1;
}
assert_eq!(
checked,
65536 - 1022,
"expected exactly the non-signaling-NaN f16 bit space to be checked"
);
}
#[test]
fn f16_bits_to_f32_signaling_nan_is_lossless_widen_independent_of_decoder() {
let mut checked = 0u32;
for bits in 0u32..=0xffff {
let bits = bits as u16;
if !is_signaling_nan_bits(bits) {
continue;
}
let sign = (bits >> 15) & 0x1;
let frac = (bits & 0x03ff) as u32;
let expected_bits = ((sign as u32) << 31) | 0x7f80_0000 | (frac << 13);
let widened = f16_bits_to_f32(bits);
assert!(
widened.is_nan(),
"signaling NaN bits {bits:#06x} must widen to NaN"
);
assert_eq!(
widened.to_bits(),
expected_bits,
"signaling NaN {bits:#06x} did not widen losslessly: \
got={:#010x} expected={expected_bits:#010x}",
widened.to_bits()
);
assert_eq!(
widened.to_bits() & 0x0040_0000,
0,
"signaling NaN {bits:#06x} must NOT be quieted by decode"
);
checked += 1;
}
assert_eq!(checked, 1022, "expected exactly 511 payloads x 2 signs");
}
#[test]
fn f32_to_f16_bits_subnormal_normal_boundary_matches_independent_oracle() {
let smallest_normal = 2f32.powi(-14);
let largest_subnormal = 2f32.powi(-14) * (1023.0 / 1024.0);
let just_below_boundary = largest_subnormal - 2f32.powi(-25); let just_above_boundary = smallest_normal + 2f32.powi(-25);
for v in [
largest_subnormal,
smallest_normal,
just_below_boundary,
just_above_boundary,
] {
let ours = f32_to_f16_bits(v);
let oracle = half::f16::from_f32(v).to_bits();
assert_eq!(
ours, oracle,
"f32_to_f16_bits({v}) diverges from the `half` crate oracle at the \
subnormal/normal boundary: ours={ours:#06x} oracle={oracle:#06x}"
);
}
assert_eq!(f32_to_f16_bits(largest_subnormal), 0x03ff);
assert_eq!(f32_to_f16_bits(smallest_normal), 0x0400);
}
#[test]
fn f32_to_f16_bits_finite_infinity_boundary_matches_independent_oracle() {
let f16_max = 65504.0f32;
let midpoint = 65520.0f32;
let just_below_midpoint = 65519.0f32; let just_above_midpoint = 65521.0f32;
for v in [f16_max, midpoint, just_below_midpoint, just_above_midpoint] {
let ours = f32_to_f16_bits(v);
let oracle = half::f16::from_f32(v).to_bits();
assert_eq!(
ours, oracle,
"f32_to_f16_bits({v}) diverges from the `half` crate oracle at the \
finite/infinity boundary: ours={ours:#06x} oracle={oracle:#06x}"
);
}
assert_eq!(f32_to_f16_bits(f16_max), 0x7bff);
assert_eq!(f32_to_f16_bits(just_below_midpoint), 0x7bff);
assert_eq!(f32_to_f16_bits(just_above_midpoint), 0x7c00);
}
#[test]
fn f16_special_values() {
assert_eq!(f16_bits_to_f32(0x0000), 0.0f32);
assert!(f16_bits_to_f32(0x8000).is_sign_negative());
assert_eq!(f16_bits_to_f32(0x8000), 0.0f32);
assert_eq!(f16_bits_to_f32(0x3c00), 1.0f32);
assert_eq!(f16_bits_to_f32(0xbc00), -1.0f32);
assert_eq!(f16_bits_to_f32(0x7c00), f32::INFINITY);
assert_eq!(f16_bits_to_f32(0xfc00), f32::NEG_INFINITY);
assert!(f16_bits_to_f32(0x7e00).is_nan());
}
#[test]
fn f16_denormals_round_trip() {
for &bits in &[0x0001u16, 0x0200, 0x03ff] {
let widened = f16_bits_to_f32(bits);
assert!(widened.is_finite() && widened != 0.0);
assert_eq!(f32_to_f16_bits(widened), bits);
}
}
#[test]
fn f32_to_f16_bits_signed_zero() {
assert_eq!(f32_to_f16_bits(0.0f32), 0x0000);
assert_eq!(f32_to_f16_bits(-0.0f32), 0x8000);
}
#[test]
fn f32_to_f16_bits_overflow_to_infinity() {
assert_eq!(f32_to_f16_bits(1.0e6), 0x7c00);
assert_eq!(f32_to_f16_bits(-1.0e6), 0xfc00);
assert_eq!(f32_to_f16_bits(f32::MAX), 0x7c00);
}
#[test]
fn f32_to_f16_bits_nan_payload_preserved_and_quiet() {
let bits = f32_to_f16_bits(f32::NAN);
assert_eq!(bits & 0x7c00, 0x7c00, "exponent field must be all-ones");
assert_ne!(bits & 0x03ff, 0, "mantissa must stay non-zero (quiet NaN)");
assert_ne!(bits & 0x0200, 0, "quiet bit must be set");
}
#[test]
fn f32_to_f16_bits_tie_to_even_rounding() {
let a = f16_bits_to_f32(0x3c00);
let b = f16_bits_to_f32(0x3c01);
let tie_low = (a + b) * 0.5;
assert_eq!(f32_to_f16_bits(tie_low), 0x3c00);
let c = f16_bits_to_f32(0x3c01);
let d = f16_bits_to_f32(0x3c02);
let tie_high = (c + d) * 0.5;
assert_eq!(f32_to_f16_bits(tie_high), 0x3c02);
}
#[test]
fn bf16_bits_to_f32_lossless_widen() {
for bits in [0x0000u16, 0x8000, 0x3f80, 0xbf80, 0x7f80, 0xff80, 0x7fc0] {
let widened = bf16_bits_to_f32(bits);
let renarrowed = f32_to_bf16_bits_reference(widened);
assert_eq!(
renarrowed, bits,
"bf16 widen must be exactly reversible for bits={bits:#06x}"
);
}
}
#[test]
fn bf16_bits_to_f32_special_values() {
assert_eq!(bf16_bits_to_f32(0x0000), 0.0f32);
assert!(bf16_bits_to_f32(0x8000).is_sign_negative());
assert_eq!(bf16_bits_to_f32(0x3f80), 1.0f32);
assert_eq!(bf16_bits_to_f32(0x7f80), f32::INFINITY);
assert_eq!(bf16_bits_to_f32(0xff80), f32::NEG_INFINITY);
assert!(bf16_bits_to_f32(0x7fc0).is_nan());
}
#[test]
fn matches_f16_weights_original_impl_golden_values() {
let cases: &[(f32, u16)] = &[
(0.0, 0x0000),
(1.0, 0x3c00),
(-1.0, 0xbc00),
(2.0, 0x4000),
(0.5, 0x3800),
(65504.0, 0x7bff), ];
for &(f, bits) in cases {
assert_eq!(f32_to_f16_bits(f), bits, "encode mismatch for {f}");
assert_eq!(f16_bits_to_f32(bits), f, "decode mismatch for {bits:#06x}");
}
}
}