1use std::collections::{HashMap, HashSet};
19
20pub fn xorshift64(state: &mut u64) -> u64 {
23 let mut x = *state;
24 x ^= x << 13;
25 x ^= x >> 7;
26 x ^= x << 17;
27 *state = x;
28 x
29}
30
31pub fn apply_penalties(
38 logits: &mut [f32],
39 prompt: &[u32],
40 emitted: &[u32],
41 presence: f32,
42 frequency: f32,
43 repetition: f32,
44) {
45 if repetition != 1.0 {
46 let mut seen = HashSet::new();
47 for &t in prompt.iter().chain(emitted) {
48 if seen.insert(t)
49 && let Some(l) = logits.get_mut(t as usize)
50 {
51 *l = if *l > 0.0 {
52 *l / repetition
53 } else {
54 *l * repetition
55 };
56 }
57 }
58 }
59 if frequency != 0.0 || presence != 0.0 {
60 let mut counts: HashMap<u32, u32> = HashMap::new();
61 for &t in emitted {
62 *counts.entry(t).or_default() += 1;
63 }
64 for (&t, &c) in &counts {
65 if let Some(l) = logits.get_mut(t as usize) {
66 *l -= frequency * c as f32;
67 *l -= presence;
68 }
69 }
70 }
71}
72
73pub fn apply_logit_bias(logits: &mut [f32], bias: &[(u32, f32)]) {
76 for &(id, b) in bias {
77 if let Some(l) = logits.get_mut(id as usize) {
78 *l += b.clamp(-100.0, 100.0);
79 }
80 }
81}
82
83pub fn argmax(logits: &[f32]) -> u32 {
85 let mut best = 0u32;
86 let mut best_v = f32::NEG_INFINITY;
87 for (i, &l) in logits.iter().enumerate() {
88 if l > best_v {
89 best_v = l;
90 best = i as u32;
91 }
92 }
93 best
94}
95
96pub fn sample(
119 logits: &[f32],
120 temperature: f32,
121 top_k: Option<usize>,
122 top_p: f32,
123 min_p: Option<f32>,
124 rng: &mut u64,
125) -> u32 {
126 let n = logits.len();
127 if n == 0 {
128 return 0;
129 }
130 let inv_t = 1.0 / temperature.max(1e-6);
131 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
132
133 let mut z = 0.0f32;
136 for &l in logits {
137 z += ((l - mx) * inv_t).exp();
138 }
139 let inv_z = if z > 0.0 { 1.0 / z } else { 0.0 };
140
141 let hard_k = top_k.map(|k| k.max(1).min(n));
144 let mut cap = hard_k.unwrap_or(64.min(n));
145 let probs = loop {
146 let probs = top_candidates(logits, cap, mx, inv_t, inv_z);
147 let covered: f32 = probs.iter().map(|(_, p)| *p).sum();
148 if hard_k.is_some() || cap >= n || covered >= top_p {
149 break probs;
150 }
151 cap = (cap * 4).min(n);
152 };
153
154 let mut end = probs.len();
156 let mut mass = 0.0f32;
158 let mut nucleus_end = 0;
159 for p in probs.iter().take(end) {
160 mass += p.1;
161 nucleus_end += 1;
162 if mass >= top_p {
163 break;
164 }
165 }
166 end = nucleus_end;
167 if let Some(mp) = min_p {
170 let thresh = mp * probs[0].1;
171 let kept = probs[1..end]
172 .iter()
173 .take_while(|(_, p)| *p >= thresh)
174 .count();
175 end = 1 + kept;
176 }
177
178 let kept_mass: f32 = probs.iter().take(end).map(|(_, p)| p).sum();
179 let draw = (xorshift64(rng) >> 11) as f32 / (1u64 << 53) as f32 * kept_mass;
180 let mut acc = 0.0f32;
181 for (id, p) in probs.iter().take(end) {
182 acc += p;
183 if draw <= acc {
184 return *id;
185 }
186 }
187 probs[end - 1].0
188}
189
190fn top_candidates(logits: &[f32], k: usize, mx: f32, inv_t: f32, inv_z: f32) -> Vec<(u32, f32)> {
196 fn worse_than(a: &(u32, f32), b: &(u32, f32)) -> bool {
199 match a.1.total_cmp(&b.1) {
200 std::cmp::Ordering::Less => true,
201 std::cmp::Ordering::Greater => false,
202 std::cmp::Ordering::Equal => a.0 > b.0,
203 }
204 }
205
206 let mut best: Vec<(u32, f32)> = Vec::with_capacity(k + 1);
209 for (i, &l) in logits.iter().enumerate() {
210 let cand = (i as u32, ((l - mx) * inv_t).exp() * inv_z);
211 if best.len() == k && !worse_than(&best[0], &cand) {
212 continue; }
214 let pos = best.partition_point(|x| worse_than(x, &cand));
215 best.insert(pos, cand);
216 if best.len() > k {
217 best.remove(0);
218 }
219 }
220 best.reverse(); best
222}
223
224#[derive(Clone, Debug, PartialEq)]
227pub struct TokenLogprobs {
228 pub chosen: u32,
229 pub chosen_logprob: f32,
230 pub top: Vec<(u32, f32)>,
232}
233
234pub fn log_softmax_at(logits: &[f32], chosen: u32, top_n: usize) -> TokenLogprobs {
236 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
237 let sumexp: f32 = logits.iter().map(|&l| (l - mx).exp()).sum();
238 let logz = mx + sumexp.ln();
239 let chosen_logprob = logits
240 .get(chosen as usize)
241 .map(|&l| l - logz)
242 .unwrap_or(f32::NEG_INFINITY);
243 let mut top: Vec<(u32, f32)> = Vec::with_capacity(top_n);
245 let sort_desc = |t: &mut Vec<(u32, f32)>| {
246 t.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal))
247 };
248 for (i, &l) in logits.iter().enumerate() {
249 let lp = l - logz;
250 if top.len() < top_n {
251 top.push((i as u32, lp));
252 if top.len() == top_n {
253 sort_desc(&mut top);
254 }
255 } else if top_n > 0 && lp > top[top_n - 1].1 {
256 top[top_n - 1] = (i as u32, lp);
257 sort_desc(&mut top);
258 }
259 }
260 if top.len() < top_n {
261 sort_desc(&mut top);
262 }
263 TokenLogprobs {
264 chosen,
265 chosen_logprob,
266 top,
267 }
268}
269
270#[cfg(test)]
271mod tests {
272 use super::*;
273
274 #[test]
275 fn presence_penalty_subtracts_once_per_emitted_token() {
276 let mut l = vec![1.0, 2.0, 3.0, 4.0];
277 apply_penalties(&mut l, &[], &[1, 1, 3], 0.5, 0.0, 1.0);
279 assert_eq!(l, vec![1.0, 1.5, 3.0, 3.5]);
280 }
281
282 #[test]
283 fn frequency_penalty_scales_with_count() {
284 let mut l = vec![0.0, 10.0, 10.0];
285 apply_penalties(&mut l, &[], &[1, 1, 1, 2], 0.0, 2.0, 1.0);
287 assert_eq!(l, vec![0.0, 4.0, 8.0]);
288 }
289
290 #[test]
291 fn repetition_penalty_is_multiplicative_over_prompt_and_output_once() {
292 let mut l = vec![2.0, -2.0, 5.0];
293 apply_penalties(&mut l, &[0], &[1], 0.0, 0.0, 2.0);
295 assert_eq!(l, vec![1.0, -4.0, 5.0]);
296 let mut l2 = vec![8.0];
298 apply_penalties(&mut l2, &[0, 0], &[0, 0], 0.0, 0.0, 2.0);
299 assert_eq!(l2, vec![4.0]);
300 }
301
302 #[test]
303 fn logit_bias_clamps_to_plus_minus_hundred() {
304 let mut l = vec![0.0, 0.0, 0.0];
305 apply_logit_bias(&mut l, &[(0, 1000.0), (1, -1000.0), (2, 5.0)]);
306 assert_eq!(l, vec![100.0, -100.0, 5.0]);
307 }
308
309 #[test]
310 fn logit_bias_minus_hundred_bans_a_token_from_greedy() {
311 let mut l = vec![1.0, 2.0, 3.0];
313 apply_logit_bias(&mut l, &[(2, -100.0)]);
314 assert_eq!(argmax(&l), 1);
315 }
316
317 #[test]
318 fn top_k_restricts_the_candidate_set() {
319 let l = vec![3.0, 2.9, 2.8, 2.7];
321 for seed in 0..8u64 {
322 let mut rng = seed ^ 0x9e37_79b9_7f4a_7c15;
323 assert_eq!(sample(&l, 1.0, Some(1), 1.0, None, &mut rng), 0);
324 }
325 }
326
327 #[test]
328 fn min_p_drops_low_probability_tail() {
329 let l = vec![10.0, 0.0, 0.0, 0.0];
331 for seed in 0..8u64 {
332 let mut rng = seed ^ 0x1234;
333 assert_eq!(sample(&l, 1.0, None, 1.0, Some(0.5), &mut rng), 0);
334 }
335 }
336
337 #[test]
338 fn sample_defaults_reduce_to_plain_nucleus_and_are_seed_reproducible() {
339 let l = vec![1.0, 2.0, 1.5, 0.5, 3.0];
341 let mut a = 42u64 ^ 0x9e37_79b9_7f4a_7c15;
342 let mut b = 42u64 ^ 0x9e37_79b9_7f4a_7c15;
343 let ta: Vec<u32> = (0..5)
344 .map(|_| sample(&l, 0.8, None, 0.9, None, &mut a))
345 .collect();
346 let tb: Vec<u32> = (0..5)
347 .map(|_| sample(&l, 0.8, None, 0.9, None, &mut b))
348 .collect();
349 assert_eq!(ta, tb);
350 }
351
352 #[test]
353 fn logprobs_are_a_normalized_log_softmax() {
354 let l = vec![0.0, 0.0]; let lp = log_softmax_at(&l, 0, 2);
356 assert!((lp.chosen_logprob - 0.5f32.ln()).abs() < 1e-5);
357 let mass: f32 = lp.top.iter().map(|(_, x)| x.exp()).sum();
359 assert!((mass - 1.0).abs() < 1e-5);
360 }
361
362 #[test]
363 fn chosen_logprob_matches_its_entry_in_top() {
364 let l = vec![1.0, 3.0, 2.0, 0.5];
366 let chosen = 1; let lp = log_softmax_at(&l, chosen, 3);
368 let in_top = lp
369 .top
370 .iter()
371 .find(|(id, _)| *id == chosen)
372 .expect("chosen in top");
373 assert_eq!(in_top.1, lp.chosen_logprob);
374 assert_eq!(lp.top[0].0, 1);
376 }
377}
378
379#[cfg(test)]
380mod selection {
381 use super::*;
382
383 fn sample_by_sorting(
390 logits: &[f32],
391 temperature: f32,
392 top_k: Option<usize>,
393 top_p: f32,
394 min_p: Option<f32>,
395 rng: &mut u64,
396 ) -> u32 {
397 let inv_t = 1.0 / temperature.max(1e-6);
398 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
399 let mut probs: Vec<(u32, f32)> = logits
400 .iter()
401 .enumerate()
402 .map(|(i, &l)| (i as u32, ((l - mx) * inv_t).exp()))
403 .collect();
404 let z: f32 = probs.iter().map(|(_, p)| p).sum();
405 for p in probs.iter_mut() {
406 p.1 /= z;
407 }
408 probs.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
409
410 let mut end = probs.len();
411 if let Some(k) = top_k {
412 end = end.min(k.max(1));
413 }
414 let mut mass = 0.0f32;
415 let mut nucleus_end = 0;
416 for p in probs.iter().take(end) {
417 mass += p.1;
418 nucleus_end += 1;
419 if mass >= top_p {
420 break;
421 }
422 }
423 end = nucleus_end;
424 if let Some(mp) = min_p {
425 let thresh = mp * probs[0].1;
426 let kept = probs[1..end]
427 .iter()
428 .take_while(|(_, p)| *p >= thresh)
429 .count();
430 end = 1 + kept;
431 }
432 let kept_mass: f32 = probs.iter().take(end).map(|(_, p)| p).sum();
433 let draw = (xorshift64(rng) >> 11) as f32 / (1u64 << 53) as f32 * kept_mass;
434 let mut acc = 0.0f32;
435 for (id, p) in probs.iter().take(end) {
436 acc += p;
437 if draw <= acc {
438 return *id;
439 }
440 }
441 probs[end - 1].0
442 }
443
444 fn logits(n: usize, seed: u64, flat: bool) -> Vec<f32> {
445 let mut r = seed;
446 (0..n)
447 .map(|_| {
448 let v = (xorshift64(&mut r) >> 40) as f32 / 1024.0;
449 if flat { v * 0.001 } else { v } })
451 .collect()
452 }
453
454 #[test]
455 fn sampler_matches_the_sorting_reference() {
456 let configs: [(Option<usize>, f32, Option<f32>, f32); 7] = [
459 (Some(20), 0.8, None, 0.7), (Some(1), 1.0, None, 1.0), (Some(50), 0.95, Some(0.05), 1.2),
462 (None, 0.8, None, 0.7), (None, 0.999, None, 1.0), (None, 1.0, Some(0.1), 0.5),
465 (Some(10_000), 0.9, None, 1.0), ];
467 for &flat in &[false, true] {
469 for &n in &[64usize, 1024, 32_000] {
470 for (ci, &(top_k, top_p, min_p, temp)) in configs.iter().enumerate() {
471 let lg = logits(n, 0xC0FFEE + n as u64 + ci as u64, flat);
472 for seed in 0..24u64 {
473 let (mut r1, mut r2) = (seed * 977 + 1, seed * 977 + 1);
474 let got = sample(&lg, temp, top_k, top_p, min_p, &mut r1);
475 let want = sample_by_sorting(&lg, temp, top_k, top_p, min_p, &mut r2);
476 assert_eq!(
477 got, want,
478 "n={n} flat={flat} cfg={ci} seed={seed}: selection drew {got}, sort drew {want}"
479 );
480 assert_eq!(r1, r2, "the RNG must advance identically");
481 }
482 }
483 }
484 }
485 }
486
487 #[test]
489 fn ties_keep_the_lowest_token_id() {
490 let lg = vec![1.0f32; 500]; for seed in 0..16u64 {
492 let (mut r1, mut r2) = (seed + 7, seed + 7);
493 assert_eq!(
494 sample(&lg, 1.0, Some(3), 1.0, None, &mut r1),
495 sample_by_sorting(&lg, 1.0, Some(3), 1.0, None, &mut r2),
496 );
497 }
498 }
499}
500
501pub fn prompt_lookup_drafts(all: &[u32], k: usize) -> Vec<u32> {
508 let n = all.len();
509 for glen in (1..=3.min(n.saturating_sub(1))).rev() {
510 let suffix = &all[n - glen..];
511 for start in (0..n - glen).rev() {
512 if &all[start..start + glen] == suffix {
513 let cont = &all[start + glen..(start + glen + k).min(n)];
514 if !cont.is_empty() {
515 let mut d = cont.to_vec();
516 while d.len() < k {
517 d.push(d[d.len() % cont.len().max(1)]);
518 }
519 return d;
520 }
521 }
522 }
523 }
524 Vec::new()
525}