use super::*;
use crate::audio::lid::Span;
const WINDOW: usize = 160_000;
fn row_from_logits(logits: &[(usize, f64)]) -> LogProbabilities {
let mut values = vec![-30.0f64; NUM_LANGUAGES];
for &(index, logit) in logits {
values[index] = logit;
}
let max = values.iter().copied().fold(f64::NEG_INFINITY, f64::max);
let log_sum = values.iter().map(|v| (v - max).exp()).sum::<f64>().ln();
LogProbabilities::new(
values
.into_iter()
.map(|v| ((v - max) - log_sum) as f32)
.collect(),
)
}
fn window(row: LogProbabilities, start: usize, len: usize) -> WindowLogProbabilities {
WindowLogProbabilities::new(row, Span::new(start, len, WINDOW))
}
fn mass(row: &LogProbabilities) -> f64 {
row.as_slice().iter().map(|v| f64::from(*v).exp()).sum()
}
fn all_poolings() -> [ScorePooling; 4] {
[
ScorePooling::MeanLogProbability,
ScorePooling::MeanProbability,
ScorePooling::Max,
ScorePooling::Vote,
]
}
#[test]
fn the_default_pooling_is_the_logarithmic_pool() {
assert_eq!(ScorePooling::default(), ScorePooling::MeanLogProbability);
}
#[test]
fn one_window_is_the_bit_exact_identity() {
let row = row_from_logits(&[(94, 6.0), (3, 1.5), (17, 0.25)]);
for pooling in all_poolings() {
let out = aggregate_windows(pooling, &[window(row.clone(), 0, WINDOW)]).expect("aggregate");
assert_eq!(out.as_slice(), row.as_slice(), "{pooling:?}");
}
let out = aggregate_windows(
ScorePooling::MeanProbability,
&[window(row.clone(), 0, 4_321)],
)
.expect("aggregate");
assert_eq!(out.as_slice(), row.as_slice());
}
#[test]
fn identical_windows_reproduce_their_common_row() {
let row = row_from_logits(&[(94, 6.0), (3, 1.5)]);
let windows = vec![
window(row.clone(), 0, WINDOW),
window(row.clone(), WINDOW, WINDOW),
];
for pooling in all_poolings() {
let out = aggregate_windows(pooling, &windows).expect("aggregate");
if pooling == ScorePooling::Vote {
assert_eq!(out.as_slice()[94], 0.0);
continue;
}
for (index, (&got, &want)) in out.as_slice().iter().zip(row.as_slice()).enumerate() {
assert!(
(got - want).abs() < 1e-5,
"{pooling:?} column {index}: {got} vs {want}"
);
}
}
}
#[test]
fn the_log_pool_and_the_linear_pool_disagree() {
let a = LogProbabilities::try_from_slice(&distribution(&[(0, 0.98), (1, 0.02)])).expect("row");
let b = LogProbabilities::try_from_slice(&distribution(&[(0, 0.30), (1, 0.70)])).expect("row");
let windows = vec![window(a, 0, WINDOW), window(b, WINDOW, WINDOW)];
let linear = aggregate_windows(ScorePooling::MeanProbability, &windows).expect("aggregate");
let logarithmic =
aggregate_windows(ScorePooling::MeanLogProbability, &windows).expect("aggregate");
assert!((f64::from(linear.as_slice()[0]).exp() - 0.64).abs() < 1e-4);
assert!((f64::from(linear.as_slice()[1]).exp() - 0.36).abs() < 1e-4);
assert!(
(f64::from(logarithmic.as_slice()[0]).exp() - 0.8209).abs() < 1e-3,
"{}",
f64::from(logarithmic.as_slice()[0]).exp()
);
assert!((f64::from(logarithmic.as_slice()[1]).exp() - 0.1791).abs() < 1e-3);
assert!(
(f64::from(linear.as_slice()[0]) - f64::from(logarithmic.as_slice()[0])).abs() > 0.2,
"the two means must not coincide"
);
}
#[test]
fn max_takes_the_per_language_peak() {
let a = LogProbabilities::try_from_slice(&distribution(&[(0, 0.90), (1, 0.10)])).expect("row");
let b = LogProbabilities::try_from_slice(&distribution(&[(0, 0.20), (2, 0.80)])).expect("row");
let out = aggregate_windows(
ScorePooling::Max,
&[window(a, 0, WINDOW), window(b, WINDOW, WINDOW)],
)
.expect("aggregate");
assert!((f64::from(out.as_slice()[0]).exp() - 0.5).abs() < 1e-3);
assert!((f64::from(out.as_slice()[1]).exp() - 0.0556).abs() < 1e-3);
assert!((f64::from(out.as_slice()[2]).exp() - 0.4444).abs() < 1e-3);
assert!((mass(&out) - 1.0).abs() < 1e-5);
}
#[test]
fn a_vote_counts_winners_and_discards_magnitude() {
let landslide =
LogProbabilities::try_from_slice(&distribution(&[(0, 0.999), (1, 0.001)])).expect("row");
let squeaker =
LogProbabilities::try_from_slice(&distribution(&[(0, 0.490), (1, 0.510)])).expect("row");
let out = aggregate_windows(
ScorePooling::Vote,
&[
window(landslide, 0, WINDOW),
window(squeaker, WINDOW, WINDOW),
],
)
.expect("aggregate");
assert_eq!(out.as_slice()[0], 0.5f32.ln());
assert_eq!(out.as_slice()[1], 0.5f32.ln());
assert_eq!(out.as_slice()[2], f32::NEG_INFINITY);
assert!((mass(&out) - 1.0).abs() < 1e-6);
let ranked = out.top_k(NUM_LANGUAGES).expect("rank");
assert_eq!(ranked[0].index(), 0, "ties break by ascending column");
assert_eq!(ranked[1].index(), 1);
assert_eq!(ranked[2].probability(), 0.0);
}
#[test]
fn a_windows_vote_agrees_with_its_own_top_1() {
let tied = LogProbabilities::try_from_slice(&distribution(&[(40, 0.5), (7, 0.5)])).expect("row");
assert_eq!(argmax(tied.as_slice()), 7);
assert_eq!(tied.top_k(1).expect("rank")[0].index(), 7);
let out = aggregate_windows(
ScorePooling::Vote,
&[
window(tied.clone(), 0, WINDOW),
window(tied, WINDOW, WINDOW),
],
)
.expect("aggregate");
assert_eq!(out.as_slice()[7], 0.0);
assert_eq!(out.as_slice()[40], f32::NEG_INFINITY);
}
#[test]
fn windows_weigh_by_the_audio_they_actually_saw() {
let full = LogProbabilities::try_from_slice(&distribution(&[(0, 1.0)])).expect("row");
let sliver = LogProbabilities::try_from_slice(&distribution(&[(1, 1.0)])).expect("row");
let weighted = aggregate_windows(
ScorePooling::MeanProbability,
&[
window(full.clone(), 0, WINDOW),
window(sliver.clone(), WINDOW, WINDOW / 10),
],
)
.expect("aggregate");
assert!((f64::from(weighted.as_slice()[0]).exp() - 10.0 / 11.0).abs() < 1e-4);
assert!((f64::from(weighted.as_slice()[1]).exp() - 1.0 / 11.0).abs() < 1e-4);
let equal = aggregate_windows(
ScorePooling::MeanProbability,
&[window(full, 0, WINDOW), window(sliver, WINDOW, WINDOW)],
)
.expect("aggregate");
assert!((f64::from(equal.as_slice()[0]).exp() - 0.5).abs() < 1e-4);
}
#[test]
fn a_vote_is_weighted_by_duration() {
let body = LogProbabilities::try_from_slice(&distribution(&[(0, 0.9), (1, 0.1)])).expect("row");
let tail = LogProbabilities::try_from_slice(&distribution(&[(1, 0.9), (0, 0.1)])).expect("row");
let out = aggregate_windows(
ScorePooling::Vote,
&[window(body, 0, WINDOW), window(tail, WINDOW, WINDOW / 4)],
)
.expect("aggregate");
assert!((f64::from(out.as_slice()[0]).exp() - 0.8).abs() < 1e-6);
assert!((f64::from(out.as_slice()[1]).exp() - 0.2).abs() < 1e-6);
}
#[test]
fn max_ignores_duration() {
let a = LogProbabilities::try_from_slice(&distribution(&[(0, 0.9), (1, 0.1)])).expect("row");
let b = LogProbabilities::try_from_slice(&distribution(&[(1, 0.7), (0, 0.3)])).expect("row");
let long = aggregate_windows(
ScorePooling::Max,
&[
window(a.clone(), 0, WINDOW),
window(b.clone(), WINDOW, WINDOW),
],
)
.expect("aggregate");
let short = aggregate_windows(
ScorePooling::Max,
&[window(a, 0, WINDOW), window(b, WINDOW, 2_000)],
)
.expect("aggregate");
assert_eq!(long.as_slice(), short.as_slice());
}
#[test]
fn every_pooling_returns_a_normalized_log_distribution() {
let windows = vec![
window(row_from_logits(&[(94, 7.0), (3, 2.0)]), 0, WINDOW),
window(
row_from_logits(&[(3, 4.0), (94, 3.5), (61, 3.0)]),
WINDOW,
WINDOW,
),
window(row_from_logits(&[(61, 9.0)]), 2 * WINDOW, WINDOW / 3),
];
for pooling in all_poolings() {
let out = aggregate_windows(pooling, &windows).expect("aggregate");
assert_eq!(out.as_slice().len(), NUM_LANGUAGES, "{pooling:?}");
assert!(
out.as_slice().iter().all(|v| !v.is_nan() && *v <= 0.0),
"{pooling:?} produced a NaN or a positive log-probability"
);
assert!(
(mass(&out) - 1.0).abs() < 1e-4,
"{pooling:?} mass {}",
mass(&out)
);
assert!(
LogProbabilities::try_from_slice(out.as_slice()).is_ok(),
"{pooling:?}"
);
}
}
#[test]
fn the_batch_and_streaming_folds_agree_bit_for_bit() {
let windows = vec![
window(row_from_logits(&[(94, 7.0), (3, 2.0)]), 0, WINDOW),
window(row_from_logits(&[(3, 4.0), (94, 3.5)]), WINDOW, WINDOW),
window(row_from_logits(&[(61, 9.0)]), 2 * WINDOW, 7_777),
];
for pooling in all_poolings() {
let batch = aggregate_windows(pooling, &windows).expect("batch");
let mut acc = Accumulator::new(pooling);
for w in &windows {
acc.push(w.value(), w.span().len()).expect("push");
}
let streamed = acc.finish().expect("streamed");
assert_eq!(batch.as_slice(), streamed.as_slice(), "{pooling:?}");
}
}
#[test]
fn an_empty_window_list_is_a_typed_refusal() {
for pooling in all_poolings() {
assert!(matches!(
aggregate_windows(pooling, &[]),
Err(Error::EmptyWindows)
));
}
}
#[test]
fn renormalize_does_not_turn_an_impossible_row_into_nan() {
let mut values = vec![f64::NEG_INFINITY; NUM_LANGUAGES];
renormalize(&mut values);
assert!(values.iter().all(|v| *v == f64::NEG_INFINITY));
}
#[test]
fn a_zero_mass_logarithmic_pool_is_refused_rather_than_returned() {
let certain_of = |index: usize| {
let mut values = vec![f32::NEG_INFINITY; NUM_LANGUAGES];
values[index] = 0.0;
LogProbabilities::try_from_slice(&values).expect("one zero among -inf normalizes exactly")
};
let windows = vec![
window(certain_of(100), 0, WINDOW),
window(certain_of(101), WINDOW, WINDOW),
];
let pooled = aggregate_windows(ScorePooling::MeanLogProbability, &windows);
assert!(
matches!(
pooled,
Err(Error::ZeroMassAggregate(ScorePooling::MeanLogProbability))
),
"a zero-mass pool must be a typed refusal, got {}",
describe(&pooled)
);
}
#[test]
fn the_linear_pool_keeps_a_finite_tail_through_exp_underflow() {
let mut values = vec![-1_000.0f32; NUM_LANGUAGES];
values[0] = 0.0;
values[100] = -800.0;
values[101] = -900.0;
let row = LogProbabilities::try_from_slice(&values).expect("a finite row, normalized to 1");
let pooled = aggregate_windows(
ScorePooling::MeanProbability,
&[
window(row.clone(), 0, WINDOW),
window(row.clone(), WINDOW, WINDOW),
],
)
.expect("aggregate");
let ranked: Vec<usize> = pooled
.top_k(3)
.expect("rank")
.iter()
.map(|score| score.index())
.collect();
assert_eq!(
ranked,
vec![0, 100, 101],
"the finite tail must keep its rank; row[100]={} row[101]={} row[1]={}",
pooled.as_slice()[100],
pooled.as_slice()[101],
pooled.as_slice()[1],
);
assert_eq!(pooled.as_slice(), row.as_slice());
}
#[test]
fn every_pooling_either_returns_a_distribution_or_refuses() {
let nothing = LogProbabilities::try_from_slice(&vec![f32::NEG_INFINITY; NUM_LANGUAGES])
.expect("an all-zero-probability row is accepted");
for pooling in all_poolings() {
let pooled = aggregate_windows(
pooling,
&[
window(nothing.clone(), 0, WINDOW),
window(nothing.clone(), WINDOW, WINDOW),
],
);
assert!(
matches!(&pooled, Err(Error::UnnormalizableWindow(0))),
"{pooling:?}: {}",
describe(&pooled)
);
}
for pooling in all_poolings() {
let pooled = aggregate_windows(pooling, &[window(nothing.clone(), 0, WINDOW)]);
assert!(
matches!(&pooled, Err(Error::UnnormalizableWindow(0))),
"{pooling:?} identity: {}",
describe(&pooled)
);
}
let certain_of = |index: usize| {
let mut values = vec![f32::NEG_INFINITY; NUM_LANGUAGES];
values[index] = 0.0;
LogProbabilities::try_from_slice(&values).expect("one zero among -inf")
};
let voted = aggregate_windows(
ScorePooling::Vote,
&[
window(certain_of(50), 0, WINDOW),
window(certain_of(3), WINDOW, WINDOW),
],
)
.expect("a vote over rows that chose something");
assert_eq!(voted.as_slice()[50], 0.5f32.ln());
assert_eq!(voted.as_slice()[3], 0.5f32.ln());
assert_eq!(voted.as_slice()[0], f32::NEG_INFINITY);
assert!((mass(&voted) - 1.0).abs() < 1e-6, "mass {}", mass(&voted));
}
#[test]
fn one_language_ruled_out_does_not_rule_out_the_row() {
let mut values = distribution(&[(0, 0.6), (1, 0.4)]);
values[5] = f32::NEG_INFINITY;
let rules_out = LogProbabilities::try_from_slice(&values).expect("row");
let votes_for =
LogProbabilities::try_from_slice(&distribution(&[(5, 0.7), (0, 0.3)])).expect("row");
let windows = vec![
window(rules_out, 0, WINDOW),
window(votes_for, WINDOW, WINDOW),
];
for pooling in all_poolings() {
let out = aggregate_windows(pooling, &windows).expect("aggregate");
assert!(
(mass(&out) - 1.0).abs() < 1e-5,
"{pooling:?} mass {}",
mass(&out)
);
let zeroed = out.as_slice()[5] == f32::NEG_INFINITY;
assert_eq!(
zeroed,
pooling == ScorePooling::MeanLogProbability,
"{pooling:?} column 5 = {}",
out.as_slice()[5]
);
}
}
fn describe(pooled: &Result<LogProbabilities>) -> String {
match pooled {
Ok(row) => format!(
"Ok(mass {}, top-3 {:?}, values[0..3] {:?})",
mass(row),
row
.top_k(3)
.expect("rank")
.iter()
.map(|score| score.index())
.collect::<Vec<_>>(),
&row.as_slice()[..3],
),
Err(error) => format!("Err({error})"),
}
}
fn distribution(entries: &[(usize, f64)]) -> Vec<f32> {
let named: f64 = entries.iter().map(|(_, p)| *p).sum();
let rest = ((1.0 - named) / (NUM_LANGUAGES - entries.len()) as f64).max(f64::MIN_POSITIVE);
let mut values = vec![rest.ln() as f32; NUM_LANGUAGES];
for &(index, p) in entries {
values[index] = p.ln() as f32;
}
values
}
#[test]
fn a_pool_far_below_zero_is_still_normalized() {
let certain_among_giants = |index: usize| {
let mut values = vec![-1e20f32; NUM_LANGUAGES];
values[index] = 0.0;
LogProbabilities::try_from_slice(&values).expect("finite and non-positive")
};
let windows = vec![
window(certain_among_giants(0), 0, WINDOW),
window(certain_among_giants(1), WINDOW, WINDOW),
];
for pooling in all_poolings() {
let pooled = aggregate_windows(pooling, &windows);
let row = match &pooled {
Ok(row) => row,
Err(error) => panic!("{pooling:?} expected a distribution, got Err({error})"),
};
assert!(
(mass(row) - 1.0).abs() < 1e-6,
"{pooling:?} mass {}, columns 0 and 1 = {} / {}",
mass(row),
row.as_slice()[0],
row.as_slice()[1]
);
assert!(
(f64::from(row.as_slice()[0]).exp() - 0.5).abs() < 1e-6,
"{pooling:?} column 0 = {}",
row.as_slice()[0]
);
assert!(
(f64::from(row.as_slice()[1]).exp() - 0.5).abs() < 1e-6,
"{pooling:?} column 1 = {}",
row.as_slice()[1]
);
}
}
#[test]
fn renormalize_does_not_lose_its_constant_against_a_huge_shift() {
let mut values = vec![-1e20f64; NUM_LANGUAGES];
values[0] = -5e19;
values[1] = -5e19;
renormalize(&mut values);
let total: f64 = values.iter().map(|v| v.exp()).sum();
assert!((total - 1.0).abs() < 1e-12, "total {total}");
assert!(
(values[0] - 0.5f64.ln()).abs() < 1e-12,
"column 0 = {}",
values[0]
);
}
#[test]
fn a_window_with_no_probability_mass_is_refused_not_folded() {
let mut one_hot = vec![f32::NEG_INFINITY; NUM_LANGUAGES];
one_hot[0] = 0.0;
let says_something = LogProbabilities::try_from_slice(&one_hot).expect("row");
let says_nothing = LogProbabilities::try_from_slice(&vec![f32::NEG_INFINITY; NUM_LANGUAGES])
.expect("an all-zero-probability row is accepted");
let windows = vec![
window(says_something, 0, WINDOW),
window(says_nothing, WINDOW, WINDOW),
];
let pooled = aggregate_windows(ScorePooling::MeanProbability, &windows);
assert!(
matches!(&pooled, Err(Error::UnnormalizableWindow(1))),
"expected the SECOND window to be named, got {}",
describe(&pooled)
);
}
#[test]
fn the_refusal_names_the_window_that_ruled_everything_out() {
let mut one_hot = vec![f32::NEG_INFINITY; NUM_LANGUAGES];
one_hot[7] = 0.0;
let says_something = LogProbabilities::try_from_slice(&one_hot).expect("row");
let says_nothing =
LogProbabilities::try_from_slice(&vec![f32::NEG_INFINITY; NUM_LANGUAGES]).expect("row");
for position in 0..3usize {
let windows: Vec<_> = (0..3usize)
.map(|i| {
let row = if i == position {
says_nothing.clone()
} else {
says_something.clone()
};
window(row, i * WINDOW, WINDOW)
})
.collect();
let pooled = aggregate_windows(ScorePooling::Max, &windows);
assert!(
matches!(&pooled, Err(Error::UnnormalizableWindow(got)) if *got == position),
"position {position}: {}",
describe(&pooled)
);
}
}
#[test]
fn a_vote_is_not_cast_by_a_window_that_ruled_everything_out() {
let mut one_hot = vec![f32::NEG_INFINITY; NUM_LANGUAGES];
one_hot[50] = 0.0;
let says_fifty = LogProbabilities::try_from_slice(&one_hot).expect("row");
let says_nothing =
LogProbabilities::try_from_slice(&vec![f32::NEG_INFINITY; NUM_LANGUAGES]).expect("row");
let pooled = aggregate_windows(
ScorePooling::Vote,
&[
window(says_fifty, 0, WINDOW),
window(says_nothing, WINDOW, WINDOW),
],
);
assert!(
matches!(&pooled, Err(Error::UnnormalizableWindow(1))),
"expected a typed refusal, got {}",
describe(&pooled)
);
}
#[test]
fn max_over_rows_of_huge_negatives_pools_to_a_distribution() {
let row = LogProbabilities::try_from_slice(&vec![-1e20f32; NUM_LANGUAGES]).expect("row");
let pooled = aggregate_windows(
ScorePooling::Max,
&[window(row.clone(), 0, WINDOW), window(row, WINDOW, WINDOW)],
);
let out = match &pooled {
Ok(out) => out,
Err(_) => panic!(
"a row with a finite maximum must fold, got {}",
describe(&pooled)
),
};
let total = mass(out);
assert!((total - 1.0).abs() < 1e-6, "mass {total}");
let uniform = (1.0 / NUM_LANGUAGES as f64).ln();
for (index, value) in out.as_slice().iter().enumerate() {
assert!(
(f64::from(*value) - uniform).abs() < 1e-6,
"column {index}: {value} against the uniform {uniform}"
);
}
}
#[test]
fn the_linear_pool_returns_a_distribution_from_rows_that_only_nearly_are() {
let short = (0.99f64 / NUM_LANGUAGES as f64).ln() as f32;
let deficient = LogProbabilities::try_from_slice(&vec![short; NUM_LANGUAGES]).expect("row");
let windows = vec![
window(deficient.clone(), 0, WINDOW),
window(deficient, WINDOW, WINDOW),
];
let pooled = aggregate_windows(ScorePooling::MeanProbability, &windows);
let row = match &pooled {
Ok(row) => row,
Err(error) => panic!("expected a distribution, got Err({error})"),
};
assert!((mass(row) - 1.0).abs() < 1e-6, "mass {}", mass(row));
}
#[test]
fn a_row_that_is_not_a_distribution_carries_the_mass_it_left() {
let error = Error::from(NotADistribution::new(ScorePooling::MeanProbability, 0.5));
let rendered = error.to_string();
assert!(rendered.contains("MeanProbability"), "{rendered}");
assert!(rendered.contains("0.5"), "{rendered}");
assert!(rendered.contains("not a distribution"), "{rendered}");
let Error::NotADistribution(payload) = error else {
panic!("wrong variant")
};
assert_eq!(payload.pooling(), ScorePooling::MeanProbability);
assert!((payload.mass() - 0.5).abs() < f64::EPSILON);
}
#[test]
fn every_folded_row_is_normalized_to_far_inside_the_tolerance() {
let windows = vec![
window(row_from_logits(&[(94, 7.0), (3, 2.0)]), 0, WINDOW),
window(
row_from_logits(&[(3, 4.0), (94, 3.5), (61, 3.0)]),
WINDOW,
WINDOW,
),
window(row_from_logits(&[(61, 9.0)]), 2 * WINDOW, WINDOW / 3),
window(
row_from_logits(&(0..NUM_LANGUAGES).map(|i| (i, 0.0)).collect::<Vec<_>>()),
3 * WINDOW,
WINDOW,
),
];
for pooling in all_poolings() {
let out = aggregate_windows(pooling, &windows).expect("aggregate");
let deviation = (mass(&out) - 1.0).abs();
assert!(
deviation < MAX_MASS_DEVIATION / 10.0,
"{pooling:?} deviation {deviation:e} is within an order of magnitude of \
the tolerance {MAX_MASS_DEVIATION:e}"
);
}
}
const CPU_ONLY_ROW_MASS: f64 = 0.99235;
fn rescaled(row: &LogProbabilities, by: f32) -> LogProbabilities {
let values: Vec<f32> = row.as_slice().iter().map(|v| v + by).collect();
LogProbabilities::try_from_slice(&values).expect("a non-positive row stays one under a shift")
}
fn max_abs_difference(a: &LogProbabilities, b: &LogProbabilities) -> f64 {
a.as_slice()
.iter()
.zip(b.as_slice())
.map(|(&x, &y)| {
if x == y {
0.0
} else {
(f64::from(x) - f64::from(y)).abs()
}
})
.fold(0.0f64, f64::max)
}
#[test]
fn a_windows_own_mass_deficit_does_not_outvote_an_equally_certain_window() {
let one_hot = |index: usize, value: f32| {
let mut values = vec![f32::NEG_INFINITY; NUM_LANGUAGES];
values[index] = value;
LogProbabilities::try_from_slice(&values).expect("one finite value among -inf")
};
let deficit = CPU_ONLY_ROW_MASS.ln() as f32;
let windows = vec![
window(one_hot(0, deficit), 0, WINDOW),
window(one_hot(1, 0.0), WINDOW, WINDOW),
];
let mut wrong = Vec::new();
for pooling in [ScorePooling::MeanProbability, ScorePooling::Max] {
let out = aggregate_windows(pooling, &windows).expect("aggregate");
let (col0, col1) = (out.as_slice()[0], out.as_slice()[1]);
let winner = argmax(out.as_slice());
println!(
"codex trigger {pooling:>18?} winner {winner} col0 {col0:.8} col1 {col1:.8} \
p0 {:.9} p1 {:.9} mass {:.10}",
f64::from(col0).exp(),
f64::from(col1).exp(),
mass(&out)
);
if winner != 0 || col0 != col1 {
wrong.push(format!(
"{pooling:?}: winner {winner}, col0 {col0} vs col1 {col1} (p0 {}, p1 {})",
f64::from(col0).exp(),
f64::from(col1).exp()
));
}
}
assert!(
wrong.is_empty(),
"a 7.7e-3 mass gap, which is fp noise and not evidence, decided the clip:\n {}",
wrong.join("\n ")
);
assert!(
matches!(
aggregate_windows(ScorePooling::MeanLogProbability, &windows),
Err(Error::ZeroMassAggregate(ScorePooling::MeanLogProbability))
),
"the logarithmic pool over disjoint supports is refused, not repaired"
);
let voted = aggregate_windows(ScorePooling::Vote, &windows).expect("aggregate");
assert_eq!(voted.as_slice()[0], 0.5f32.ln());
assert_eq!(voted.as_slice()[1], 0.5f32.ln());
assert_eq!(argmax(voted.as_slice()), 0);
}
#[test]
fn the_fold_is_invariant_to_a_rows_own_scale() {
let baseline =
LogProbabilities::try_from_slice(&distribution(&[(94, 0.62), (3, 0.23), (61, 0.15)]))
.expect("row");
let other = LogProbabilities::try_from_slice(&distribution(&[(3, 0.44), (94, 0.31), (61, 0.25)]))
.expect("row");
let deficit = CPU_ONLY_ROW_MASS.ln() as f32;
let scaled_other = rescaled(&other, deficit);
assert!(deficit < 0.0, "the deficit must be a real shift");
assert_ne!(scaled_other.as_slice(), other.as_slice());
let plain = vec![
window(baseline.clone(), 0, WINDOW),
window(other, WINDOW, WINDOW),
];
let scaled = vec![
window(baseline, 0, WINDOW),
window(scaled_other, WINDOW, WINDOW),
];
let mut moved = Vec::new();
for pooling in all_poolings() {
let from_plain = aggregate_windows(pooling, &plain).expect("aggregate");
let from_scaled = aggregate_windows(pooling, &scaled).expect("aggregate");
let deviation = max_abs_difference(&from_plain, &from_scaled);
println!("scale-invariance {pooling:>18?} {deviation:.4e}");
if deviation >= SCALE_INVARIANCE_FLOOR {
moved.push(format!("{pooling:?}: {deviation:e}"));
}
}
assert!(
moved.is_empty(),
"rescaling one window's row by a constant moved the pooled row further than the \
f32 narrowing floor {SCALE_INVARIANCE_FLOOR:e}:\n {}",
moved.join("\n ")
);
}
const SCALE_INVARIANCE_FLOOR: f64 = 1e-6;
#[test]
fn a_lone_window_keeps_its_own_mass_deficit() {
let deficit = CPU_ONLY_ROW_MASS.ln() as f32;
let row = rescaled(
&LogProbabilities::try_from_slice(&distribution(&[(94, 0.62), (3, 0.23), (61, 0.15)]))
.expect("row"),
deficit,
);
let remaining = mass(&row);
assert!(
(remaining - CPU_ONLY_ROW_MASS).abs() < 1e-4,
"the fixture must not already be a distribution, or this proves nothing: mass {remaining}"
);
for pooling in all_poolings() {
let out = aggregate_windows(pooling, &[window(row.clone(), 0, WINDOW)]).expect("aggregate");
assert_eq!(out.as_slice(), row.as_slice(), "{pooling:?}");
}
}
fn ramp(top: f32) -> LogProbabilities {
let values: Vec<f32> = (0..NUM_LANGUAGES).map(|i| top - i as f32).collect();
LogProbabilities::try_from_slice(&values).expect("a descending non-positive ramp")
}
fn fold_pair(pooling: ScorePooling, row: &LogProbabilities) -> Result<LogProbabilities> {
aggregate_windows(
pooling,
&[
window(row.clone(), 0, WINDOW),
window(row.clone(), WINDOW, WINDOW),
],
)
}
#[test]
fn a_rows_own_scale_does_not_decide_whether_the_door_accepts_it() {
let low = ramp(-800.0);
let high = ramp(0.0);
assert_ne!(low.as_slice(), high.as_slice());
for (&a, &b) in low.as_slice().iter().zip(high.as_slice()) {
assert_eq!(a - low.as_slice()[0], b - high.as_slice()[0]);
}
assert_eq!(
as_distribution(low.as_slice()),
as_distribution(high.as_slice())
);
assert_eq!(f64::from(low.as_slice()[0]).exp(), 0.0);
assert!(f64::from(high.as_slice()[0]).exp() > 0.0);
for pooling in all_poolings() {
let from_low = fold_pair(pooling, &low);
let from_high = fold_pair(pooling, &high);
match (&from_low, &from_high) {
(Ok(a), Ok(b)) => assert_eq!(
a.as_slice(),
b.as_slice(),
"{pooling:?}: two rows carrying identical evidence must fold alike"
),
_ => panic!(
"{pooling:?}: the SAME evidence at two scales got two verdicts — \
low {} / high {}",
describe(&from_low),
describe(&from_high)
),
}
}
}
#[test]
fn the_door_is_invariant_to_a_rows_own_scale() {
let reference: Vec<Vec<f32>> = all_poolings()
.into_iter()
.map(|pooling| {
fold_pair(pooling, &ramp(0.0))
.expect("the unshifted ramp folds")
.as_slice()
.to_vec()
})
.collect();
for top in [
-1.0f32,
-100.0,
-744.0,
-745.0,
-746.0,
-800.0,
-10_000.0,
-16_000_000.0,
] {
let row = ramp(top);
for (pooling, want) in all_poolings().into_iter().zip(&reference) {
let got = fold_pair(pooling, &row);
match &got {
Ok(folded) => assert_eq!(folded.as_slice(), want.as_slice(), "top {top}, {pooling:?}"),
Err(_) => panic!("top {top}, {pooling:?}: {}", describe(&got)),
}
}
}
let flattened = ramp(-3.0e38);
assert!(flattened.as_slice().iter().all(|v| *v == -3.0e38));
for pooling in all_poolings() {
let got = fold_pair(pooling, &flattened);
let folded = match &got {
Ok(folded) => folded,
Err(_) => panic!("{pooling:?}: {}", describe(&got)),
};
let total = mass(folded);
assert!((total - 1.0).abs() < 1e-6, "{pooling:?}: mass {total}");
}
}
#[test]
fn a_lone_low_scale_window_comes_back_verbatim_and_ranks_as_identify_would() {
let low = ramp(-800.0);
assert_eq!(
mass(&low),
0.0,
"the fixture must be in the underflow regime"
);
for pooling in all_poolings() {
let out = aggregate_windows(pooling, &[window(low.clone(), 0, WINDOW)]);
let row = match &out {
Ok(row) => row,
Err(_) => panic!("{pooling:?}: {}", describe(&out)),
};
assert_eq!(row.as_slice(), low.as_slice(), "{pooling:?}");
assert_eq!(
row.top_k(3).expect("rank"),
low.top_k(3).expect("rank"),
"{pooling:?}"
);
}
}
#[test]
fn a_window_whose_maximum_is_not_finite_upward_is_refused_at_the_door() {
let mut values = vec![-1.0f32; NUM_LANGUAGES];
values[7] = f32::INFINITY;
assert!(
matches!(
LogProbabilities::try_from_slice(&values),
Err(Error::InvalidLogProbability(detail)) if detail.index() == 7
),
"the public constructor must still be the first line of defence"
);
let row = LogProbabilities::new(values);
for pooling in all_poolings() {
let lone = aggregate_windows(pooling, &[window(row.clone(), 0, WINDOW)]);
assert!(
matches!(&lone, Err(Error::UnnormalizableWindow(0))),
"{pooling:?}, lone window: {}",
describe(&lone)
);
let pair = fold_pair(pooling, &row);
assert!(
matches!(&pair, Err(Error::UnnormalizableWindow(0))),
"{pooling:?}, two windows: {}",
describe(&pair)
);
}
}
#[test]
fn a_fold_that_leaves_a_nan_is_refused_rather_than_returned() {
let row = row_from_logits(&[(94, 6.0), (3, 1.5)]);
for pooling in all_poolings() {
let mut fold = Accumulator::new(pooling);
fold.push(&row, WINDOW).expect("push");
fold.push(&row, WINDOW).expect("push");
fold.acc = vec![f64::NAN; NUM_LANGUAGES];
let finished = fold.finish();
assert!(
matches!(&finished, Err(Error::NotADistribution(detail)) if detail.mass().is_nan()),
"{pooling:?}: {}",
describe(&finished)
);
let rendered = finished.expect_err("refused").to_string();
assert!(
rendered.contains("sum to NaN, not 1"),
"{pooling:?}: {rendered}"
);
}
}
#[test]
fn a_window_holding_a_nan_is_refused_at_the_door() {
let mut values = vec![-1.0f32; NUM_LANGUAGES];
values[7] = f32::NAN;
assert!(
matches!(
LogProbabilities::try_from_slice(&values),
Err(Error::InvalidLogProbability(detail)) if detail.index() == 7
),
"the public constructor must still be the first line of defence"
);
assert!(!has_a_finite_maximum(&values));
let row = LogProbabilities::new(values);
for pooling in all_poolings() {
let lone = aggregate_windows(pooling, &[window(row.clone(), 0, WINDOW)]);
assert!(
matches!(&lone, Err(Error::UnnormalizableWindow(0))),
"{pooling:?}, lone window: {}",
describe(&lone)
);
let pair = fold_pair(pooling, &row);
assert!(
matches!(&pair, Err(Error::UnnormalizableWindow(0))),
"{pooling:?}, two windows: {}",
describe(&pair)
);
let good = row_from_logits(&[(94, 6.0), (3, 1.5)]);
let mixed = aggregate_windows(
pooling,
&[window(row.clone(), 0, WINDOW), window(good, WINDOW, WINDOW)],
);
assert!(
matches!(&mixed, Err(Error::UnnormalizableWindow(0))),
"{pooling:?}, NaN beside a good window: {}",
describe(&mixed)
);
}
}
#[test]
fn a_fold_over_admissible_rows_cannot_produce_a_nan() {
let certain_of = |index: usize| {
let mut values = vec![f32::NEG_INFINITY; NUM_LANGUAGES];
values[index] = 0.0;
LogProbabilities::try_from_slice(&values).expect("one zero among -inf")
};
let mut mostly_ruled_out = distribution(&[(0, 0.6), (1, 0.4)]);
mostly_ruled_out[5] = f32::NEG_INFINITY;
let rows = [
vec![certain_of(0), certain_of(1)],
vec![
LogProbabilities::try_from_slice(&mostly_ruled_out).expect("row"),
certain_of(0),
],
vec![ramp(-800.0), ramp(0.0)],
vec![certain_of(3), certain_of(3)],
];
for pooling in all_poolings() {
for (case, pair) in rows.iter().enumerate() {
let windows = vec![
window(pair[0].clone(), 0, if case == 3 { 1 } else { WINDOW }),
window(pair[1].clone(), WINDOW, WINDOW),
];
let pooled = aggregate_windows(pooling, &windows);
match &pooled {
Ok(out) => assert!(
out.as_slice().iter().all(|v| !v.is_nan()),
"{pooling:?} case {case}: {}",
describe(&pooled)
),
Err(error) => assert!(
matches!(error, Error::ZeroMassAggregate(_)),
"{pooling:?} case {case}: {}",
describe(&pooled)
),
}
}
}
}
#[test]
fn score_pooling_display_pins_the_wire_word() {
for pooling in all_poolings() {
let expected = match pooling {
ScorePooling::MeanLogProbability => "mean_log_probability",
ScorePooling::MeanProbability => "mean_probability",
ScorePooling::Max => "max",
ScorePooling::Vote => "vote",
};
assert_eq!(pooling.to_string(), expected);
}
assert_eq!(ScorePooling::default().to_string(), "mean_log_probability");
}
#[cfg(feature = "serde")]
#[test]
fn score_pooling_serde_wire_values_are_snake_case() {
for pooling in all_poolings() {
let expected = match pooling {
ScorePooling::MeanLogProbability => "\"mean_log_probability\"",
ScorePooling::MeanProbability => "\"mean_probability\"",
ScorePooling::Max => "\"max\"",
ScorePooling::Vote => "\"vote\"",
};
assert_eq!(serde_json::to_string(&pooling).unwrap(), expected);
}
}
#[cfg(feature = "serde")]
#[test]
fn score_pooling_serde_round_trips() {
for pooling in all_poolings() {
let json = serde_json::to_string(&pooling).unwrap();
let back: ScorePooling = serde_json::from_str(&json).unwrap();
assert_eq!(back, pooling);
}
}