use super::*;
use crate::audio::lid::NUM_LANGUAGES;
fn row(base: f32, overrides: &[(usize, f32)]) -> Vec<f32> {
let mut values = vec![base; NUM_LANGUAGES];
for &(index, value) in overrides {
values[index] = value;
}
values
}
fn top_k(values: &[f32], k: usize) -> Vec<LanguageScore> {
top_k_from_scores(values.iter().copied().enumerate(), k).expect("ranking")
}
#[test]
fn ranking_is_descending_and_carries_the_roster_row() {
let values = row(-25.0, &[(94, -0.0101), (55, -4.6038), (52, -18.7132)]);
let ranked = top_k(&values, 3);
assert_eq!(ranked.len(), 3);
assert_eq!(ranked[0].index(), 94);
assert_eq!(ranked[0].code(), "th");
assert_eq!(ranked[0].name(), "Thai");
assert_eq!(ranked[1].code(), "lo");
assert_eq!(ranked[2].code(), "la");
for pair in ranked.windows(2) {
assert!(
pair[0].log_probability() >= pair[1].log_probability(),
"output must be descending"
);
}
assert_eq!(
ranked[0].language(),
crate::audio::lid::Language::from_index(94).expect("Thai")
);
}
#[test]
fn probability_is_exp_of_the_log_score() {
let values = row(-30.0, &[(94, -0.0101), (55, -4.6038), (52, -18.7132)]);
let ranked = top_k(&values, 3);
assert!((ranked[0].probability() - 0.98995).abs() < 1e-4);
assert!((ranked[1].probability() - 0.01001).abs() < 1e-4);
assert_eq!(ranked[0].probability(), ranked[0].log_probability().exp());
assert!(ranked[2].probability() < 1e-6);
assert!(ranked[2].log_probability() - (-30.0) > 10.0);
}
#[test]
fn log_space_ranking_matches_probability_space_ranking() {
let mut values = row(-9.0, &[]);
for (i, slot) in values.iter_mut().enumerate() {
*slot = -((i % 17) as f32) - 0.25 * (i as f32 % 5.0);
}
let by_log: Vec<usize> = top_k(&values, NUM_LANGUAGES)
.iter()
.map(LanguageScore::index)
.collect();
let mut by_probability: Vec<(usize, f32)> = values.iter().map(|v| v.exp()).enumerate().collect();
by_probability.sort_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
assert_eq!(
by_log,
by_probability.iter().map(|(i, _)| *i).collect::<Vec<_>>()
);
}
#[test]
fn ties_break_by_ascending_language_index() {
let flat = row(-4.5, &[]);
let ranked = top_k(&flat, 5);
assert_eq!(
ranked.iter().map(LanguageScore::index).collect::<Vec<_>>(),
vec![0, 1, 2, 3, 4]
);
let values = row(-9.0, &[(80, -1.0), (12, -1.0), (44, -2.0)]);
let ranked = top_k(&values, 3);
assert_eq!(
ranked.iter().map(LanguageScore::index).collect::<Vec<_>>(),
vec![12, 80, 44]
);
}
#[test]
fn k_zero_and_oversized_k_are_safe() {
let values = row(-3.0, &[(7, -0.5)]);
assert!(top_k(&values, 0).is_empty());
assert_eq!(top_k(&values, 1).len(), 1);
assert_eq!(top_k(&values, NUM_LANGUAGES).len(), NUM_LANGUAGES);
assert_eq!(top_k(&values, NUM_LANGUAGES + 50).len(), NUM_LANGUAGES);
assert_eq!(top_k(&values, usize::MAX).len(), NUM_LANGUAGES);
assert_eq!(top_k(&values, usize::MAX)[0].index(), 7);
}
#[test]
fn ranking_everything_is_a_permutation_of_the_roster() {
let values: Vec<f32> = (0..NUM_LANGUAGES)
.map(|i| -((i as f32) * 0.37 % 11.0))
.collect();
let mut indices: Vec<usize> = top_k(&values, NUM_LANGUAGES)
.iter()
.map(LanguageScore::index)
.collect();
assert_eq!(indices.len(), NUM_LANGUAGES);
indices.sort_unstable();
assert_eq!(indices, (0..NUM_LANGUAGES).collect::<Vec<_>>());
}
#[test]
fn an_out_of_roster_index_is_a_typed_error() {
let error = top_k_from_scores([(NUM_LANGUAGES, -1.0f32)], 1).expect_err("must reject");
assert!(matches!(error, Error::UnknownLanguageIndex(i) if i == NUM_LANGUAGES));
}
#[test]
fn a_hand_built_row_is_copied_verbatim_and_ranks_like_the_model_path() {
let mut values = vec![-14.0f32; NUM_LANGUAGES];
values[94] = -0.01;
values[3] = -5.5;
let row = LogProbabilities::try_from_slice(&values).expect("valid row");
assert_eq!(row.as_slice(), values.as_slice());
assert_eq!(row.as_slice().len(), NUM_LANGUAGES);
let ranked = row.top_k(2).expect("rank");
assert_eq!(ranked[0].index(), 94);
assert_eq!(ranked[0].log_probability(), -0.01);
assert_eq!(ranked[1].index(), 3);
let direct = top_k_from_scores(values.iter().copied().enumerate(), 2).expect("rank");
assert_eq!(direct, ranked);
}
#[test]
fn the_row_invariant_admits_negative_infinity_and_nothing_positive() {
assert!(LogProbabilities::try_from_slice(&vec![f32::NEG_INFINITY; NUM_LANGUAGES]).is_ok());
assert!(LogProbabilities::try_from_slice(&vec![0.0f32; NUM_LANGUAGES]).is_ok());
for (index, bad) in [(0usize, 1e-7f32), (94, f32::NAN), (106, f32::INFINITY)] {
let mut values = vec![-1.0f32; NUM_LANGUAGES];
values[index] = bad;
let error = LogProbabilities::try_from_slice(&values).expect_err("must reject");
let Error::InvalidLogProbability(detail) = error else {
panic!("expected InvalidLogProbability for {bad} at {index}, got {error:?}");
};
assert_eq!(detail.index(), index);
}
}
#[test]
fn a_wrong_width_row_is_a_typed_refusal() {
for width in [0usize, NUM_LANGUAGES - 1, NUM_LANGUAGES + 1] {
assert!(matches!(
LogProbabilities::try_from_slice(&vec![-1.0f32; width]),
Err(Error::LanguageCountMismatch(got)) if got == width
));
}
}
#[test]
fn a_zero_probability_ranks_last_and_reads_as_zero() {
let mut values = vec![f32::NEG_INFINITY; NUM_LANGUAGES];
values[7] = 0.0;
let ranked = LogProbabilities::try_from_slice(&values)
.expect("valid row")
.top_k(3)
.expect("rank");
assert_eq!(ranked[0].index(), 7);
assert_eq!(ranked[0].probability(), 1.0);
assert_eq!(ranked[1].probability(), 0.0);
assert_eq!(ranked[2].probability(), 0.0);
assert!(ranked[1].index() < ranked[2].index());
}