use std::collections::{HashMap, HashSet};
pub fn xorshift64(state: &mut u64) -> u64 {
let mut x = *state;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
*state = x;
x
}
pub fn apply_penalties(
logits: &mut [f32],
prompt: &[u32],
emitted: &[u32],
presence: f32,
frequency: f32,
repetition: f32,
) {
if repetition != 1.0 {
let mut seen = HashSet::new();
for &t in prompt.iter().chain(emitted) {
if seen.insert(t)
&& let Some(l) = logits.get_mut(t as usize)
{
*l = if *l > 0.0 {
*l / repetition
} else {
*l * repetition
};
}
}
}
if frequency != 0.0 || presence != 0.0 {
let mut counts: HashMap<u32, u32> = HashMap::new();
for &t in emitted {
*counts.entry(t).or_default() += 1;
}
for (&t, &c) in &counts {
if let Some(l) = logits.get_mut(t as usize) {
*l -= frequency * c as f32;
*l -= presence;
}
}
}
}
pub fn apply_logit_bias(logits: &mut [f32], bias: &[(u32, f32)]) {
for &(id, b) in bias {
if let Some(l) = logits.get_mut(id as usize) {
*l += b.clamp(-100.0, 100.0);
}
}
}
pub fn argmax(logits: &[f32]) -> u32 {
let mut best = 0u32;
let mut best_v = f32::NEG_INFINITY;
for (i, &l) in logits.iter().enumerate() {
if l > best_v {
best_v = l;
best = i as u32;
}
}
best
}
pub fn sample(
logits: &[f32],
temperature: f32,
top_k: Option<usize>,
top_p: f32,
min_p: Option<f32>,
rng: &mut u64,
) -> u32 {
let n = logits.len();
if n == 0 {
return 0;
}
let inv_t = 1.0 / temperature.max(1e-6);
let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut z = 0.0f32;
for &l in logits {
z += ((l - mx) * inv_t).exp();
}
let inv_z = if z > 0.0 { 1.0 / z } else { 0.0 };
let hard_k = top_k.map(|k| k.max(1).min(n));
let mut cap = hard_k.unwrap_or(64.min(n));
let probs = loop {
let probs = top_candidates(logits, cap, mx, inv_t, inv_z);
let covered: f32 = probs.iter().map(|(_, p)| *p).sum();
if hard_k.is_some() || cap >= n || covered >= top_p {
break probs;
}
cap = (cap * 4).min(n);
};
let mut end = probs.len();
let mut mass = 0.0f32;
let mut nucleus_end = 0;
for p in probs.iter().take(end) {
mass += p.1;
nucleus_end += 1;
if mass >= top_p {
break;
}
}
end = nucleus_end;
if let Some(mp) = min_p {
let thresh = mp * probs[0].1;
let kept = probs[1..end]
.iter()
.take_while(|(_, p)| *p >= thresh)
.count();
end = 1 + kept;
}
let kept_mass: f32 = probs.iter().take(end).map(|(_, p)| p).sum();
let draw = (xorshift64(rng) >> 11) as f32 / (1u64 << 53) as f32 * kept_mass;
let mut acc = 0.0f32;
for (id, p) in probs.iter().take(end) {
acc += p;
if draw <= acc {
return *id;
}
}
probs[end - 1].0
}
fn top_candidates(logits: &[f32], k: usize, mx: f32, inv_t: f32, inv_z: f32) -> Vec<(u32, f32)> {
fn worse_than(a: &(u32, f32), b: &(u32, f32)) -> bool {
match a.1.total_cmp(&b.1) {
std::cmp::Ordering::Less => true,
std::cmp::Ordering::Greater => false,
std::cmp::Ordering::Equal => a.0 > b.0,
}
}
let mut best: Vec<(u32, f32)> = Vec::with_capacity(k + 1);
for (i, &l) in logits.iter().enumerate() {
let cand = (i as u32, ((l - mx) * inv_t).exp() * inv_z);
if best.len() == k && !worse_than(&best[0], &cand) {
continue; }
let pos = best.partition_point(|x| worse_than(x, &cand));
best.insert(pos, cand);
if best.len() > k {
best.remove(0);
}
}
best.reverse(); best
}
#[derive(Clone, Debug, PartialEq)]
pub struct TokenLogprobs {
pub chosen: u32,
pub chosen_logprob: f32,
pub top: Vec<(u32, f32)>,
}
pub fn log_softmax_at(logits: &[f32], chosen: u32, top_n: usize) -> TokenLogprobs {
let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let sumexp: f32 = logits.iter().map(|&l| (l - mx).exp()).sum();
let logz = mx + sumexp.ln();
let chosen_logprob = logits
.get(chosen as usize)
.map(|&l| l - logz)
.unwrap_or(f32::NEG_INFINITY);
let mut top: Vec<(u32, f32)> = Vec::with_capacity(top_n);
let sort_desc = |t: &mut Vec<(u32, f32)>| {
t.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal))
};
for (i, &l) in logits.iter().enumerate() {
let lp = l - logz;
if top.len() < top_n {
top.push((i as u32, lp));
if top.len() == top_n {
sort_desc(&mut top);
}
} else if top_n > 0 && lp > top[top_n - 1].1 {
top[top_n - 1] = (i as u32, lp);
sort_desc(&mut top);
}
}
if top.len() < top_n {
sort_desc(&mut top);
}
TokenLogprobs {
chosen,
chosen_logprob,
top,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn presence_penalty_subtracts_once_per_emitted_token() {
let mut l = vec![1.0, 2.0, 3.0, 4.0];
apply_penalties(&mut l, &[], &[1, 1, 3], 0.5, 0.0, 1.0);
assert_eq!(l, vec![1.0, 1.5, 3.0, 3.5]);
}
#[test]
fn frequency_penalty_scales_with_count() {
let mut l = vec![0.0, 10.0, 10.0];
apply_penalties(&mut l, &[], &[1, 1, 1, 2], 0.0, 2.0, 1.0);
assert_eq!(l, vec![0.0, 4.0, 8.0]);
}
#[test]
fn repetition_penalty_is_multiplicative_over_prompt_and_output_once() {
let mut l = vec![2.0, -2.0, 5.0];
apply_penalties(&mut l, &[0], &[1], 0.0, 0.0, 2.0);
assert_eq!(l, vec![1.0, -4.0, 5.0]);
let mut l2 = vec![8.0];
apply_penalties(&mut l2, &[0, 0], &[0, 0], 0.0, 0.0, 2.0);
assert_eq!(l2, vec![4.0]);
}
#[test]
fn logit_bias_clamps_to_plus_minus_hundred() {
let mut l = vec![0.0, 0.0, 0.0];
apply_logit_bias(&mut l, &[(0, 1000.0), (1, -1000.0), (2, 5.0)]);
assert_eq!(l, vec![100.0, -100.0, 5.0]);
}
#[test]
fn logit_bias_minus_hundred_bans_a_token_from_greedy() {
let mut l = vec![1.0, 2.0, 3.0];
apply_logit_bias(&mut l, &[(2, -100.0)]);
assert_eq!(argmax(&l), 1);
}
#[test]
fn top_k_restricts_the_candidate_set() {
let l = vec![3.0, 2.9, 2.8, 2.7];
for seed in 0..8u64 {
let mut rng = seed ^ 0x9e37_79b9_7f4a_7c15;
assert_eq!(sample(&l, 1.0, Some(1), 1.0, None, &mut rng), 0);
}
}
#[test]
fn min_p_drops_low_probability_tail() {
let l = vec![10.0, 0.0, 0.0, 0.0];
for seed in 0..8u64 {
let mut rng = seed ^ 0x1234;
assert_eq!(sample(&l, 1.0, None, 1.0, Some(0.5), &mut rng), 0);
}
}
#[test]
fn sample_defaults_reduce_to_plain_nucleus_and_are_seed_reproducible() {
let l = vec![1.0, 2.0, 1.5, 0.5, 3.0];
let mut a = 42u64 ^ 0x9e37_79b9_7f4a_7c15;
let mut b = 42u64 ^ 0x9e37_79b9_7f4a_7c15;
let ta: Vec<u32> = (0..5)
.map(|_| sample(&l, 0.8, None, 0.9, None, &mut a))
.collect();
let tb: Vec<u32> = (0..5)
.map(|_| sample(&l, 0.8, None, 0.9, None, &mut b))
.collect();
assert_eq!(ta, tb);
}
#[test]
fn logprobs_are_a_normalized_log_softmax() {
let l = vec![0.0, 0.0]; let lp = log_softmax_at(&l, 0, 2);
assert!((lp.chosen_logprob - 0.5f32.ln()).abs() < 1e-5);
let mass: f32 = lp.top.iter().map(|(_, x)| x.exp()).sum();
assert!((mass - 1.0).abs() < 1e-5);
}
#[test]
fn chosen_logprob_matches_its_entry_in_top() {
let l = vec![1.0, 3.0, 2.0, 0.5];
let chosen = 1; let lp = log_softmax_at(&l, chosen, 3);
let in_top = lp
.top
.iter()
.find(|(id, _)| *id == chosen)
.expect("chosen in top");
assert_eq!(in_top.1, lp.chosen_logprob);
assert_eq!(lp.top[0].0, 1);
}
}
#[cfg(test)]
mod selection {
use super::*;
fn sample_by_sorting(
logits: &[f32],
temperature: f32,
top_k: Option<usize>,
top_p: f32,
min_p: Option<f32>,
rng: &mut u64,
) -> u32 {
let inv_t = 1.0 / temperature.max(1e-6);
let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut probs: Vec<(u32, f32)> = logits
.iter()
.enumerate()
.map(|(i, &l)| (i as u32, ((l - mx) * inv_t).exp()))
.collect();
let z: f32 = probs.iter().map(|(_, p)| p).sum();
for p in probs.iter_mut() {
p.1 /= z;
}
probs.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
let mut end = probs.len();
if let Some(k) = top_k {
end = end.min(k.max(1));
}
let mut mass = 0.0f32;
let mut nucleus_end = 0;
for p in probs.iter().take(end) {
mass += p.1;
nucleus_end += 1;
if mass >= top_p {
break;
}
}
end = nucleus_end;
if let Some(mp) = min_p {
let thresh = mp * probs[0].1;
let kept = probs[1..end]
.iter()
.take_while(|(_, p)| *p >= thresh)
.count();
end = 1 + kept;
}
let kept_mass: f32 = probs.iter().take(end).map(|(_, p)| p).sum();
let draw = (xorshift64(rng) >> 11) as f32 / (1u64 << 53) as f32 * kept_mass;
let mut acc = 0.0f32;
for (id, p) in probs.iter().take(end) {
acc += p;
if draw <= acc {
return *id;
}
}
probs[end - 1].0
}
fn logits(n: usize, seed: u64, flat: bool) -> Vec<f32> {
let mut r = seed;
(0..n)
.map(|_| {
let v = (xorshift64(&mut r) >> 40) as f32 / 1024.0;
if flat { v * 0.001 } else { v } })
.collect()
}
#[test]
fn sampler_matches_the_sorting_reference() {
let configs: [(Option<usize>, f32, Option<f32>, f32); 7] = [
(Some(20), 0.8, None, 0.7), (Some(1), 1.0, None, 1.0), (Some(50), 0.95, Some(0.05), 1.2),
(None, 0.8, None, 0.7), (None, 0.999, None, 1.0), (None, 1.0, Some(0.1), 0.5),
(Some(10_000), 0.9, None, 1.0), ];
for &flat in &[false, true] {
for &n in &[64usize, 1024, 32_000] {
for (ci, &(top_k, top_p, min_p, temp)) in configs.iter().enumerate() {
let lg = logits(n, 0xC0FFEE + n as u64 + ci as u64, flat);
for seed in 0..24u64 {
let (mut r1, mut r2) = (seed * 977 + 1, seed * 977 + 1);
let got = sample(&lg, temp, top_k, top_p, min_p, &mut r1);
let want = sample_by_sorting(&lg, temp, top_k, top_p, min_p, &mut r2);
assert_eq!(
got, want,
"n={n} flat={flat} cfg={ci} seed={seed}: selection drew {got}, sort drew {want}"
);
assert_eq!(r1, r2, "the RNG must advance identically");
}
}
}
}
}
#[test]
fn ties_keep_the_lowest_token_id() {
let lg = vec![1.0f32; 500]; for seed in 0..16u64 {
let (mut r1, mut r2) = (seed + 7, seed + 7);
assert_eq!(
sample(&lg, 1.0, Some(3), 1.0, None, &mut r1),
sample_by_sorting(&lg, 1.0, Some(3), 1.0, None, &mut r2),
);
}
}
}
pub fn prompt_lookup_drafts(all: &[u32], k: usize) -> Vec<u32> {
let n = all.len();
for glen in (1..=3.min(n.saturating_sub(1))).rev() {
let suffix = &all[n - glen..];
for start in (0..n - glen).rev() {
if &all[start..start + glen] == suffix {
let cont = &all[start + glen..(start + glen + k).min(n)];
if !cont.is_empty() {
let mut d = cont.to_vec();
while d.len() < k {
d.push(d[d.len() % cont.len().max(1)]);
}
return d;
}
}
}
}
Vec::new()
}