use statrs::distribution::{Beta, ContinuousCDF, Continuous};
use std::collections::HashMap;
use std::sync::Mutex;
static MEMO: Mutex<Option<HashMap<(usize, usize), (Vec<f32>, Vec<f32>)>>> = Mutex::new(None);
pub(crate) fn codebook(bits: usize, dim: usize) -> (Vec<f32>, Vec<f32>) {
if let Ok(memo) = MEMO.try_lock() {
if let Some(hit) = memo.as_ref().and_then(|m| m.get(&(bits, dim))) {
return hit.clone();
}
}
let computed = lloyd_max(bits, dim, 200, 1e-12);
if let Ok(mut memo) = MEMO.try_lock() {
memo.get_or_insert_with(HashMap::new)
.insert((bits, dim), computed.clone());
}
computed
}
fn lloyd_max(bits: usize, dim: usize, max_iter: usize, tol: f64) -> (Vec<f32>, Vec<f32>) {
let a = (dim as f64 - 1.0) / 2.0;
let beta = Beta::new(a, a).unwrap();
let n_levels = 1usize << bits;
let std_dev = (2.0 * a / ((2.0 * a + 1.0) * 4.0 * a)).sqrt(); let spread = 3.0 * std_dev;
let mut centroids: Vec<f64> = (0..n_levels)
.map(|i| -spread + 2.0 * spread * i as f64 / (n_levels as f64 - 1.0))
.collect();
for _ in 0..max_iter {
let boundaries: Vec<f64> = (0..n_levels - 1)
.map(|i| (centroids[i] + centroids[i + 1]) / 2.0)
.collect();
let mut edges = Vec::with_capacity(n_levels + 1);
edges.push(-1.0);
edges.extend_from_slice(&boundaries);
edges.push(1.0);
let mut new_centroids = vec![0.0f64; n_levels];
for i in 0..n_levels {
let lo = edges[i];
let hi = edges[i + 1];
let cdf_lo = beta.cdf((lo + 1.0) / 2.0);
let cdf_hi = beta.cdf((hi + 1.0) / 2.0);
let prob = cdf_hi - cdf_lo;
if prob < 1e-15 {
new_centroids[i] = centroids[i];
} else {
let mean = adaptive_simpson(
|x| {
let t = (x + 1.0) / 2.0;
x * beta.pdf(t) / 2.0
},
lo,
hi,
1e-14,
50,
);
new_centroids[i] = mean / prob;
}
}
let max_change = centroids
.iter()
.zip(new_centroids.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f64, f64::max);
centroids = new_centroids;
if max_change < tol {
break;
}
}
let centroids_f32: Vec<f32> = centroids.iter().map(|&c| c as f32).collect();
let boundaries: Vec<f32> = (0..n_levels - 1)
.map(|i| (centroids_f32[i] + centroids_f32[i + 1]) * 0.5)
.collect();
(boundaries, centroids_f32)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn boundaries_are_f32_midpoints_of_f32_centroids() {
for bits in 2..=4usize {
for dim in [8usize, 128, 200, 768, 1000, 1024, 1536, 3072] {
let (boundaries, centroids) = codebook(bits, dim);
assert_eq!(centroids.len(), 1 << bits);
assert_eq!(boundaries.len(), (1 << bits) - 1);
for i in 0..boundaries.len() {
let expect = (centroids[i] + centroids[i + 1]) * 0.5;
assert_eq!(
boundaries[i].to_bits(),
expect.to_bits(),
"bits={bits} dim={dim} boundary {i} is not the f32 \
midpoint of centroids {i} and {}",
i + 1,
);
}
for w in boundaries.windows(2) {
assert!(w[0] < w[1], "bits={bits} dim={dim}: boundaries not ascending");
}
for w in centroids.windows(2) {
assert!(w[0] < w[1], "bits={bits} dim={dim}: centroids not ascending");
}
}
}
}
}
fn adaptive_simpson<F: Fn(f64) -> f64>(f: F, a: f64, b: f64, tol: f64, max_depth: usize) -> f64 {
let mid = (a + b) / 2.0;
let fa = f(a);
let fb = f(b);
let fm = f(mid);
let whole = (b - a) / 6.0 * (fa + 4.0 * fm + fb);
adaptive_simpson_rec(&f, a, b, fa, fb, fm, whole, tol, max_depth)
}
fn adaptive_simpson_rec<F: Fn(f64) -> f64>(
f: &F,
a: f64,
b: f64,
fa: f64,
fb: f64,
fm: f64,
whole: f64,
tol: f64,
depth: usize,
) -> f64 {
let mid = (a + b) / 2.0;
let m1 = (a + mid) / 2.0;
let m2 = (mid + b) / 2.0;
let fm1 = f(m1);
let fm2 = f(m2);
let left = (mid - a) / 6.0 * (fa + 4.0 * fm1 + fm);
let right = (b - mid) / 6.0 * (fm + 4.0 * fm2 + fb);
let refined = left + right;
if depth == 0 || (refined - whole).abs() < 15.0 * tol {
refined + (refined - whole) / 15.0
} else {
adaptive_simpson_rec(f, a, mid, fa, fm, fm1, left, tol / 2.0, depth - 1)
+ adaptive_simpson_rec(f, mid, b, fm, fb, fm2, right, tol / 2.0, depth - 1)
}
}
#[cfg(test)]
mod fork_safety_tests {
use super::MEMO;
static SERIAL: std::sync::Mutex<()> = std::sync::Mutex::new(());
#[test]
fn codebook_does_not_block_on_a_held_memo_lock() {
let _serial = SERIAL.lock().unwrap_or_else(|e| e.into_inner());
let held = MEMO.lock().unwrap_or_else(|e| e.into_inner());
let (tx, rx) = std::sync::mpsc::channel();
let worker = std::thread::spawn(move || {
let (boundaries, centroids) = super::codebook(2, 1237);
let _ = tx.send((boundaries.len(), centroids.len()));
});
let got = rx.recv_timeout(std::time::Duration::from_secs(30));
drop(held);
let _ = worker.join();
assert_eq!(
got.ok(),
Some((3, 4)),
"codebook() blocked on the memoisation mutex — a process that forks \
while that lock is held would deadlock in the child on its first \
load (#147/#288/#321/#364)",
);
}
#[test]
fn codebook_repeats_are_served_from_the_memo() {
let _serial = SERIAL.lock().unwrap_or_else(|e| e.into_inner());
let cold = std::time::Instant::now();
let first = super::codebook(4, 1553);
let cold = cold.elapsed();
let warm = std::time::Instant::now();
let second = super::codebook(4, 1553);
let warm = warm.elapsed();
assert_eq!(first, second);
assert!(
warm * 20 < cold,
"repeat codebook() took {warm:?} against a cold solve of {cold:?} — \
the memo is not being hit",
);
}
}