Skip to main content

rlx_moshi/
generate.rs

1use crate::config::GenerateConfig;
2use crate::lm::LmModel;
3use crate::sampling::LogitsProcessor;
4use anyhow::{Result, ensure};
5
6pub const UNGENERATED: u32 = u32::MAX;
7
8/// Multistream autoregressive state (ported from kyutai `lm_generate_multistream`).
9pub struct GenerateState {
10    audio_tokens: Vec<Vec<u32>>,
11    text_tokens: Vec<u32>,
12    text_lp: LogitsProcessor,
13    audio_lp: LogitsProcessor,
14    step_idx: usize,
15    forced_audio_tokens: ForcedAudioTokens,
16    cfg: GenerateConfig,
17}
18
19#[derive(Debug, Clone)]
20pub(crate) struct ForcedAudioTokens {
21    delay: usize,
22    pad: u32,
23    pattern: Vec<usize>,
24}
25
26impl ForcedAudioTokens {
27    pub(crate) fn new(delay: usize, pad: u32, pattern: &[usize]) -> Self {
28        Self {
29            delay,
30            pad,
31            pattern: pattern.to_vec(),
32        }
33    }
34
35    pub(crate) fn forced_tokens(&self, step: usize) -> Vec<Option<u32>> {
36        if step >= self.delay {
37            return vec![None; self.pattern.len()];
38        }
39        self.pattern
40            .iter()
41            .map(|&v| if v == 0 { None } else { Some(self.pad) })
42            .collect()
43    }
44}
45
46impl GenerateState {
47    pub fn new(
48        max_steps: usize,
49        text_lp: LogitsProcessor,
50        audio_lp: LogitsProcessor,
51        cfg: GenerateConfig,
52    ) -> Self {
53        let buf = max_steps + cfg.acoustic_delay;
54        let audio_tokens = vec![vec![UNGENERATED; cfg.total_audio_codebooks()]; buf];
55        let text_tokens = vec![UNGENERATED; buf];
56        let forced = ForcedAudioTokens::new(cfg.acoustic_delay, cfg.audio_pad_token(), &[8, 8]);
57        Self {
58            audio_tokens,
59            text_tokens,
60            text_lp,
61            audio_lp,
62            step_idx: 0,
63            forced_audio_tokens: forced,
64            cfg,
65        }
66    }
67
68    pub fn config(&self) -> &GenerateConfig {
69        &self.cfg
70    }
71
72    pub fn step_idx(&self) -> usize {
73        self.step_idx
74    }
75
76    pub fn text_tokens(&self) -> &[u32] {
77        let n = self.step_idx.min(self.text_tokens.len());
78        &self.text_tokens[..n]
79    }
80
81    /// Advance one 12.5 Hz frame. `input_audio` is user codebooks (empty for one-way).
82    pub fn step(&mut self, lm: &mut LmModel, text_token: u32, input_audio: &[u32]) -> Result<u32> {
83        ensure!(
84            input_audio.len() == self.cfg.input_audio_codebooks,
85            "expected {} user codebooks, got {}",
86            self.cfg.input_audio_codebooks,
87            input_audio.len()
88        );
89        for (ci, &t) in input_audio.iter().enumerate() {
90            let idx = ci + self.cfg.generated_audio_codebooks;
91            self.audio_tokens[self.step_idx][idx] = t;
92        }
93        let pad = self.cfg.audio_pad_token();
94        let mut delayed = Vec::with_capacity(self.cfg.total_audio_codebooks());
95        for codebook in 0..self.cfg.total_audio_codebooks() {
96            let t = if codebook == 0 || codebook == self.cfg.generated_audio_codebooks {
97                if self.step_idx == 0 {
98                    pad
99                } else {
100                    self.audio_tokens[self.step_idx - 1][codebook]
101                }
102            } else if self.step_idx <= self.cfg.acoustic_delay {
103                pad
104            } else {
105                self.audio_tokens[self.step_idx - self.cfg.acoustic_delay - 1][codebook]
106            };
107            ensure!(
108                t != UNGENERATED,
109                "internal: ungenerated audio at step {}",
110                self.step_idx
111            );
112            delayed.push(Some(t));
113        }
114        let (text_logits, hidden) = lm.forward_step(Some(text_token), &delayed)?;
115        let sampled_text = self.text_lp.sample(text_logits.view())?;
116        self.text_tokens[self.step_idx] = sampled_text;
117        let forced = self.forced_audio_tokens.forced_tokens(self.step_idx);
118        if let Some(tokens) =
119            lm.depformer_sample(&hidden, Some(sampled_text), &forced, &mut self.audio_lp)?
120        {
121            for (ci, &tok) in tokens.iter().enumerate() {
122                let delay = if ci == 0 { 0 } else { self.cfg.acoustic_delay };
123                let pos = self.step_idx.saturating_sub(delay);
124                self.audio_tokens[pos][ci] = tok;
125            }
126        }
127        self.step_idx += 1;
128        Ok(sampled_text)
129    }
130
131    /// Moshi output codebooks ready for Mimi decode (past acoustic delay).
132    pub fn last_audio_frame(&self) -> Option<Vec<u32>> {
133        if self.step_idx <= self.cfg.acoustic_delay {
134            return None;
135        }
136        let pos = self.step_idx - self.cfg.acoustic_delay - 1;
137        let frame = &self.audio_tokens[pos];
138        let pad = self.cfg.audio_pad_token();
139        if frame[..self.cfg.generated_audio_codebooks]
140            .iter()
141            .any(|&t| t >= pad)
142        {
143            return None;
144        }
145        Some(frame[..self.cfg.generated_audio_codebooks].to_vec())
146    }
147
148    pub fn reset(&mut self, lm: &mut LmModel) {
149        lm.reset_state();
150        self.step_idx = 0;
151        let buf = self.audio_tokens.len();
152        let tc = self.cfg.total_audio_codebooks();
153        self.audio_tokens = vec![vec![UNGENERATED; tc]; buf];
154        self.text_tokens = vec![UNGENERATED; self.text_tokens.len()];
155    }
156}