1use crate::audio::{load_wav_mono, write_wav_mono};
2use crate::codes::MimiCodes;
3use crate::config::MimiConfig;
4use crate::graph::{CodecGraph, DecodeWeights, EncodeWeights};
5use crate::layout::{ct_to_tc, tc_to_ct};
6use crate::rvq::SplitRvq;
7use crate::rvq::build_split_rvq;
8use crate::seanet::{
9 FrameRateDownsample, FrameRateUpsample, SeanetDecoder, SeanetEncoder, build_decoder,
10 build_encoder,
11};
12use crate::transformer::{MimiTransformer, build_transformer};
13use anyhow::{Context, Result, ensure};
14use ndarray::Array2;
15use rlx_core::safetensors_checkpoint::SafetensorsCheckpoint;
16use rlx_runtime::{Device, is_available};
17use std::cell::RefCell;
18use std::collections::HashMap;
19use std::path::{Path, PathBuf};
20use std::time::Instant;
21
22pub const SAMPLE_RATE: u32 = 24_000;
23pub const FRAME_RATE: f32 = 12.5;
24
25struct EagerCodec {
26 cfg: MimiConfig,
27 encoder: SeanetEncoder,
28 encoder_transformer: MimiTransformer,
29 downsample: FrameRateDownsample,
30 quantizer: SplitRvq,
31 upsample: FrameRateUpsample,
32 decoder_transformer: MimiTransformer,
33 decoder: SeanetDecoder,
34}
35
36pub struct MimiCodec {
37 model_dir: PathBuf,
38 cfg: MimiConfig,
39 device: Device,
40 eager: EagerCodec,
41 enc_graphs: RefCell<HashMap<usize, CodecGraph>>,
44 dec_graphs: RefCell<HashMap<usize, CodecGraph>>,
45}
46
47#[derive(Debug, Clone)]
48pub struct RoundtripStats {
49 pub encode_ms: f64,
50 pub decode_ms: f64,
51 pub num_frames: usize,
52 pub pcm_samples: usize,
53}
54
55impl MimiCodec {
56 pub fn open(model_dir: impl AsRef<Path>) -> Result<Self> {
57 Self::open_on(model_dir, Device::Cpu)
58 }
59
60 pub fn open_on(model_dir: impl AsRef<Path>, device: Device) -> Result<Self> {
61 Self::open_on_with_moshi(model_dir, None, device, None)
62 }
63
64 pub fn open_on_with_moshi(
65 model_dir: impl AsRef<Path>,
66 _moshi_dir: Option<&Path>,
67 device: Device,
68 _mimi_codebooks: Option<usize>,
69 ) -> Result<Self> {
70 let model_dir = model_dir.as_ref().to_path_buf();
71 let cfg = MimiConfig::load(&model_dir)?;
72 let actual_device = if device == Device::Cpu || is_available(device) {
75 device
76 } else {
77 eprintln!("mimi: {device:?} not available — using CPU");
78 Device::Cpu
79 };
80 let eager = EagerCodec::open(&model_dir, &cfg)?;
81 Ok(Self {
82 model_dir,
83 cfg,
84 device: actual_device,
85 eager,
86 enc_graphs: RefCell::new(HashMap::new()),
87 dec_graphs: RefCell::new(HashMap::new()),
88 })
89 }
90
91 fn encode_weights(&self) -> EncodeWeights<'_> {
92 EncodeWeights {
93 encoder: &self.eager.encoder,
94 transformer: &self.eager.encoder_transformer,
95 downsample: &self.eager.downsample,
96 audio_channels: self.cfg.audio_channels,
97 hidden_size: self.cfg.hidden_size,
98 }
99 }
100
101 fn decode_weights(&self) -> DecodeWeights<'_> {
102 DecodeWeights {
103 upsample: &self.eager.upsample,
104 transformer: &self.eager.decoder_transformer,
105 decoder: &self.eager.decoder,
106 hidden_size: self.cfg.hidden_size,
107 }
108 }
109
110 fn run_encode(&self, pcm: &[f32]) -> Result<Array2<f32>> {
113 let in_len = pcm.len();
114 let mut cache = self.enc_graphs.borrow_mut();
115 let graph = match cache.get_mut(&in_len) {
116 Some(g) => g,
117 None => {
118 let g = CodecGraph::encoder(self.device, &self.encode_weights(), in_len)?;
119 cache.entry(in_len).or_insert(g)
120 }
121 };
122 graph.run(pcm)
123 }
124
125 fn run_decode(&self, emb: &Array2<f32>) -> Result<Array2<f32>> {
128 let in_t = emb.shape()[1];
129 let mut cache = self.dec_graphs.borrow_mut();
130 let graph = match cache.get_mut(&in_t) {
131 Some(g) => g,
132 None => {
133 let g = CodecGraph::decoder(self.device, &self.decode_weights(), in_t)?;
134 cache.entry(in_t).or_insert(g)
135 }
136 };
137 let flat: Vec<f32> = emb.iter().copied().collect();
138 graph.run(&flat)
139 }
140
141 pub fn config(&self) -> &MimiConfig {
142 &self.cfg
143 }
144
145 pub fn model_dir(&self) -> &Path {
146 &self.model_dir
147 }
148
149 pub fn device(&self) -> Device {
150 self.device
151 }
152
153 pub fn encode_pcm(&self, pcm: &[f32], num_quantizers: Option<usize>) -> Result<MimiCodes> {
155 ensure!(!pcm.is_empty(), "empty PCM");
156 if self.device == Device::Cpu {
157 return self.eager.encode_pcm(pcm, num_quantizers);
158 }
159 let ds = self.run_encode(pcm)?;
160 let nq = num_quantizers.unwrap_or(self.cfg.num_quantizers);
161 let frames = self.eager.quantizer.encode_frames(&ds, Some(nq));
162 Ok(MimiCodes {
163 frames,
164 num_quantizers: nq,
165 })
166 }
167
168 pub fn decode_codes(&self, codes: &MimiCodes) -> Result<Vec<f32>> {
170 ensure!(!codes.frames.is_empty(), "empty codec frames");
171 if self.device == Device::Cpu {
172 return self.eager.decode_codes(codes);
173 }
174 let emb = self.eager.quantizer.decode_frames(&codes.frames);
175 let wav = self.run_decode(&emb)?;
176 ensure!(wav.dim().0 >= 1, "decoder produced no channels");
177 Ok(wav.row(0).to_vec())
178 }
179
180 pub fn encode_wav(
181 &self,
182 wav: impl AsRef<Path>,
183 num_quantizers: Option<usize>,
184 ) -> Result<MimiCodes> {
185 let pcm = load_wav_mono(wav.as_ref(), SAMPLE_RATE)?;
186 self.encode_pcm(&pcm, num_quantizers)
187 }
188
189 pub fn decode_wav(
190 &self,
191 codes: &MimiCodes,
192 out: impl AsRef<Path>,
193 trim_to_samples: Option<usize>,
194 ) -> Result<()> {
195 let mut pcm = self.decode_codes(codes)?;
196 if let Some(n) = trim_to_samples {
197 pcm.truncate(n.min(pcm.len()));
198 }
199 write_wav_mono(out.as_ref(), &pcm, SAMPLE_RATE)
200 }
201
202 pub fn roundtrip_pcm(
203 &self,
204 pcm: &[f32],
205 num_quantizers: Option<usize>,
206 ) -> Result<(MimiCodes, Vec<f32>, RoundtripStats)> {
207 let t0 = Instant::now();
208 let codes = self.encode_pcm(pcm, num_quantizers)?;
209 let num_frames = codes.num_frames();
210 let encode_ms = t0.elapsed().as_secs_f64() * 1000.0;
211 let t1 = Instant::now();
212 let mut recon = self.decode_codes(&codes)?;
213 let decode_ms = t1.elapsed().as_secs_f64() * 1000.0;
214 recon.truncate(pcm.len().min(recon.len()));
215 Ok((
216 codes,
217 recon,
218 RoundtripStats {
219 encode_ms,
220 decode_ms,
221 num_frames,
222 pcm_samples: pcm.len(),
223 },
224 ))
225 }
226
227 pub fn roundtrip_wav(
228 &self,
229 in_wav: impl AsRef<Path>,
230 out_wav: impl AsRef<Path>,
231 num_quantizers: Option<usize>,
232 ) -> Result<MimiCodes> {
233 let pcm = load_wav_mono(in_wav.as_ref(), SAMPLE_RATE)?;
234 let (codes, recon, _) = self.roundtrip_pcm(&pcm, num_quantizers)?;
235 write_wav_mono(out_wav.as_ref(), &recon, SAMPLE_RATE)?;
236 Ok(codes)
237 }
238}
239
240impl EagerCodec {
241 fn open(model_dir: &Path, cfg: &MimiConfig) -> Result<Self> {
242 let map = load_weight_map(model_dir)?;
243 let mut aux: HashMap<String, _> = map
244 .iter()
245 .filter(|(k, _)| k.starts_with("downsample.") || k.starts_with("upsample."))
246 .map(|(k, v)| (k.clone(), v.clone()))
247 .collect();
248 Ok(Self {
249 cfg: cfg.clone(),
250 encoder: build_encoder(cfg, subset_prefix(&map, "encoder."))?,
251 encoder_transformer: build_transformer(
252 cfg,
253 "encoder_transformer",
254 subset_prefix(&map, "encoder_transformer."),
255 )?,
256 downsample: FrameRateDownsample::from_weights(cfg, &mut aux)?,
257 quantizer: build_split_rvq(cfg, subset_prefix(&map, "quantizer."))?,
258 upsample: FrameRateUpsample::from_weights(cfg, &mut aux)?,
259 decoder_transformer: build_transformer(
260 cfg,
261 "decoder_transformer",
262 subset_prefix(&map, "decoder_transformer."),
263 )?,
264 decoder: build_decoder(cfg, subset_prefix(&map, "decoder."))?,
265 })
266 }
267
268 fn encode_pcm(&self, pcm: &[f32], num_quantizers: Option<usize>) -> Result<MimiCodes> {
269 let mut input = Array2::<f32>::zeros((self.cfg.audio_channels, pcm.len()));
270 for (i, &s) in pcm.iter().enumerate() {
271 input[[0, i]] = s;
272 }
273 let conv = self.encoder.forward(input.view());
274 let pre_tf = ct_to_tc(conv.view());
275 let post_tf = self.encoder_transformer.forward(pre_tf.view());
276 let pre_ds = tc_to_ct(post_tf.view());
277 let ds = self.downsample.forward(pre_ds.view());
278 let nq = num_quantizers.unwrap_or(self.cfg.num_quantizers);
279 let frames = self.quantizer.encode_frames(&ds, Some(nq));
280 Ok(MimiCodes {
281 frames,
282 num_quantizers: nq,
283 })
284 }
285
286 fn decode_codes(&self, codes: &MimiCodes) -> Result<Vec<f32>> {
287 let emb = self.quantizer.decode_frames(&codes.frames);
288 let up = self.upsample.forward(emb.view());
289 let pre_tf = ct_to_tc(up.view());
290 let post_tf = self.decoder_transformer.forward(pre_tf.view());
291 let pre_dec = tc_to_ct(post_tf.view());
292 let wav = self.decoder.forward(pre_dec.view());
293 ensure!(wav.dim().0 >= 1, "decoder produced no channels");
294 Ok(wav.row(0).to_vec())
295 }
296}
297
298type RawTensorMap = HashMap<String, (Vec<f32>, Vec<usize>)>;
300
301fn load_weight_map(model_dir: &Path) -> Result<RawTensorMap> {
302 let ckpt = SafetensorsCheckpoint::open(model_dir)?;
303 let keys: std::collections::HashSet<String> = ckpt.keys().map(str::to_string).collect();
304 let mut wm = ckpt.load_selected(&keys)?;
305 let mut map = HashMap::with_capacity(keys.len());
306 for key in keys {
307 let (data, shape) = wm
308 .take(&key)
309 .with_context(|| format!("tensor {key} missing after load"))?;
310 map.insert(key, (data, shape));
311 }
312 Ok(map)
313}
314
315fn subset_prefix(map: &RawTensorMap, prefix: &str) -> RawTensorMap {
316 map.iter()
317 .filter(|(k, _)| k.starts_with(prefix))
318 .map(|(k, v)| (k.clone(), v.clone()))
319 .collect()
320}
321
322impl rlx_core::AudioCodec for MimiCodec {
325 fn info(&self) -> rlx_core::CodecInfo {
326 rlx_core::CodecInfo {
327 sample_rate: self.cfg.sampling_rate,
328 frame_rate: self.cfg.frame_rate,
329 hop_length: self.cfg.samples_per_codec_frame(),
330 channels: self.cfg.audio_channels,
331 max_quantizers: self.cfg.num_quantizers,
332 codebook_size: self.cfg.codebook_size,
333 }
334 }
335
336 fn device(&self) -> Device {
337 self.device
338 }
339
340 fn encode_pcm(&self, pcm: &[f32], num_quantizers: Option<usize>) -> Result<rlx_core::RvqCodes> {
341 let codes = MimiCodec::encode_pcm(self, pcm, num_quantizers)?;
344 Ok(rlx_core::RvqCodes::new(codes.frames, codes.num_quantizers))
345 }
346
347 fn decode_codes(&self, codes: &rlx_core::RvqCodes) -> Result<Vec<f32>> {
348 let mc = MimiCodes {
349 frames: codes.frames.clone(),
350 num_quantizers: codes.num_quantizers,
351 };
352 MimiCodec::decode_codes(self, &mc)
353 }
354}