Skip to main content

rlx_moshi/stream/
duplex.rs

1use crate::backend::{MoshiGenState, MoshiLm};
2use crate::config::GenerateConfig;
3use crate::sampling::LogitsProcessor;
4use crate::session::{GenerationConfig, MoshiSession};
5use crate::tokenizer::MoshiTokenizer;
6use anyhow::{Result, ensure};
7use rlx_mimi::{MimiCodec, MimiCodes};
8use rlx_runtime::Device;
9
10/// One streaming step's output.
11#[derive(Debug, Clone)]
12pub struct StreamStepOutput {
13    pub step: usize,
14    pub text_token: u32,
15    pub moshi_pcm: Vec<f32>,
16    pub transcript_delta: Option<String>,
17}
18
19/// Incremental full-duplex engine — feed 24 kHz PCM, receive Moshi PCM chunks.
20pub struct DuplexStreamEngine {
21    lm: MoshiLm,
22    mimi: MimiCodec,
23    tokenizer: MoshiTokenizer,
24    state: MoshiGenState,
25    gen_cfg: GenerateConfig,
26    run_cfg: GenerationConfig,
27    text_frames: Vec<u32>,
28    pcm_buf: Vec<f32>,
29    frame_samples: usize,
30    frame_idx: usize,
31    finished: bool,
32    device: Device,
33}
34
35impl DuplexStreamEngine {
36    pub fn from_session(
37        session: MoshiSession,
38        prompt: &str,
39        run_cfg: &GenerationConfig,
40    ) -> Result<Self> {
41        let parts = session.into_parts()?;
42        ensure!(
43            parts.gen_cfg.input_audio_codebooks > 0,
44            "DuplexStreamEngine requires a full-duplex variant (Moshiko or Moshika)"
45        );
46        let max = run_cfg.max_steps;
47        let text_frames = parts.tokenizer.prompt_frame_tokens(prompt, max)?;
48        let text_lp = LogitsProcessor::new(
49            run_cfg.text_temperature,
50            run_cfg.text_top_k,
51            run_cfg.text_seed,
52        );
53        let audio_lp = LogitsProcessor::new(
54            run_cfg.audio_temperature,
55            run_cfg.audio_top_k,
56            run_cfg.audio_seed,
57        );
58        let mut state = parts
59            .lm
60            .new_gen_state(max, text_lp, audio_lp, parts.gen_cfg.clone())?;
61        let mut lm = parts.lm;
62        state.reset(
63            &mut lm,
64            max,
65            LogitsProcessor::new(
66                run_cfg.text_temperature,
67                run_cfg.text_top_k,
68                run_cfg.text_seed,
69            ),
70            LogitsProcessor::new(
71                run_cfg.audio_temperature,
72                run_cfg.audio_top_k,
73                run_cfg.audio_seed,
74            ),
75        )?;
76        let frame_samples = parts.mimi.config().samples_per_codec_frame();
77        Ok(Self {
78            lm,
79            mimi: parts.mimi,
80            tokenizer: parts.tokenizer,
81            state,
82            gen_cfg: parts.gen_cfg,
83            run_cfg: run_cfg.clone(),
84            text_frames,
85            pcm_buf: Vec::new(),
86            frame_samples,
87            frame_idx: 0,
88            finished: false,
89            device: parts.device,
90        })
91    }
92
93    pub fn device(&self) -> Device {
94        self.device
95    }
96
97    pub fn frame_samples(&self) -> usize {
98        self.frame_samples
99    }
100
101    pub fn steps_done(&self) -> usize {
102        self.frame_idx
103    }
104
105    /// Append PCM (mono f32 @ 24 kHz). Returns zero or more completed step outputs.
106    pub fn feed_pcm(&mut self, pcm: &[f32]) -> Result<Vec<StreamStepOutput>> {
107        ensure!(!self.finished, "stream already finished");
108        self.pcm_buf.extend_from_slice(pcm);
109        let mut outs = Vec::new();
110        while self.pcm_buf.len() >= self.frame_samples && self.frame_idx < self.run_cfg.max_steps {
111            let frame_pcm: Vec<f32> = self.pcm_buf.drain(..self.frame_samples).collect();
112            outs.push(self.step_frame(&frame_pcm)?);
113        }
114        Ok(outs)
115    }
116
117    /// Pad tail, run remaining buffered audio, mark finished.
118    pub fn finish(&mut self) -> Result<Vec<StreamStepOutput>> {
119        if self.finished {
120            return Ok(Vec::new());
121        }
122        let mut outs = Vec::new();
123        if !self.pcm_buf.is_empty() && self.frame_idx < self.run_cfg.max_steps {
124            let mut frame_pcm = std::mem::take(&mut self.pcm_buf);
125            frame_pcm.resize(self.frame_samples, 0.0);
126            outs.push(self.step_frame(&frame_pcm)?);
127        }
128        self.finished = true;
129        Ok(outs)
130    }
131
132    pub fn collected_text_tokens(&self) -> Vec<u32> {
133        self.state.text_tokens().to_vec()
134    }
135
136    fn step_frame(&mut self, frame_pcm: &[f32]) -> Result<StreamStepOutput> {
137        let user_codes = self
138            .mimi
139            .encode_pcm(frame_pcm, Some(self.run_cfg.mimi_codebooks))?;
140        ensure!(
141            user_codes.num_frames() >= 1,
142            "mimi encode produced no frames"
143        );
144        let user_frame = user_codes.frames[0].clone();
145        let text_tok = self.text_frames[self.frame_idx];
146        let sampled = self.state.step(&mut self.lm, text_tok, &user_frame)?;
147        let mut moshi_pcm = Vec::new();
148        if let Some(frame) = self.state.last_audio_frame() {
149            moshi_pcm = self.decode_frame(&frame)?;
150        }
151        let transcript_delta = self.token_delta(sampled);
152        let step = self.frame_idx;
153        self.frame_idx += 1;
154        Ok(StreamStepOutput {
155            step,
156            text_token: sampled,
157            moshi_pcm,
158            transcript_delta,
159        })
160    }
161
162    fn decode_frame(&mut self, frame: &[u32]) -> Result<Vec<f32>> {
163        let codes = MimiCodes {
164            frames: vec![frame.to_vec()],
165            num_quantizers: self.run_cfg.mimi_codebooks,
166        };
167        self.mimi.decode_codes(&codes)
168    }
169
170    fn token_delta(&self, token: u32) -> Option<String> {
171        let g = &self.gen_cfg;
172        if token == g.text_start_token || token == g.text_pad_token || token == g.text_eop_token {
173            return None;
174        }
175        self.tokenizer.decode_piece(token).ok()
176    }
177}