pub fn f16_to_f32(bits: u16) -> f32 {
let sign: f32 = if bits & 0x8000 != 0 { -1.0 } else { 1.0 };
let exponent = (bits >> 10) & 0x1F;
let mantissa = f32::from(bits & 0x3FF);
match exponent {
0 => sign * mantissa * 2f32.powi(-24),
0x1F if mantissa == 0.0 => sign * f32::INFINITY,
0x1F => f32::NAN,
e => sign * (1.0 + mantissa / 1024.0) * 2f32.powi(i32::from(e) - 15),
}
}
pub fn f32_to_f16(value: f32) -> u16 {
let bits = value.to_bits();
let sign = ((bits >> 16) & 0x8000) as u16;
let exponent = ((bits >> 23) & 0xFF) as i32;
let mantissa = bits & 0x007F_FFFF;
if exponent == 0xFF {
return if mantissa == 0 {
sign | 0x7C00
} else {
sign | 0x7E00
};
}
if exponent == 0 {
return sign;
}
let half_exponent = exponent - 127 + 15;
if half_exponent >= 0x1F {
return sign | 0x7C00;
}
if half_exponent <= 0 {
return sign | encode_f16_subnormal(mantissa, half_exponent);
}
let mut half_mantissa = (mantissa >> 13) as u16;
if round_to_even(mantissa, 13) {
half_mantissa += 1;
}
if half_mantissa == 0x0400 {
return finish_f16_normal(sign, half_exponent + 1, 0);
}
finish_f16_normal(sign, half_exponent, half_mantissa)
}
fn finish_f16_normal(sign: u16, exponent: i32, mantissa: u16) -> u16 {
if exponent >= 0x1F {
return sign | 0x7C00;
}
sign | ((exponent as u16) << 10) | mantissa
}
fn encode_f16_subnormal(mantissa: u32, half_exponent: i32) -> u16 {
let shift = 14 - half_exponent;
if shift > 24 {
return 0;
}
let significand = 0x0080_0000 | mantissa; let mut half_mantissa = (significand >> shift) as u16;
if round_to_even(significand, shift) {
half_mantissa += 1;
}
half_mantissa
}
fn round_to_even(value: u32, shift: i32) -> bool {
if shift <= 0 {
return false;
}
let shift = shift as u32;
let round_bit = 1u32 << (shift - 1);
let round_bit_set = value & round_bit != 0;
let sticky = value & (round_bit - 1) != 0;
let kept_is_odd = (value >> shift) & 1 != 0;
round_bit_set && (sticky || kept_is_odd)
}
pub fn bf16_to_f32(bits: u16) -> f32 {
f32::from_bits(u32::from(bits) << 16)
}
pub fn f32_to_bf16(value: f32) -> u16 {
let bits = value.to_bits();
if value.is_nan() {
return ((bits >> 16) as u16) | 0x0040;
}
let round = 0x7FFF_u32 + ((bits >> 16) & 1);
(bits.wrapping_add(round) >> 16) as u16
}
#[cfg(test)]
mod tests {
use super::*;
fn assert_close(a: f32, b: f32, epsilon: f32) {
assert!(
(a - b).abs() <= epsilon,
"expected {b} within {epsilon}, got {a}"
);
}
#[test]
fn f16_known_constants_decode_exactly() {
assert_eq!(f16_to_f32(0x0000), 0.0);
assert!(f16_to_f32(0x8000).is_sign_negative());
assert_eq!(f16_to_f32(0x8000), -0.0);
assert_eq!(f16_to_f32(0x3C00), 1.0);
assert_eq!(f16_to_f32(0xBC00), -1.0);
assert_eq!(f16_to_f32(0x4000), 2.0);
assert_eq!(f16_to_f32(0x3800), 0.5);
assert_eq!(f16_to_f32(0xC800), -8.0);
assert_eq!(f16_to_f32(0x7BFF), 65504.0);
assert_eq!(f16_to_f32(0x0001), 2f32.powi(-24));
}
#[test]
fn f16_infinity_and_nan_decode_correctly() {
assert_eq!(f16_to_f32(0x7C00), f32::INFINITY);
assert_eq!(f16_to_f32(0xFC00), f32::NEG_INFINITY);
assert!(f16_to_f32(0x7E00).is_nan());
assert!(f16_to_f32(0xFE00).is_nan());
}
#[test]
fn f32_to_f16_known_constants_encode_exactly() {
assert_eq!(f32_to_f16(1.0), 0x3C00);
assert_eq!(f32_to_f16(2.0), 0x4000);
assert_eq!(f32_to_f16(0.5), 0x3800);
assert_eq!(f32_to_f16(-8.0), 0xC800);
assert_eq!(f32_to_f16(65504.0), 0x7BFF);
assert_eq!(f32_to_f16(0.0), 0x0000);
assert_eq!(f32_to_f16(-0.0), 0x8000);
}
#[test]
fn f32_to_f16_overflow_saturates_to_infinity() {
assert_eq!(f32_to_f16(70_000.0), 0x7C00);
assert_eq!(f32_to_f16(-70_000.0), 0xFC00);
assert_eq!(f32_to_f16(3e38), 0x7C00);
}
#[test]
fn f32_to_f16_underflow_flushes_to_zero() {
assert_eq!(f32_to_f16(1e-9), 0x0000);
assert_eq!(f32_to_f16(2f32.powi(-25)), 0x0000);
}
#[test]
fn f32_to_f16_infinity_and_nan_round_trip() {
assert_eq!(f32_to_f16(f32::INFINITY), 0x7C00);
assert_eq!(f32_to_f16(f32::NEG_INFINITY), 0xFC00);
assert!(f16_to_f32(f32_to_f16(f32::NAN)).is_nan());
}
#[test]
fn f32_to_f16_matches_numpy_oracle_for_irrational_values() {
assert_eq!(f32_to_f16(1.0 / 3.0), 0x3555);
assert_eq!(f32_to_f16(100.5), 0x5648);
assert_eq!(f32_to_f16(-100.5), 0xD648);
assert_eq!(f32_to_f16(std::f32::consts::PI), 0x4248);
assert_eq!(f32_to_f16(6.10352e-05), 0x0400);
assert_eq!(f32_to_f16(0.00012), 0x07DD);
}
#[test]
fn f16_round_trips_every_possible_bit_pattern() {
for bits in 0..=u16::MAX {
let value = f16_to_f32(bits);
let back = f32_to_f16(value);
if value.is_nan() {
assert!(f16_to_f32(back).is_nan(), "0x{bits:04x} -> NaN -> not NaN");
} else {
assert_eq!(
back, bits,
"0x{bits:04x} -> {value} -> 0x{back:04x} (round-trip mismatch)"
);
}
}
}
#[test]
fn bf16_known_constants_round_trip_exactly() {
for v in [1.0f32, 2.0, 0.5, 100.5, -1.5, 0.0, -0.0] {
let bits = f32_to_bf16(v);
assert_eq!(bf16_to_f32(bits), v);
}
}
#[test]
fn bf16_matches_ggml_rounding_oracle() {
assert_eq!(f32_to_bf16(1.0 / 3.0), 0x3EAB);
assert_eq!(f32_to_bf16(std::f32::consts::PI), 0x4049);
assert_eq!(f32_to_bf16(100.5), 0x42C9);
}
#[test]
fn bf16_infinity_and_nan_round_trip() {
assert_eq!(f32_to_bf16(f32::INFINITY), 0x7F80);
assert_eq!(f32_to_bf16(f32::NEG_INFINITY), 0xFF80);
assert!(bf16_to_f32(f32_to_bf16(f32::NAN)).is_nan());
}
#[test]
fn bf16_denormals_are_preserved_not_flushed() {
let tiny = f32::from_bits(1); let bits = f32_to_bf16(tiny);
assert!(bf16_to_f32(bits) >= 0.0);
}
#[test]
fn bf16_round_trips_every_possible_bit_pattern() {
for bits in 0..=u16::MAX {
let value = bf16_to_f32(bits);
let back = f32_to_bf16(value);
if value.is_nan() {
assert!(bf16_to_f32(back).is_nan());
} else {
assert_eq!(back, bits, "0x{bits:04x} -> {value} -> 0x{back:04x}");
}
}
}
#[test]
fn conversions_are_reasonably_close_for_arbitrary_values() {
for v in [1.23456f32, -9.8765, 12345.6, 0.001234] {
let f16_bits = f32_to_f16(v);
assert_close(f16_to_f32(f16_bits), v, v.abs() * 0.001 + 1e-6);
let bf16_bits = f32_to_bf16(v);
assert_close(bf16_to_f32(bf16_bits), v, v.abs() * 0.01 + 1e-6);
}
}
}