Skip to main content

rlx_mimi/
codec.rs

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    // Compiled encode/decode graphs, cached by input length. Unused when
42    // `device == Cpu` (which runs the exact ndarray path).
43    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        // Non-CPU devices run the rlx-runtime graph; fall back to CPU eager if
73        // the requested backend isn't compiled in / available on this host.
74        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    /// Run SEANet-encoder + transformer + downsample on the selected backend,
111    /// returning the pre-quantization latent `[hidden, t_ds]`.
112    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    /// Run upsample + transformer + SEANet-decoder on the selected backend,
126    /// returning the waveform `[1, out_len]`.
127    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    /// Encode mono PCM @ [`SAMPLE_RATE`].
154    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    /// Decode codec frames → mono PCM @ [`SAMPLE_RATE`].
169    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
298/// Raw safetensors map: tensor name -> `(data, shape)`.
299type 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
322/// Unified [`rlx_core::AudioCodec`] view so TTS/ASR consumers can use Mimi
323/// interchangeably with other codecs (bitrate control + resampling for free).
324impl 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        // Inherent `MimiCodec::encode_pcm` (inherent methods take resolution
342        // priority over the trait method of the same name).
343        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}