use std::cell::RefCell;
use std::collections::HashSet;
thread_local! {
static SAMPLE_INDEXED_SCRATCH: RefCell<Vec<(usize, f32)>> =
RefCell::new(Vec::new());
}
#[derive(Debug, Clone)]
pub struct SamplingParams {
pub temperature: f64,
pub top_p: f64,
pub top_k: usize,
pub min_p: f64,
pub repetition_penalty: f64,
pub max_tokens: usize,
}
impl Default for SamplingParams {
fn default() -> Self {
Self {
temperature: 0.7,
top_p: 0.9,
top_k: 50,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 2048,
}
}
}
pub const SAMPLING_EPS: f64 = 1e-5;
pub fn sample_greedy(logits: &[f32]) -> u32 {
if logits.is_empty() {
return 0;
}
let mut best_idx: usize = 0;
let mut best_val: f32 = logits[0];
for (i, &v) in logits.iter().enumerate().skip(1) {
if v > best_val {
best_val = v;
best_idx = i;
}
}
best_idx as u32
}
pub fn sample_token(logits: &mut [f32], params: &SamplingParams, previous_tokens: &[u32]) -> u32 {
if params.repetition_penalty != 1.0 && !previous_tokens.is_empty() {
apply_repetition_penalty(logits, previous_tokens, params.repetition_penalty);
}
if params.temperature < SAMPLING_EPS {
return sample_greedy(logits);
}
let result = SAMPLE_INDEXED_SCRATCH.with(|cell| -> Option<u32> {
let mut indexed = cell.borrow_mut();
indexed.clear();
indexed.reserve(logits.len());
for (i, &l) in logits.iter().enumerate() {
indexed.push((i, l));
}
if indexed.is_empty() {
return Some(sample_greedy(logits));
}
sample_token_indexed(&mut indexed, params)
});
if let Some(out) = result {
return out;
}
sample_greedy(logits)
}
pub fn sample_token_with_logprob(
logits: &mut [f32],
params: &SamplingParams,
previous_tokens: &[u32],
) -> (u32, f32) {
let max_logit = logits
.iter()
.copied()
.fold(f32::NEG_INFINITY, |acc, v| if v > acc { v } else { acc });
if !max_logit.is_finite() {
let token = sample_greedy(logits);
return (token, f32::NEG_INFINITY);
}
let mut sum_exp = 0.0f32;
for &v in logits.iter() {
sum_exp += (v - max_logit).exp();
}
let log_z = max_logit + sum_exp.ln();
let raw_logprobs: Vec<f32> = logits.iter().map(|&v| v - log_z).collect();
let token = sample_token(logits, params, previous_tokens);
let logprob = raw_logprobs
.get(token as usize)
.copied()
.unwrap_or(f32::NEG_INFINITY);
(token, logprob)
}
pub fn sample_token_from_topk(
top_indices: &[u32],
top_values: &[f32],
params: &SamplingParams,
) -> u32 {
debug_assert_eq!(
top_indices.len(),
top_values.len(),
"sample_token_from_topk: top_indices.len()={} != top_values.len()={}",
top_indices.len(),
top_values.len(),
);
if top_indices.is_empty() {
return 0;
}
debug_assert!(
(params.repetition_penalty - 1.0).abs() < 1e-9,
"sample_token_from_topk: repetition_penalty={} != 1.0 — caller must \
route through the full-vocab sample_token path when rep_penalty is \
engaged",
params.repetition_penalty
);
if (params.repetition_penalty - 1.0).abs() >= 1e-9 {
let mut best_i = 0usize;
let mut best_v = top_values[0];
for (i, &v) in top_values.iter().enumerate().skip(1) {
if v > best_v {
best_v = v;
best_i = i;
}
}
return top_indices[best_i];
}
if params.temperature < SAMPLING_EPS {
let mut best_i = 0usize;
let mut best_v = top_values[0];
for (i, &v) in top_values.iter().enumerate().skip(1) {
if v > best_v {
best_v = v;
best_i = i;
}
}
return top_indices[best_i];
}
let mut indexed: Vec<(usize, f32)> = top_indices
.iter()
.zip(top_values.iter())
.map(|(&i, &l)| (i as usize, l))
.collect();
match sample_token_indexed(&mut indexed, params) {
Some(tok) => tok,
None => {
let mut best_i = 0usize;
let mut best_v = top_values[0];
for (i, &v) in top_values.iter().enumerate().skip(1) {
if v > best_v {
best_v = v;
best_i = i;
}
}
top_indices[best_i]
}
}
}
fn sample_token_indexed(indexed: &mut Vec<(usize, f32)>, params: &SamplingParams) -> Option<u32> {
let cmp_desc = |a: &(usize, f32), b: &(usize, f32)| {
b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal)
};
if params.top_k > 0 && params.top_k < indexed.len() {
let top_k = params.top_k;
indexed.select_nth_unstable_by(top_k - 1, cmp_desc);
indexed.truncate(top_k);
indexed.sort_by(cmp_desc);
} else {
indexed.sort_by(cmp_desc);
}
if params.top_p < 1.0 && indexed.len() > 1 {
let probs = softmax_logits_to_probs(&indexed);
let mut cumsum = 0.0f32;
let mut cutoff = indexed.len();
for (i, &p) in probs.iter().enumerate() {
cumsum += p;
if cumsum >= params.top_p as f32 {
cutoff = i + 1;
break;
}
}
indexed.truncate(cutoff);
}
if params.min_p > 0.0 && indexed.len() > 1 {
let max_logit = indexed[0].1;
let min_logit_threshold = max_logit + (params.min_p as f32).ln();
let mut cutoff = 1;
for (i, &(_, l)) in indexed.iter().enumerate().skip(1) {
if l >= min_logit_threshold {
cutoff = i + 1;
} else {
break;
}
}
indexed.truncate(cutoff);
}
let inv_temp = 1.0 / params.temperature as f32;
for (_, l) in indexed.iter_mut() {
*l *= inv_temp;
}
let probs = softmax_logits_to_probs(indexed);
let sum: f32 = probs.iter().sum();
if sum <= 0.0 || !sum.is_finite() {
return Some(indexed.first().map(|&(idx, _)| idx as u32).unwrap_or(0));
}
let mut rng_val = rand_f32();
for (i, &p) in probs.iter().enumerate() {
let normalized = p / sum;
if rng_val < normalized {
return Some(indexed[i].0 as u32);
}
rng_val -= normalized;
}
Some(indexed.last().map(|&(idx, _)| idx as u32).unwrap_or(0))
}
fn softmax_logits_to_probs(indexed: &[(usize, f32)]) -> Vec<f32> {
if indexed.is_empty() {
return Vec::new();
}
let max = indexed
.iter()
.map(|&(_, l)| l)
.fold(f32::NEG_INFINITY, f32::max);
let mut probs: Vec<f32> = indexed.iter().map(|&(_, l)| (l - max).exp()).collect();
let sum: f32 = probs.iter().sum();
if sum > 0.0 && sum.is_finite() {
let inv = 1.0 / sum;
for p in probs.iter_mut() {
*p *= inv;
}
}
probs
}
pub fn apply_repetition_penalty(logits: &mut [f32], previous_tokens: &[u32], penalty: f64) {
let vocab = logits.len();
let penalty_f = penalty as f32;
let inv_penalty = 1.0f32 / penalty_f;
let mut seen = HashSet::new();
for &id in previous_tokens {
let idx = id as usize;
if idx < vocab && seen.insert(id) {
if logits[idx] >= 0.0 {
logits[idx] *= inv_penalty;
} else {
logits[idx] *= penalty_f;
}
}
}
}
fn softmax_in_place(probs: &mut [f32]) {
let max = probs.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
for p in probs.iter_mut() {
*p = (*p - max).exp();
sum += *p;
}
if sum > 0.0 {
let inv_sum = 1.0 / sum;
for p in probs.iter_mut() {
*p *= inv_sum;
}
}
}
fn softmax_pairs_in_place(indexed: &mut [(usize, f32)]) {
let max = indexed
.iter()
.map(|&(_, v)| v)
.fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
for (_, v) in indexed.iter_mut() {
*v = (*v - max).exp();
sum += *v;
}
if sum > 0.0 {
let inv_sum = 1.0 / sum;
for (_, v) in indexed.iter_mut() {
*v *= inv_sum;
}
}
}
fn rand_f32() -> f32 {
use std::time::SystemTime;
thread_local! {
static STATE: std::cell::Cell<u64> = const { std::cell::Cell::new(0) };
static SEEDED: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
}
SEEDED.with(|seeded| {
if !seeded.get() {
let t = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap_or_default()
.as_nanos() as u64;
let tid = {
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
let mut h = DefaultHasher::new();
std::thread::current().id().hash(&mut h);
h.finish()
};
let mut seed = t ^ tid;
seed = seed.wrapping_add(0x9e3779b97f4a7c15);
seed = (seed ^ (seed >> 30)).wrapping_mul(0xbf58476d1ce4e5b9);
seed = (seed ^ (seed >> 27)).wrapping_mul(0x94d049bb133111eb);
seed ^= seed >> 31;
STATE.with(|s| s.set(if seed == 0 { 0x1234567890abcdef } else { seed }));
seeded.set(true);
}
});
STATE.with(|s| {
let mut x = s.get();
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
s.set(x);
let u = x.wrapping_mul(0x2545f4914f6cdd1d) >> 32;
u as f32 / (u32::MAX as f32 + 1.0)
})
}
#[cfg(test)]
mod tests {
use super::*;
fn raw_softmax(logits: &[f32]) -> Vec<f64> {
let max = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max) as f64;
let mut p: Vec<f64> = logits.iter().map(|l| ((*l as f64) - max).exp()).collect();
let s: f64 = p.iter().sum();
for v in p.iter_mut() {
*v /= s;
}
p
}
fn sample_frequencies(base_logits: &[f32], params: &SamplingParams, n: usize) -> Vec<f64> {
let mut counts = vec![0usize; base_logits.len()];
for _ in 0..n {
let mut buf = base_logits.to_vec();
let tok = sample_token(&mut buf, params, &[]);
counts[tok as usize] += 1;
}
counts.iter().map(|&c| c as f64 / n as f64).collect()
}
#[test]
fn sampler_all_pass_matches_raw_softmax() {
let logits = vec![2.0_f32, 1.0, 0.0, -1.0, -2.0];
let expected = raw_softmax(&logits);
let params = SamplingParams {
temperature: 1.0,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 1,
};
let freq = sample_frequencies(&logits, ¶ms, 20_000);
for i in 0..logits.len() {
let diff = (freq[i] - expected[i]).abs();
assert!(
diff < 0.02,
"all-pass sampler diverges from softmax at idx {i}: \
expected={:.4} got={:.4} diff={:.4}",
expected[i],
freq[i],
diff,
);
}
}
#[test]
fn sampler_temperature_concentrates_winner() {
let logits = vec![3.0_f32, 2.0, 1.0, 0.0];
let cold = SamplingParams {
temperature: 0.3,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 1,
};
let hot = SamplingParams {
temperature: 2.0,
..cold.clone()
};
let cold_freq = sample_frequencies(&logits, &cold, 10_000);
let hot_freq = sample_frequencies(&logits, &hot, 10_000);
assert!(
cold_freq[0] - hot_freq[0] > 0.20,
"temperature does not concentrate winner: cold={:.3} hot={:.3} \
diff={:.3}",
cold_freq[0],
hot_freq[0],
cold_freq[0] - hot_freq[0],
);
assert!(
cold_freq[0] > 0.80,
"cold winner-rate too low: {:.3} (expected > 0.80)",
cold_freq[0],
);
}
#[test]
fn sampler_min_p_filters_distant_tokens() {
let mut logits = vec![0.5_f32; 100];
logits[7] = 10.0; let params = SamplingParams {
temperature: 1.0,
top_p: 1.0,
top_k: 0,
min_p: 0.05,
repetition_penalty: 1.0,
max_tokens: 1,
};
let freq = sample_frequencies(&logits, ¶ms, 5_000);
assert!(
freq[7] > 0.99,
"min-p failed to isolate dominant token: freq[7]={:.4}",
freq[7],
);
}
#[test]
fn sampler_top_k_truncates_to_k_candidates() {
let logits = vec![1.0_f32, 2.0, 3.0, 0.0, -1.0];
let params = SamplingParams {
temperature: 1.0,
top_p: 1.0,
top_k: 2,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 1,
};
let freq = sample_frequencies(&logits, ¶ms, 5_000);
assert_eq!(freq[3], 0.0, "top-k=2 leaked to idx 3 (logit=0.0)");
assert_eq!(freq[4], 0.0, "top-k=2 leaked to idx 4 (logit=-1.0)");
assert_eq!(freq[0], 0.0, "top-k=2 leaked to idx 0 (logit=1.0)");
}
#[test]
fn sampler_default_chain_preserves_dominant_winner() {
let mut logits = vec![0.0_f32; 50];
logits[7] = 5.0; let params = SamplingParams {
temperature: 0.8,
top_p: 0.95,
top_k: 40,
min_p: 0.05,
repetition_penalty: 1.0,
max_tokens: 1,
};
let freq = sample_frequencies(&logits, ¶ms, 5_000);
assert!(
freq[7] > 0.70,
"regression: dominant-winner pick rate too low under default \
chain: freq[7]={:.4} (pre-fix bug returned ~0.025)",
freq[7],
);
}
#[test]
fn greedy_path_unchanged_at_temp_zero() {
let mut logits = vec![0.0_f32, 5.0, 3.0, 1.0];
let params = SamplingParams {
temperature: 0.0,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 1,
};
assert_eq!(sample_token(&mut logits, ¶ms, &[]), 1);
}
#[test]
fn topk_with_k1_returns_single_index_regardless_of_temperature() {
let top_indices = vec![42u32];
let top_values = vec![3.7_f32];
for &temp in &[0.0_f64, 0.5, 0.8, 1.5, 5.0] {
for &top_p in &[0.0_f64, 0.5, 0.95, 1.0] {
let params = SamplingParams {
temperature: temp,
top_p,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 1,
};
assert_eq!(
sample_token_from_topk(&top_indices, &top_values, ¶ms),
42,
"K=1 must always return the single index (temp={}, top_p={})",
temp,
top_p,
);
}
}
}
#[test]
fn topk_empty_returns_zero() {
let params = SamplingParams::default();
assert_eq!(sample_token_from_topk(&[], &[], ¶ms), 0);
}
#[test]
fn topk_temp_zero_returns_max_value_index() {
let top_indices = vec![100u32, 200, 300, 400];
let top_values = vec![1.0_f32, 5.0, 3.0, 2.0];
let params = SamplingParams {
temperature: 0.0,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 1,
};
assert_eq!(
sample_token_from_topk(&top_indices, &top_values, ¶ms),
200,
);
}
#[test]
fn topk_matches_full_v_path_within_sampling_noise() {
let v: usize = 1000;
let mut full_logits = vec![-10.0_f32; v];
let top_pos = [3usize, 17, 42, 99, 250, 500, 700, 999];
let top_vals = [5.0_f32, 4.5, 4.0, 3.5, 3.0, 2.5, 2.0, 1.5];
for (&pos, &val) in top_pos.iter().zip(top_vals.iter()) {
full_logits[pos] = val;
}
let params = SamplingParams {
temperature: 0.5,
top_p: 1.0,
top_k: 8,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 1,
};
let n = 4_000usize;
let mut full_counts = vec![0usize; v];
for _ in 0..n {
let mut buf = full_logits.clone();
let tok = sample_token(&mut buf, ¶ms, &[]);
full_counts[tok as usize] += 1;
}
let topk_indices: Vec<u32> = top_pos.iter().map(|&p| p as u32).collect();
let topk_values: Vec<f32> = top_vals.to_vec();
let mut topk_counts = vec![0usize; v];
for _ in 0..n {
let tok = sample_token_from_topk(&topk_indices, &topk_values, ¶ms);
topk_counts[tok as usize] += 1;
}
for &pos in &top_pos {
let f_full = full_counts[pos] as f64 / n as f64;
let f_topk = topk_counts[pos] as f64 / n as f64;
assert!(
(f_full - f_topk).abs() < 0.05,
"top-K path diverges from full-V path at idx {}: \
full={:.4} topk={:.4} diff={:.4}",
pos,
f_full,
f_topk,
(f_full - f_topk).abs(),
);
}
for i in 0..v {
if !top_pos.contains(&i) {
assert_eq!(
full_counts[i], 0,
"full-V path leaked to tail idx {} ({} hits)",
i, full_counts[i]
);
assert_eq!(
topk_counts[i], 0,
"top-K path leaked to tail idx {} ({} hits)",
i, topk_counts[i]
);
}
}
}
#[test]
fn topk_unsorted_input_preserves_distribution() {
let top_indices_a = vec![10u32, 20, 30, 40];
let top_values_a = vec![1.0_f32, 5.0, 3.0, 2.0];
let top_indices_b = vec![10u32, 30, 20, 40];
let top_values_b = vec![1.0_f32, 3.0, 5.0, 2.0];
let params = SamplingParams {
temperature: 0.0,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 1,
};
let a = sample_token_from_topk(&top_indices_a, &top_values_a, ¶ms);
let b = sample_token_from_topk(&top_indices_b, &top_values_b, ¶ms);
assert_eq!(a, 20, "expected max-value idx 20, got {}", a);
assert_eq!(b, 20, "scrambled order changed greedy result: {}", b);
}
#[test]
fn sample_token_with_logprob_uniform_distribution() {
let n = 64usize;
let mut logits = vec![0.5_f32; n];
let params = SamplingParams {
temperature: 0.0,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 1,
};
let (_token, logprob) = sample_token_with_logprob(&mut logits, ¶ms, &[]);
let expected = -(n as f32).ln();
assert!(
(logprob - expected).abs() < 1e-4,
"uniform logprob: expected {expected:.6}, got {logprob:.6}"
);
}
#[test]
fn sample_token_with_logprob_concentrated_distribution() {
let n = 64usize;
let mut logits = vec![-100.0_f32; n];
logits[42] = 100.0;
let params = SamplingParams {
temperature: 0.0,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 1,
};
let (token, logprob) = sample_token_with_logprob(&mut logits, ¶ms, &[]);
assert_eq!(token, 42);
assert!(
logprob > -1e-3,
"concentrated logprob: expected ≈ 0, got {logprob:.6}"
);
}
#[test]
fn sample_token_with_logprob_known_two_token_distribution() {
let mut logits = vec![0.0_f32, 2.0_f32.ln()];
let params = SamplingParams {
temperature: 0.0,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 1,
};
let (token, logprob) = sample_token_with_logprob(&mut logits, ¶ms, &[]);
assert_eq!(token, 1, "greedy should pick the larger logit");
let expected = (2.0_f32 / 3.0).ln();
assert!(
(logprob - expected).abs() < 1e-5,
"two-token logprob: expected {expected:.6}, got {logprob:.6}"
);
}
#[test]
fn sample_token_with_logprob_all_neg_inf_returns_inf() {
let mut logits = vec![f32::NEG_INFINITY; 16];
let params = SamplingParams {
temperature: 0.0,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: 1,
};
let (_token, logprob) = sample_token_with_logprob(&mut logits, ¶ms, &[]);
assert!(logprob.is_infinite() && logprob < 0.0);
}
}