#[inline]
pub fn f16_bits_to_f32(bits: u16) -> f32 {
let u = bits as u32;
let sign = (u >> 15) << 31;
let exp = (u >> 10) & 0x1F;
let mant = u & 0x3FF;
if exp == 0 {
if mant == 0 {
f32::from_bits(sign) } else {
let z = mant.leading_zeros().wrapping_sub(22); let e = (112u32).wrapping_sub(z);
let frac = mant.wrapping_shl(z + 14) & 0x7F_FFFF;
f32::from_bits(sign | (e << 23) | frac)
}
} else if exp == 0x1F {
if mant == 0 {
f32::from_bits(sign | 0x7F80_0000) } else {
f32::from_bits(sign | 0x7FC0_0000 | (mant << 13)) }
} else {
let e = ((exp as i32) - 15 + 127).wrapping_shl(23) as u32;
f32::from_bits(sign | e | (mant << 13))
}
}
#[inline]
pub fn f32_to_f16_bits(f: f32) -> u16 {
let bits = f.to_bits();
let sign = ((bits >> 16) as u16) & 0x8000;
let f32_exp = ((bits >> 23) & 0xFF) as i32;
let f32_mant = bits & 0x7F_FFFF;
if f32_exp == 0xFF {
return if f32_mant == 0 {
sign | 0x7C00 } else {
sign | 0x7E00 | ((f32_mant >> 13) as u16 & 0x1FF) };
}
let (exp_adj, fraction) = if f32_exp == 0 {
if f32_mant == 0 {
return sign; }
let lz = f32_mant.leading_zeros().wrapping_sub(9); (-127 - lz as i32, f32_mant.wrapping_shl(lz + 1) & 0x7F_FFFF)
} else {
(f32_exp - 127, f32_mant)
};
if exp_adj > 15 {
return sign | 0x7C00;
}
if exp_adj < -14 {
if exp_adj < -25 {
return sign; }
let shift = (-1 - exp_adj) as u32;
let mant_full = 0x80_0000u32 | fraction;
let mut m = (mant_full >> shift) as u16;
let dropped = mant_full & ((1u32 << shift).wrapping_sub(1));
let guard = 1u32 << (shift - 1);
if dropped > guard || (dropped == guard && (m & 1) != 0) {
m += 1;
if m >= 0x400 {
return sign | 0x0400; }
}
return sign | (m & 0x3FF);
}
let e = ((exp_adj + 15) as u16) << 10;
let mut m = ((fraction >> 13) as u16) & 0x3FF;
let dropped = fraction & 0x1FFF;
if dropped > 0x1000 || (dropped == 0x1000 && (m & 1) != 0) {
m += 1;
if m >= 0x400 {
if e + 0x0400 >= 0x7C00 {
return sign | 0x7C00; }
return sign | (e + 0x0400);
}
}
sign | e | m
}
#[inline]
#[target_feature(enable = "f16c")]
pub unsafe fn f16_bits_to_f32_f16c(bits: u16) -> f32 {
use std::arch::x86_64::_mm_cvtph_ps;
use std::arch::x86_64::_mm_cvtsi32_si128;
use std::arch::x86_64::_mm_cvtss_f32;
_mm_cvtss_f32(_mm_cvtph_ps(_mm_cvtsi32_si128(bits as i32)))
}
#[inline]
#[target_feature(enable = "f16c")]
pub unsafe fn f32_to_f16_bits_f16c(f: f32) -> u16 {
use std::arch::x86_64::_MM_FROUND_TO_NEAREST_INT;
use std::arch::x86_64::_mm_cvtps_ph;
use std::arch::x86_64::_mm_cvtsi128_si32;
use std::arch::x86_64::_mm_set_ss;
(_mm_cvtsi128_si32(_mm_cvtps_ph(_mm_set_ss(f), _MM_FROUND_TO_NEAREST_INT)) as u32 & 0xFFFF)
as u16
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn roundtrip_f32_f16_f32() {
let edge = [
0.0f32,
-0.0,
f32::INFINITY,
f32::NEG_INFINITY,
f32::NAN,
f32::MIN_POSITIVE,
1.0,
-1.0,
65504.0, -65504.0,
6.1035156e-5, 5.9604645e-8, ];
for &v in &edge {
let bits = f32_to_f16_bits(v);
let back = f16_bits_to_f32(bits);
let bits2 = f32_to_f16_bits(back);
assert_eq!(bits, bits2, "idempotent failed for {:e}", v);
}
for e in (-14i32..=15i32).rev() {
for m in 0..1024 {
let v = (1.0 + m as f32 / 1024.0) * 2.0f32.powi(e);
let bits = f32_to_f16_bits(v);
let back = f16_bits_to_f32(bits);
let bits2 = f32_to_f16_bits(back);
assert_eq!(bits, bits2, "roundtrip failed for v={:e}", v);
}
}
}
#[test]
fn zero_roundtrip() {
assert_eq!(f32_to_f16_bits(0.0), 0);
assert_eq!(f32_to_f16_bits(-0.0), 0x8000);
assert_eq!(f16_bits_to_f32(0), 0.0);
assert_eq!(f16_bits_to_f32(0x8000), -0.0);
}
#[test]
fn infinity_roundtrip() {
assert_eq!(f32_to_f16_bits(f32::INFINITY), 0x7C00);
assert_eq!(f32_to_f16_bits(f32::NEG_INFINITY), 0xFC00);
assert!(f16_bits_to_f32(0x7C00).is_infinite());
assert!(f16_bits_to_f32(0xFC00).is_infinite());
}
#[test]
fn nan_roundtrip() {
assert!(f16_bits_to_f32(0x7C01).is_nan());
assert!(f16_bits_to_f32(0xFC01).is_nan());
let bits = f32_to_f16_bits(f32::NAN);
assert!(f16_bits_to_f32(bits).is_nan());
}
#[test]
fn exhaustive_decode_f16c_vs_software() {
for bits in 0u16..=0xFFFFu16 {
let expected = f16_bits_to_f32(bits);
let got = unsafe { f16_bits_to_f32_f16c(bits) };
assert_eq!(
got.to_bits(),
expected.to_bits(),
"F16C decode mismatch for bits=0x{:04X}",
bits
);
}
}
#[test]
fn exhaustive_encode_f16c_vs_software() {
for bits in 0u16..=0xFFFFu16 {
let f = f16_bits_to_f32(bits);
let expected = f32_to_f16_bits(f);
let got = unsafe { f32_to_f16_bits_f16c(f) };
assert_eq!(
got, expected,
"F16C encode mismatch for f={:e} (f16 bits 0x{:04X})",
f, bits
);
}
}
#[test]
fn encode_rounding_overflow_carry_vs_f16c() {
for exp in -14i32..=15 {
let boundary = 2.0f32.powi(exp + 1);
for k in 1..=4096u32 {
let f = boundary * (1.0 - (k as f32) * 1e-7);
for &v in &[f, -f] {
let expected = unsafe { f32_to_f16_bits_f16c(v) };
let got = f32_to_f16_bits(v);
assert_eq!(
got, expected,
"encoder carry mismatch for v={:e}: software=0x{:04X} hardware=0x{:04X}",
v, got, expected
);
}
}
}
}
}