pub trait FixedWidth: Copy {
const MAX_CODE: u32;
fn from_code(code: u32) -> Self;
fn code(self) -> u32;
}
impl FixedWidth for u8 {
const MAX_CODE: u32 = u8::MAX as u32;
fn from_code(code: u32) -> Self {
code as u8
}
fn code(self) -> u32 {
u32::from(self)
}
}
impl FixedWidth for u16 {
const MAX_CODE: u32 = u16::MAX as u32;
fn from_code(code: u32) -> Self {
code as u16
}
fn code(self) -> u32 {
u32::from(self)
}
}
impl FixedWidth for u32 {
const MAX_CODE: u32 = u32::MAX;
fn from_code(code: u32) -> Self {
code
}
fn code(self) -> u32 {
self
}
}
#[must_use]
pub fn quantize_dist<Q: FixedWidth>(probs: &[f32]) -> Vec<Q> {
let scale = Q::MAX_CODE as f32;
probs
.iter()
.map(|&x| {
let clamped = x.clamp(0.0, 1.0);
Q::from_code((clamped * scale).round() as u32)
})
.collect()
}
#[must_use]
pub fn dequantize_dist<Q: FixedWidth>(codes: &[Q]) -> Vec<f32> {
let scale = Q::MAX_CODE as f32;
let mut out: Vec<f32> = codes.iter().map(|&c| c.code() as f32 / scale).collect();
crate::probability::normalize_inplace(&mut out);
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn u16_roundtrip_within_one_quantum() {
let probs = [0.1f32, 0.2, 0.3, 0.4];
let back = dequantize_dist::<u16>(&quantize_dist::<u16>(&probs));
for (a, b) in probs.iter().zip(&back) {
assert!((a - b).abs() < 1.0 / 65535.0 + 1e-7, "{a} vs {b}");
}
assert!((back.iter().sum::<f32>() - 1.0).abs() < 1e-6);
}
#[test]
fn one_hot_roundtrips() {
let back = dequantize_dist::<u16>(&quantize_dist::<u16>(&[0.0f32, 1.0, 0.0]));
assert!((back[1] - 1.0).abs() < 1e-6);
assert!(back[0].abs() < 1e-6 && back[2].abs() < 1e-6);
}
#[test]
fn all_zero_codes_decode_to_uniform() {
let back = dequantize_dist::<u16>(&[0u16, 0, 0, 0]);
assert!(back.iter().all(|&v| (v - 0.25).abs() < 1e-6));
}
#[test]
fn u8_and_u32_widths_supported() {
let probs = [0.25f32, 0.25, 0.5];
let b8 = dequantize_dist::<u8>(&quantize_dist::<u8>(&probs));
assert!((b8.iter().sum::<f32>() - 1.0).abs() < 1e-6);
let b32 = dequantize_dist::<u32>(&quantize_dist::<u32>(&probs));
for (a, b) in probs.iter().zip(&b32) {
assert!((a - b).abs() < 1e-6, "{a} vs {b}");
}
}
#[test]
fn out_of_range_inputs_are_clamped() {
let codes = quantize_dist::<u16>(&[2.0, -1.0]);
assert_eq!(codes[0], 65535);
assert_eq!(codes[1], 0);
}
}