Skip to main content

rlx_moshi/
session.rs

1use crate::backend::{MoshiLm, resolve_lm_device};
2use crate::checkpoint::MoshiCheckpoint;
3use crate::config::{GenerateConfig, MoshiVariant};
4use crate::download::{default_mimi_dir, ensure_weights_checkpoint, tokenizer_path};
5use crate::sampling::LogitsProcessor;
6use crate::tokenizer::MoshiTokenizer;
7use anyhow::{Result, ensure};
8use rlx_mimi::{MimiCodec, MimiCodes, SAMPLE_RATE as MIMI_RATE};
9use rlx_runtime::Device;
10use std::path::{Path, PathBuf};
11
12/// Sampling overrides for Moshi generation.
13#[derive(Debug, Clone)]
14pub struct GenerationConfig {
15    pub max_steps: usize,
16    pub text_temperature: f64,
17    pub text_top_k: usize,
18    pub audio_temperature: f64,
19    pub audio_top_k: usize,
20    pub text_seed: u64,
21    pub audio_seed: u64,
22    pub mimi_codebooks: usize,
23}
24
25impl Default for GenerationConfig {
26    fn default() -> Self {
27        Self {
28            max_steps: 25,
29            text_temperature: 0.8,
30            text_top_k: 250,
31            audio_temperature: 0.8,
32            audio_top_k: 250,
33            text_seed: 42,
34            audio_seed: 43,
35            mimi_codebooks: 8,
36        }
37    }
38}
39
40/// Synthesis output: mono PCM @ 24 kHz + token trace.
41#[derive(Debug, Clone)]
42pub struct GenerationResult {
43    pub samples: Vec<f32>,
44    pub sample_rate: u32,
45    pub text_tokens: Vec<u32>,
46    pub audio_frames: Vec<Vec<u32>>,
47    pub transcript: String,
48}
49
50/// Decomposed session for streaming worker ownership.
51pub struct MoshiSessionParts {
52    pub lm: MoshiLm,
53    pub mimi: MimiCodec,
54    pub tokenizer: MoshiTokenizer,
55    pub gen_cfg: GenerateConfig,
56    pub device: Device,
57    pub variant: MoshiVariant,
58    pub checkpoint: MoshiCheckpoint,
59    pub moshi_dir: PathBuf,
60}
61
62/// Moshi session — LM + Mimi codec + tokenizer.
63pub struct MoshiSession {
64    variant: MoshiVariant,
65    lm: MoshiLm,
66    mimi: MimiCodec,
67    tokenizer: MoshiTokenizer,
68    gen_cfg: GenerateConfig,
69    moshi_dir: PathBuf,
70    device: Device,
71    checkpoint: MoshiCheckpoint,
72}
73
74impl MoshiSession {
75    pub fn open(
76        moshi_dir: impl AsRef<Path>,
77        mimi_dir: impl AsRef<Path>,
78        variant: MoshiVariant,
79    ) -> Result<Self> {
80        Self::open_on(moshi_dir, mimi_dir, variant, Device::Cpu)
81    }
82
83    pub fn open_on(
84        moshi_dir: impl AsRef<Path>,
85        mimi_dir: impl AsRef<Path>,
86        variant: MoshiVariant,
87        device: Device,
88    ) -> Result<Self> {
89        Self::open_with_checkpoint(
90            moshi_dir,
91            mimi_dir,
92            variant,
93            device,
94            MoshiCheckpoint::from_env_or_default(),
95        )
96    }
97
98    pub fn open_with_checkpoint(
99        moshi_dir: impl AsRef<Path>,
100        mimi_dir: impl AsRef<Path>,
101        variant: MoshiVariant,
102        device: Device,
103        checkpoint: MoshiCheckpoint,
104    ) -> Result<Self> {
105        let moshi_dir = moshi_dir.as_ref().to_path_buf();
106        ensure_weights_checkpoint(&moshi_dir, variant, checkpoint)?;
107        rlx_mimi::ensure_weights(mimi_dir.as_ref())?;
108        let device = resolve_lm_device(device, checkpoint);
109        let lm = MoshiLm::open(&moshi_dir, variant, checkpoint, device)?;
110        let mimi =
111            MimiCodec::open_on_with_moshi(mimi_dir.as_ref(), Some(&moshi_dir), device, Some(8))?;
112        let tokenizer = MoshiTokenizer::open(tokenizer_path(&moshi_dir))?;
113        Ok(Self {
114            variant,
115            lm,
116            mimi,
117            tokenizer,
118            gen_cfg: variant.generate_config(),
119            moshi_dir,
120            device,
121            checkpoint,
122        })
123    }
124
125    pub fn open_default(variant: MoshiVariant) -> Result<Self> {
126        Self::open_default_on(variant, Device::Cpu)
127    }
128
129    pub fn open_default_on(variant: MoshiVariant, device: Device) -> Result<Self> {
130        Self::open_on(
131            crate::download::default_moshi_dir(),
132            default_mimi_dir(),
133            variant,
134            device,
135        )
136    }
137
138    pub fn variant(&self) -> MoshiVariant {
139        self.variant
140    }
141
142    pub fn device(&self) -> Device {
143        self.device
144    }
145
146    pub fn checkpoint(&self) -> MoshiCheckpoint {
147        self.checkpoint
148    }
149
150    pub fn moshi_dir(&self) -> &Path {
151        &self.moshi_dir
152    }
153
154    pub fn gen_cfg_internal(&self) -> &GenerateConfig {
155        &self.gen_cfg
156    }
157
158    pub fn into_parts(self) -> Result<MoshiSessionParts> {
159        Ok(MoshiSessionParts {
160            lm: self.lm,
161            mimi: self.mimi,
162            tokenizer: self.tokenizer,
163            gen_cfg: self.gen_cfg,
164            device: self.device,
165            variant: self.variant,
166            checkpoint: self.checkpoint,
167            moshi_dir: self.moshi_dir,
168        })
169    }
170
171    /// One-way TTS from a text prompt (blank user audio).
172    pub fn generate_one_way(
173        &mut self,
174        prompt: &str,
175        cfg: &GenerationConfig,
176    ) -> Result<GenerationResult> {
177        ensure!(
178            self.gen_cfg.input_audio_codebooks == 0,
179            "generate_one_way requires a one-way variant (MoshikoOneWay or MoshikaOneWay)"
180        );
181        self.run_generation(prompt, cfg)
182    }
183
184    /// Full-duplex: encode user WAV with Mimi, condition generation, decode Moshi reply.
185    pub fn generate_duplex(
186        &mut self,
187        user_wav: impl AsRef<Path>,
188        cfg: &GenerationConfig,
189    ) -> Result<GenerationResult> {
190        ensure!(
191            self.gen_cfg.input_audio_codebooks > 0,
192            "generate_duplex requires a full-duplex variant (Moshiko or Moshika)"
193        );
194        let user_codes = self
195            .mimi
196            .encode_wav(user_wav.as_ref(), Some(cfg.mimi_codebooks))?;
197        let num_frames = user_codes.num_frames().min(cfg.max_steps);
198        self.run_generation_with_user("", &user_codes, num_frames, cfg)
199    }
200
201    fn run_generation(&mut self, prompt: &str, cfg: &GenerationConfig) -> Result<GenerationResult> {
202        let text_frames = self.tokenizer.prompt_frame_tokens(prompt, cfg.max_steps)?;
203        let text_lp = LogitsProcessor::new(cfg.text_temperature, cfg.text_top_k, cfg.text_seed);
204        let audio_lp = LogitsProcessor::new(cfg.audio_temperature, cfg.audio_top_k, cfg.audio_seed);
205        let mut state = self.lm.new_gen_state(
206            cfg.max_steps,
207            text_lp.clone(),
208            audio_lp.clone(),
209            self.gen_cfg.clone(),
210        )?;
211        state.reset(&mut self.lm, cfg.max_steps, text_lp, audio_lp)?;
212        let empty_user: Vec<u32> = vec![];
213        let mut decoded_pcm = Vec::new();
214        let mut audio_trace = Vec::new();
215        for step in 0..cfg.max_steps {
216            let tt = text_frames[step];
217            state.step(&mut self.lm, tt, &empty_user)?;
218            if let Some(frame) = state.last_audio_frame() {
219                audio_trace.push(frame.clone());
220                let pcm = self.decode_frame(&frame, cfg.mimi_codebooks)?;
221                decoded_pcm.extend(pcm);
222            }
223        }
224        let text_tokens = state.text_tokens().to_vec();
225        let transcript = self.tokens_to_text(&text_tokens)?;
226        Ok(GenerationResult {
227            samples: decoded_pcm,
228            sample_rate: MIMI_RATE,
229            text_tokens,
230            audio_frames: audio_trace,
231            transcript,
232        })
233    }
234
235    fn run_generation_with_user(
236        &mut self,
237        prompt: &str,
238        user_codes: &MimiCodes,
239        num_frames: usize,
240        cfg: &GenerationConfig,
241    ) -> Result<GenerationResult> {
242        let text_frames = self.tokenizer.prompt_frame_tokens(prompt, num_frames)?;
243        let text_lp = LogitsProcessor::new(cfg.text_temperature, cfg.text_top_k, cfg.text_seed);
244        let audio_lp = LogitsProcessor::new(cfg.audio_temperature, cfg.audio_top_k, cfg.audio_seed);
245        let mut state = self.lm.new_gen_state(
246            num_frames,
247            text_lp.clone(),
248            audio_lp.clone(),
249            self.gen_cfg.clone(),
250        )?;
251        state.reset(&mut self.lm, num_frames, text_lp, audio_lp)?;
252        let pad = state.config().audio_pad_token();
253        let mut decoded_pcm = Vec::new();
254        let mut audio_trace = Vec::new();
255        for step in 0..num_frames {
256            let user_frame: Vec<u32> = if step < user_codes.frames.len() {
257                user_codes.frames[step].clone()
258            } else {
259                vec![pad; cfg.mimi_codebooks]
260            };
261            state.step(&mut self.lm, text_frames[step], &user_frame)?;
262            if let Some(frame) = state.last_audio_frame() {
263                audio_trace.push(frame.clone());
264                let pcm = self.decode_frame(&frame, cfg.mimi_codebooks)?;
265                decoded_pcm.extend(pcm);
266            }
267        }
268        let text_tokens = state.text_tokens().to_vec();
269        let transcript = self.tokens_to_text(&text_tokens)?;
270        Ok(GenerationResult {
271            samples: decoded_pcm,
272            sample_rate: MIMI_RATE,
273            text_tokens,
274            audio_frames: audio_trace,
275            transcript,
276        })
277    }
278
279    fn decode_frame(&mut self, frame: &[u32], nq: usize) -> Result<Vec<f32>> {
280        let codes = MimiCodes {
281            frames: vec![frame.to_vec()],
282            num_quantizers: nq,
283        };
284        self.mimi.decode_codes(&codes)
285    }
286
287    fn tokens_to_text(&self, tokens: &[u32]) -> Result<String> {
288        let mut out = String::new();
289        let cfg = &self.gen_cfg;
290        let mut prev = cfg.text_pad_token;
291        for &t in tokens {
292            if t != cfg.text_start_token
293                && t != cfg.text_pad_token
294                && t != cfg.text_eop_token
295                && prev == cfg.text_start_token
296            {
297                out.push_str(&self.tokenizer.decode_piece(t)?);
298            }
299            prev = t;
300        }
301        Ok(out)
302    }
303}