1use anyhow::{Context, Result};
2use serde::Deserialize;
3use std::path::Path;
4
5#[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 pub fn seanet_hop(&self) -> usize {
73 self.upsampling_ratios.iter().product()
74 }
75
76 pub fn samples_per_codec_frame(&self) -> usize {
78 self.seanet_hop() * 2
79 }
80
81 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 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 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}