pub(crate) struct Candidates {
ids: Vec<usize>,
logits: Vec<f32>,
probs: Vec<f32>,
}
impl Candidates {
pub(crate) fn new(logits: &[f32]) -> Self {
let mut ids: Vec<usize> = (0..logits.len()).collect();
ids.sort_unstable_by(|&a, &b| logits[b].total_cmp(&logits[a]).then(a.cmp(&b)));
let ordered: Vec<f32> = ids.iter().map(|&i| logits[i]).collect();
let probs = vec![0.0f32; ordered.len()];
Candidates {
ids,
logits: ordered,
probs,
}
}
pub(crate) fn len(&self) -> usize {
self.ids.len()
}
fn truncate(&mut self, n: usize) {
self.ids.truncate(n);
self.logits.truncate(n);
self.probs.truncate(n);
}
fn softmax(&mut self) {
if self.logits.is_empty() {
return;
}
let max = self.logits[0];
let mut sum = 0.0f32;
for (p, &l) in self.probs.iter_mut().zip(self.logits.iter()) {
*p = (l - max).exp();
sum += *p;
}
if sum <= 0.0 || !sum.is_finite() {
let uniform = 1.0 / self.probs.len() as f32;
self.probs.fill(uniform);
return;
}
for p in self.probs.iter_mut() {
*p /= sum;
}
}
pub(crate) fn top_k(&mut self, k: usize) {
if k == 0 {
return;
}
self.truncate(k.min(self.len()));
}
pub(crate) fn top_p(&mut self, p: f32) {
if p >= 1.0 {
return;
}
self.softmax();
let mut cum_sum = 0.0f32;
let mut last_idx = self.len();
for i in 0..self.len() {
cum_sum += self.probs[i];
if cum_sum >= p {
last_idx = i + 1;
break;
}
}
self.truncate(last_idx);
}
pub(crate) fn min_p(&mut self, p: f32) {
if p <= 0.0 || self.logits.is_empty() {
return;
}
let min_logit = self.logits[0] + p.ln();
let mut i = 1;
while i < self.len() && self.logits[i] >= min_logit {
i += 1;
}
self.truncate(i);
}
pub(crate) fn temperature(&mut self, temp: f32) {
if temp <= 0.0 {
return;
}
for l in self.logits.iter_mut() {
*l /= temp;
}
}
pub(crate) fn into_distribution(mut self, vocab: usize) -> Vec<f32> {
self.softmax();
let mut out = vec![0.0f32; vocab];
for (&id, &p) in self.ids.iter().zip(self.probs.iter()) {
if let Some(slot) = out.get_mut(id) {
*slot = p;
}
}
out
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn min_p_keeps_candidates_within_ln_p_of_the_top_logit() {
let mut c = Candidates::new(&[4.0, 3.0, 2.0, 1.0]);
c.min_p(0.2);
assert_eq!(c.ids, vec![0, 1], "threshold is 4 + ln(0.2) = 2.3905");
let mut c = Candidates::new(&[0.0, (0.5f32).ln(), -5.0]);
c.min_p(0.5);
assert_eq!(c.ids, vec![0, 1], "p_i == p * p_max must be kept");
let mut c = Candidates::new(&[1.0, 0.9, 0.8]);
c.min_p(2.0);
assert_eq!(c.ids, vec![0]);
let mut c = Candidates::new(&[4.0, 3.0, 2.0, 1.0]);
c.min_p(0.0);
assert_eq!(c.len(), 4);
}
#[test]
fn top_p_renormalises_over_what_top_k_left() {
let logits = vec![3.0f32, 2.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0];
let mut c = Candidates::new(&logits);
c.top_k(2);
c.top_p(0.72);
assert_eq!(
c.ids,
vec![0],
"renormalised over the top 2, token 0 already holds 0.731"
);
let mut c = Candidates::new(&logits);
c.top_p(0.72);
assert_eq!(c.ids, vec![0, 1]);
}
#[test]
fn top_p_includes_the_candidate_that_crosses_the_threshold() {
let mut c = Candidates::new(&[1.0f32, 1.0]);
c.top_p(0.5);
assert_eq!(c.ids, vec![0]);
let mut c = Candidates::new(&[1.0f32, 1.0]);
c.top_p(0.6);
assert_eq!(c.ids, vec![0, 1]);
}
#[test]
fn equal_logits_break_ties_on_the_token_id() {
let c = Candidates::new(&[1.0f32; 6]);
assert_eq!(c.ids, vec![0, 1, 2, 3, 4, 5]);
}
#[test]
fn the_published_distribution_is_normalised_and_zero_outside_the_survivors() {
let mut c = Candidates::new(&[3.0f32, 2.0, 1.0, 0.0]);
c.top_k(2);
c.temperature(0.5);
let probs = c.into_distribution(4);
assert!((probs.iter().sum::<f32>() - 1.0).abs() < 1e-6);
assert_eq!(probs[2], 0.0);
assert_eq!(probs[3], 0.0);
assert!((probs[0] - 0.880_797).abs() < 1e-5, "got {}", probs[0]);
}
}