Skip to main content

rlx_mimi/
config.rs

1use anyhow::{Context, Result};
2use serde::Deserialize;
3use std::path::Path;
4
5/// Flat HF `config.json` for [`kyutai/mimi`](https://huggingface.co/kyutai/mimi).
6#[derive(Debug, Clone, Deserialize)]
7pub struct MimiConfig {
8    pub audio_channels: usize,
9    pub num_filters: usize,
10    pub kernel_size: usize,
11    pub last_kernel_size: usize,
12    pub hidden_size: usize,
13    pub upsampling_ratios: Vec<usize>,
14    pub num_residual_layers: usize,
15    pub residual_kernel_size: usize,
16    pub dilation_growth_rate: usize,
17    pub compress: usize,
18    pub codebook_dim: usize,
19    pub codebook_size: usize,
20    #[serde(default = "default_num_quantizers")]
21    pub num_quantizers: usize,
22    pub num_semantic_quantizers: usize,
23    pub vector_quantization_hidden_dimension: usize,
24    pub upsample_groups: usize,
25    pub num_hidden_layers: usize,
26    pub intermediate_size: usize,
27    pub num_attention_heads: usize,
28    pub num_key_value_heads: usize,
29    pub head_dim: usize,
30    pub sliding_window: usize,
31    #[serde(default = "default_norm_eps")]
32    pub norm_eps: f32,
33    #[serde(default = "default_rope_theta")]
34    pub rope_theta: f64,
35    #[serde(default = "default_layer_scale")]
36    pub layer_scale_initial_scale: f32,
37    pub sampling_rate: u32,
38    #[serde(default = "default_frame_rate")]
39    pub frame_rate: f32,
40    #[serde(default = "default_trim_right_ratio")]
41    pub trim_right_ratio: f32,
42}
43
44fn default_num_quantizers() -> usize {
45    32
46}
47fn default_norm_eps() -> f32 {
48    1e-5
49}
50fn default_rope_theta() -> f64 {
51    10_000.0
52}
53fn default_layer_scale() -> f32 {
54    0.01
55}
56fn default_frame_rate() -> f32 {
57    12.5
58}
59fn default_trim_right_ratio() -> f32 {
60    1.0
61}
62
63impl MimiConfig {
64    pub fn load(model_dir: &Path) -> Result<Self> {
65        let path = model_dir.join("config.json");
66        let text =
67            std::fs::read_to_string(&path).with_context(|| format!("read {}", path.display()))?;
68        serde_json::from_str(&text).with_context(|| format!("parse {}", path.display()))
69    }
70
71    /// SEANet hop length before the stride-2 frame-rate downsample.
72    pub fn seanet_hop(&self) -> usize {
73        self.upsampling_ratios.iter().product()
74    }
75
76    /// PCM samples per one codec frame @ [`Self::frame_rate`] Hz.
77    pub fn samples_per_codec_frame(&self) -> usize {
78        self.seanet_hop() * 2
79    }
80
81    /// Approximate encoded frame count for `pcm_len` samples (matches HF `get_encoded_length`).
82    pub fn encoded_frame_count(&self, pcm_len: usize) -> usize {
83        if pcm_len == 0 {
84            return 0;
85        }
86        pcm_len.div_ceil(self.samples_per_codec_frame())
87    }
88
89    pub fn num_acoustic_quantizers(&self) -> usize {
90        self.num_quantizers
91            .saturating_sub(self.num_semantic_quantizers)
92    }
93
94    /// Nominal bitrate in bits per second (codebooks × log2(codebook_size) × frame_rate).
95    pub fn bitrate_bps(&self) -> f32 {
96        let bits_per_frame = self.num_quantizers as f32 * (self.codebook_size as f32).log2();
97        bits_per_frame * self.frame_rate
98    }
99
100    /// Kernel size for the frame-rate `downsample` / `upsample` convs.
101    pub fn frame_rate_kernel(&self) -> usize {
102        let encodec_frame_rate = self.sampling_rate as f32 / self.seanet_hop() as f32;
103        let ratio = (encodec_frame_rate / self.frame_rate) as usize;
104        2 * ratio.max(1)
105    }
106}