use crate::ops::SoftMax;
use nalgebra::{Dyn, StorageMut, Vector};
use ordered_float::OrderedFloat;
use rand::random;
#[derive(Default, Copy, Clone)]
pub struct ProbIndex {
prob: f32,
index: usize,
}
pub struct Sampler {
vocab_size: usize,
prob_index: Vec<ProbIndex>, temperature: f32,
topp: f32,
}
impl Sampler {
pub fn new(vocab_size: usize, temperature: f32, topp: f32) -> Self {
Self {
vocab_size,
prob_index: vec![ProbIndex::default(); vocab_size],
temperature,
topp,
}
}
pub fn sample<S: StorageMut<f32, Dyn>>(&mut self, logits: &mut Vector<f32, Dyn, S>) -> usize {
if self.temperature == 0.0 {
Self::sample_argmax(logits)
} else {
*logits /= self.temperature;
SoftMax::run_cpu(logits);
let coin = random();
if self.topp <= 0.0 || self.topp >= 1.0 {
Self::sample_mult(logits, coin)
} else {
self.sample_topp(logits, coin)
}
}
}
pub fn sample_argmax<S: StorageMut<f32, Dyn>>(probabilities: &Vector<f32, Dyn, S>) -> usize {
probabilities.imax()
}
pub fn sample_mult<S: StorageMut<f32, Dyn>>(
probabilities: &Vector<f32, Dyn, S>,
coin: f32,
) -> usize {
let mut cdf = 0.0;
for (i, prob) in probabilities.iter().enumerate() {
cdf += *prob;
if coin < cdf {
return i;
}
}
probabilities.len() - 1
}
pub fn sample_topp<S: StorageMut<f32, Dyn>>(
&mut self,
probabilities: &Vector<f32, Dyn, S>,
coin: f32,
) -> usize {
let mut n0 = 0;
let cutoff = (1.0 - self.topp) / (self.vocab_size as f32 - 1.0);
for i in 0..probabilities.len() {
if probabilities[i] >= cutoff {
self.prob_index[n0].index = i;
self.prob_index[n0].prob = probabilities[i];
n0 += 1;
}
}
self.prob_index[..n0].sort_by_key(|pid| OrderedFloat(-pid.prob));
let mut cumulative_prob = 0.0;
let mut last_idx = n0 - 1;
for i in 0..n0 {
cumulative_prob += self.prob_index[i].prob;
if cumulative_prob > self.topp {
last_idx = i;
break;
}
}
let r = coin * cumulative_prob;
let mut cdf = 0.0;
for i in 0..=last_idx {
cdf += self.prob_index[i].prob;
if r < cdf {
return self.prob_index[i].index;
}
}
self.prob_index[last_idx].index }
}