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 cfg.temperature > 0.0 && cfg.temperature != 1.0 {
let inv_temp = 1.0 / cfg.temperature;
for v in &mut adjusted {
*v *= inv_temp;
}
}
if cfg.temperature <= 0.0 {
return greedy_token(&adjusted);
}
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)
});
indices.truncate(k);
}
let mut probs = build_softmax_probs(&adjusted, &indices);
probs.sort_unstable_by(|(_, a), (_, b)| b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal));
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();
for &id in previous_ids {
let idx = id as usize;
if idx < vocab_size {
if adjusted[idx] > 0.0 {
adjusted[idx] /= penalty;
} else {
adjusted[idx] *= penalty;
}
}
}
}
fn greedy_token(adjusted: &[f32]) -> u32 {
adjusted
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map(|(i, _)| i as u32)
.unwrap_or(0)
}
fn build_softmax_probs(adjusted: &[f32], indices: &[usize]) -> Vec<(usize, f32)> {
let max_logit = indices
.iter()
.map(|&i| adjusted[i])
.fold(f32::NEG_INFINITY, f32::max);
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();
for (_, p) in &mut probs {
*p /= sum;
}
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 {
let r = xorshift64(rng_state);
let mut cumsum = 0.0f32;
for &(idx, p) in probs {
cumsum += p;
if r <= cumsum {
return idx as u32;
}
}
probs[0].0 as u32
}
pub(crate) fn xorshift64(state: &mut u64) -> f32 {
let mut s = *state;
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
*state = s;
(s >> 11) as f32 / (1u64 << 53) as f32
}