use crate::sampler_chain::Candidates;
#[derive(Debug, Clone)]
pub struct SamplingParams {
pub temperature: f32,
pub top_p: f32,
pub min_p: f32,
pub top_k: usize,
pub repetition_penalty: f32,
pub penalty_last_n: usize,
pub presence_penalty: f32,
pub frequency_penalty: f32,
}
impl Default for SamplingParams {
fn default() -> Self {
SamplingParams {
temperature: 0.0,
top_p: 1.0,
min_p: 0.0,
top_k: 0,
repetition_penalty: 1.0,
penalty_last_n: 64,
presence_penalty: 0.0,
frequency_penalty: 0.0,
}
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq)]
pub struct RecommendedSampling {
pub temperature: Option<f32>,
pub top_p: Option<f32>,
pub top_k: Option<usize>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq)]
pub struct RequestedSampling {
pub temperature: Option<f32>,
pub top_p: Option<f32>,
pub top_k: Option<usize>,
}
impl RecommendedSampling {
pub fn is_empty(&self) -> bool {
*self == RecommendedSampling::default()
}
pub fn resolve(
&self,
requested: RequestedSampling,
framework: SamplingParams,
) -> SamplingParams {
SamplingParams {
temperature: requested
.temperature
.or(self.temperature)
.unwrap_or(framework.temperature),
top_p: requested.top_p.or(self.top_p).unwrap_or(framework.top_p),
top_k: requested.top_k.or(self.top_k).unwrap_or(framework.top_k),
..framework
}
}
pub fn from_generation_config(json: &str) -> Self {
let Ok(serde_json::Value::Object(map)) = serde_json::from_str::<serde_json::Value>(json)
else {
return RecommendedSampling::default();
};
if map.get("do_sample").and_then(|v| v.as_bool()) == Some(false) {
return RecommendedSampling {
temperature: Some(0.0),
..RecommendedSampling::default()
};
}
RecommendedSampling {
temperature: map
.get("temperature")
.and_then(|v| v.as_f64())
.map(|v| v as f32),
top_p: map.get("top_p").and_then(|v| v.as_f64()).map(|v| v as f32),
top_k: map
.get("top_k")
.and_then(|v| v.as_u64())
.map(|v| v as usize),
}
}
pub fn from_model_dir(dir: &std::path::Path) -> Self {
match std::fs::read_to_string(dir.join("generation_config.json")) {
Ok(text) => Self::from_generation_config(&text),
Err(_) => RecommendedSampling::default(),
}
}
}
pub type LogitMask<'a> = &'a mut dyn FnMut(&mut [f32]);
pub struct Sampler {
state: u64,
}
impl Sampler {
pub fn new(seed: u64) -> Self {
Sampler {
state: if seed == 0 { 0x9E3779B97F4A7C15 } else { seed },
}
}
fn next_u64(&mut self) -> u64 {
self.state ^= self.state << 13;
self.state ^= self.state >> 7;
self.state ^= self.state << 17;
self.state.wrapping_mul(0x2545F491_4F6CDD1D)
}
fn next_f32(&mut self) -> f32 {
(self.next_u64() >> 40) as f32 / (1u64 << 24) as f32
}
pub fn sample(&mut self, logits: &[f32], params: &SamplingParams, history: &[usize]) -> usize {
self.sample_with_mask(logits, params, history, None)
}
pub fn sample_with_mask(
&mut self,
logits: &[f32],
params: &SamplingParams,
history: &[usize],
mut mask: Option<LogitMask<'_>>,
) -> usize {
if params.temperature <= 0.0 && mask.is_none() {
if logits.len() == 1 {
return logits[0] as usize;
}
let mut scores = logits.to_vec();
apply_history_penalties(&mut scores, params, history);
return argmax(&scores);
}
let mut scores: Vec<f32> = logits.to_vec();
apply_history_penalties(&mut scores, params, history);
if let Some(m) = mask.as_mut() {
m(&mut scores);
}
if params.temperature <= 0.0 {
if scores.len() == 1 {
return scores[0] as usize;
}
return argmax(&scores);
}
let probs = filtered_distribution(scores, params);
self.sample_from(&probs)
}
pub fn uniform(&mut self) -> f32 {
self.next_f32()
}
pub fn sample_from(&mut self, probs: &[f32]) -> usize {
let draw = self.next_f32();
let mut cumulative = 0.0f32;
for (i, &p) in probs.iter().enumerate() {
cumulative += p;
if draw < cumulative {
return i;
}
}
probs
.iter()
.enumerate()
.rev()
.find(|&(_, &p)| p > 0.0)
.map(|(i, _)| i)
.unwrap_or(0)
}
}
pub fn sampling_distribution(
logits: &[f32],
params: &SamplingParams,
history: &[usize],
) -> Vec<f32> {
let mut scores = logits.to_vec();
apply_history_penalties(&mut scores, params, history);
if params.temperature <= 0.0 {
let mut probs = vec![0.0f32; scores.len()];
if let Some(p) = probs.get_mut(argmax(&scores)) {
*p = 1.0;
}
return probs;
}
filtered_distribution(scores, params)
}
fn filtered_distribution(scores: Vec<f32>, params: &SamplingParams) -> Vec<f32> {
let vocab = scores.len();
let mut candidates = Candidates::new(&scores);
candidates.top_k(params.top_k);
candidates.top_p(params.top_p);
candidates.min_p(params.min_p);
candidates.temperature(params.temperature);
candidates.into_distribution(vocab)
}
fn apply_history_penalties(scores: &mut [f32], params: &SamplingParams, history: &[usize]) {
if params.repetition_penalty == 1.0
&& params.presence_penalty == 0.0
&& params.frequency_penalty == 0.0
{
return;
}
if params.penalty_last_n == 0 {
return;
}
let window = history.len().saturating_sub(params.penalty_last_n);
let mut counts = std::collections::HashMap::<usize, usize>::new();
for &tok in &history[window..] {
*counts.entry(tok).or_insert(0) += 1;
}
for (tok, count) in counts {
let Some(s) = scores.get_mut(tok) else {
continue;
};
if params.repetition_penalty != 1.0 {
*s = if *s > 0.0 {
*s / params.repetition_penalty
} else {
*s * params.repetition_penalty
};
}
*s -= params.frequency_penalty * count as f32;
*s -= params.presence_penalty;
}
}
fn argmax(logits: &[f32]) -> usize {
logits
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i)
.unwrap_or(0)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_repetition_penalty_does_not_compound_with_repeats() {
let params = SamplingParams {
temperature: 1.0,
top_p: 1.0,
top_k: 0,
repetition_penalty: 2.0,
..SamplingParams::default()
};
let logits = vec![4.0f32, 1.0, 1.0];
let mut scores = logits.clone();
apply_history_penalties(&mut scores, ¶ms, &[0, 0, 0, 0, 0]);
assert!(
(scores[0] - 2.0).abs() < 1e-6,
"expected one division (2.0), got {} -- {} would be 2^5",
scores[0],
4.0f32 / 32.0
);
let mut once = logits.clone();
apply_history_penalties(&mut once, ¶ms, &[0]);
assert_eq!(once[0].to_bits(), scores[0].to_bits());
let mut negative = vec![-4.0f32];
apply_history_penalties(&mut negative, ¶ms, &[0, 0, 0]);
assert!((negative[0] + 8.0).abs() < 1e-6, "got {}", negative[0]);
}
#[test]
fn the_penalties_only_see_the_last_n_tokens() {
let params = SamplingParams {
repetition_penalty: 2.0,
penalty_last_n: 2,
..SamplingParams::default()
};
let mut scores = vec![8.0f32, 8.0, 8.0];
apply_history_penalties(&mut scores, ¶ms, &[0, 1, 2]);
assert_eq!(
scores[0].to_bits(),
8.0f32.to_bits(),
"token 0 is outside the window"
);
assert!((scores[1] - 4.0).abs() < 1e-6, "got {}", scores[1]);
assert!((scores[2] - 4.0).abs() < 1e-6, "got {}", scores[2]);
let off = SamplingParams {
penalty_last_n: 0,
..params
};
let mut untouched = vec![8.0f32; 3];
apply_history_penalties(&mut untouched, &off, &[0, 1, 2]);
assert_eq!(untouched, vec![8.0f32; 3]);
let wide = SamplingParams {
penalty_last_n: 1000,
..params
};
let mut short = vec![8.0f32];
apply_history_penalties(&mut short, &wide, &[0]);
assert!((short[0] - 4.0).abs() < 1e-6);
}
#[test]
fn the_frequency_penalty_still_counts_repeats() {
let params = SamplingParams {
frequency_penalty: 0.5,
presence_penalty: 0.25,
..SamplingParams::default()
};
let mut scores = vec![10.0f32];
apply_history_penalties(&mut scores, ¶ms, &[0, 0, 0, 0]);
assert!((scores[0] - 7.75).abs() < 1e-6, "got {}", scores[0]);
}
#[test]
fn temperature_does_not_change_which_candidates_top_p_keeps() {
let logits = vec![3.0f32, 2.0, 1.0, 0.0];
let at = |temperature: f32| -> Vec<bool> {
let params = SamplingParams {
temperature,
top_p: 0.9,
top_k: 0,
..SamplingParams::default()
};
sampling_distribution(&logits, ¶ms, &[])
.iter()
.map(|&p| p > 0.0)
.collect()
};
let cold = at(0.5);
let hot = at(4.0);
assert_eq!(
cold, hot,
"the surviving set must not depend on the temperature: \
cold={cold:?} hot={hot:?}"
);
assert!(
cold.iter().any(|&k| !k),
"top_p = 0.9 must drop at least one of these four candidates"
);
}
#[test]
fn min_p_truncates_at_ln_p_below_the_top_logit() {
let logits = vec![4.0f32, 3.0, 2.0, 1.0];
let params = SamplingParams {
temperature: 1.0,
min_p: 0.2,
..SamplingParams::default()
};
let probs = sampling_distribution(&logits, ¶ms, &[]);
assert!(probs[0] > 0.0 && probs[1] > 0.0);
assert_eq!(probs[2], 0.0, "2.0 is below 4 + ln(0.2) = 2.3905");
assert_eq!(probs[3], 0.0);
assert!((probs.iter().sum::<f32>() - 1.0).abs() < 1e-6);
assert!((probs[0] - 0.731_059).abs() < 1e-5, "got {}", probs[0]);
let off = SamplingParams {
min_p: 0.0,
..params.clone()
};
let unfiltered = sampling_distribution(&logits, &off, &[]);
assert!(unfiltered.iter().all(|&p| p > 0.0));
}
#[test]
fn temperature_does_not_change_which_candidates_min_p_keeps() {
let logits = vec![3.0f32, 2.0, 1.0, 0.0];
let survivors = |temperature: f32| -> Vec<bool> {
let params = SamplingParams {
temperature,
min_p: 0.2,
..SamplingParams::default()
};
sampling_distribution(&logits, ¶ms, &[])
.iter()
.map(|&p| p > 0.0)
.collect()
};
let cold = survivors(0.5);
let warm = survivors(1.0);
let hot = survivors(2.0);
assert_eq!(cold, warm, "cold={cold:?} warm={warm:?}");
assert_eq!(warm, hot, "warm={warm:?} hot={hot:?}");
assert_eq!(warm, vec![true, true, false, false]);
}
#[test]
fn top_p_and_min_p_both_apply() {
let logits = vec![3.0f32, 2.0, 1.0, 0.0];
let params = SamplingParams {
temperature: 1.0,
top_p: 0.95,
min_p: 0.2,
..SamplingParams::default()
};
let probs = sampling_distribution(&logits, ¶ms, &[]);
assert_eq!(
probs.iter().map(|&p| p > 0.0).collect::<Vec<_>>(),
vec![true, true, false, false]
);
let top_p_only = SamplingParams {
min_p: 0.0,
..params.clone()
};
assert_eq!(
sampling_distribution(&logits, &top_p_only, &[])
.iter()
.filter(|&&p| p > 0.0)
.count(),
3
);
}
#[test]
fn temperature_zero_accepts_precomputed_argmax_singleton() {
let mut sampler = Sampler::new(1);
let params = SamplingParams::default();
assert_eq!(sampler.sample(&[42.0], ¶ms, &[]), 42);
let sampled = SamplingParams {
temperature: 0.8,
..SamplingParams::default()
};
assert_eq!(sampler.sample(&[42.0], &sampled, &[]), 0);
}
#[test]
fn temperature_zero_is_deterministic_greedy_argmax() {
let logits = vec![0.1, 0.9, 0.3, -0.2];
let params = SamplingParams::default();
let mut sampler = Sampler::new(42);
assert_eq!(sampler.sample(&logits, ¶ms, &[]), 1);
assert_eq!(sampler.sample(&logits, ¶ms, &[]), 1);
}
#[test]
fn high_temperature_can_pick_a_non_argmax_token_over_many_draws() {
let logits = vec![1.0, 1.0, 1.0, 1.0];
let params = SamplingParams {
temperature: 1.0,
..SamplingParams::default()
};
let mut sampler = Sampler::new(7);
let mut seen = std::collections::HashSet::new();
for _ in 0..200 {
seen.insert(sampler.sample(&logits, ¶ms, &[]));
}
assert!(
seen.len() > 1,
"uniform logits at temperature=1.0 must produce more than one distinct token across 200 draws"
);
}
#[test]
fn top_k_one_is_equivalent_to_greedy() {
let logits = vec![0.1, 0.9, 0.3, -0.2];
let params = SamplingParams {
temperature: 1.0,
top_k: 1,
..SamplingParams::default()
};
let mut sampler = Sampler::new(123);
for _ in 0..20 {
assert_eq!(sampler.sample(&logits, ¶ms, &[]), 1);
}
}
#[test]
fn top_p_near_zero_is_equivalent_to_greedy() {
let logits = vec![0.1, 5.0, 0.3, -0.2];
let params = SamplingParams {
temperature: 1.0,
top_p: 0.001,
..SamplingParams::default()
};
let mut sampler = Sampler::new(9);
for _ in 0..20 {
assert_eq!(sampler.sample(&logits, ¶ms, &[]), 1);
}
}
#[test]
fn presence_and_frequency_penalties_reduce_seen_token_logits() {
let logits = vec![0.0, 5.0, 0.0];
let params = SamplingParams {
temperature: 1.0,
presence_penalty: 10.0,
frequency_penalty: 0.0,
..SamplingParams::default()
};
let mut sampler = Sampler::new(1);
let mut counts = [0usize; 3];
for _ in 0..500 {
counts[sampler.sample(&logits, ¶ms, &[1])] += 1;
}
assert!(
counts[1] < 250,
"presence_penalty should discourage token 1; counts={counts:?}"
);
let params = SamplingParams {
temperature: 1.0,
presence_penalty: 0.0,
frequency_penalty: 10.0,
..SamplingParams::default()
};
let mut sampler = Sampler::new(2);
counts = [0; 3];
for _ in 0..500 {
counts[sampler.sample(&logits, ¶ms, &[1, 1, 1])] += 1;
}
assert!(
counts[1] < 250,
"frequency_penalty should discourage repeated token 1; counts={counts:?}"
);
}
#[test]
fn repetition_penalty_reduces_probability_of_recently_seen_token() {
let logits = vec![0.0, 5.0, 0.0];
let params = SamplingParams {
temperature: 1.0,
repetition_penalty: 1000.0,
..SamplingParams::default()
};
let mut sampler = Sampler::new(3);
let mut counts = [0usize; 3];
for _ in 0..500 {
counts[sampler.sample(&logits, ¶ms, &[1])] += 1;
}
assert!(
counts[1] < 250,
"heavily penalizing token 1 (already in history) should make it far less likely than its raw logit alone would suggest; got counts={counts:?}"
);
}
#[test]
fn low_seeds_do_not_bias_the_first_draw() {
let vocab = 8;
let logits = vec![0.0f32; vocab];
let params = SamplingParams {
temperature: 1.0,
..SamplingParams::default()
};
let seeds = 4_000u64;
let mut counts = vec![0usize; vocab];
for seed in 1..=seeds {
counts[Sampler::new(seed).sample(&logits, ¶ms, &[])] += 1;
}
let expected = seeds as f64 / vocab as f64;
for (token, &c) in counts.iter().enumerate() {
assert!(
(c as f64 - expected).abs() < expected * 0.25,
"uniform logits: token {token} came up {c} times across {seeds} seeds, \
expected about {expected:.0} (counts={counts:?})"
);
}
}
#[test]
fn the_published_distribution_is_the_one_sample_actually_draws_from() {
let logits = vec![0.4, 2.0, -1.0, 1.2, 0.9, -0.3];
let params = SamplingParams {
temperature: 0.8,
top_p: 0.9,
top_k: 4,
repetition_penalty: 1.3,
..SamplingParams::default()
};
let history = [1usize, 4];
let claimed = sampling_distribution(&logits, ¶ms, &history);
assert!((claimed.iter().sum::<f32>() - 1.0).abs() < 1e-5);
let draws = 100_000;
let mut counts = vec![0usize; logits.len()];
let mut sampler = Sampler::new(0xC0FFEE);
for _ in 0..draws {
counts[sampler.sample(&logits, ¶ms, &history)] += 1;
}
for (i, &c) in counts.iter().enumerate() {
let empirical = c as f64 / draws as f64;
assert!(
(empirical - claimed[i] as f64).abs() < 0.01,
"token {i}: sample() draws it {empirical:.4} of the time but \
sampling_distribution claims {:.4}",
claimed[i]
);
}
}
#[test]
fn greedy_is_published_as_a_point_mass_not_a_special_case() {
let logits = vec![0.1, 0.9, 0.3, -0.2];
let probs = sampling_distribution(&logits, &SamplingParams::default(), &[]);
assert_eq!(probs, vec![0.0, 1.0, 0.0, 0.0]);
let penalized = sampling_distribution(
&logits,
&SamplingParams {
repetition_penalty: 100.0,
..SamplingParams::default()
},
&[1],
);
assert_eq!(penalized[1], 0.0);
assert_eq!(penalized.iter().sum::<f32>(), 1.0);
}
#[test]
fn degenerate_all_zero_probability_falls_back_to_greedy() {
let logits = vec![0.1, 0.9, 0.3, -0.2];
let params = SamplingParams {
temperature: 1.0,
top_k: 1,
top_p: 1.0,
..SamplingParams::default()
};
let mut sampler = Sampler::new(1);
assert_eq!(sampler.sample(&logits, ¶ms, &[]), 1);
}
#[test]
fn an_absent_generation_config_key_stays_absent_rather_than_taking_a_default() {
let recommended = RecommendedSampling::from_generation_config(r#"{"temperature": 0.6}"#);
assert_eq!(recommended.temperature, Some(0.6));
assert_eq!(recommended.top_p, None, "top_p was not in the file");
assert_eq!(recommended.top_k, None, "top_k was not in the file");
let nulled = RecommendedSampling::from_generation_config(r#"{"top_p": null}"#);
assert_eq!(nulled, RecommendedSampling::default());
}
#[test]
fn every_generation_config_key_present_is_recommended() {
let recommended = RecommendedSampling::from_generation_config(
r#"{"do_sample": true, "temperature": 1.0, "top_k": 20, "top_p": 0.95}"#,
);
assert_eq!(
recommended,
RecommendedSampling {
temperature: Some(1.0),
top_p: Some(0.95),
top_k: Some(20),
}
);
}
#[test]
fn do_sample_false_recommends_greedy_and_no_other_field() {
let recommended = RecommendedSampling::from_generation_config(
r#"{"do_sample": false, "temperature": 0.7, "top_k": 50, "top_p": 0.9}"#,
);
assert_eq!(recommended.temperature, Some(0.0));
assert_eq!(recommended.top_p, None);
assert_eq!(recommended.top_k, None);
}
#[test]
fn a_malformed_generation_config_recommends_nothing() {
for text in ["", "not json", "[1, 2, 3]", "null"] {
assert!(
RecommendedSampling::from_generation_config(text).is_empty(),
"{text:?} must recommend nothing"
);
}
}
#[test]
fn a_request_outranks_the_recommendation_which_outranks_the_framework_default() {
let recommended = RecommendedSampling {
temperature: Some(1.0),
top_p: Some(0.95),
top_k: Some(20),
};
let resolved = recommended.resolve(
RequestedSampling {
temperature: Some(0.0),
..RequestedSampling::default()
},
SamplingParams::default(),
);
assert_eq!(resolved.temperature, 0.0, "the request asked for greedy");
assert_eq!(resolved.top_p, 0.95, "the request said nothing about top_p");
assert_eq!(resolved.top_k, 20, "the request said nothing about top_k");
assert_eq!(resolved.repetition_penalty, 1.0);
}
#[test]
fn a_checkpoint_that_recommends_nothing_leaves_the_framework_defaults_alone() {
let resolved = RecommendedSampling::default()
.resolve(RequestedSampling::default(), SamplingParams::default());
let default = SamplingParams::default();
assert_eq!(resolved.temperature, default.temperature);
assert_eq!(resolved.top_p, default.top_p);
assert_eq!(resolved.top_k, default.top_k);
}
#[test]
fn a_model_directory_without_a_generation_config_recommends_nothing() {
let dir = std::env::temp_dir().join(format!(
"ferrox_test_no_generation_config_{}",
std::process::id()
));
std::fs::create_dir_all(&dir).unwrap();
assert!(RecommendedSampling::from_model_dir(&dir).is_empty());
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn a_model_directory_generation_config_is_read_from_beside_the_weights() {
let dir = std::env::temp_dir().join(format!(
"ferrox_test_generation_config_dir_{}",
std::process::id()
));
std::fs::create_dir_all(&dir).unwrap();
std::fs::write(
dir.join("generation_config.json"),
r#"{"temperature": 0.6, "top_p": 0.95}"#,
)
.unwrap();
let recommended = RecommendedSampling::from_model_dir(&dir);
std::fs::remove_dir_all(&dir).ok();
assert_eq!(recommended.temperature, Some(0.6));
assert_eq!(recommended.top_p, Some(0.95));
assert_eq!(recommended.top_k, None);
}
}