1use serde::{Deserialize, Serialize};
8
9#[derive(Debug, Clone)]
11pub struct SplitMix64 {
12 state: u64,
13}
14
15#[derive(Debug, Default)]
19pub struct SamplerScratch {
20 seen_epoch: Vec<u32>,
21 epoch: u32,
22 presence_seen: std::collections::HashSet<u32>,
24 probs: Vec<f32>,
28 topk: Vec<f32>,
33 cand: Vec<(u32, f32)>,
35 cand_parts: Vec<Vec<(u32, f32)>>,
36 sum_parts: Vec<f32>,
37 sparse: Sparse,
38}
39
40impl SamplerScratch {
41 fn begin_seen(&mut self, vocab_size: usize) -> u32 {
42 if self.seen_epoch.len() < vocab_size {
43 self.seen_epoch.resize(vocab_size, 0);
44 }
45 self.epoch = self.epoch.wrapping_add(1);
46 if self.epoch == 0 {
47 self.seen_epoch.fill(0);
48 self.epoch = 1;
49 }
50 self.epoch
51 }
52}
53
54impl SplitMix64 {
55 pub fn new(seed: u64) -> Self {
56 Self { state: seed }
57 }
58
59 pub fn from_entropy() -> Self {
61 let t = std::time::SystemTime::now()
62 .duration_since(std::time::UNIX_EPOCH)
63 .unwrap_or_default();
64 let addr = Box::into_raw(Box::new(0u8)) as u64;
65 unsafe { drop(Box::from_raw(addr as *mut u8)) };
67 Self::new(t.as_nanos() as u64 ^ addr.rotate_left(17) ^ 0x9E3779B97F4A7C15)
68 }
69
70 #[inline]
71 pub fn next_u64(&mut self) -> u64 {
72 self.state = self.state.wrapping_add(0x9E3779B97F4A7C15);
73 let mut z = self.state;
74 z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
75 z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
76 z ^ (z >> 31)
77 }
78
79 #[inline]
81 pub fn next_f32(&mut self) -> f32 {
82 (self.next_u64() >> 40) as f32 / (1u64 << 24) as f32
83 }
84}
85
86#[derive(Debug, Clone, Serialize, Deserialize)]
88pub struct SamplerConfig {
89 pub temperature: f32,
90 pub top_p: f32,
91 pub top_k: u32,
92 pub repetition_penalty: f32,
93 pub min_p: f32,
94 #[serde(default)]
99 pub presence_penalty: f32,
100 #[serde(default)]
102 pub seed: Option<u64>,
103 #[serde(default)]
105 pub suppress_tokens: Vec<u32>,
106 #[serde(default)]
112 pub penalty_window: usize,
113}
114
115pub const BOUNDED_PENALTY_WINDOW: usize = 128;
118
119impl SamplerConfig {
120 pub fn penalty_past<'a>(&self, past: &'a [u32], bounded_native: bool) -> &'a [u32] {
124 let mut w = self.penalty_window;
125 if w == 0 && bounded_native {
126 w = std::env::var("CMF_PENALTY_WINDOW")
127 .ok()
128 .and_then(|v| v.parse::<usize>().ok())
129 .filter(|&v| v > 0)
130 .unwrap_or(BOUNDED_PENALTY_WINDOW);
131 }
132 if w == 0 {
133 past
134 } else {
135 &past[past.len().saturating_sub(w)..]
136 }
137 }
138}
139
140impl Default for SamplerConfig {
141 fn default() -> Self {
142 Self {
143 temperature: 0.7,
144 top_p: 0.9,
145 top_k: 40,
146 repetition_penalty: 1.1,
147 presence_penalty: 0.0,
148 min_p: 0.05,
149 seed: None,
150 suppress_tokens: Vec::new(),
151 penalty_window: 0,
152 }
153 }
154}
155
156pub fn sample(
159 logits: &[f32],
160 config: &SamplerConfig,
161 past_tokens: &[u32],
162 rng: &mut SplitMix64,
163) -> u32 {
164 let mut scratch = SamplerScratch::default();
165 sample_with_scratch(logits, config, past_tokens, rng, &mut scratch)
166}
167
168pub fn sample_with_scratch(
170 logits: &[f32],
171 config: &SamplerConfig,
172 past_tokens: &[u32],
173 rng: &mut SplitMix64,
174 scratch: &mut SamplerScratch,
175) -> u32 {
176 sample_with_scratch_pool(logits, config, past_tokens, rng, scratch, None)
177}
178
179pub fn sample_with_scratch_pool(
195 logits: &[f32],
196 config: &SamplerConfig,
197 past_tokens: &[u32],
198 rng: &mut SplitMix64,
199 scratch: &mut SamplerScratch,
200 pool: Option<&crate::pool::Pool>,
201) -> u32 {
202 if config.temperature < 1e-6
203 && config.repetition_penalty == 1.0
204 && config.presence_penalty == 0.0
205 && config.suppress_tokens.is_empty()
206 {
207 return argmax(logits);
208 }
209 if config.temperature < 1e-6 {
215 return argmax_penalized(logits, config, past_tokens, scratch, pool);
217 }
218 if sparse_ok(config) {
219 let mut sp = std::mem::take(&mut scratch.sparse);
221 let ok = sparse_distribution_into(logits, config, past_tokens, scratch, pool, &mut sp);
222 let t = if ok {
223 draw_sparse(&sp, rng)
224 } else {
225 argmax(logits)
226 };
227 scratch.sparse = sp;
228 return t;
229 }
230 let mut probs = std::mem::take(&mut scratch.probs);
231 let normalized = chain(logits, config, past_tokens, scratch, pool, &mut probs);
232
233 let mut done = |probs: Vec<f32>, tok: u32| -> u32 {
234 scratch.probs = probs;
235 tok
236 };
237
238 if !normalized {
239 let t = argmax(logits);
241 return done(probs, t);
242 }
243 let t = categorical_sample(&probs, rng.next_f32());
244 done(probs, t)
245}
246
247pub fn argmax_penalized(
257 logits: &[f32],
258 config: &SamplerConfig,
259 past_tokens: &[u32],
260 scratch: &mut SamplerScratch,
261 pool: Option<&crate::pool::Pool>,
262) -> u32 {
263 let n = logits.len();
264 if config.repetition_penalty == 1.0
265 && config.presence_penalty == 0.0
266 && config.suppress_tokens.is_empty()
267 {
268 return argmax(logits);
269 }
270 let epoch = scratch.begin_seen(n);
272 for &tok in past_tokens {
273 let idx = tok as usize;
274 if idx < n {
275 scratch.seen_epoch[idx] = epoch;
276 }
277 }
278 let rep = config.repetition_penalty;
279 let pres = config.presence_penalty;
280 let suppress = &config.suppress_tokens;
282 let seen = &scratch.seen_epoch;
283 let value = |i: usize| -> f32 {
284 let mut v = logits[i];
285 if suppress.iter().any(|&t| t as usize == i) {
286 return f32::NEG_INFINITY;
287 }
288 if seen[i] == epoch {
289 if rep != 1.0 {
290 if v > 0.0 {
291 v /= rep;
292 } else {
293 v *= rep;
294 }
295 }
296 if pres != 0.0 {
297 v -= pres;
298 }
299 }
300 v
301 };
302 let best_in = |s: usize, e: usize| -> (usize, f32) {
305 let mut bi = s;
306 let mut bv = f32::NEG_INFINITY;
307 for i in s..e {
308 let v = value(i);
309 if v >= bv {
310 bv = v;
311 bi = i;
312 }
313 }
314 (bi, bv)
315 };
316 match pool {
317 Some(p) if n >= PAR_MIN && suppress.is_empty() => {
318 let m = std::sync::Mutex::new(Vec::<(usize, f32)>::new());
319 p.run_rows(n, &|s, e| {
320 let r = best_in(s, e);
321 m.lock().unwrap().push(r);
322 });
323 let mut parts = m.into_inner().unwrap();
324 parts.sort_by_key(|(i, _)| *i);
328 let mut bi = 0usize;
329 let mut bv = f32::NEG_INFINITY;
330 for (i, v) in parts {
331 if v >= bv {
332 bv = v;
333 bi = i;
334 }
335 }
336 bi as u32
337 }
338 _ => best_in(0, n).0 as u32,
339 }
340}
341
342fn chain(
348 logits: &[f32],
349 config: &SamplerConfig,
350 past_tokens: &[u32],
351 scratch: &mut SamplerScratch,
352 pool: Option<&crate::pool::Pool>,
353 probs: &mut Vec<f32>,
354) -> bool {
355 probs.clear();
356 probs.extend_from_slice(logits);
357 apply_penalties(probs, config, past_tokens, scratch);
358
359 if config.temperature < 1e-6 {
360 return true;
361 }
362 if config.temperature != 1.0 {
363 let t = config.temperature;
364 par_map(pool, probs, &move |p| p / t);
365 }
366
367 softmax_inplace_pool(pool, probs);
368
369 if config.min_p > 0.0 {
370 let max_prob = par_max(pool, probs, 0.0);
371 let threshold = max_prob * config.min_p;
372 par_map(pool, probs, &move |p| if p < threshold { 0.0 } else { p });
373 }
374
375 if config.top_k > 0 && (config.top_k as usize) < probs.len() {
376 apply_top_k_pool(pool, probs, config.top_k as usize);
377 }
378
379 if config.top_p < 1.0 && config.top_p > 0.0 {
380 apply_top_p(probs, config.top_p);
381 }
382
383 let sum: f32 = probs.iter().sum();
384 if sum > 0.0 {
385 par_map(pool, probs, &move |p| p / sum);
386 true
387 } else {
388 false
389 }
390}
391
392fn apply_penalties(
397 probs: &mut [f32],
398 config: &SamplerConfig,
399 past_tokens: &[u32],
400 scratch: &mut SamplerScratch,
401) {
402 for &tok in &config.suppress_tokens {
403 if (tok as usize) < probs.len() {
404 probs[tok as usize] = f32::NEG_INFINITY;
405 }
406 }
407 if config.repetition_penalty != 1.0 {
408 apply_repetition_penalty(probs, past_tokens, config.repetition_penalty, scratch);
409 }
410 if config.presence_penalty != 0.0 {
411 let mut seen = std::mem::take(&mut scratch.presence_seen);
416 seen.clear();
417 seen.extend(past_tokens.iter().copied());
418 for &tok in &seen {
419 if (tok as usize) < probs.len() {
420 probs[tok as usize] -= config.presence_penalty;
421 }
422 }
423 scratch.presence_seen = seen;
424 }
425}
426
427fn config_penalized(config: &SamplerConfig) -> bool {
428 config.repetition_penalty != 1.0
429 || config.presence_penalty != 0.0
430 || !config.suppress_tokens.is_empty()
431}
432
433pub const SPARSE_TOPK_MAX: usize = 256;
436
437pub fn sparse_ok(config: &SamplerConfig) -> bool {
441 config.temperature >= 1e-6 && config.top_k > 0 && (config.top_k as usize) <= SPARSE_TOPK_MAX
442}
443
444pub type Sparse = Vec<(u32, f32)>;
447
448pub fn sparse_distribution_into(
474 logits: &[f32],
475 config: &SamplerConfig,
476 past_tokens: &[u32],
477 scratch: &mut SamplerScratch,
478 pool: Option<&crate::pool::Pool>,
479 out: &mut Sparse,
480) -> bool {
481 debug_assert!(sparse_ok(config));
482 out.clear();
483 let k = (config.top_k as usize).min(logits.len());
484 if k == 0 {
485 return false;
486 }
487 let penalized = config_penalized(config);
488 let mut probs = std::mem::take(&mut scratch.probs);
489 if penalized {
490 probs.clear();
491 probs.extend_from_slice(logits);
492 apply_penalties(&mut probs, config, past_tokens, scratch);
493 }
494 let src: &[f32] = if penalized { &probs } else { logits };
495 let t = if config.temperature > 0.0 {
496 config.temperature
497 } else {
498 1.0
499 };
500 let mut cand = std::mem::take(&mut scratch.cand);
501 par_topk(pool, src, k, &mut cand, &mut scratch.cand_parts);
502 let ok = if let Some(&(_, lmax)) = cand.first().filter(|c| c.1.is_finite()) {
505 let sum_all = par_sum_exp(pool, src, lmax, t, &mut scratch.sum_parts);
506 let min_p = config.min_p;
509 let mut cum = 0.0f32;
510 let mut cut = false;
511 for &(id, l) in cand.iter() {
512 if cut {
513 break;
514 }
515 let e = ((l - lmax) / t).exp();
516 if min_p > 0.0 && e < min_p {
517 continue;
518 }
519 let pr = e / sum_all;
520 if pr <= 0.0 {
521 continue;
522 }
523 out.push((id, pr));
524 cum += pr;
525 if config.top_p < 1.0 && config.top_p > 0.0 && cum >= config.top_p {
526 cut = true;
527 }
528 }
529 let sum: f32 = out.iter().map(|c| c.1).sum();
531 if sum > 0.0 {
532 for c in out.iter_mut() {
533 c.1 /= sum;
534 }
535 out.sort_unstable_by_key(|c| c.0);
536 true
537 } else {
538 out.clear();
539 false
540 }
541 } else {
542 false
543 };
544 scratch.cand = cand;
545 scratch.probs = probs;
546 ok
547}
548
549fn par_topk(
556 pool: Option<&crate::pool::Pool>,
557 src: &[f32],
558 k: usize,
559 out: &mut Vec<(u32, f32)>,
560 parts: &mut Vec<Vec<(u32, f32)>>,
561) {
562 let better = |a: (u32, f32), b: (u32, f32)| a.1 > b.1 || (a.1 == b.1 && a.0 < b.0);
563 let scan = |s: usize, e: usize, best: &mut Vec<(u32, f32)>| {
566 best.clear();
567 for i in s..e {
568 let c = (i as u32, src[i]);
569 if best.len() < k {
570 let pos = best
571 .iter()
572 .position(|&b| better(c, b))
573 .unwrap_or(best.len());
574 best.insert(pos, c);
575 } else if better(c, best[k - 1]) {
576 let pos = best.iter().position(|&b| better(c, b)).unwrap_or(k - 1);
577 best.pop();
578 best.insert(pos, c);
579 }
580 }
581 };
582 let by_value_desc = |a: &(u32, f32), b: &(u32, f32)| {
583 b.1.partial_cmp(&a.1)
584 .unwrap_or(std::cmp::Ordering::Equal)
585 .then(a.0.cmp(&b.0))
586 };
587 let cap = k * 4 + 64;
591 let gather = |s: usize, e: usize, kth: f32, slot: &mut Vec<(u32, f32)>| {
593 slot.clear();
594 for i in s..e {
595 let v = src[i];
596 if v >= kth {
597 slot.push((i as u32, v));
598 if slot.len() >= cap {
599 break;
600 }
601 }
602 }
603 };
604 out.clear();
605 match pool {
606 Some(p) if src.len() >= PAR_MIN => {
607 let n = src.len();
608 let grain = crate::pool::grain_for(n, p.n_workers() + 1);
609 let ng = n.div_ceil(grain);
610 parts.resize_with(ng, Vec::new);
611 let pp = crate::pool::SendMutT::new(parts.as_mut_ptr());
612 p.run_rows(n, &|s, e| {
613 let slot = unsafe { &mut *pp.at(s / grain) };
616 scan(s, e, slot);
617 });
618 for g in 0..ng {
619 out.extend_from_slice(&parts[g]);
620 }
621 out.sort_unstable_by(by_value_desc);
622 out.truncate(k);
623 let Some(&(_, kth)) = out.last() else {
624 return;
625 };
626 if !kth.is_finite() {
627 return; }
629 p.run_rows(n, &|s, e| {
630 let slot = unsafe { &mut *pp.at(s / grain) };
631 gather(s, e, kth, slot);
632 });
633 out.clear();
634 for g in 0..ng {
635 out.extend_from_slice(&parts[g]);
636 if out.len() >= cap {
637 break;
638 }
639 }
640 out.sort_unstable_by(by_value_desc);
641 out.truncate(cap);
642 }
643 _ => {
644 scan(0, src.len(), out);
645 let Some(&(_, kth)) = out.last() else {
646 return;
647 };
648 if !kth.is_finite() {
649 return;
650 }
651 let mut all = std::mem::take(out);
652 gather(0, src.len(), kth, &mut all);
653 all.sort_unstable_by(by_value_desc);
654 all.truncate(cap);
655 *out = all;
656 }
657 }
658}
659
660fn par_sum_exp(
663 pool: Option<&crate::pool::Pool>,
664 src: &[f32],
665 lmax: f32,
666 t: f32,
667 parts: &mut Vec<f32>,
668) -> f32 {
669 let term = |s: usize, e: usize| -> f32 {
670 let mut acc = 0.0f32;
671 for &l in &src[s..e] {
672 acc += ((l - lmax) / t).exp();
673 }
674 acc
675 };
676 match pool {
677 Some(p) if src.len() >= PAR_MIN => {
678 let n = src.len();
679 let grain = crate::pool::grain_for(n, p.n_workers() + 1);
680 let ng = n.div_ceil(grain);
681 parts.clear();
682 parts.resize(ng, 0.0);
683 let pp = crate::pool::SendMut::new(parts.as_mut_ptr());
684 p.run_rows(n, &|s, e| {
685 unsafe { *pp.at(s / grain) = term(s, e) };
687 });
688 parts.iter().sum()
689 }
690 _ => term(0, src.len()),
691 }
692}
693
694pub fn draw_sparse(p: &[(u32, f32)], rng: &mut SplitMix64) -> u32 {
698 let r = rng.next_f32();
699 let mut cum = 0.0f32;
700 for &(id, pr) in p {
701 cum += pr;
702 if r < cum {
703 return id;
704 }
705 }
706 p.iter().rev().find(|c| c.1 > 0.0).map(|c| c.0).unwrap_or(0)
707}
708
709fn sparse_get(p: &[(u32, f32)], id: u32) -> f32 {
710 p.binary_search_by_key(&id, |c| c.0)
711 .map(|i| p[i].1)
712 .unwrap_or(0.0)
713}
714
715pub fn spec_accept_or_correct_sparse(
720 p: &[(u32, f32)],
721 q: &[(u32, f32)],
722 d: u32,
723 rng: &mut SplitMix64,
724 res: &mut Sparse,
725) -> Option<u32> {
726 let (pd, qd) = (sparse_get(p, d), sparse_get(q, d));
727 let r = rng.next_f32();
728 if qd > 0.0 && r * qd < pd {
729 return None;
730 }
731 res.clear();
732 let mut total = 0.0f32;
733 for &(id, pi) in p {
734 let ri = pi - sparse_get(q, id);
735 if ri > 0.0 {
736 res.push((id, ri));
737 total += ri;
738 }
739 }
740 if total <= 0.0 {
741 return Some(draw_sparse(p, rng));
742 }
743 for c in res.iter_mut() {
744 c.1 /= total;
745 }
746 Some(draw_sparse(res, rng))
747}
748
749pub fn distribution_into(
756 logits: &[f32],
757 config: &SamplerConfig,
758 past_tokens: &[u32],
759 scratch: &mut SamplerScratch,
760 pool: Option<&crate::pool::Pool>,
761 out: &mut Vec<f32>,
762) {
763 let one_hot = |out: &mut Vec<f32>, t: usize, n: usize| {
764 out.clear();
765 out.resize(n, 0.0);
766 if t < n {
767 out[t] = 1.0;
768 }
769 };
770 if config.temperature < 1e-6
771 && config.repetition_penalty == 1.0
772 && config.presence_penalty == 0.0
773 && config.suppress_tokens.is_empty()
774 {
775 return one_hot(out, argmax(logits) as usize, logits.len());
776 }
777 let mut probs = std::mem::take(&mut scratch.probs);
778 let normalized = chain(logits, config, past_tokens, scratch, pool, &mut probs);
779 if config.temperature < 1e-6 {
780 let t = argmax(&probs) as usize;
781 scratch.probs = probs;
782 return one_hot(out, t, logits.len());
783 }
784 if !normalized {
785 scratch.probs = probs;
786 return one_hot(out, argmax(logits) as usize, logits.len());
787 }
788 out.clear();
789 out.extend_from_slice(&probs);
790 scratch.probs = probs;
791}
792
793pub fn draw(probs: &[f32], rng: &mut SplitMix64) -> u32 {
795 categorical_sample(probs, rng.next_f32())
796}
797
798pub fn spec_accept_or_correct(
808 p: &[f32],
809 q: &[f32],
810 d: u32,
811 rng: &mut SplitMix64,
812 scratch: &mut Vec<f32>,
813 pool: Option<&crate::pool::Pool>,
814) -> Option<u32> {
815 let di = d as usize;
816 let (pd, qd) = (
817 p.get(di).copied().unwrap_or(0.0),
818 q.get(di).copied().unwrap_or(0.0),
819 );
820 let r = rng.next_f32();
821 if qd > 0.0 && r * qd < pd {
823 return None;
824 }
825 let n = p.len().min(q.len());
826 scratch.clear();
827 scratch.extend_from_slice(&p[..n]);
828 {
830 let qp = q.as_ptr() as usize;
831 let sm = crate::pool::SendMut::new(scratch.as_mut_ptr());
832 let body = move |s: usize, e: usize| {
833 let qs = unsafe { std::slice::from_raw_parts(qp as *const f32, n) };
835 for i in s..e {
836 unsafe {
837 let x = sm.at(i);
838 *x = (*x - qs[i]).max(0.0);
839 }
840 }
841 };
842 match pool {
843 Some(pl) if n >= PAR_MIN => pl.run_rows(n, &body),
844 _ => body(0, n),
845 }
846 }
847 let sum: f32 = scratch.iter().sum();
848 if sum > 0.0 {
849 let inv = 1.0 / sum;
850 par_map(pool, scratch, &move |v| v * inv);
851 Some(categorical_sample(scratch, rng.next_f32()))
852 } else {
853 Some(categorical_sample(&p[..n], rng.next_f32()))
854 }
855}
856
857pub fn argmax(values: &[f32]) -> u32 {
869 if values.is_empty() {
870 return 0;
871 }
872 let n = values.len();
873 let mut best = [(0usize, f32::NEG_INFINITY); 4];
874 for (l, b) in best.iter_mut().enumerate() {
875 b.0 = l.min(n - 1);
876 }
877 let mut i = 0;
878 while i + 4 <= n {
879 for l in 0..4 {
880 let v = values[i + l];
881 if v >= best[l].1 {
882 best[l] = (i + l, v);
883 }
884 }
885 i += 4;
886 }
887 let mut bi = best[0].0;
888 let mut bv = best[0].1;
889 for b in &best[1..] {
890 if b.1 > bv || (b.1 == bv && b.0 > bi) {
891 bi = b.0;
892 bv = b.1;
893 }
894 }
895 while i < n {
896 if values[i] >= bv {
897 bv = values[i];
898 bi = i;
899 }
900 i += 1;
901 }
902 bi as u32
903}
904
905fn softmax_inplace(logits: &mut [f32]) {
906 let max_val = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
907 let mut sum = 0.0f32;
908 for v in logits.iter_mut() {
909 *v = (*v - max_val).exp();
910 sum += *v;
911 }
912 if sum > 0.0 {
913 for v in logits.iter_mut() {
914 *v /= sum;
915 }
916 }
917}
918
919const PAR_MIN: usize = 1 << 14;
921
922fn par_map(pool: Option<&crate::pool::Pool>, buf: &mut [f32], f: &(dyn Fn(f32) -> f32 + Sync)) {
926 match pool {
927 Some(p) if buf.len() >= PAR_MIN => {
928 let out = crate::pool::SendMut::new(buf.as_mut_ptr());
929 p.run_rows(buf.len(), &move |s, e| {
930 for i in s..e {
931 unsafe {
934 let q = out.at(i);
935 *q = f(*q);
936 }
937 }
938 });
939 }
940 _ => {
941 for v in buf.iter_mut() {
942 *v = f(*v);
943 }
944 }
945 }
946}
947
948fn par_max(pool: Option<&crate::pool::Pool>, buf: &[f32], init: f32) -> f32 {
951 match pool {
952 Some(p) if buf.len() >= PAR_MIN => {
953 let m = std::sync::Mutex::new(init);
954 p.run_rows(buf.len(), &|s, e| {
955 let local = buf[s..e].iter().cloned().fold(init, f32::max);
956 let mut g = m.lock().unwrap();
957 *g = g.max(local);
958 });
959 m.into_inner().unwrap()
960 }
961 _ => buf.iter().cloned().fold(init, f32::max),
962 }
963}
964
965fn softmax_inplace_pool(pool: Option<&crate::pool::Pool>, logits: &mut [f32]) {
970 if pool.is_none() || logits.len() < PAR_MIN {
971 return softmax_inplace(logits);
972 }
973 let max_val = par_max(pool, logits, f32::NEG_INFINITY);
974 par_map(pool, logits, &move |v| (v - max_val).exp());
975 let sum: f32 = logits.iter().sum();
976 if sum > 0.0 {
977 par_map(pool, logits, &move |v| v / sum);
978 }
979}
980
981fn kth_largest(probs: &[f32], k: usize) -> f32 {
986 use std::cmp::Ordering;
987 let mut heap: Vec<f32> = Vec::with_capacity(k);
990 let desc = |a: f32, b: f32| b.partial_cmp(&a).unwrap_or(Ordering::Equal);
991 let sift_down = |h: &mut [f32], mut i: usize| {
992 let n = h.len();
993 loop {
994 let (l, r) = (2 * i + 1, 2 * i + 2);
995 let mut m = i;
996 if l < n && desc(h[l], h[m]) == Ordering::Greater {
998 m = l;
999 }
1000 if r < n && desc(h[r], h[m]) == Ordering::Greater {
1001 m = r;
1002 }
1003 if m == i {
1004 break;
1005 }
1006 h.swap(i, m);
1007 i = m;
1008 }
1009 };
1010 let sift_up = |h: &mut [f32], mut i: usize| {
1011 while i > 0 {
1012 let parent = (i - 1) / 2;
1013 if desc(h[i], h[parent]) == Ordering::Greater {
1014 h.swap(i, parent);
1015 i = parent;
1016 } else {
1017 break;
1018 }
1019 }
1020 };
1021 for &v in probs {
1022 if heap.len() < k {
1023 heap.push(v);
1024 let n = heap.len();
1025 sift_up(&mut heap, n - 1);
1026 } else if desc(v, heap[0]) == Ordering::Less {
1027 heap[0] = v;
1029 sift_down(&mut heap, 0);
1030 }
1031 }
1032 heap[0]
1033}
1034
1035fn apply_top_k_pool(pool: Option<&crate::pool::Pool>, probs: &mut [f32], k: usize) {
1039 if k == 0 || k >= probs.len() {
1040 return;
1041 }
1042 let threshold = kth_largest(probs, k);
1043 par_map(pool, probs, &move |p| if p < threshold { 0.0 } else { p });
1044}
1045
1046pub fn top1_prob_pool(
1051 pool: Option<&crate::pool::Pool>,
1052 scratch: &mut SamplerScratch,
1053 logits: &[f32],
1054 id: u32,
1055 temp: f32,
1056) -> f32 {
1057 let t = if temp > 1e-3 { temp } else { 1.0 };
1058 let max = par_max(pool, logits, f32::NEG_INFINITY);
1059 let mut e = std::mem::take(&mut scratch.topk);
1060 e.clear();
1061 e.extend_from_slice(logits);
1062 par_map(pool, &mut e, &move |v| ((v - max) / t).exp());
1063 let sum: f32 = e.iter().sum();
1064 let out = if sum > 0.0 {
1065 (((logits[id as usize] - max) / t).exp()) / sum
1066 } else {
1067 0.0
1068 };
1069 scratch.topk = e;
1070 out
1071}
1072
1073fn apply_repetition_penalty(
1074 logits: &mut [f32],
1075 past_tokens: &[u32],
1076 penalty: f32,
1077 scratch: &mut SamplerScratch,
1078) {
1079 let epoch = scratch.begin_seen(logits.len());
1080 for &tok in past_tokens {
1081 let idx = tok as usize;
1082 if idx < logits.len() && scratch.seen_epoch[idx] != epoch {
1083 scratch.seen_epoch[idx] = epoch;
1084 if logits[idx] > 0.0 {
1085 logits[idx] /= penalty;
1086 } else {
1087 logits[idx] *= penalty;
1088 }
1089 }
1090 }
1091}
1092
1093fn apply_top_k(probs: &mut [f32], k: usize, sel: &mut Vec<f32>) {
1098 if k == 0 || k >= probs.len() {
1099 return;
1100 }
1101 sel.clear();
1104 sel.extend_from_slice(probs);
1105 let (_, kth, _) = sel.select_nth_unstable_by(k - 1, |a, b| {
1107 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
1108 });
1109 let threshold = *kth;
1110 for p in probs.iter_mut() {
1111 if *p < threshold {
1112 *p = 0.0;
1113 }
1114 }
1115}
1116
1117fn apply_top_p(probs: &mut [f32], top_p: f32) {
1122 let mut indexed: Vec<(usize, f32)> = probs
1123 .iter()
1124 .copied()
1125 .enumerate()
1126 .filter(|&(_, p)| p > 0.0)
1127 .collect();
1128 indexed.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
1129
1130 let mut cumsum = 0.0f32;
1131 let mut cutoff_idx = indexed.len();
1132 for (i, &(_, prob)) in indexed.iter().enumerate() {
1133 cumsum += prob;
1134 if cumsum >= top_p {
1135 cutoff_idx = i + 1;
1136 break;
1137 }
1138 }
1139
1140 for &(i, _) in &indexed[cutoff_idx..] {
1142 probs[i] = 0.0;
1143 }
1144}
1145
1146fn categorical_sample(probs: &[f32], r: f32) -> u32 {
1148 let mut cumsum = 0.0f32;
1149 for (i, &p) in probs.iter().enumerate() {
1150 cumsum += p;
1151 if r < cumsum {
1152 return i as u32;
1153 }
1154 }
1155 probs.iter().rposition(|&p| p > 0.0).unwrap_or(0) as u32
1156}
1157
1158#[cfg(test)]
1159mod tests {
1160 #[test]
1165 fn argmax_lanes_match_the_scalar_one_ties_and_all() {
1166 let scalar = |v: &[f32]| -> u32 {
1169 v.iter()
1170 .enumerate()
1171 .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
1172 .map(|(i, _)| i as u32)
1173 .unwrap_or(0)
1174 };
1175 for n in 0..40usize {
1176 for seed in 0..8u64 {
1177 let mut r = super::SplitMix64::new(seed * 7 + n as u64);
1178 let v: Vec<f32> = (0..n).map(|_| ((r.next_u64() % 5) as f32) - 2.0).collect();
1181 assert_eq!(super::argmax(&v), scalar(&v), "n={n} seed={seed} {v:?}");
1182 }
1183 }
1184 let flat = vec![f32::NEG_INFINITY; 13];
1185 assert_eq!(super::argmax(&flat), scalar(&flat));
1186 }
1187
1188 use super::*;
1189
1190 #[test]
1191 fn test_argmax() {
1192 let logits = vec![0.1, 0.5, 0.3, 0.9, 0.2];
1193 assert_eq!(argmax(&logits), 3);
1194 }
1195
1196 #[test]
1197 fn test_greedy_sampling() {
1198 let logits = vec![1.0, 5.0, 2.0, 3.0];
1199 let config = SamplerConfig {
1200 temperature: 0.0,
1201 ..Default::default()
1202 };
1203 let mut rng = SplitMix64::new(1);
1204 assert_eq!(sample(&logits, &config, &[], &mut rng), 1);
1205 }
1206
1207 #[test]
1210 fn argmax_penalized_matches_chain_argmax() {
1211 let pool = crate::pool::Pool::new(3);
1212 let n = 40_000usize;
1213 for seed in 0..8u64 {
1214 let mut r = SplitMix64::new(seed + 3);
1215 let logits: Vec<f32> = (0..n)
1217 .map(|_| ((r.next_u64() % 41) as f32 - 20.0) / 4.0)
1218 .collect();
1219 let past: Vec<u32> = (0..2000)
1220 .map(|_| (r.next_u64() % n as u64) as u32)
1221 .collect();
1222 for cfg in [
1223 SamplerConfig {
1224 temperature: 0.0,
1225 repetition_penalty: 1.1,
1226 ..Default::default()
1227 },
1228 SamplerConfig {
1229 temperature: 0.0,
1230 repetition_penalty: 1.0,
1231 presence_penalty: 1.5,
1232 ..Default::default()
1233 },
1234 SamplerConfig {
1235 temperature: 0.0,
1236 repetition_penalty: 1.3,
1237 presence_penalty: 0.7,
1238 suppress_tokens: vec![5, 77, 3000],
1239 ..Default::default()
1240 },
1241 ] {
1242 let mut s1 = SamplerScratch::default();
1243 let mut probs = Vec::new();
1244 chain(&logits, &cfg, &past, &mut s1, None, &mut probs);
1245 let want = argmax(&probs);
1246 let mut s2 = SamplerScratch::default();
1247 let got_serial = argmax_penalized(&logits, &cfg, &past, &mut s2, None);
1248 let got_pool = argmax_penalized(&logits, &cfg, &past, &mut s2, Some(&pool));
1249 assert_eq!(want, got_serial, "serial, seed {seed} cfg {cfg:?}");
1250 assert_eq!(want, got_pool, "pool, seed {seed} cfg {cfg:?}");
1251 }
1252 }
1253 }
1254
1255 #[test]
1258 fn pool_chain_matches_serial_bit_for_bit() {
1259 let pool = crate::pool::Pool::new(3);
1260 let n = 40_000usize; for seed in 0..6u64 {
1262 let mut r = SplitMix64::new(seed + 11);
1263 let logits: Vec<f32> = (0..n)
1264 .map(|i| ((r.next_u64() % 2001) as f32 - 1000.0) / 90.0 + (i % 7) as f32 * 0.01)
1265 .collect();
1266 let past: Vec<u32> = (0..500).map(|_| (r.next_u64() % n as u64) as u32).collect();
1267 for cfg in [
1268 SamplerConfig::default(),
1269 SamplerConfig {
1270 temperature: 0.7,
1271 top_p: 0.8,
1272 top_k: 20,
1273 min_p: 0.0,
1274 presence_penalty: 1.5,
1275 repetition_penalty: 1.0,
1276 ..Default::default()
1277 },
1278 SamplerConfig {
1279 temperature: 1.0,
1280 top_p: 0.95,
1281 top_k: 20,
1282 min_p: 0.0,
1283 ..Default::default()
1284 },
1285 SamplerConfig {
1286 temperature: 1.3,
1287 top_p: 1.0,
1288 top_k: 0,
1289 min_p: 0.02,
1290 ..Default::default()
1291 },
1292 ] {
1293 let mut s1 = SamplerScratch::default();
1294 let mut s2 = SamplerScratch::default();
1295 for step in 0..5u64 {
1296 let mut r1 = SplitMix64::new(seed * 100 + step);
1297 let mut r2 = r1.clone();
1298 let a = sample_with_scratch(&logits, &cfg, &past, &mut r1, &mut s1);
1299 let b = sample_with_scratch_pool(
1300 &logits,
1301 &cfg,
1302 &past,
1303 &mut r2,
1304 &mut s2,
1305 Some(&pool),
1306 );
1307 assert_eq!(a, b, "seed {seed} step {step} cfg {cfg:?}");
1308 assert_eq!(s1.probs, s2.probs, "probs differ seed {seed} step {step}");
1310 }
1311 }
1312 }
1313 }
1314
1315 #[test]
1320 fn spec_accept_or_correct_reproduces_the_target() {
1321 let n = 40usize;
1322 let mk = |seed: u64, sharp: f32| -> Vec<f32> {
1323 let mut r = SplitMix64::new(seed);
1324 let mut v: Vec<f32> = (0..n)
1325 .map(|_| ((r.next_u64() % 1000) as f32 / 1000.0).powf(sharp))
1326 .collect();
1327 for i in 0..n {
1329 if (i * 7 + seed as usize) % 5 == 0 {
1330 v[i] = 0.0;
1331 }
1332 }
1333 let s: f32 = v.iter().sum();
1334 v.iter().map(|x| x / s).collect()
1335 };
1336 let p = mk(3, 3.0);
1337 for (qi, q) in [mk(3, 3.0), mk(11, 1.0), mk(29, 6.0)]
1338 .into_iter()
1339 .enumerate()
1340 {
1341 let mut rng = SplitMix64::new(77 + qi as u64);
1342 let mut counts = vec![0u64; n];
1343 let mut scratch = Vec::new();
1344 let trials = 400_000u64;
1345 for _ in 0..trials {
1346 let d = categorical_sample(&q, rng.next_f32());
1347 let t = match spec_accept_or_correct(&p, &q, d, &mut rng, &mut scratch, None) {
1348 None => d,
1349 Some(c) => c,
1350 };
1351 counts[t as usize] += 1;
1352 }
1353 let l1: f64 = (0..n)
1354 .map(|i| (counts[i] as f64 / trials as f64 - p[i] as f64).abs())
1355 .sum();
1356 eprintln!("spec q#{qi}: L1(empirical, p) = {l1:.4}");
1357 assert!(
1358 l1 < 0.01,
1359 "q#{qi}: empirical distribution drifted from p, L1 {l1}"
1360 );
1361 for i in 0..n {
1363 if p[i] == 0.0 {
1364 assert_eq!(counts[i], 0, "q#{qi}: token {i} outside p emitted");
1365 }
1366 }
1367 }
1368 }
1369
1370 #[test]
1374 fn sparse_chain_matches_the_dense_chain() {
1375 let pool = crate::pool::Pool::new(3);
1376 let n = 40_000usize; for seed in 0..5u64 {
1378 let mut r = SplitMix64::new(100 + seed);
1379 let logits: Vec<f32> = (0..n)
1380 .map(|_| ((r.next_u64() % 3000) as f32 - 1500.0) / 120.0)
1381 .collect();
1382 let past: Vec<u32> = (0..400).map(|_| (r.next_u64() % n as u64) as u32).collect();
1383 for cfg in [
1384 SamplerConfig {
1385 temperature: 0.7,
1386 top_p: 0.8,
1387 top_k: 20,
1388 min_p: 0.0,
1389 presence_penalty: 1.5,
1390 repetition_penalty: 1.0,
1391 ..Default::default()
1392 },
1393 SamplerConfig {
1394 temperature: 1.0,
1395 top_p: 0.95,
1396 top_k: 40,
1397 min_p: 0.05,
1398 presence_penalty: 0.0,
1399 repetition_penalty: 1.1,
1400 ..Default::default()
1401 },
1402 SamplerConfig {
1403 temperature: 0.6,
1404 top_p: 1.0,
1405 top_k: 3,
1406 min_p: 0.0,
1407 presence_penalty: 0.0,
1408 repetition_penalty: 1.0,
1409 suppress_tokens: vec![5, 6, 7],
1410 ..Default::default()
1411 },
1412 ] {
1413 assert!(sparse_ok(&cfg));
1414 let mut sd = SamplerScratch::default();
1415 let mut dense = Vec::new();
1416 distribution_into(&logits, &cfg, &past, &mut sd, None, &mut dense);
1417 for pl in [None, Some(&pool)] {
1418 let mut ss = SamplerScratch::default();
1419 let mut sp = Vec::new();
1420 let ok = sparse_distribution_into(&logits, &cfg, &past, &mut ss, pl, &mut sp);
1421 assert!(ok, "seed {seed} cfg {cfg:?}");
1422 let dense_nz: Vec<(u32, f32)> = dense
1423 .iter()
1424 .enumerate()
1425 .filter(|&(_, &v)| v > 0.0)
1426 .map(|(i, &v)| (i as u32, v))
1427 .collect();
1428 assert_eq!(
1429 dense_nz.len(),
1430 sp.len(),
1431 "seed {seed} pool {} cfg {cfg:?}: support {:?} vs {:?}",
1432 pl.is_some(),
1433 dense_nz,
1434 sp
1435 );
1436 for (a, b) in dense_nz.iter().zip(sp.iter()) {
1437 assert_eq!(a.0, b.0, "seed {seed} cfg {cfg:?}: ids differ");
1438 assert!(
1439 (a.1 - b.1).abs() <= 2e-5 * a.1.max(1e-3),
1440 "seed {seed} cfg {cfg:?}: prob {} vs {}",
1441 a.1,
1442 b.1
1443 );
1444 }
1445 let mut agree = 0usize;
1448 let trials = 400usize;
1449 for k in 0..trials as u64 {
1450 let mut r1 = SplitMix64::new(500 + k);
1451 let mut r2 = SplitMix64::new(500 + k);
1452 let a = categorical_sample(&dense, r1.next_f32());
1453 let b = draw_sparse(&sp, &mut r2);
1454 agree += (a == b) as usize;
1455 }
1456 assert!(
1457 agree >= trials - 2,
1458 "seed {seed} cfg {cfg:?}: agree {agree}/{trials}"
1459 );
1460 let mut r1 = SplitMix64::new(9);
1462 let mut r2 = SplitMix64::new(9);
1463 let a = categorical_sample(&dense, r1.next_f32());
1464 let mut s3 = SamplerScratch::default();
1465 let b = sample_with_scratch_pool(&logits, &cfg, &past, &mut r2, &mut s3, pl);
1466 assert_eq!(a, b, "seed {seed} cfg {cfg:?}: entry draw");
1467 }
1468 }
1469 }
1470 }
1471
1472 #[test]
1475 fn spec_accept_or_correct_sparse_reproduces_the_target() {
1476 let n = 40usize;
1477 let mk = |seed: u64, sharp: f32| -> Vec<(u32, f32)> {
1478 let mut r = SplitMix64::new(seed);
1479 let mut v: Vec<f32> = (0..n)
1480 .map(|_| ((r.next_u64() % 1000) as f32 / 1000.0).powf(sharp))
1481 .collect();
1482 for i in 0..n {
1483 if (i * 7 + seed as usize) % 5 == 0 {
1484 v[i] = 0.0;
1485 }
1486 }
1487 let s: f32 = v.iter().sum();
1488 v.iter()
1489 .enumerate()
1490 .filter(|&(_, &x)| x > 0.0)
1491 .map(|(i, &x)| (i as u32, x / s))
1492 .collect()
1493 };
1494 let p = mk(3, 3.0);
1495 for (qi, q) in [mk(3, 3.0), mk(11, 1.0), mk(29, 6.0)]
1496 .into_iter()
1497 .enumerate()
1498 {
1499 let mut rng = SplitMix64::new(77 + qi as u64);
1500 let mut counts = vec![0u64; n];
1501 let mut res = Vec::new();
1502 let trials = 400_000u64;
1503 for _ in 0..trials {
1504 let d = draw_sparse(&q, &mut rng);
1505 let t = match spec_accept_or_correct_sparse(&p, &q, d, &mut rng, &mut res) {
1506 None => d,
1507 Some(c) => c,
1508 };
1509 counts[t as usize] += 1;
1510 }
1511 let l1: f64 = (0..n)
1512 .map(|i| (counts[i] as f64 / trials as f64 - sparse_get(&p, i as u32) as f64).abs())
1513 .sum();
1514 eprintln!("sparse spec q#{qi}: L1(empirical, p) = {l1:.4}");
1515 assert!(l1 < 0.01, "q#{qi}: drifted, L1 {l1}");
1516 for i in 0..n {
1517 if sparse_get(&p, i as u32) == 0.0 {
1518 assert_eq!(counts[i], 0, "q#{qi}: token {i} outside p emitted");
1519 }
1520 }
1521 }
1522 }
1523
1524 #[test]
1527 fn distribution_matches_the_sampler_draw() {
1528 let n = 20_000usize;
1529 let mut r = SplitMix64::new(9);
1530 let logits: Vec<f32> = (0..n)
1531 .map(|_| ((r.next_u64() % 3000) as f32 - 1500.0) / 120.0)
1532 .collect();
1533 let past: Vec<u32> = (0..300).map(|_| (r.next_u64() % n as u64) as u32).collect();
1534 for cfg in [
1535 SamplerConfig::default(),
1536 SamplerConfig {
1537 temperature: 0.7,
1538 top_p: 0.8,
1539 top_k: 20,
1540 min_p: 0.0,
1541 presence_penalty: 1.5,
1542 repetition_penalty: 1.0,
1543 ..Default::default()
1544 },
1545 SamplerConfig {
1546 temperature: 0.0,
1547 repetition_penalty: 1.1,
1548 ..Default::default()
1549 },
1550 SamplerConfig {
1551 temperature: 0.0,
1552 repetition_penalty: 1.0,
1553 presence_penalty: 0.0,
1554 ..Default::default()
1555 },
1556 ] {
1557 let mut s1 = SamplerScratch::default();
1558 let mut s2 = SamplerScratch::default();
1559 for step in 0..4u64 {
1560 let mut r1 = SplitMix64::new(step + 1);
1561 let mut r2 = r1.clone();
1562 let a = sample_with_scratch(&logits, &cfg, &past, &mut r1, &mut s1);
1563 let mut dist = Vec::new();
1564 distribution_into(&logits, &cfg, &past, &mut s2, None, &mut dist);
1565 assert!(
1566 (dist.iter().sum::<f32>() - 1.0).abs() < 1e-3,
1567 "not normalized: {}",
1568 dist.iter().sum::<f32>()
1569 );
1570 let b = if cfg.temperature < 1e-6 {
1571 argmax(&dist)
1572 } else {
1573 draw(&dist, &mut r2)
1574 };
1575 assert_eq!(a, b, "cfg {cfg:?} step {step}");
1576 }
1577 }
1578 }
1579
1580 #[test]
1583 fn kth_largest_equals_select_nth() {
1584 for seed in 0..20u64 {
1585 let mut r = SplitMix64::new(seed);
1586 let n = 50 + (r.next_u64() % 3000) as usize;
1587 let v: Vec<f32> = (0..n)
1588 .map(|_| {
1589 if r.next_u64() % 3 == 0 {
1590 0.0
1591 } else {
1592 (r.next_u64() % 97) as f32 / 97.0
1593 }
1594 })
1595 .collect();
1596 for k in [1usize, 2, 5, 20, 40, n / 2, n - 1] {
1597 let mut sel = v.clone();
1598 let (_, kth, _) = sel.select_nth_unstable_by(k - 1, |a, b| {
1599 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
1600 });
1601 assert_eq!(kth_largest(&v, k), *kth, "seed {seed} k {k}");
1602 }
1603 }
1604 }
1605
1606 #[test]
1608 fn top1_prob_pool_matches_serial() {
1609 let pool = crate::pool::Pool::new(2);
1610 let n = 30_000usize;
1611 let mut r = SplitMix64::new(5);
1612 let logits: Vec<f32> = (0..n)
1613 .map(|_| ((r.next_u64() % 1000) as f32) / 37.0)
1614 .collect();
1615 let serial = |id: u32, temp: f32| -> f32 {
1616 let t = if temp > 1e-3 { temp } else { 1.0 };
1617 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
1618 let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
1619 (((logits[id as usize] - max) / t).exp()) / sum
1620 };
1621 let mut sc = SamplerScratch::default();
1622 for (id, t) in [(3u32, 1.0f32), (777, 0.7), (29_999, 2.0), (12, 0.0)] {
1623 let a = serial(id, t);
1624 let b = top1_prob_pool(Some(&pool), &mut sc, &logits, id, t);
1625 assert_eq!(a.to_bits(), b.to_bits(), "id {id} t {t}: {a} vs {b}");
1626 }
1627 }
1628
1629 #[test]
1630 fn test_softmax() {
1631 let mut logits = vec![1.0, 2.0, 3.0];
1632 softmax_inplace(&mut logits);
1633 let sum: f32 = logits.iter().sum();
1634 assert!((sum - 1.0).abs() < 1e-5);
1635 assert!(logits[2] > logits[1] && logits[1] > logits[0]);
1636 }
1637
1638 #[test]
1639 fn test_repetition_penalty() {
1640 let mut logits = vec![1.0, 2.0, 3.0, 4.0];
1641 let mut scratch = SamplerScratch::default();
1642 apply_repetition_penalty(&mut logits, &[1, 3], 2.0, &mut scratch);
1643 assert_eq!(logits, vec![1.0, 1.0, 3.0, 2.0]);
1644 }
1645
1646 #[test]
1647 fn repetition_penalty_applies_once_per_unique_token() {
1648 let mut logits = vec![1.0, 4.0, -6.0];
1649 let mut scratch = SamplerScratch::default();
1650 apply_repetition_penalty(&mut logits, &[1, 1, 2, 1, 2], 2.0, &mut scratch);
1651 assert_eq!(logits, vec![1.0, 2.0, -12.0]);
1652 }
1653
1654 #[test]
1655 fn top_k_keeps_exactly_k() {
1656 let mut probs = vec![0.1, 0.4, 0.05, 0.3, 0.15];
1657 apply_top_k(&mut probs, 2, &mut Vec::new());
1658 let kept = probs.iter().filter(|&&p| p > 0.0).count();
1659 assert_eq!(kept, 2, "top-k must keep exactly k (was k+1 in v1)");
1660 assert!(probs[1] > 0.0 && probs[3] > 0.0);
1661 }
1662
1663 #[test]
1664 fn rng_reaches_full_cdf() {
1665 let probs = vec![0.25f32; 4];
1668 let mut rng = SplitMix64::new(42);
1669 let mut hits = [0usize; 4];
1670 for _ in 0..4000 {
1671 let i = categorical_sample(&probs, rng.next_f32()) as usize;
1672 hits[i] += 1;
1673 }
1674 for (i, &h) in hits.iter().enumerate() {
1675 assert!(h > 700, "index {i} sampled only {h}/4000 — biased RNG");
1676 }
1677 }
1678
1679 #[test]
1680 fn same_seed_same_sequence() {
1681 let logits: Vec<f32> = (0..32).map(|i| (i as f32 * 0.37).sin()).collect();
1682 let config = SamplerConfig {
1683 temperature: 1.0,
1684 seed: Some(7),
1685 ..Default::default()
1686 };
1687 let run = |seed: u64| -> Vec<u32> {
1688 let mut rng = SplitMix64::new(seed);
1689 (0..16)
1690 .map(|_| sample(&logits, &config, &[], &mut rng))
1691 .collect()
1692 };
1693 assert_eq!(run(7), run(7), "same seed must reproduce");
1694 assert_ne!(run(7), run(8), "different seed must differ");
1695 }
1696}