#[must_use]
pub fn largest_remainder(p: &[f64], digits: u32) -> Vec<u64> {
assert!(digits <= 6, "precision is 0 to 6 digits");
assert!(!p.is_empty(), "a distribution has at least one value");
assert!(p.iter().all(|x| x.is_finite() && *x >= 0.0), "probabilities are finite and >= 0");
let total: f64 = p.iter().sum();
let scale = 10u64.pow(digits);
let scaled: Vec<f64> = if total > 0.0 {
p.iter().map(|x| x / total * scale as f64).collect()
} else {
vec![scale as f64 / p.len() as f64; p.len()]
};
let mut units: Vec<u64> = scaled.iter().map(|x| x.floor() as u64).collect();
let used: u64 = units.iter().sum();
let left = scale.saturating_sub(used) as usize;
let mut order: Vec<usize> = (0..p.len()).collect();
order.sort_by(|&a, &b| {
let ra = scaled[a] - scaled[a].floor();
let rb = scaled[b] - scaled[b].floor();
rb.total_cmp(&ra).then(scaled[b].total_cmp(&scaled[a])).then(a.cmp(&b))
});
for &i in order.iter().take(left) {
units[i] += 1;
}
units
}
#[must_use]
pub fn round_distribution(p: &[f64], digits: u32) -> Vec<f64> {
let scale = 10u64.pow(digits) as f64;
largest_remainder(p, digits).into_iter().map(|u| u as f64 / scale).collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_worked_example_is_already_exact() {
assert_eq!(largest_remainder(&[0.91, 0.02, 0.03, 0.04], 2), vec![91, 2, 3, 4]);
}
#[test]
fn thirds_sum_to_one() {
let units = largest_remainder(&[1.0 / 3.0; 3], 2);
assert_eq!(units.iter().sum::<u64>(), 100);
assert_eq!(units, vec![34, 33, 33]);
}
#[test]
fn many_small_values_still_sum_to_one() {
for k in [7, 32, 255] {
for digits in 0..=6 {
let p = vec![1.0 / k as f64; k];
let units = largest_remainder(&p, digits);
assert_eq!(units.iter().sum::<u64>(), 10u64.pow(digits), "k={k} digits={digits}");
}
}
}
#[test]
fn the_argmax_survives_rounding() {
let p = [0.3349, 0.3351, 0.33];
let units = largest_remainder(&p, 2);
assert_eq!(units.iter().sum::<u64>(), 100);
let top =
units.iter().enumerate().max_by_key(|&(i, u)| (*u, std::cmp::Reverse(i))).unwrap();
assert_eq!(top.0, 1);
}
#[test]
fn an_unnormalized_softmax_is_renormalized() {
let units = largest_remainder(&[0.4999999, 0.4999999], 2);
assert_eq!(units, vec![50, 50]);
}
#[test]
fn floats_come_back_on_the_grid() {
let v = round_distribution(&[0.123456, 0.876544], 3);
assert_eq!(v, vec![0.123, 0.877]);
}
}