use super::*;
use crate::audio::whisper::options::DecodingOptions;
fn greedy(temperature: f32) -> GreedyTokenSampler {
GreedyTokenSampler::new(temperature, 3, &DecodingOptions::new()).with_seed(42)
}
#[test]
fn argmax_at_zero_temperature_with_exact_logprob() {
let logits = [1.0f32, 3.0, 2.0, 0.0];
let result = greedy(0.0).sample(&logits);
assert_eq!(result.token(), 1);
assert!(!result.completed());
let log_z = logits.iter().map(|v| v.exp()).sum::<f32>().ln();
assert!((result.logprob() - (3.0 - log_z)).abs() < 1e-5);
}
#[test]
fn eot_completes() {
let logits = [0.0f32, 0.0, 0.0, 5.0]; let result = greedy(0.0).sample(&logits);
assert_eq!(result.token(), 3);
assert!(result.completed());
}
#[test]
fn nonzero_temperature_is_seed_deterministic_and_top_k_bounded() {
let mut logits = vec![0.0f32; 16];
for (i, v) in [9.0, 8.0, 7.0, 6.0, 5.0].iter().enumerate() {
logits[i + 8] = *v; }
let mut a = greedy(0.7);
let mut b = greedy(0.7);
for _ in 0..20 {
let (ra, rb) = (a.sample(&logits), b.sample(&logits));
assert_eq!(ra.token(), rb.token(), "same seed, same draw");
assert!(
(8..13).contains(&(ra.token() as usize)),
"outside top-k drawn"
);
assert!(ra.logprob() <= 0.0);
}
}
#[test]
fn fully_masked_logits_degenerate_without_panic_or_nan() {
let masked = [f32::NEG_INFINITY; 8];
for temperature in [0.0, 0.7] {
let result = greedy(temperature).sample(&masked);
assert_eq!(result.token(), 0, "t={temperature}");
assert_eq!(result.logprob(), f32::NEG_INFINITY, "t={temperature}");
assert!(!result.completed(), "t={temperature}");
}
let eot_zero = GreedyTokenSampler::new(0.7, 0, &DecodingOptions::new())
.with_seed(42)
.sample(&masked);
assert!(eot_zero.completed());
}
#[test]
fn fully_masked_sample_does_not_consume_rng() {
let logits: Vec<f32> = (0..16).map(|i| i as f32 * 0.25).collect();
let mut interrupted = greedy(0.7);
let mut fresh = greedy(0.7);
interrupted.sample(&[f32::NEG_INFINITY; 16]);
for _ in 0..10 {
assert_eq!(
interrupted.sample(&logits).token(),
fresh.sample(&logits).token()
);
}
}
#[test]
fn drew_from_rng_tracks_real_rng_draws() {
let logits = [1.0f32, 3.0, 2.0, 0.0];
let mut argmax_sampler = greedy(0.0);
assert!(
!argmax_sampler.drew_from_rng(),
"a fresh sampler has not drawn"
);
argmax_sampler.sample(&logits);
assert!(
!argmax_sampler.drew_from_rng(),
"argmax decoding does not consult the RNG"
);
let mut sampling = greedy(0.7);
sampling.sample(&logits);
assert!(
sampling.drew_from_rng(),
"a non-zero-temperature sample draws from the RNG"
);
let mut masked = greedy(0.7);
masked.sample(&[f32::NEG_INFINITY; 4]);
assert!(
!masked.drew_from_rng(),
"the all-masked degenerate path must not consult the RNG"
);
}
#[test]
fn negative_temperature_wide_logits_sample_without_panic() {
let result = greedy(-0.2).sample(&[-10.0f32, 10.0]);
assert!(result.logprob().is_finite(), "logprob must be finite");
assert!((result.token() as usize) < 2, "token must index the logits");
let wide = [-10.0f32, 10.0, -8.0, 5.0, -3.0, 2.0, 9.0, -1.0];
let inv_t = 1.0f32 / -0.2;
let scaled: Vec<f32> = wide.iter().map(|&v| v * inv_t).collect();
let scaled_max = scaled.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let reference: Vec<f32> = scaled.iter().map(|&s| (s - scaled_max).exp()).collect();
assert!(
reference.iter().all(|p| p.is_finite()),
"the reference stable softmax over scaled logits must stay finite"
);
let mut order: Vec<usize> = (0..wide.len()).collect();
order.sort_by(|&a, &b| reference[b].total_cmp(&reference[a]));
let top_k: std::collections::HashSet<usize> = order.into_iter().take(5).collect();
let mut sampler = greedy(-0.2);
for _ in 0..50 {
let r = sampler.sample(&wide);
assert!(
r.logprob().is_finite(),
"finite logprob under negative temperature"
);
assert!(
top_k.contains(&(r.token() as usize)),
"drawn token {} fell outside the stable-softmax top-k {top_k:?}",
r.token()
);
}
}
#[test]
fn negative_temperature_masked_logits_sample_without_panic() {
let mut sampler = greedy(-0.2);
for _ in 0..200 {
let r = sampler.sample(&[0.0f32, f32::NEG_INFINITY]);
assert!(
r.logprob().is_finite(),
"masked negative-temperature draw must have a finite log-prob"
);
assert_eq!(
r.token(),
0,
"the masked index (1) must never be drawn -- a mask is a mask at any temperature sign"
);
}
assert!(
sampler.drew_from_rng(),
"a non-zero-temperature draw on a non-masked-max buffer consults the RNG"
);
let mixed = [
f32::NEG_INFINITY,
-4.0,
f32::NEG_INFINITY,
7.0,
2.0,
f32::NEG_INFINITY,
];
let masked_indices = [0usize, 2, 5];
let mut sampler = greedy(-0.35);
for _ in 0..200 {
let r = sampler.sample(&mixed);
assert!(
r.logprob().is_finite(),
"finite log-prob under the mask + negative temperature"
);
assert!(
!masked_indices.contains(&(r.token() as usize)),
"a masked (-inf) index {} was drawn under negative temperature",
r.token()
);
}
}
#[test]
fn tiny_temperature_overflow_sample_without_panic_or_nan() {
for &temperature in &[1e-40f32, -1e-40, f32::MIN_POSITIVE, 1e-30] {
let mut sampler = GreedyTokenSampler::new(temperature, 3, &DecodingOptions::new()).with_seed(1);
for _ in 0..50 {
let r = sampler.sample(&[0.0f32, 20.0, -20.0, 5.0]);
assert!(
r.logprob().is_finite() || r.logprob() == f32::NEG_INFINITY,
"t={temperature}: log-prob must be finite or the degenerate -inf, never NaN"
);
assert!(
(r.token() as usize) < 4,
"t={temperature}: token must index the logits"
);
}
}
}
#[test]
fn tiny_temperature_preserves_rank_and_logprob() {
let logits = [5.0f32, 20.0, -3.0, 12.0]; for &(temperature, winner) in &[(1e-40f32, 1u32), (-1e-40f32, 2)] {
for seed in 0..8u64 {
let mut sampler =
GreedyTokenSampler::new(temperature, 3, &DecodingOptions::new()).with_seed(seed);
let result = sampler.sample(&logits);
assert_eq!(
result.token(),
winner,
"t={temperature}, seed={seed}: the order-preserving extreme must win, not a uniform draw"
);
assert!(
result.logprob().abs() < 1e-6,
"t={temperature}, seed={seed}: the point-mass winner's logprob → 0, got {}",
result.logprob()
);
assert!(!result.completed(), "t={temperature}, seed={seed}");
}
}
}
#[test]
#[should_panic(expected = "non-empty logits")]
fn empty_logits_panic() {
greedy(0.0).sample(&[]);
}
#[test]
fn finalize_appends_eot_once() {
let sampler = greedy(0.0);
let (mut tokens, mut logprobs) = (vec![1u32, 2], vec![-0.5f32, -0.25]);
sampler.finalize(&mut tokens, &mut logprobs);
assert_eq!(tokens, vec![1, 2, 3]);
assert_eq!(logprobs, vec![-0.5, -0.25, 0.0]);
sampler.finalize(&mut tokens, &mut logprobs); assert_eq!(tokens.len(), 3);
}
#[test]
fn derive_attempt_seed_is_pure_and_deterministic() {
assert_eq!(
derive_attempt_seed(1, 2, 3, 4),
derive_attempt_seed(1, 2, 3, 4),
"pure function: identical inputs must reproduce identical output"
);
assert_ne!(
derive_attempt_seed(1, 2, 3, 4),
derive_attempt_seed(2, 2, 3, 4),
"the base seed must change the derived seed"
);
assert_ne!(
derive_attempt_seed(1, 2, 3, 4),
derive_attempt_seed(1, 9, 3, 4),
"worker_index must change the derived seed"
);
assert_ne!(
derive_attempt_seed(1, 2, 3, 4),
derive_attempt_seed(1, 2, 9, 4),
"window_index must change the derived seed"
);
assert_ne!(
derive_attempt_seed(1, 2, 3, 4),
derive_attempt_seed(1, 2, 3, 9),
"attempt_index must change the derived seed"
);
}
#[test]
fn derive_attempt_seed_domain_separates_worker_and_window() {
for seed in [0u64, 1, 7, u64::MAX] {
for attempt in 0..=5u64 {
assert_ne!(
derive_attempt_seed(seed, 0, 1, attempt),
derive_attempt_seed(seed, 1, 0, attempt),
"(worker=0, window=1) must not alias (worker=1, window=0) \
[seed={seed} attempt={attempt}]"
);
}
}
}
#[test]
fn derive_attempt_seed_has_no_zero_collapse_across_base_seeds() {
assert_ne!(
derive_attempt_seed(0, 0, 0, 0),
0,
"the all-zero tuple must not collapse to 0"
);
assert_ne!(
derive_attempt_seed(0, 0, 0, 0),
derive_attempt_seed(1, 0, 1, 0),
"(seed=0, window=0) must not alias (seed=1, window=1)"
);
assert_ne!(
derive_attempt_seed(0, 0, 1, 0),
derive_attempt_seed(1, 0, 0, 0),
"(seed=0, window=1) must not alias (seed=1, window=0)"
);
}
#[test]
fn derive_attempt_seed_has_no_collisions_over_realistic_ranges() {
let mut seen = std::collections::HashSet::new();
for worker in 0..64u64 {
for window in 0..64u64 {
for attempt in 0..=8u64 {
let derived = derive_attempt_seed(0xABCD_1234_5678_9ABC, worker, window, attempt);
assert!(
seen.insert(derived),
"collision at worker={worker} window={window} attempt={attempt}"
);
}
}
}
}
fn seeded_sampler(seed: u64) -> GreedyTokenSampler {
GreedyTokenSampler::new(0.7, 999, &DecodingOptions::new()).with_seed(seed)
}
fn draw_sequence(sampler: &mut GreedyTokenSampler, logits: &[f32], n: usize) -> Vec<u32> {
(0..n).map(|_| sampler.sample(logits).token()).collect()
}
#[test]
fn attempt_seed_derivation_changes_sampled_draws_across_attempts() {
let seed = 0xC0FFEE_u64;
let worker = 2u64;
let window = 3u64;
let logits: Vec<f32> = (0..32).map(|i| i as f32 * 0.1 - 1.6).collect();
let mut attempt0 = seeded_sampler(derive_attempt_seed(seed, worker, window, 0));
let mut attempt1 = seeded_sampler(derive_attempt_seed(seed, worker, window, 1));
let draws0 = draw_sequence(&mut attempt0, &logits, 20);
let draws1 = draw_sequence(&mut attempt1, &logits, 20);
assert_ne!(
draws0, draws1,
"different attempt_index must decorrelate the sampled stream"
);
let mut replay0 = seeded_sampler(derive_attempt_seed(seed, worker, window, 0));
let replay_draws0 = draw_sequence(&mut replay0, &logits, 20);
assert_eq!(draws0, replay_draws0);
}
#[test]
fn attempt_seed_derivation_changes_sampled_draws_across_windows() {
let seed = 0xC0FFEE_u64;
let worker = 2u64;
let attempt = 1u64;
let logits: Vec<f32> = (0..32).map(|i| i as f32 * 0.1 - 1.6).collect();
let mut window0 = seeded_sampler(derive_attempt_seed(seed, worker, 0, attempt));
let mut window1 = seeded_sampler(derive_attempt_seed(seed, worker, 1, attempt));
let draws0 = draw_sequence(&mut window0, &logits, 20);
let draws1 = draw_sequence(&mut window1, &logits, 20);
assert_ne!(
draws0, draws1,
"different window_index must decorrelate the sampled stream"
);
}
#[test]
fn attempt_seed_derivation_changes_sampled_draws_across_workers() {
let seed = 0xC0FFEE_u64;
let attempt = 0u64;
let logits: Vec<f32> = (0..32).map(|i| i as f32 * 0.1 - 1.6).collect();
let mut worker0_window1 = seeded_sampler(derive_attempt_seed(seed, 0, 1, attempt));
let mut worker1_window0 = seeded_sampler(derive_attempt_seed(seed, 1, 0, attempt));
let draws_a = draw_sequence(&mut worker0_window1, &logits, 20);
let draws_b = draw_sequence(&mut worker1_window0, &logits, 20);
assert_ne!(
draws_a, draws_b,
"(worker=0, window=1) and (worker=1, window=0) must not share a stream"
);
}
#[test]
fn argmax_tie_keeps_first_index() {
let mut distant = [0.0f32; 16];
distant[2] = 5.0;
distant[5] = 5.0;
assert_eq!(
argmax(&distant),
2,
"distant tie -> first (tie_2_and_5_small)"
);
let mut adjacent = [0.0f32; 1024];
adjacent[100] = 5.0;
adjacent[101] = 5.0;
assert_eq!(
argmax(&adjacent),
100,
"adjacent tie -> first (adjacent_tie_100_101)"
);
let mut zero_and_last = [0.0f32; 16];
zero_and_last[0] = 5.0;
zero_and_last[15] = 5.0;
assert_eq!(
argmax(&zero_and_last),
0,
"tie at 0 & last -> 0 (tie_0_and_last_small)"
);
let mut manyway = [0.0f32; 16];
for v in manyway.iter_mut().step_by(5) {
*v = 7.0;
}
assert_eq!(
argmax(&manyway),
0,
"many-way tie -> first (vocab_manyway_tie_every_5000)"
);
}
#[test]
fn argmax_all_equal_returns_zero() {
assert_eq!(argmax(&[0.0f32; 16]), 0);
assert_eq!(argmax(&[1.0f32; 16]), 0);
}
#[test]
fn argmax_neg_infinity_floor_tie() {
let mut v = [f32::NEG_INFINITY; 16];
v[2] = 1.5;
v[5] = 1.5;
assert_eq!(
argmax(&v),
2,
"first finite of the tied pair over an -inf floor"
);
}
#[test]
fn argmax_signed_zero_ties_keep_first() {
let mut neg_then_pos = [-1.0f32; 16];
neg_then_pos[2] = -0.0;
neg_then_pos[5] = 0.0;
assert_eq!(argmax(&neg_then_pos), 2, "-0.0@2, +0.0@5 -> 2");
let mut pos_then_neg = [-1.0f32; 16];
pos_then_neg[2] = 0.0;
pos_then_neg[5] = -0.0;
assert_eq!(argmax(&pos_then_neg), 2, "+0.0@2, -0.0@5 -> 2");
}
#[test]
fn argmax_skips_nan_anywhere() {
let mut nan_mid = [0.0f32; 16];
nan_mid[4] = f32::NAN;
nan_mid[7] = 5.0;
assert_eq!(argmax(&nan_mid), 7, "NaN@4 skipped, max@7 wins");
let mut nan_first = [0.0f32; 16];
nan_first[0] = f32::NAN;
nan_first[7] = 5.0;
assert_eq!(
argmax(&nan_first),
7,
"NaN@0 must not seed best, max@7 wins"
);
}
#[test]
fn argmax_all_nan_pins_zero() {
assert_eq!(argmax(&[f32::NAN; 16]), 0);
}