#[derive(Clone, Debug)]
pub struct SamplerConfig {
pub temperature: f32, pub top_k: usize, pub top_p: f32, pub min_p: f32, pub penalty_last_n: usize, pub penalty_repeat: f32, pub penalty_freq: f32, pub penalty_present: f32, pub seed: u64,
}
impl Default for SamplerConfig {
fn default() -> Self {
SamplerConfig {
temperature: 0.0,
top_k: 0,
top_p: 1.0,
min_p: 0.0,
penalty_last_n: 0,
penalty_repeat: 1.0,
penalty_freq: 0.0,
penalty_present: 0.0,
seed: 0,
}
}
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub struct SamplerIdentity {
greedy: bool,
temp_bits: u32,
seed: u64,
top_k: usize,
top_p_bits: u32,
min_p_bits: u32,
penalty_last_n: usize,
penalty_repeat_bits: u32,
penalty_freq_bits: u32,
penalty_present_bits: u32,
}
impl SamplerIdentity {
pub fn of(cfg: &SamplerConfig) -> Self {
let greedy = cfg.temperature <= 0.0;
let pen_on = cfg.penalty_last_n > 0
&& (cfg.penalty_repeat != 1.0 || cfg.penalty_freq != 0.0 || cfg.penalty_present != 0.0);
SamplerIdentity {
greedy,
temp_bits: if greedy { 0.0f32 } else { cfg.temperature }.to_bits(),
seed: cfg.seed,
top_k: cfg.top_k,
top_p_bits: if cfg.top_p >= 1.0 { 1.0f32 } else { cfg.top_p }.to_bits(),
min_p_bits: if cfg.min_p <= 0.0 { 0.0f32 } else { cfg.min_p }.to_bits(),
penalty_last_n: if pen_on { cfg.penalty_last_n } else { 0 },
penalty_repeat_bits: if pen_on { cfg.penalty_repeat } else { 1.0f32 }.to_bits(),
penalty_freq_bits: if pen_on { cfg.penalty_freq } else { 0.0f32 }.to_bits(),
penalty_present_bits: if pen_on { cfg.penalty_present } else { 0.0f32 }.to_bits(),
}
}
pub fn seed(&self) -> u64 {
self.seed
}
pub fn mismatch(&self, parked: &Self) -> Option<&'static str> {
if self.greedy != parked.greedy {
return Some("regime");
}
if self.temp_bits != parked.temp_bits {
return Some("temperature");
}
if self.top_k != parked.top_k {
return Some("top_k");
}
if self.top_p_bits != parked.top_p_bits {
return Some("top_p");
}
if self.min_p_bits != parked.min_p_bits {
return Some("min_p");
}
if self.penalty_last_n != parked.penalty_last_n {
return Some("penalty_last_n");
}
if self.penalty_repeat_bits != parked.penalty_repeat_bits {
return Some("penalty_repeat");
}
if self.penalty_freq_bits != parked.penalty_freq_bits {
return Some("penalty_freq");
}
if self.penalty_present_bits != parked.penalty_present_bits {
return Some("penalty_present");
}
None
}
pub fn legacy_admits(&self, _parked: &Self) -> bool {
true
}
}
pub struct Sampler {
cfg: SamplerConfig,
rng: SplitMix64,
history: Vec<u32>, }
impl Sampler {
pub fn new(cfg: SamplerConfig) -> Self {
let rng = SplitMix64::new(cfg.seed);
Sampler {
cfg,
rng,
history: Vec::new(),
}
}
pub fn is_greedy(&self) -> bool {
self.cfg.temperature <= 0.0
}
pub fn is_spec_sampling(&self) -> bool {
self.cfg.temperature > 0.0
&& self.cfg.penalty_repeat == 1.0
&& self.cfg.penalty_freq == 0.0
&& self.cfg.penalty_present == 0.0
&& self.cfg.top_k == 0
&& self.cfg.top_p >= 1.0
&& self.cfg.min_p <= 0.0
}
pub fn top_k(&self) -> usize {
self.cfg.top_k
}
pub fn penalty_last_n(&self) -> usize {
self.cfg.penalty_last_n
}
pub fn penalty_repeat(&self) -> f32 {
self.cfg.penalty_repeat
}
pub fn penalty_freq(&self) -> f32 {
self.cfg.penalty_freq
}
pub fn penalty_present(&self) -> f32 {
self.cfg.penalty_present
}
pub fn top_p(&self) -> f32 {
self.cfg.top_p
}
pub fn min_p(&self) -> f32 {
self.cfg.min_p
}
pub fn temperature(&self) -> f32 {
self.cfg.temperature
}
pub fn seed(&self) -> u64 {
self.cfg.seed
}
pub fn identity(&self) -> SamplerIdentity {
SamplerIdentity::of(&self.cfg)
}
pub fn accept(&mut self, token: u32) {
self.history.push(token);
}
pub fn sample(&mut self, logits: &[f32]) -> u32 {
if self.is_greedy()
&& self.cfg.penalty_repeat == 1.0
&& self.cfg.penalty_freq == 0.0
&& self.cfg.penalty_present == 0.0
{
return argmax_u32(logits);
}
let mut cand: Vec<(u32, f32)> = logits
.iter()
.enumerate()
.map(|(i, &l)| (i as u32, l))
.collect();
self.apply_penalties(&mut cand);
if self.is_greedy() {
let mut best = cand[0];
for &c in &cand[1..] {
if c.1 > best.1 {
best = c;
}
}
return best.0;
}
if self.cfg.temperature > 0.0 && self.cfg.temperature != 1.0 {
let inv = 1.0 / self.cfg.temperature;
for c in cand.iter_mut() {
c.1 *= inv;
}
}
if self.cfg.top_k > 0 && self.cfg.top_k < cand.len() {
cand.sort_unstable_by(|a, b| b.1.total_cmp(&a.1));
cand.truncate(self.cfg.top_k);
}
softmax_inplace(&mut cand);
if self.cfg.top_p < 1.0 {
cand.sort_unstable_by(|a, b| b.1.total_cmp(&a.1));
let mut cum = 0.0f32;
let mut keep = 0usize;
for (i, c) in cand.iter().enumerate() {
cum += c.1;
keep = i + 1;
if cum >= self.cfg.top_p {
break;
}
}
cand.truncate(keep.max(1));
}
if self.cfg.min_p > 0.0 {
let maxp = cand.iter().map(|c| c.1).fold(0.0f32, f32::max);
let thresh = self.cfg.min_p * maxp;
cand.retain(|c| c.1 >= thresh);
if cand.is_empty() {
return argmax_u32(logits);
} }
let sum: f32 = cand.iter().map(|c| c.1).sum();
let r = self.rng.next_f32() * sum;
let mut acc = 0.0f32;
for c in &cand {
acc += c.1;
if acc >= r {
return c.0;
}
}
cand.last().unwrap().0
}
fn apply_penalties(&self, cand: &mut [(u32, f32)]) {
let n = self.cfg.penalty_last_n;
if n == 0 {
return;
}
if self.cfg.penalty_repeat == 1.0
&& self.cfg.penalty_freq == 0.0
&& self.cfg.penalty_present == 0.0
{
return;
}
let start = self.history.len().saturating_sub(n);
let window = &self.history[start..];
if window.is_empty() {
return;
}
use std::collections::HashMap;
let mut counts: HashMap<u32, i32> = HashMap::new();
for &t in window {
*counts.entry(t).or_insert(0) += 1;
}
for c in cand.iter_mut() {
if let Some(&cnt) = counts.get(&c.0) {
if self.cfg.penalty_repeat != 1.0 {
if c.1 > 0.0 {
c.1 /= self.cfg.penalty_repeat;
} else {
c.1 *= self.cfg.penalty_repeat;
}
}
c.1 -= cnt as f32 * self.cfg.penalty_freq;
c.1 -= self.cfg.penalty_present; }
}
}
}
fn argmax_u32(logits: &[f32]) -> u32 {
let mut best = 0u32;
let mut bv = f32::NEG_INFINITY;
for (i, &v) in logits.iter().enumerate() {
if v > bv {
bv = v;
best = i as u32;
}
}
best
}
fn softmax_inplace(cand: &mut [(u32, f32)]) {
let maxl = cand.iter().map(|c| c.1).fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
for c in cand.iter_mut() {
let e = (c.1 - maxl).exp();
c.1 = e;
sum += e;
}
let inv = if sum > 0.0 { 1.0 / sum } else { 0.0 };
for c in cand.iter_mut() {
c.1 *= inv;
}
}
struct SplitMix64 {
state: u64,
}
impl SplitMix64 {
fn new(seed: u64) -> Self {
SplitMix64 {
state: seed.wrapping_add(0x9E3779B97F4A7C15),
}
}
fn next_u64(&mut self) -> u64 {
self.state = self.state.wrapping_add(0x9E3779B97F4A7C15);
let mut z = self.state;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
z ^ (z >> 31)
}
fn next_f32(&mut self) -> f32 {
((self.next_u64() >> 40) as f32) / (1u32 << 24) as f32
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn greedy_is_argmax() {
let mut s = Sampler::new(SamplerConfig::default()); let logits = vec![0.1, 5.0, 2.0, -1.0];
assert_eq!(s.sample(&logits), 1);
}
#[test]
fn temp_sampling_deterministic_with_seed() {
let cfg = SamplerConfig {
temperature: 1.0,
seed: 42,
..Default::default()
};
let logits = vec![1.0, 2.0, 3.0, 0.5];
let a = Sampler::new(cfg.clone()).sample(&logits);
let b = Sampler::new(cfg).sample(&logits);
assert_eq!(a, b, "same seed must reproduce the draw");
assert!(a < 4);
}
#[test]
fn top_k_one_is_argmax() {
let cfg = SamplerConfig {
temperature: 1.0,
top_k: 1,
seed: 7,
..Default::default()
};
let logits = vec![0.1, 5.0, 2.0, -1.0];
assert_eq!(
Sampler::new(cfg).sample(&logits),
1,
"top_k=1 collapses to argmax"
);
}
#[test]
fn min_p_keeps_only_high_prob() {
let cfg = SamplerConfig {
temperature: 1.0,
min_p: 0.5,
seed: 3,
..Default::default()
};
let logits = vec![0.0, 0.0, 10.0, 0.0];
for _ in 0..16 {
assert_eq!(Sampler::new(cfg.clone()).sample(&logits), 2);
}
}
#[test]
fn repeat_penalty_suppresses_recent() {
let mut cfg = SamplerConfig::default();
cfg.penalty_last_n = 8;
cfg.penalty_repeat = 100.0;
let mut s = Sampler::new(cfg);
s.accept(1); let logits = vec![4.0, 5.0, 4.5, 1.0]; let got = s.sample(&logits);
assert_ne!(
got, 1,
"recent token must be penalized out of greedy argmax"
);
assert_eq!(got, 2, "next-highest after penalizing 1");
}
}
#[cfg(test)]
mod resume_sampler_predicate_tests {
use super::*;
fn vendor() -> SamplerConfig {
SamplerConfig {
temperature: 0.7,
top_k: 20,
top_p: 0.95,
seed: 20260820,
..Default::default()
}
}
fn pure_temp() -> SamplerConfig {
SamplerConfig {
temperature: 0.7,
seed: 20260820,
..Default::default()
}
}
fn id(cfg: &SamplerConfig) -> SamplerIdentity {
SamplerIdentity::of(cfg)
}
#[test]
fn identical_sampler_resumes() {
for cfg in [pure_temp(), vendor(), SamplerConfig::default()] {
assert_eq!(
id(&cfg).mismatch(&id(&cfg)),
None,
"a request must resume a session its own sampler shaped: {cfg:?}"
);
}
}
#[test]
fn disabled_sentinels_are_the_same_program() {
let a = SamplerConfig {
temperature: 0.7,
top_p: 1.0,
min_p: 0.0,
..Default::default()
};
let b = SamplerConfig {
temperature: 0.7,
top_p: 1.5,
min_p: -1.0,
..Default::default()
};
assert_eq!(id(&a).mismatch(&id(&b)), None, "off spelled two ways");
}
#[test]
fn greedy_temperature_encodings_are_one_program() {
let a = SamplerConfig {
temperature: 0.0,
..Default::default()
};
let b = SamplerConfig {
temperature: -1.0,
..Default::default()
};
assert_eq!(id(&a).mismatch(&id(&b)), None, "temp<=0 is one regime");
}
#[test]
fn neutral_penalty_coefficients_equal_penalties_absent() {
let a = SamplerConfig {
temperature: 0.7,
penalty_last_n: 64,
penalty_repeat: 1.0,
penalty_freq: 0.0,
penalty_present: 0.0,
..Default::default()
};
let b = SamplerConfig {
temperature: 0.7,
penalty_last_n: 0,
..Default::default()
};
assert_eq!(
id(&a).mismatch(&id(&b)),
None,
"an inert penalty window is not a penalty change"
);
}
#[test]
fn the_reproduced_collision_pair_refuses_and_names_a_filter() {
let parked = id(&pure_temp());
let incoming = id(&vendor());
let field = incoming
.mismatch(&parked)
.expect("the reproduced collision pair must refuse");
assert_eq!(field, "top_k", "coarsest-first order names top_k here");
assert!(
incoming.legacy_admits(&parked),
"legacy must admit the collision pair, or this test proves nothing"
);
}
#[test]
fn every_compared_field_refuses_on_its_own_and_names_itself() {
let base = pure_temp();
let pen_base = SamplerConfig {
penalty_last_n: 64,
penalty_repeat: 1.1,
penalty_freq: 0.5,
penalty_present: 0.5,
..base.clone()
};
let cases: [(&str, SamplerConfig, SamplerConfig); 9] = [
(
"regime",
base.clone(),
SamplerConfig {
temperature: 0.0,
..base.clone()
},
),
(
"temperature",
base.clone(),
SamplerConfig {
temperature: 0.8,
..base.clone()
},
),
(
"top_k",
base.clone(),
SamplerConfig {
top_k: 20,
..base.clone()
},
),
(
"top_p",
base.clone(),
SamplerConfig {
top_p: 0.95,
..base.clone()
},
),
(
"min_p",
base.clone(),
SamplerConfig {
min_p: 0.05,
..base.clone()
},
),
(
"penalty_last_n",
pen_base.clone(),
SamplerConfig {
penalty_last_n: 128,
..pen_base.clone()
},
),
(
"penalty_repeat",
pen_base.clone(),
SamplerConfig {
penalty_repeat: 1.2,
..pen_base.clone()
},
),
(
"penalty_freq",
pen_base.clone(),
SamplerConfig {
penalty_freq: 0.6,
..pen_base.clone()
},
),
(
"penalty_present",
pen_base.clone(),
SamplerConfig {
penalty_present: 0.6,
..pen_base.clone()
},
),
];
for (expect, parked_cfg, cfg) in cases {
let parked = id(&parked_cfg);
let incoming = id(&cfg);
assert_eq!(
incoming.mismatch(&parked),
Some(expect),
"changing {expect} alone must refuse and name {expect} ({cfg:?})"
);
assert!(
incoming.legacy_admits(&parked),
"legacy must admit the {expect} change, or the refusal test is tautological"
);
}
assert_eq!(
id(&pen_base).mismatch(&id(&base)),
Some("penalty_last_n"),
"penalties on vs off is named at the window, not at a coefficient"
);
}
#[test]
fn greedy_to_sampled_and_back_both_refuse_as_regime() {
let g = id(&SamplerConfig::default());
let s = id(&pure_temp());
assert_eq!(s.mismatch(&g), Some("regime"));
assert_eq!(g.mismatch(&s), Some("regime"));
}
#[test]
fn seed_alone_does_not_refuse() {
let a = pure_temp();
let b = SamplerConfig {
seed: 999,
..a.clone()
};
assert_eq!(
id(&a).mismatch(&id(&b)),
None,
"seed is carried but not compared"
);
assert_ne!(id(&a).seed(), id(&b).seed(), "the seed is still recorded");
}
#[test]
fn mismatch_is_symmetric_and_identity_is_an_equivalence() {
let a = id(&pure_temp());
let b = id(&vendor());
assert_eq!(a.mismatch(&b).is_some(), b.mismatch(&a).is_some());
assert_eq!(a.mismatch(&a), None);
assert_eq!(b.mismatch(&b), None);
}
}