1use anyhow::{Context, Result, ensure};
19use rlx_llama32::Llama32Config;
20use serde::Deserialize;
21use std::path::Path;
22
23#[derive(Debug, Clone, Deserialize)]
25pub struct VoxtralAudioConfig {
26 pub num_mel_bins: usize,
27 pub max_source_positions: usize,
28 #[serde(rename = "hidden_size", alias = "d_model")]
29 pub d_model: usize,
30 #[serde(rename = "num_attention_heads", alias = "encoder_attention_heads")]
31 pub encoder_attention_heads: usize,
32 #[serde(rename = "num_hidden_layers", alias = "encoder_layers")]
33 pub encoder_layers: usize,
34 pub intermediate_size: usize,
35 #[serde(default)]
36 pub scale_embedding: bool,
37}
38
39impl VoxtralAudioConfig {
40 pub fn head_dim(&self) -> usize {
41 self.d_model / self.encoder_attention_heads
42 }
43
44 pub fn encoder_seq_len(&self, mel_frames: usize) -> usize {
46 let after_conv1 = mel_frames;
47 let pad = 1usize;
48 let k = 3usize;
49 let stride2 = 2usize;
50 (after_conv1 + 2 * pad - k) / stride2 + 1
51 }
52
53 pub fn audio_token_count(&self, mel_frames: usize) -> usize {
55 self.encoder_seq_len(mel_frames) / 4
56 }
57
58 pub fn tiny_synthetic() -> Self {
59 Self {
60 num_mel_bins: 4,
61 max_source_positions: 16,
62 d_model: 8,
63 encoder_attention_heads: 2,
64 encoder_layers: 1,
65 intermediate_size: 32,
66 scale_embedding: false,
67 }
68 }
69
70 pub fn mini_3b() -> Self {
71 Self {
72 num_mel_bins: 128,
73 max_source_positions: 1500,
74 d_model: 1280,
75 encoder_attention_heads: 20,
76 encoder_layers: 32,
77 intermediate_size: 5120,
78 scale_embedding: false,
79 }
80 }
81}
82
83#[derive(Debug, Clone, Deserialize)]
85pub struct VoxtralConfig {
86 pub audio_config: VoxtralAudioConfig,
87 pub text_config: Llama32Config,
88 #[serde(default = "default_audio_token_id")]
89 pub audio_token_id: u32,
90 #[serde(default = "default_eos_token_id")]
92 pub eos_token_id: u32,
93 #[serde(default = "default_projector_act")]
94 pub projector_hidden_act: String,
95 pub vocab_size: usize,
96}
97
98fn default_audio_token_id() -> u32 {
99 24
100}
101
102fn default_eos_token_id() -> u32 {
103 2
104}
105
106fn default_projector_act() -> String {
107 "gelu".into()
108}
109
110impl VoxtralConfig {
111 pub fn from_file(path: &Path) -> Result<Self> {
112 let data = std::fs::read_to_string(path)?;
113 serde_json::from_str(&data).with_context(|| format!("parse Voxtral config {path:?}"))
114 }
115
116 pub fn llama_config(&self) -> &Llama32Config {
117 &self.text_config
118 }
119
120 pub fn validate(&self) -> Result<()> {
121 ensure!(
122 self.text_config.hidden_size > 0,
123 "text_config.hidden_size must be > 0"
124 );
125 ensure!(
126 self.audio_config.intermediate_size == self.audio_config.d_model * 4,
127 "audio_config.intermediate_size should be 4× d_model for the projector reshape"
128 );
129 Ok(())
130 }
131
132 pub fn tiny_synthetic() -> Self {
133 Self {
134 audio_config: VoxtralAudioConfig::tiny_synthetic(),
135 text_config: Llama32Config {
136 vocab_size: 32,
137 hidden_size: 16,
138 intermediate_size: 32,
139 num_hidden_layers: 1,
140 num_attention_heads: 4,
141 num_key_value_heads: 2,
142 max_position_embeddings: 16,
143 rms_norm_eps: 1e-5,
144 rope_theta: 100_000_000.0,
145 hidden_act: "silu".into(),
146 tie_word_embeddings: true,
147 attention_bias: false,
148 head_dim: Some(4),
149 rope_scaling: None,
150 },
151 audio_token_id: 24,
152 eos_token_id: 2,
153 projector_hidden_act: "gelu".into(),
154 vocab_size: 32,
155 }
156 }
157
158 pub fn mini_3b() -> Self {
159 Self {
160 audio_config: VoxtralAudioConfig::mini_3b(),
161 text_config: Llama32Config {
162 vocab_size: 131_072,
163 hidden_size: 3072,
164 intermediate_size: 8192,
165 num_hidden_layers: 30,
166 num_attention_heads: 32,
167 num_key_value_heads: 8,
168 max_position_embeddings: 131_072,
169 rms_norm_eps: 1e-5,
170 rope_theta: 100_000_000.0,
171 hidden_act: "silu".into(),
172 tie_word_embeddings: true,
173 attention_bias: false,
174 head_dim: Some(128),
175 rope_scaling: None,
176 },
177 audio_token_id: 24,
178 eos_token_id: 2,
179 projector_hidden_act: "gelu".into(),
180 vocab_size: 131_072,
181 }
182 }
183}