rlx_moshi/stream/
duplex.rs1use 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#[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
19pub 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 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 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}