#[derive(Debug, Clone, Copy, PartialEq)]
pub struct MixedPrecisionConfig {
pub init_scale: f32,
pub growth_factor: f32,
pub backoff_factor: f32,
pub growth_interval: u32,
pub min_scale: f32,
pub max_scale: f32,
pub use_bfloat16: bool,
}
impl Default for MixedPrecisionConfig {
fn default() -> Self {
Self {
init_scale: 65536.0,
growth_factor: 2.0,
backoff_factor: 0.5,
growth_interval: 2000,
min_scale: 1.0,
max_scale: 65536.0 * 128.0,
use_bfloat16: false,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct OverflowStats {
pub total_steps: usize,
pub overflow_count: usize,
pub overflow_rate: f32,
pub current_scale: f32,
}
#[derive(Debug, Clone)]
pub struct DynamicLossScaler {
scale: f32,
config: MixedPrecisionConfig,
growth_tracker: u32,
window: Vec<bool>,
}
impl DynamicLossScaler {
pub const WINDOW: usize = 100;
pub fn new(config: MixedPrecisionConfig) -> Self {
Self {
scale: config.init_scale.clamp(config.min_scale, config.max_scale),
config,
growth_tracker: 0,
window: Vec::with_capacity(Self::WINDOW),
}
}
pub fn scale(&self) -> f32 {
self.scale
}
pub fn inv_scale(&self) -> f32 {
1.0 / self.scale
}
pub fn update(&mut self, has_overflow: bool) {
if self.window.len() == Self::WINDOW {
self.window.remove(0);
}
self.window.push(has_overflow);
if has_overflow {
self.scale = (self.scale * self.config.backoff_factor).max(self.config.min_scale);
self.growth_tracker = 0;
} else {
self.growth_tracker = self.growth_tracker.saturating_add(1);
if self.growth_tracker >= self.config.growth_interval {
self.scale = (self.scale * self.config.growth_factor).min(self.config.max_scale);
self.growth_tracker = 0;
}
}
}
pub fn overflow_stats(&self) -> OverflowStats {
let total = self.window.len();
let overflows = self.window.iter().filter(|&&x| x).count();
OverflowStats {
total_steps: total,
overflow_count: overflows,
overflow_rate: if total > 0 {
overflows as f32 / total as f32
} else {
0.0
},
current_scale: self.scale,
}
}
pub fn unscale_and_check(&mut self, values: &mut [f32]) -> bool {
let inv = self.inv_scale();
let mut overflow = false;
for v in values.iter_mut() {
if !v.is_finite() {
overflow = true;
}
*v *= inv;
}
self.update(overflow);
overflow
}
}
pub fn f32_to_f16_bits(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 | 0x7c00 | ((mantissa >> 13) as u16) | 0x0200
};
}
let unbiased = exponent - 127;
let half_exp = unbiased + 15;
if half_exp >= 0x1f {
return sign | 0x7c00;
}
if half_exp <= 0 {
if half_exp < -10 {
return sign;
}
let significand = mantissa | 0x0080_0000;
let shift = (14 - half_exp) as u32; let result = significand >> shift;
let round_bit = 1u32 << (shift - 1);
let remainder = significand & (round_bit.saturating_mul(2) - 1);
let mut half = result as u16;
if remainder > round_bit || (remainder == round_bit && (result & 1) == 1) {
half = half.wrapping_add(1);
}
return sign | half;
}
let mut half = ((half_exp as u16) << 10) | ((mantissa >> 13) as u16);
let remainder = mantissa & 0x1fff;
if remainder > 0x1000 || (remainder == 0x1000 && (half & 1) == 1) {
half = half.wrapping_add(1);
}
sign | half
}
pub fn f16_bits_to_f32(bits: u16) -> f32 {
let sign = ((bits as u32) & 0x8000) << 16;
let exponent = ((bits >> 10) & 0x1f) as u32;
let mantissa = ((bits & 0x03ff) as u32) << 13;
if exponent == 0 {
if mantissa == 0 {
return f32::from_bits(sign);
}
let mut mant = mantissa;
let mut shifts: i32 = 0;
while mant & 0x0080_0000 == 0 {
mant <<= 1;
shifts += 1;
}
mant &= 0x007f_ffff;
let f32_exp = ((113 - shifts) as u32) << 23;
return f32::from_bits(sign | f32_exp | mant);
}
if exponent == 0x1f {
return f32::from_bits(sign | 0x7f80_0000 | mantissa);
}
let f32_exp = (exponent + (127 - 15)) << 23;
f32::from_bits(sign | f32_exp | mantissa)
}
pub fn f32_slice_to_f16_bits(values: &[f32]) -> Vec<u16> {
values.iter().copied().map(f32_to_f16_bits).collect()
}
pub fn f16_bits_slice_to_f32(bits: &[u16]) -> Vec<f32> {
bits.iter().copied().map(f16_bits_to_f32).collect()
}
pub const F16_MAX: f32 = 65504.0;
pub fn saturate_to_f16_range(value: f32) -> f32 {
if value.is_nan() {
value
} else {
value.clamp(-F16_MAX, F16_MAX)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn round_trips_exact_values() {
for &v in &[
0.0f32, -0.0, 1.0, -1.0, 0.5, 2.0, 65504.0, -65504.0, 0.125, 1024.0,
] {
let back = f16_bits_to_f32(f32_to_f16_bits(v));
assert_eq!(back.to_bits(), v.to_bits(), "value {v} did not round-trip");
}
}
#[test]
fn round_trips_every_finite_f16() {
for bits in 0u16..=u16::MAX {
let exponent = (bits >> 10) & 0x1f;
let value = f16_bits_to_f32(bits);
if exponent == 0x1f {
let back = f32_to_f16_bits(value);
assert_eq!(
(back >> 10) & 0x1f,
0x1f,
"bits {bits:#06x} lost its inf/NaN class"
);
assert_eq!(back & 0x8000, bits & 0x8000, "bits {bits:#06x} lost sign");
continue;
}
let back = f32_to_f16_bits(value);
assert_eq!(
back, bits,
"bits {bits:#06x} did not round-trip (f32 {value})"
);
}
}
#[test]
fn subnormals_are_not_flushed_to_zero() {
let smallest = f16_bits_to_f32(1);
assert!(smallest > 0.0);
assert!((smallest - 2f32.powi(-24)).abs() < f32::EPSILON * smallest);
assert_eq!(f32_to_f16_bits(smallest), 1);
}
#[test]
fn rounds_half_to_even() {
let halfway = 1.0f32 + 2f32.powi(-11);
assert_eq!(f32_to_f16_bits(halfway), f32_to_f16_bits(1.0));
let halfway_up = 1.0f32 + 3.0 * 2f32.powi(-11);
assert_eq!(f32_to_f16_bits(halfway_up), f32_to_f16_bits(1.0) + 2);
}
#[test]
fn overflow_saturates_to_infinity() {
assert_eq!(f32_to_f16_bits(1.0e30), 0x7c00);
assert_eq!(f32_to_f16_bits(-1.0e30), 0xfc00);
assert!(f16_bits_to_f32(0x7c00).is_infinite());
}
#[test]
fn nan_stays_nan() {
assert!(f16_bits_to_f32(f32_to_f16_bits(f32::NAN)).is_nan());
}
#[test]
fn loss_scaler_backs_off_and_grows() {
let config = MixedPrecisionConfig {
growth_interval: 4,
..MixedPrecisionConfig::default()
};
let mut scaler = DynamicLossScaler::new(config);
assert_eq!(scaler.scale(), 65536.0);
scaler.update(true);
assert_eq!(scaler.scale(), 32768.0);
for _ in 0..4 {
scaler.update(false);
}
assert_eq!(scaler.scale(), 65536.0);
}
#[test]
fn loss_scaler_detects_overflow_while_unscaling() {
let mut scaler = DynamicLossScaler::new(MixedPrecisionConfig {
init_scale: 4.0,
..MixedPrecisionConfig::default()
});
let mut grads = [8.0f32, -4.0, 2.0];
assert!(!scaler.unscale_and_check(&mut grads));
assert_eq!(grads, [2.0, -1.0, 0.5]);
assert_eq!(scaler.scale(), 4.0);
let mut bad = [f32::INFINITY, 1.0];
assert!(scaler.unscale_and_check(&mut bad));
assert_eq!(scaler.scale(), 2.0);
let stats = scaler.overflow_stats();
assert_eq!(stats.total_steps, 2);
assert_eq!(stats.overflow_count, 1);
}
#[test]
fn saturation_keeps_values_finite() {
assert_eq!(saturate_to_f16_range(1.0e30), F16_MAX);
assert_eq!(saturate_to_f16_range(-1.0e30), -F16_MAX);
assert!(saturate_to_f16_range(f32::NAN).is_nan());
}
}