Skip to main content

rlx_moshi/
sampling.rs

1use ndarray::ArrayView1;
2use rand::Rng;
3use rand::SeedableRng;
4
5/// Greedy / top-k / temperature sampling over a 1-D logit vector.
6#[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}