1use ndarray::ArrayView1;
2use rand::Rng;
3use rand::SeedableRng;
4
5#[derive(Clone)]
7pub struct LogitsProcessor {
8 pub temperature: f64,
9 pub top_k: usize,
10 pub seed: u64,
11}
12
13impl LogitsProcessor {
14 pub fn new(temperature: f64, top_k: usize, seed: u64) -> Self {
15 Self {
16 temperature,
17 top_k,
18 seed,
19 }
20 }
21
22 pub fn temperature(&self) -> f64 {
23 self.temperature
24 }
25
26 pub fn top_k(&self) -> usize {
27 self.top_k
28 }
29
30 pub fn seed(&self) -> u64 {
31 self.seed
32 }
33
34 pub fn sample(&mut self, logits: ArrayView1<f32>) -> anyhow::Result<u32> {
35 let v = logits.len();
36 let mut idx_logits: Vec<(usize, f32)> = (0..v).map(|i| (i, logits[i])).collect();
37 if self.top_k > 0 && self.top_k < v {
38 idx_logits.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
39 idx_logits.truncate(self.top_k);
40 }
41 if self.temperature <= 0.0 || (self.temperature - 1.0).abs() < 1e-6 && self.top_k == 1 {
42 let best = idx_logits
43 .iter()
44 .max_by(|a, b| a.1.partial_cmp(&b.1).unwrap())
45 .map(|(i, _)| *i)
46 .unwrap_or(0);
47 return Ok(best as u32);
48 }
49 let inv_t = 1.0 / self.temperature as f32;
50 let max = idx_logits
51 .iter()
52 .map(|(_, l)| *l)
53 .fold(f32::NEG_INFINITY, f32::max);
54 let mut probs: Vec<f32> = idx_logits
55 .iter()
56 .map(|(_, l)| ((l - max) * inv_t).exp())
57 .collect();
58 let sum: f32 = probs.iter().sum();
59 let mut rng = rand::rngs::StdRng::seed_from_u64(self.seed);
60 self.seed = rng.r#gen();
61 let r: f32 = rng.r#gen::<f32>() * sum;
62 let mut acc = 0.0f32;
63 for (i, p) in probs.iter_mut().enumerate() {
64 acc += *p;
65 if r <= acc {
66 return Ok(idx_logits[i].0 as u32);
67 }
68 }
69 Ok(idx_logits.last().map(|(i, _)| *i).unwrap_or(0) as u32)
70 }
71}