use crate::quantization::turboquant::TQBits;
const CENTROIDS_1BIT: [f32; 2] = [-0.797_884_6, 0.797_884_6];
const CENTROIDS_1BIT_BOUNDARIES: [f32; 1] = calculate_boundaries(CENTROIDS_1BIT);
const CENTROIDS_2BIT: [f32; 4] = [-1.510, -0.4528, 0.4528, 1.510];
const CENTROIDS_2BIT_BOUNDARIES: [f32; 3] = calculate_boundaries(CENTROIDS_2BIT);
const CENTROIDS_4BIT: [f32; 16] = [
-2.733, -2.069, -1.618, -1.256, -0.9424, -0.6568, -0.3881, -0.1284, 0.1284, 0.3881, 0.6568,
0.9424, 1.256, 1.618, 2.069, 2.733,
];
const CENTROIDS_4BIT_BOUNDARIES: [f32; 15] = calculate_boundaries(CENTROIDS_4BIT);
const fn calculate_boundaries<const N: usize, const B: usize>(centroids: [f32; N]) -> [f32; B] {
assert!(B + 1 == N, "B must equal N - 1");
let mut out = [0.0; B];
let mut i = 0;
while i < B {
out[i] = (centroids[i] + centroids[i + 1]) / 2.0;
i += 1;
}
out
}
impl TQBits {
#[inline]
pub fn get_centroids(&self) -> &'static [f32] {
match self {
TQBits::Bits1 => &CENTROIDS_1BIT,
TQBits::Bits1_5 => &CENTROIDS_1BIT,
TQBits::Bits2 => &CENTROIDS_2BIT,
TQBits::Bits4 => &CENTROIDS_4BIT,
}
}
#[inline]
pub fn get_centroid_boundaries(&self) -> &'static [f32] {
match self {
TQBits::Bits1 => &CENTROIDS_1BIT_BOUNDARIES,
TQBits::Bits1_5 => &CENTROIDS_1BIT_BOUNDARIES,
TQBits::Bits2 => &CENTROIDS_2BIT_BOUNDARIES,
TQBits::Bits4 => &CENTROIDS_4BIT_BOUNDARIES,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::quantization::turboquant::math::std_normal_cdf;
fn std_normal_pdf(x: f64) -> f64 {
(-0.5 * x * x).exp() / (2.0 * std::f64::consts::PI).sqrt()
}
fn inv_std_normal_cdf(p: f64) -> f64 {
if p <= 0.0 {
return f64::NEG_INFINITY;
}
if p >= 1.0 {
return f64::INFINITY;
}
let t = if p < 0.5 {
(-2.0 * p.ln()).sqrt()
} else {
(-2.0 * (1.0 - p).ln()).sqrt()
};
let result = t
- (2.515517 + 0.802853 * t + 0.010328 * t * t)
/ (1.0 + 1.432788 * t + 0.189269 * t * t + 0.001308 * t * t * t);
if p < 0.5 { -result } else { result }
}
fn gaussian_conditional_expectation(sigma: f64, a: f64, b: f64) -> f64 {
let a_std = if a.is_finite() { a / sigma } else { a };
let b_std = if b.is_finite() { b / sigma } else { b };
let prob = if a_std.is_infinite() && a_std < 0.0 {
std_normal_cdf(b_std)
} else if b_std.is_infinite() && b_std > 0.0 {
1.0 - std_normal_cdf(a_std)
} else {
std_normal_cdf(b_std) - std_normal_cdf(a_std)
};
if prob < 1e-15 {
return if a.is_finite() && b.is_finite() {
(a + b) / 2.0
} else if a.is_finite() {
a + sigma
} else if b.is_finite() {
b - sigma
} else {
0.0
};
}
let pdf_diff = std_normal_pdf(a_std) - std_normal_pdf(b_std);
sigma * pdf_diff / prob
}
fn solve_lloyd_max(d: usize, bits: u32) -> Vec<f32> {
const MAX_ITER: usize = 200;
const TOL: f64 = 1e-10;
let n_levels = 1usize << bits;
let sigma = 1.0 / (d as f64).sqrt();
let mut boundaries: Vec<f64> = (1..n_levels)
.map(|i| sigma * inv_std_normal_cdf(i as f64 / n_levels as f64))
.collect();
let mut centroids = vec![0.0f64; n_levels];
for _iter in 0..MAX_ITER {
let mut max_shift: f64 = 0.0;
for i in 0..n_levels {
let a = if i == 0 {
f64::NEG_INFINITY
} else {
boundaries[i - 1]
};
let b = if i == n_levels - 1 {
f64::INFINITY
} else {
boundaries[i]
};
let new_c = gaussian_conditional_expectation(sigma, a, b);
max_shift = max_shift.max((new_c - centroids[i]).abs());
centroids[i] = new_c;
}
if max_shift < TOL {
break;
}
for i in 0..n_levels - 1 {
boundaries[i] = (centroids[i] + centroids[i + 1]) / 2.0;
}
}
centroids.iter().map(|&c| c as f32).collect()
}
#[test]
fn test_matches_hardcoded_centroids() {
let cases: &[(u32, &[f32])] = &[
(1, &CENTROIDS_1BIT),
(2, &CENTROIDS_2BIT),
(4, &CENTROIDS_4BIT),
];
for &(bits, expected) in cases {
let centroids = solve_lloyd_max(1, bits);
assert_eq!(
centroids.len(),
expected.len(),
"bits={bits}: wrong number of centroids"
);
for (i, (&got, &want)) in centroids.iter().zip(expected).enumerate() {
assert!(
(got - want).abs() < 1e-3,
"bits={bits}, centroid[{i}]: got {got}, expected {want}"
);
}
}
}
}