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#[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#[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
50pub 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
62pub 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 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 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}