Skip to main content

cortiq_engine/
sampler.rs

1//! Token sampling — temperature, top-p, top-k, min-p, repetition penalty.
2//!
3//! Randomness comes from an explicit SplitMix64 PRNG carried by the
4//! caller: reproducible with a seed, unbiased across the whole CDF
5//! (the v1 `subsec_nanos` source could never pick past ~23% of it).
6
7use serde::{Deserialize, Serialize};
8
9/// SplitMix64 — tiny, fast, statistically solid for sampling.
10#[derive(Debug, Clone)]
11pub struct SplitMix64 {
12    state: u64,
13}
14
15/// Reusable per-pipeline sampling workspace. The epoch table lets the
16/// repetition penalty visit each token id once without allocating a HashSet
17/// or clearing a vocab-sized boolean vector on every decode step.
18#[derive(Debug, Default)]
19pub struct SamplerScratch {
20    seen_epoch: Vec<u32>,
21    epoch: u32,
22    /// Distinct-token set for the presence penalty; reused per token.
23    presence_seen: std::collections::HashSet<u32>,
24    /// The working copy of the logits. At a 129k vocab that is half a
25    /// megabyte allocated, filled and dropped per token; the struct that
26    /// exists to hold scratch may as well hold this one too.
27    probs: Vec<f32>,
28    /// The SECOND whole-vocab copy — top-k's partition buffer. Qwen3.6's
29    /// vocab is 248320, so this was another megabyte allocated, filled
30    /// and dropped per token, on the same hot path and for the same
31    /// reason. Same fix.
32    topk: Vec<f32>,
33    /// The sparse chain's candidate list and its per-grain partials.
34    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    /// Seed from OS entropy (address-space + time mix) when none given.
60    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        // SAFETY: pointer came from Box::into_raw just above.
66        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    /// Uniform f32 in [0, 1).
80    #[inline]
81    pub fn next_f32(&mut self) -> f32 {
82        (self.next_u64() >> 40) as f32 / (1u64 << 24) as f32
83    }
84}
85
86/// Sampling configuration.
87#[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    /// Flat additive penalty on every token that has appeared at least
95    /// once (OpenAI-style presence penalty). Qwen3.8's instruct sampling
96    /// asks for 1.5 here — the multiplicative repetition_penalty is a
97    /// different curve and cannot stand in for it.
98    #[serde(default)]
99    pub presence_penalty: f32,
100    /// Fixed seed for reproducible generation (None = entropy).
101    #[serde(default)]
102    pub seed: Option<u64>,
103    /// Token IDs to suppress (force logit to -inf).
104    #[serde(default)]
105    pub suppress_tokens: Vec<u32>,
106    /// How many of the most recent ids the repetition / presence
107    /// penalties look at. 0 = the whole sequence (the historical
108    /// behaviour). A natively bounded model (Embryo-O1 anchor) never
109    /// scans unbounded history: the pipeline substitutes
110    /// [`BOUNDED_PENALTY_WINDOW`] there when this is 0.
111    #[serde(default)]
112    pub penalty_window: usize,
113}
114
115/// Penalty window a bounded-state model falls back to when
116/// `penalty_window == 0` (`CMF_PENALTY_WINDOW` overrides).
117pub const BOUNDED_PENALTY_WINDOW: usize = 128;
118
119impl SamplerConfig {
120    /// The slice of `past` the penalties may scan: the last
121    /// `penalty_window` ids, or all of them when the window is 0 and the
122    /// model is not bounded-native.
123    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
156/// Sample next token from logits. Chain order is fixed:
157/// rep-penalty → temperature → softmax → min-p → top-k → top-p → sample.
158pub 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
168/// Sampling entry point for hot decode loops with reusable scratch storage.
169pub 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
179/// The same chain with the whole-vocab passes spread over the CPU pool.
180///
181/// WHY: at Qwen3.8's 248 320-entry vocab the serial sampler is ~14 passes
182/// over a megabyte plus 248k `exp` and a `select_nth` over a second copy
183/// — measured on the RTX 5090 pod as the gap between `bench --core` and
184/// the production loop (50.9 against 46.8 tok/s with penalty+confidence
185/// alone; the temperature path pays the softmax and the partition on top).
186/// The GPU graph owns the token, so during decode the pool sits idle —
187/// this is free work.
188///
189/// WHAT IS PRESERVED: every value. The parallel passes are elementwise
190/// (each output depends on its own input only), the sums that feed
191/// divisions stay sequential in index order, and top-k's threshold is
192/// the k-th largest VALUE — the same number `select_nth` returned. The
193/// sampled token is bit-identical to the serial chain for the same seed.
194pub 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    // Borrowed from the scratch and handed back at the single exit: at a
210    // 129k vocab this copy is half a megabyte allocated, filled and dropped
211    // per token, and the struct that exists to hold scratch may as well
212    // hold it. Every early return goes through `done` so the buffer never
213    // leaks back to the allocator.
214    if config.temperature < 1e-6 {
215        // greedy over the penalized logits, no working copy
216        return argmax_penalized(logits, config, past_tokens, scratch, pool);
217    }
218    if sparse_ok(config) {
219        // The sparse chain: same distribution, a tenth of the passes.
220        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        // Everything filtered out — fall back to greedy over original logits.
240        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
247/// Greedy over the PENALIZED logits without the working copy: one pass
248/// that applies the repetition / presence penalty and the suppress list
249/// on the fly (a membership table over the vocab, built from the past
250/// tokens) and keeps the argmax with the same tie rule as `argmax`
251/// (highest index among equal maxima). Bit-identical to
252/// `chain` + `argmax` for temperature 0 — the values compared are the
253/// same expressions — and it is what a greedy decode with penalties pays
254/// per token, and what a speculative round pays per draft and per
255/// verified row (nine such passes a round at k=4).
256pub 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    // Membership: seen_epoch[i] == epoch for past tokens (each once).
271    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    // Suppressed ids get -inf; a second, rarer set — keep it exact.
281    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    // The suppress list is scanned per element above; keep that path
303    // serial and rare. The common (no suppress) case runs the pool.
304    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            // Same rule across chunks: the max value, and among equal
325            // maxima the HIGHEST index — chunk order does not matter once
326            // sorted by index.
327            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
342/// The chain up to the draw, into `probs`: penalties, temperature,
343/// softmax, min-p, top-k, top-p, renormalize. Returns false when the
344/// filters left nothing (the caller's greedy fallback); for a greedy
345/// config it stops after the penalties (`probs` then holds penalized
346/// logits, and argmax over them is the token).
347fn 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
392/// The chain's first stage — suppress, repetition and presence penalties
393/// — in place on a working copy of the logits. Every penalty only LOWERS
394/// a logit, which is what lets the sparse chain below bound its
395/// candidates.
396fn 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        // Once per DISTINCT seen token — presence, not frequency. The
412        // scratch set the repetition penalty uses would serve, but it is
413        // only built on its own branch; a local pass stays correct when
414        // rep-penalty is 1.0 (Qwen3.8's recommended pairing).
415        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
433/// Largest top-k the sparse chain serves. Past this the dense chain is
434/// the better tool anyway.
435pub const SPARSE_TOPK_MAX: usize = 256;
436
437/// Whether `config` can go through the sparse chain: a real temperature
438/// and a top-k in 1..=256. Qwen's recommended instruct settings
439/// (0.7 / top-p 0.8 / top-k 20 / presence 1.5) do.
440pub 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
444/// A distribution over at most `SPARSE_TOPK_MAX` tokens: `(id, prob)`
445/// sorted by id, probs summing to 1.
446pub type Sparse = Vec<(u32, f32)>;
447
448/// The sampler chain's distribution as a SPARSE list — the same
449/// distribution `chain` builds over the whole vocab, for configs with a
450/// top-k, at a fraction of the cost. The dense chain copies the vocab,
451/// exponentiates it, selects, filters and normalises it — six or seven
452/// passes over 248k floats — and every one of them past the selection
453/// touches only the k survivors. Here: penalties on a copy ONLY when
454/// there are penalties, one pooled pass that selects the top-k penalized
455/// logits, one pooled pass for the vocab-wide softmax denominator (top-p
456/// is defined against the FULL normalisation, so the denominator must
457/// see every token), and the rest over k entries.
458///
459/// Why it is the same distribution: softmax → min-p → top-k → top-p →
460/// renormalise, in the dense order. Softmax is monotone in the logit, so
461/// the top-k SET is the top-k of the penalized logits; min-p drops
462/// tokens below `max_prob·min_p`, i.e. below `exp((l − l_max)/T) <
463/// min_p` — every token outside the top-k that fails it is dropped
464/// either way, and the ones inside are tested exactly as the dense chain
465/// tests them; top-p cuts on the cumulative FULL-normalised probs of the
466/// survivors sorted descending, computed here from the same terms. What
467/// differs is floating-point: the denominator's summation order and the
468/// exp of `(l − l_max)/T` against `l/T − max(l/T)`. Ties at the k-th
469/// place resolve by lower id here where the dense chain keeps them all.
470///
471/// Returns false when the chain filtered everything (the dense chain's
472/// `!normalized`) — the caller falls back to greedy the same way.
473pub 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    // `cand` is by value descending, ties by id — the top-k AND every tie
503    // at the k-th place, as the dense chain keeps them.
504    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        // e_i = exp((l_i − l_max)/T); min-p against e_i < min_p (max_prob
507        // is e = 1 over the same denominator); probs e_i / sum_all.
508        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        // renormalise over the survivors and order by id
530        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
549/// The k largest of `src` as `(id, value)` sorted by value descending,
550/// ties by id ascending — INCLUDING every value tied with the k-th, which
551/// is what the dense chain's `p < threshold → 0` keeps. Two pooled passes:
552/// a per-grain k-slot selection merged in grain order for the k-th value,
553/// then a gather of everything at or above it. Nothing here depends on
554/// scheduling, so a seed reproduces.
555fn 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    // sorted insertion into a fixed k-slot list: the compare against the
564    // current k-th is what nearly every element pays, and nothing else.
565    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    // A degenerate row (thousands tied at the k-th place) is capped: the
588    // dense chain would keep them all; nobody samples such a row on
589    // purpose.
590    let cap = k * 4 + 64;
591    // gather everything ≥ kth into `slot`, at most `cap` entries
592    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                // SAFETY: grain g is written by exactly one range (start =
614                // g·grain) and `parts` outlives the joined dispatch.
615                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; // -inf ties are the filtered-out set; keep the k
628            }
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
660/// Σ exp((l − lmax)/t) over the vocab, per-grain partials summed in
661/// grain order (deterministic across runs, so a seed reproduces).
662fn 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                // SAFETY: one writer per grain slot; joined before read.
686                unsafe { *pp.at(s / grain) = term(s, e) };
687            });
688            parts.iter().sum()
689        }
690        _ => term(0, src.len()),
691    }
692}
693
694/// Draw from a sparse distribution: inverse CDF in id order — the same
695/// walk the dense `categorical_sample` makes over the vocab, so a seed
696/// lands on the same token when the survivor set and probs agree.
697pub 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
715/// `spec_accept_or_correct` over sparse distributions: accept the draft
716/// `d` with min(1, p[d]/q[d]); on rejection draw the correction from the
717/// residual max(0, p − q) over p's support (q's support outside p
718/// contributes nothing to the residual). Empty residual → a draw from p.
719pub 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
749/// The distribution the sampler would draw from — the whole chain minus
750/// the draw — as a normalized vector over the vocab, in `out`. Greedy
751/// configs (and the filtered-out fallback) come back as a one-hot, so a
752/// caller can treat every configuration uniformly. This is what
753/// speculative SAMPLING needs from both the draft head and the verify:
754/// accept-with-min(1, p/q), correct from max(0, p − q).
755pub 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
793/// Draw from a normalized distribution with the caller's RNG.
794pub fn draw(probs: &[f32], rng: &mut SplitMix64) -> u32 {
795    categorical_sample(probs, rng.next_f32())
796}
797
798/// One step of speculative sampling (Leviathan et al. / Chen et al.):
799/// the draft `d` was drawn from `q`; the target distribution at the same
800/// position is `p`. Returns `None` when `d` is accepted (with probability
801/// min(1, p[d]/q[d])) and `Some(c)` when it is rejected, `c` drawn from
802/// the residual max(0, p − q) renormalized — which is exactly what makes
803/// the emitted token stream distributed as `p`, draft or no draft. When
804/// the residual is empty (p ⊆ q, so p == q on the support) the correction
805/// falls back to a draw from `p` itself. `scratch` holds the residual;
806/// the pool spreads the vocab-wide pass.
807pub 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    // accept iff r < min(1, pd/qd)  ⇔  r·qd < pd (qd > 0 since d was drawn from q)
822    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    // residual = max(0, p − q), elementwise over the pool
829    {
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            // SAFETY: disjoint ranges; q outlives the (joined) dispatch.
834            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
857/// Greedy: index of the maximum value.
858///
859/// Four running maxima instead of one: the scalar `max_by` carried a loop
860/// dependency through the comparison, which at a 129k vocab is a tenth of a
861/// millisecond of pure serial work per token.
862///
863/// Ties resolve to the HIGHEST index — not an arbitrary choice, it is what
864/// `Iterator::max_by` does (it keeps the last of several equal maxima) and
865/// therefore what this has always returned. `explain`'s preview compares
866/// its own argmax against what greedy emits, and that test is what catches
867/// the flip.
868pub 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
919/// Below this length the pool's dispatch costs more than the pass.
920const PAR_MIN: usize = 1 << 14;
921
922/// Elementwise `buf[i] = f(buf[i])` over the pool (serial without one, or
923/// for short buffers). Each output depends on its own input alone, so the
924/// chunking cannot change a single bit.
925fn 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                    // SAFETY: ranges from run_rows are disjoint and the
932                    // buffer outlives the (joined) dispatch.
933                    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
948/// `fold(init, f32::max)` over the pool. Max is order-free on non-NaN
949/// input, so per-chunk maxima combined give the serial fold's answer.
950fn 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
965/// `softmax_inplace` with the exp and the normalisation spread over the
966/// pool. The max is order-free, the exp is elementwise, and the SUM stays
967/// a sequential index-order fold — exactly the serial loop's accumulation
968/// — so the probabilities are bit-identical.
969fn 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
981/// The k-th largest value of `probs` (k ≥ 1, k ≤ len), by the same
982/// descending `partial_cmp` order `select_nth_unstable_by` used — one
983/// streaming pass with a k-slot min-heap instead of a whole-vocab copy
984/// and partition. Values, not indices, so ties give the same threshold.
985fn kth_largest(probs: &[f32], k: usize) -> f32 {
986    use std::cmp::Ordering;
987    // Min-heap on the k largest seen so far: `heap[0]` is the smallest of
988    // them, i.e. the running k-th largest.
989    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            // child "smaller" in the descending order = later in it
997            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            // v is larger than the current k-th largest: replace it.
1028            heap[0] = v;
1029            sift_down(&mut heap, 0);
1030        }
1031    }
1032    heap[0]
1033}
1034
1035/// `apply_top_k` without the second vocab-sized copy: the threshold is
1036/// the k-th largest value from one streaming pass, the zeroing is an
1037/// elementwise pass over the pool. Same kept set, same values.
1038fn 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
1046/// Top-1 probability of `id` under a softmax at temperature `temp` — the
1047/// per-token confidence — with the exp pass over the pool and the sum
1048/// sequential in index order (bit-identical to the serial fold). Uses
1049/// the scratch's partition buffer, idle now that top-k streams.
1050pub 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
1093/// Keep the k highest-probability tokens (plus exact ties at the
1094/// threshold), zero the rest. Selection, not a full vocab sort — the
1095/// old double `sort_by` over ~150k probs was ~1ms of pure per-token
1096/// overhead (roadmap §3 P0).
1097fn apply_top_k(probs: &mut [f32], k: usize, sel: &mut Vec<f32>) {
1098    if k == 0 || k >= probs.len() {
1099        return;
1100    }
1101    // `select_nth_unstable` permutes, so the partition needs its own
1102    // buffer — but not a FRESH one each token.
1103    sel.clear();
1104    sel.extend_from_slice(probs);
1105    // k-th largest = (k-1)-th index in a descending partition.
1106    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
1117/// Nucleus: keep the smallest prefix of tokens whose cumulative
1118/// probability reaches top_p. Only surviving (non-zero) candidates are
1119/// sorted — after top-k that is ≤ k elements, not the whole vocab; the
1120/// kept set is marked in-place instead of a per-token HashSet.
1121fn 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    // Zero the dropped tail directly — indices, not membership tests.
1141    for &(i, _) in &indexed[cutoff_idx..] {
1142        probs[i] = 0.0;
1143    }
1144}
1145
1146/// Inverse-CDF sampling with an externally supplied uniform r ∈ [0, 1).
1147fn 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    /// The four-lane argmax against the obvious scalar one, including the
1161    /// tie rule: `max_by` keeps the LAST of several equal maxima, and a
1162    /// lane split that quietly picked the first would move greedy output on
1163    /// any model with two equally-likely tokens.
1164    #[test]
1165    fn argmax_lanes_match_the_scalar_one_ties_and_all() {
1166        // The reference IS the old implementation, `max_by` and all — the
1167        // point is that nothing observable changed, tie rule included.
1168        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                // Quantized to few distinct values on purpose: ties are the
1179                // case the lanes can get wrong and random floats never hit.
1180                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    /// `argmax_penalized` must equal the copy-and-penalize chain's argmax
1208    /// — same values, same tie rule — with and without the pool.
1209    #[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            // coarse values so ties happen
1216            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    /// The pool chain must sample the SAME token as the serial one on the
1256    /// same seed — that is the whole contract of the parallel passes.
1257    #[test]
1258    fn pool_chain_matches_serial_bit_for_bit() {
1259        let pool = crate::pool::Pool::new(3);
1260        let n = 40_000usize; // above PAR_MIN so the pool arm is exercised
1261        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                    // and the working copies agree value for value
1309                    assert_eq!(s1.probs, s2.probs, "probs differ seed {seed} step {step}");
1310                }
1311            }
1312        }
1313    }
1314
1315    /// Speculative sampling must reproduce the TARGET distribution however
1316    /// good or bad the draft is: draw d ~ q, accept with min(1, p/q), else
1317    /// correct from max(0, p − q). Empirical law over many trials against
1318    /// p itself, for a sharp draft, a flat draft and a wrong draft.
1319    #[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            // a few exact zeros, like a top-k'd distribution
1328            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            // and nothing outside p's support was ever emitted
1362            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    /// The sparse chain is the dense chain: same survivor set, same
1371    /// probabilities (to fp), serial and pooled, with and without
1372    /// penalties, min-p, top-p — and a seed draws the same token.
1373    #[test]
1374    fn sparse_chain_matches_the_dense_chain() {
1375        let pool = crate::pool::Pool::new(3);
1376        let n = 40_000usize; // above PAR_MIN: the pooled arms run
1377        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                    // the seed lands on the same token (both walk the ids
1446                    // in order); allow a boundary rounding miss or two
1447                    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                    // and the public entry uses it
1461                    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    /// The sparse accept/correct emits the target distribution, like its
1473    /// dense twin — the same 400k-trial law test over sparse p and q.
1474    #[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    /// `distribution_into` is the sampler's own chain: drawing from it
1525    /// with the same uniform lands on the same token as `sample_*`.
1526    #[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    /// The streaming k-th largest against `select_nth_unstable_by` — the
1581    /// value the old partition returned, ties and zeros included.
1582    #[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    /// The pooled confidence equals the serial formula, bit for bit.
1607    #[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        // v1 bug: r < 0.233 always, so the CDF tail was unreachable.
1666        // With uniform probs the LAST index must be sampled sometimes.
1667        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}