#[derive(Debug, Clone)]
pub struct SamplingParams {
pub temperature: f32,
pub top_p: f32,
pub top_k: usize,
pub repetition_penalty: f32,
pub presence_penalty: f32,
pub frequency_penalty: f32,
}
impl Default for SamplingParams {
fn default() -> Self {
SamplingParams {
temperature: 0.0,
top_p: 1.0,
top_k: 0,
repetition_penalty: 1.0,
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, logits, 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, logits, params)
}
fn filtered_distribution(
mut scores: Vec<f32>,
raw_logits: &[f32],
params: &SamplingParams,
) -> Vec<f32> {
for s in scores.iter_mut() {
*s /= params.temperature;
}
let mut probs = softmax(&scores);
if params.top_k > 0 && params.top_k < probs.len() {
let mut idx: Vec<usize> = (0..probs.len()).collect();
idx.sort_unstable_by(|&a, &b| probs[b].partial_cmp(&probs[a]).unwrap());
for &i in idx.iter().skip(params.top_k) {
probs[i] = 0.0;
}
}
if params.top_p < 1.0 {
let mut idx: Vec<usize> = (0..probs.len()).collect();
idx.sort_unstable_by(|&a, &b| probs[b].partial_cmp(&probs[a]).unwrap());
let mut cumulative = 0.0f32;
let mut cutoff = idx.len();
for (rank, &i) in idx.iter().enumerate() {
cumulative += probs[i];
if cumulative >= params.top_p {
cutoff = rank + 1;
break;
}
}
for &i in idx.iter().skip(cutoff) {
probs[i] = 0.0;
}
}
let total: f32 = probs.iter().sum();
if total <= 0.0 {
let mut point = vec![0.0f32; probs.len()];
if let Some(p) = point.get_mut(argmax(raw_logits)) {
*p = 1.0;
}
return point;
}
for p in probs.iter_mut() {
*p /= total;
}
probs
}
fn apply_history_penalties(scores: &mut [f32], params: &SamplingParams, history: &[usize]) {
if params.repetition_penalty != 1.0 {
for &tok in history {
if let Some(s) = scores.get_mut(tok) {
*s = if *s > 0.0 {
*s / params.repetition_penalty
} else {
*s * params.repetition_penalty
};
}
}
}
if params.presence_penalty != 0.0 || params.frequency_penalty != 0.0 {
let mut counts = std::collections::HashMap::<usize, usize>::new();
for &tok in history {
*counts.entry(tok).or_insert(0) += 1;
}
for (tok, count) in counts {
if let Some(s) = scores.get_mut(tok) {
*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)
}
fn softmax(logits: &[f32]) -> Vec<f32> {
let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let exps: Vec<f32> = logits.iter().map(|&l| (l - max).exp()).collect();
let sum: f32 = exps.iter().sum();
if sum <= 0.0 {
vec![1.0 / logits.len().max(1) as f32; logits.len()]
} else {
exps.into_iter().map(|e| e / sum).collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[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);
}
}