use crate::model::qwen35_config::GenerateConfig;
pub(crate) fn sample_token(
logits: &[f32],
cfg: &GenerateConfig,
previous_ids: &[u32],
rng_state: &mut u64,
) -> u32 {
let vocab_size = logits.len();
let mut adjusted = logits.to_vec();
if cfg.repetition_penalty != 1.0 {
apply_repetition_penalty(&mut adjusted, previous_ids, cfg.repetition_penalty);
}
if crate::sampling::temperature_degenerate(cfg.temperature) {
return greedy_token(&adjusted);
}
if cfg.temperature != 1.0 {
let inv_temp = 1.0 / cfg.temperature;
for v in &mut adjusted {
*v *= inv_temp;
}
}
let mut indices: Vec<usize> = (0..vocab_size).collect();
if cfg.top_k > 0 && cfg.top_k < vocab_size {
let k = cfg.top_k;
indices.select_nth_unstable_by(k - 1, |&a, &b| {
adjusted[b]
.partial_cmp(&adjusted[a])
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.cmp(&b))
});
indices.truncate(k);
}
indices.sort_unstable_by(|&a, &b| {
adjusted[b]
.partial_cmp(&adjusted[a])
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.cmp(&b))
});
let Some(mut probs) = build_softmax_probs(&adjusted, &indices) else {
return greedy_token(&adjusted);
};
if cfg.top_p < 1.0 {
apply_top_p(&mut probs, cfg.top_p);
}
draw_from_distribution(&probs, rng_state)
}
fn apply_repetition_penalty(adjusted: &mut [f32], previous_ids: &[u32], penalty: f32) {
let vocab_size = adjusted.len();
let mut seen = std::collections::HashSet::with_capacity(previous_ids.len());
for &id in previous_ids {
let idx = id as usize;
if idx < vocab_size && seen.insert(id) {
adjusted[idx] = crate::sampling::penalized_logit(adjusted[idx], penalty);
}
}
}
fn greedy_token(adjusted: &[f32]) -> u32 {
let mut best_idx = 0u32;
let mut best_val = f32::NEG_INFINITY;
for (i, &v) in adjusted.iter().enumerate() {
if v > best_val {
best_val = v;
best_idx = i as u32;
}
}
best_idx
}
fn build_softmax_probs(adjusted: &[f32], indices: &[usize]) -> Option<Vec<(usize, f32)>> {
let max_logit = indices
.iter()
.map(|&i| adjusted[i])
.fold(f32::NEG_INFINITY, f32::max);
if !max_logit.is_finite() {
return None;
}
let mut probs: Vec<(usize, f32)> = indices
.iter()
.map(|&i| (i, (adjusted[i] - max_logit).exp()))
.collect();
let sum: f32 = probs.iter().map(|(_, p)| p).sum();
if !sum.is_finite() || sum <= 0.0 {
return None;
}
for (_, p) in &mut probs {
*p /= sum;
}
Some(probs)
}
fn apply_top_p(probs: &mut Vec<(usize, f32)>, top_p: f32) {
let mut cumsum = 0.0f32;
let mut cutoff = probs.len();
for (i, (_, p)) in probs.iter().enumerate() {
cumsum += p;
if cumsum >= top_p {
cutoff = i + 1;
break;
}
}
probs.truncate(cutoff);
let new_sum: f32 = probs.iter().map(|(_, p)| p).sum();
for (_, p) in probs.iter_mut() {
*p /= new_sum;
}
}
fn draw_from_distribution(probs: &[(usize, f32)], rng_state: &mut u64) -> u32 {
draw_index(probs, xorshift64(rng_state))
}
fn draw_index(probs: &[(usize, f32)], r: f32) -> u32 {
let mut cumsum = 0.0f32;
for &(idx, p) in probs {
cumsum += p;
if r < cumsum {
return idx as u32;
}
}
probs[probs.len() - 1].0 as u32
}
pub(crate) fn xorshift64(state: &mut u64) -> f32 {
let x = crate::sampling::xorshift64_next(state);
crate::sampling::uniform_f32_from_u64(x)
}
#[cfg(test)]
mod tests {
use super::*;
fn cfg(temperature: f32, top_k: usize) -> GenerateConfig {
GenerateConfig {
temperature,
top_k,
top_p: 1.0,
repetition_penalty: 1.0,
..Default::default()
}
}
#[test]
fn test_sample_token_degenerate_temperature_falls_back_to_argmax() {
let logits = [10.0_f32, 11.0, 9.0];
for bad in [
f32::NAN,
f32::INFINITY,
-1.0,
0.0,
1e-45,
1e-39,
f32::MIN_POSITIVE,
1e-37,
] {
let mut rng = 7u64;
let token = sample_token(&logits, &cfg(bad, 0), &[], &mut rng);
assert_eq!(
token, 1,
"degenerate temperature {bad} must fall back to argmax (token 1)"
);
}
}
#[test]
fn test_greedy_tie_break_is_first_wins() {
let logits = [0.0_f32, 1.0, 1.0];
assert_eq!(greedy_token(&logits), 1, "first-wins on tied max");
let mut rng = 3u64;
let token = sample_token(&logits, &cfg(0.0, 0), &[], &mut rng);
assert_eq!(token, 1, "temperature=0 greedy must be first-wins on ties");
let with_nan = [f32::NAN, 2.0_f32, 2.0];
assert_eq!(
greedy_token(&with_nan),
1,
"NaN skipped, first finite max wins"
);
}
#[test]
fn test_sample_token_valid_temperature_unchanged() {
let logits = [10.0_f32, 11.0, 9.0];
let mut rng = 7u64;
let token = sample_token(&logits, &cfg(0.7, 0), &[], &mut rng);
assert!(token < 3, "valid temperature must return an in-range token");
}
#[test]
fn test_inf_logit_routes_to_argmax_not_nan_draw() {
let logits = [0.0_f32, 1.0, f32::INFINITY];
let mut rng = 42u64;
let token = sample_token(&logits, &cfg(1.0, 0), &[], &mut rng);
assert_eq!(
token, 2,
"infinite-logit token must win via argmax fallback"
);
}
#[test]
fn test_nan_in_nonmax_position_routes_to_argmax() {
let logits = [5.0_f32, f32::NAN, 1.0];
let mut rng = 7u64;
let token = sample_token(&logits, &cfg(1.0, 0), &[], &mut rng);
assert_eq!(
token, 0,
"finite argmax must win when a non-max logit is NaN"
);
}
#[test]
fn test_empty_logits_returns_zero_without_panic() {
let logits: [f32; 0] = [];
let mut rng = 1u64;
let token = sample_token(&logits, &cfg(1.0, 0), &[], &mut rng);
assert_eq!(token, 0, "empty logits must return 0, never panic");
}
#[test]
fn test_invalid_repetition_penalty_is_noop_not_signflip() {
let logits = [5.0_f32, 1.0, 0.5];
let mut cfg_rp = cfg(0.0, 0); cfg_rp.repetition_penalty = -1.0;
let mut rng = 7u64;
let token = sample_token(&logits, &cfg_rp, &[0], &mut rng);
assert_eq!(
token, 0,
"invalid penalty must be a no-op, argmax stays index 0"
);
}
#[test]
fn test_nan_repetition_penalty_is_noop() {
let logits = [5.0_f32, 1.0, 0.5];
let mut cfg_rp = cfg(0.0, 0);
cfg_rp.repetition_penalty = f32::NAN;
let mut rng = 7u64;
let token = sample_token(&logits, &cfg_rp, &[0], &mut rng);
assert_eq!(token, 0, "NaN penalty must be a no-op");
}
#[test]
fn cross_path_parity_sampler_vs_sample_token() {
use crate::model::qwen35_config::GenerateConfig;
use crate::sampling::{Sampler, SamplingConfig};
let logits: Vec<f32> = (0..64u64)
.map(|i| {
let h = i
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(h as f32 / u64::MAX as f32) * 20.0 - 10.0
})
.collect();
let temperature = 0.8f32;
let top_k = 40usize;
let top_p = 0.95f32;
let seed = 0xdead_beef_cafe_babe_u64;
let n = 200usize;
let config_a = SamplingConfig {
temperature,
top_k,
top_p,
repetition_penalty: 1.0,
};
let mut sampler = Sampler::new(config_a).with_seed(seed);
let tokens_a: Vec<u32> = (0..n).map(|_| sampler.sample(&logits)).collect();
let config_b = GenerateConfig {
temperature,
top_k,
top_p,
repetition_penalty: 1.0,
..Default::default()
};
let mut rng_state = seed;
let tokens_b: Vec<u32> = (0..n)
.map(|_| sample_token(&logits, &config_b, &[], &mut rng_state))
.collect();
assert_eq!(
tokens_a, tokens_b,
"Sampler::sample and sample_token must produce identical token streams \
for the same logits, config, and seed (consolidation parity proof)"
);
}
#[test]
fn canonical_rng_streams_match_across_sampling_paths() {
let seed = 0x1234_5678_9abc_def0u64;
let n = 32usize;
let mut state_q = seed;
let q35: Vec<f32> = (0..n).map(|_| xorshift64(&mut state_q)).collect();
let mut state_c = seed;
let canonical: Vec<f32> = (0..n)
.map(|_| {
let x = crate::sampling::xorshift64_next(&mut state_c);
crate::sampling::uniform_f32_from_u64(x)
})
.collect();
assert_eq!(
q35, canonical,
"xorshift64 and canonical primitives must produce identical streams"
);
}
#[test]
fn draw_index_uses_strict_less_than_at_exact_boundary() {
let probs: Vec<(usize, f32)> = vec![(0, 0.5), (1, 0.5)];
assert_eq!(
draw_index(&probs, 0.5),
1,
"at the exact boundary (cumsum == r), strict `r < cumsum` must skip \
token 0 and select token 1 (a `r <= cumsum` regression selects token 0)"
);
assert_eq!(
draw_index(&probs, 0.25),
0,
"r below first cumsum picks token 0"
);
assert_eq!(
draw_index(&probs, 0.75),
1,
"r between the two bucket boundaries picks token 1"
);
}
#[test]
fn draw_index_fallback_returns_last_bucket_not_first() {
let probs: Vec<(usize, f32)> = vec![(7, 0.3), (8, 0.3), (9, 0.3)]; assert_eq!(
draw_index(&probs, 0.95),
9,
"when cumsum never reaches r (float error), the draw belongs to the \
LAST bucket (idx 9); the old `probs[0]` behaviour returns idx 7"
);
}
#[test]
fn cross_path_parity_with_logit_ties() {
use crate::model::qwen35_config::GenerateConfig;
use crate::sampling::{Sampler, SamplingConfig};
let logits: Vec<f32> = vec![10.0, 10.0, 9.5, 9.5, 9.0, 8.0, 7.0, 6.0];
let temperature = 1.0f32;
let top_k = 3usize; let top_p = 1.0f32; let seed = 0xdead_beef_cafe_babe_u64;
let n = 200usize;
let config_a = SamplingConfig {
temperature,
top_k,
top_p,
repetition_penalty: 1.0,
};
let mut sampler = Sampler::new(config_a).with_seed(seed);
let tokens_a: Vec<u32> = (0..n).map(|_| sampler.sample(&logits)).collect();
let config_b = GenerateConfig {
temperature,
top_k,
top_p,
repetition_penalty: 1.0,
..Default::default()
};
let mut rng_state = seed;
let tokens_b: Vec<u32> = (0..n)
.map(|_| sample_token(&logits, &config_b, &[], &mut rng_state))
.collect();
assert_eq!(
tokens_a, tokens_b,
"Sampler::sample and sample_token must produce identical token streams \
with high-probability ties at the max (idx 0 & 1) and the top_k=3 \
boundary (idx 2 & 3). A backwards tie-break in either the selection or \
the pre-softmax sort changes which token wins the contested rank or the \
order of the two max-probability buckets, flipping dozens of the 200 \
draws and failing this assertion."
);
}
}