1use serde::de::DeserializeOwned;
2use std::{
3 collections::{HashMap, HashSet},
4 f32::consts::PI,
5 path::{Path, PathBuf},
6};
7use thiserror::Error;
8
9#[derive(Debug, Error)]
10pub enum SpeechLoaderError {
11 #[error("io error: {0}")]
12 Io(String),
13 #[error("invalid model metadata: {0}")]
14 Metadata(String),
15 #[error("model weights are missing: {0}")]
16 MissingWeights(String),
17 #[error("audio input is invalid: {0}")]
18 InvalidAudio(String),
19 #[error("MLX runtime failed: {0}")]
20 Mlx(String),
21}
22
23#[derive(Debug, Clone, PartialEq, Eq)]
24pub struct SpeechModelFiles {
25 pub model_dir: PathBuf,
26 pub config_path: PathBuf,
27 pub tokenizer_path: Option<PathBuf>,
28 pub safetensors_files: Vec<PathBuf>,
29}
30
31impl SpeechModelFiles {
32 pub fn discover(model_dir: impl Into<PathBuf>) -> Result<Self, SpeechLoaderError> {
33 let model_dir = model_dir.into();
34 let config_path = model_dir.join("config.json");
35 if !config_path.is_file() {
36 return Err(SpeechLoaderError::Metadata(format!(
37 "missing config.json in {}",
38 model_dir.display()
39 )));
40 }
41
42 let tokenizer_json = model_dir.join("tokenizer.json");
43 let tokenizer_path = tokenizer_json.is_file().then_some(tokenizer_json);
44 let safetensors_files = collect_safetensors_files(&model_dir)?;
45
46 Ok(Self {
47 model_dir,
48 config_path,
49 tokenizer_path,
50 safetensors_files,
51 })
52 }
53}
54
55#[derive(Debug, Clone, PartialEq, Eq)]
56pub struct WeightMapIndex {
57 pub weight_map: HashMap<String, String>,
58}
59
60impl serde::Serialize for WeightMapIndex {
61 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
62 where
63 S: serde::Serializer,
64 {
65 let mut root = serde_json::Map::new();
66 root.insert(
67 "weight_map".to_string(),
68 serde_json::to_value(&self.weight_map).map_err(serde::ser::Error::custom)?,
69 );
70 root.serialize(serializer)
71 }
72}
73
74impl<'de> serde::Deserialize<'de> for WeightMapIndex {
75 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
76 where
77 D: serde::Deserializer<'de>,
78 {
79 let value = serde_json::Value::deserialize(deserializer)?;
80 let weight_map = value
81 .get("weight_map")
82 .ok_or_else(|| serde::de::Error::custom("missing weight_map"))?;
83 Ok(Self {
84 weight_map: serde_json::from_value(weight_map.clone())
85 .map_err(serde::de::Error::custom)?,
86 })
87 }
88}
89
90pub fn read_json<T: DeserializeOwned>(path: &Path) -> Result<T, SpeechLoaderError> {
91 let text = std::fs::read_to_string(path).map_err(|error| {
92 SpeechLoaderError::Io(format!("failed to read {}: {error}", path.display()))
93 })?;
94 serde_json::from_str(&text).map_err(|error| {
95 SpeechLoaderError::Metadata(format!("failed to parse {}: {error}", path.display()))
96 })
97}
98
99pub fn collect_safetensors_files(model_dir: &Path) -> Result<Vec<PathBuf>, SpeechLoaderError> {
100 let index_path = model_dir.join("model.safetensors.index.json");
101 if index_path.is_file() {
102 let index: WeightMapIndex = read_json(&index_path)?;
103 let mut names: Vec<String> = index
104 .weight_map
105 .values()
106 .collect::<HashSet<_>>()
107 .into_iter()
108 .cloned()
109 .collect();
110 names.sort();
111 return Ok(names.into_iter().map(|name| model_dir.join(name)).collect());
112 }
113
114 let single_path = model_dir.join("model.safetensors");
115 if single_path.is_file() {
116 return Ok(vec![single_path]);
117 }
118
119 Err(SpeechLoaderError::MissingWeights(format!(
120 "no model.safetensors or model.safetensors.index.json in {}",
121 model_dir.display()
122 )))
123}
124
125#[cfg(feature = "native-mlx")]
126pub fn load_safetensors(
127 files: &[PathBuf],
128) -> Result<HashMap<String, mlx_rs::Array>, SpeechLoaderError> {
129 use std::io::{Read, Seek};
130
131 #[derive(Debug, serde::Deserialize)]
132 struct SafeTensorHeaderEntry {
133 dtype: String,
134 shape: Vec<i32>,
135 data_offsets: [u64; 2],
136 }
137
138 fn expected_len(shape: &[i32], name: &str, path: &Path) -> Result<usize, SpeechLoaderError> {
139 shape
140 .iter()
141 .try_fold(1_usize, |total, dim| {
142 usize::try_from(*dim)
143 .ok()
144 .and_then(|dim| total.checked_mul(dim))
145 })
146 .ok_or_else(|| {
147 SpeechLoaderError::MissingWeights(format!(
148 "invalid safetensors shape for {name} in {}",
149 path.display()
150 ))
151 })
152 }
153
154 let mut weights = HashMap::new();
155 for file in files {
156 let mut reader = std::fs::File::open(file).map_err(|error| {
157 SpeechLoaderError::Io(format!("failed to open {}: {error}", file.display()))
158 })?;
159 let mut len_bytes = [0_u8; 8];
160 reader.read_exact(&mut len_bytes).map_err(|error| {
161 SpeechLoaderError::Io(format!(
162 "failed to read safetensors header length from {}: {error}",
163 file.display()
164 ))
165 })?;
166 let header_len = u64::from_le_bytes(len_bytes);
167 if header_len > 128 * 1024 * 1024 {
168 return Err(SpeechLoaderError::Metadata(format!(
169 "safetensors header in {} is unexpectedly large",
170 file.display()
171 )));
172 }
173
174 let mut header = vec![0_u8; header_len as usize];
175 reader.read_exact(&mut header).map_err(|error| {
176 SpeechLoaderError::Io(format!(
177 "failed to read safetensors header from {}: {error}",
178 file.display()
179 ))
180 })?;
181 let header = serde_json::from_slice::<serde_json::Map<String, serde_json::Value>>(&header)
182 .map_err(|error| {
183 SpeechLoaderError::Metadata(format!(
184 "failed to parse safetensors header in {}: {error}",
185 file.display()
186 ))
187 })?;
188
189 for (name, value) in header {
190 if name == "__metadata__" {
191 continue;
192 }
193 let entry =
194 serde_json::from_value::<SafeTensorHeaderEntry>(value).map_err(|error| {
195 SpeechLoaderError::Metadata(format!(
196 "failed to parse safetensors entry {name} in {}: {error}",
197 file.display()
198 ))
199 })?;
200 let byte_len = entry.data_offsets[1]
201 .checked_sub(entry.data_offsets[0])
202 .ok_or_else(|| {
203 SpeechLoaderError::MissingWeights(format!(
204 "invalid safetensors offsets for {name} in {}",
205 file.display()
206 ))
207 })?;
208 reader
209 .seek(std::io::SeekFrom::Start(
210 8 + header_len + entry.data_offsets[0],
211 ))
212 .map_err(|error| {
213 SpeechLoaderError::Io(format!(
214 "failed to seek tensor {name} in {}: {error}",
215 file.display()
216 ))
217 })?;
218 let mut data = vec![0_u8; byte_len as usize];
219 reader.read_exact(&mut data).map_err(|error| {
220 SpeechLoaderError::Io(format!(
221 "failed to read tensor {name} in {}: {error}",
222 file.display()
223 ))
224 })?;
225 let expected = expected_len(&entry.shape, &name, file)?;
226 let array = match entry.dtype.as_str() {
227 "F32" => {
228 let values = read_le_chunks::<4, f32>(&data, expected, &name, file, |bytes| {
229 f32::from_le_bytes(bytes)
230 })?;
231 mlx_rs::Array::from_slice(&values, &entry.shape)
232 }
233 "F16" => {
234 let values =
235 read_le_chunks::<2, half::f16>(&data, expected, &name, file, |bytes| {
236 half::f16::from_bits(u16::from_le_bytes(bytes))
237 })?;
238 mlx_rs::Array::from_slice(&values, &entry.shape)
239 }
240 "BF16" => {
241 let values =
242 read_le_chunks::<2, half::bf16>(&data, expected, &name, file, |bytes| {
243 half::bf16::from_bits(u16::from_le_bytes(bytes))
244 })?;
245 mlx_rs::Array::from_slice(&values, &entry.shape)
246 }
247 "I32" => {
248 let values = read_le_chunks::<4, i32>(&data, expected, &name, file, |bytes| {
249 i32::from_le_bytes(bytes)
250 })?;
251 mlx_rs::Array::from_slice(&values, &entry.shape)
252 }
253 "I64" => {
254 let values = read_le_chunks::<8, i64>(&data, expected, &name, file, |bytes| {
255 i64::from_le_bytes(bytes)
256 })?;
257 mlx_rs::Array::from_slice(&values, &entry.shape)
258 }
259 "U32" => {
260 let values = read_le_chunks::<4, u32>(&data, expected, &name, file, |bytes| {
261 u32::from_le_bytes(bytes)
262 })?;
263 mlx_rs::Array::from_slice(&values, &entry.shape)
264 }
265 "U8" => {
266 if data.len() != expected {
267 return Err(SpeechLoaderError::MissingWeights(format!(
268 "safetensors tensor {name} in {} has {} values, expected {expected}",
269 file.display(),
270 data.len()
271 )));
272 }
273 mlx_rs::Array::from_slice(&data, &entry.shape)
274 }
275 other => {
276 return Err(SpeechLoaderError::MissingWeights(format!(
277 "unsupported safetensors dtype {other} for {name} in {}",
278 file.display()
279 )));
280 }
281 };
282 weights.insert(name, array);
283 }
284 }
285 Ok(weights)
286}
287
288#[cfg(feature = "native-mlx")]
289fn read_le_chunks<const N: usize, T>(
290 data: &[u8],
291 expected: usize,
292 name: &str,
293 path: &Path,
294 convert: impl Fn([u8; N]) -> T,
295) -> Result<Vec<T>, SpeechLoaderError> {
296 if data.len() != expected.checked_mul(N).unwrap_or(usize::MAX) {
297 return Err(SpeechLoaderError::MissingWeights(format!(
298 "safetensors tensor {name} in {} has {} bytes, expected {}",
299 path.display(),
300 data.len(),
301 expected * N
302 )));
303 }
304 Ok(data
305 .chunks_exact(N)
306 .map(|chunk| convert(chunk.try_into().expect("chunk size is fixed")))
307 .collect())
308}
309
310#[cfg(feature = "native-mlx")]
311pub fn safetensors_file_contains_dtype(
312 path: &Path,
313 dtype: &str,
314) -> Result<bool, SpeechLoaderError> {
315 use std::io::Read;
316
317 let mut file = std::fs::File::open(path).map_err(|error| {
318 SpeechLoaderError::Io(format!("failed to open {}: {error}", path.display()))
319 })?;
320 let mut len_bytes = [0_u8; 8];
321 file.read_exact(&mut len_bytes).map_err(|error| {
322 SpeechLoaderError::Io(format!(
323 "failed to read safetensors header length from {}: {error}",
324 path.display()
325 ))
326 })?;
327 let header_len = u64::from_le_bytes(len_bytes);
328 if header_len > 128 * 1024 * 1024 {
329 return Err(SpeechLoaderError::Metadata(format!(
330 "safetensors header in {} is unexpectedly large",
331 path.display()
332 )));
333 }
334 let mut header = vec![0_u8; header_len as usize];
335 file.read_exact(&mut header).map_err(|error| {
336 SpeechLoaderError::Io(format!(
337 "failed to read safetensors header from {}: {error}",
338 path.display()
339 ))
340 })?;
341 let header = serde_json::from_slice::<serde_json::Value>(&header).map_err(|error| {
342 SpeechLoaderError::Metadata(format!(
343 "failed to parse safetensors header in {}: {error}",
344 path.display()
345 ))
346 })?;
347 let Some(tensors) = header.as_object() else {
348 return Ok(false);
349 };
350 Ok(tensors.iter().any(|(key, value)| {
351 key != "__metadata__"
352 && value
353 .get("dtype")
354 .and_then(serde_json::Value::as_str)
355 .is_some_and(|value| value == dtype)
356 }))
357}
358
359#[cfg(feature = "native-mlx")]
360pub struct MlxWeightView<'a> {
361 weights: &'a HashMap<String, mlx_rs::Array>,
362}
363
364#[cfg(feature = "native-mlx")]
365impl<'a> MlxWeightView<'a> {
366 pub fn new(weights: &'a HashMap<String, mlx_rs::Array>) -> Self {
367 Self { weights }
368 }
369
370 pub fn get(&self, key: &str) -> Result<&'a mlx_rs::Array, SpeechLoaderError> {
371 self.weights
372 .get(key)
373 .ok_or_else(|| SpeechLoaderError::MissingWeights(format!("missing weight {key}")))
374 }
375
376 pub fn get_any(&self, keys: &[&str]) -> Result<&'a mlx_rs::Array, SpeechLoaderError> {
377 keys.iter()
378 .find_map(|key| self.weights.get(*key))
379 .ok_or_else(|| {
380 SpeechLoaderError::MissingWeights(format!(
381 "missing any of weights: {}",
382 keys.join(", ")
383 ))
384 })
385 }
386
387 pub fn optional(&self, key: &str) -> Option<&'a mlx_rs::Array> {
388 self.weights.get(key)
389 }
390}
391
392#[cfg(feature = "native-mlx")]
393pub fn linear(
394 input: &mlx_rs::Array,
395 weight: &mlx_rs::Array,
396 bias: Option<&mlx_rs::Array>,
397) -> Result<mlx_rs::Array, SpeechLoaderError> {
398 let mut output = mlx_rs::ops::matmul(input, weight.t())
399 .map_err(|error| SpeechLoaderError::Mlx(error.to_string()))?;
400 if let Some(bias) = bias {
401 output += bias;
402 }
403 Ok(output)
404}
405
406#[cfg(feature = "native-mlx")]
407pub fn embedding(
408 weight: &mlx_rs::Array,
409 token_ids: &[i32],
410) -> Result<mlx_rs::Array, SpeechLoaderError> {
411 let len = i32::try_from(token_ids.len()).map_err(|_| {
412 SpeechLoaderError::InvalidAudio("token sequence exceeds MLX shape limits".to_string())
413 })?;
414 let indices = mlx_rs::Array::from_slice(token_ids, &[len]);
415 weight
416 .take_axis(&indices, 0)
417 .map_err(|error| SpeechLoaderError::Mlx(error.to_string()))
418}
419
420#[cfg(feature = "native-mlx")]
421pub fn conv1d_pytorch(
422 input_nlc: &mlx_rs::Array,
423 weight_okc: &mlx_rs::Array,
424 bias: Option<&mlx_rs::Array>,
425 stride: i32,
426 padding: i32,
427) -> Result<mlx_rs::Array, SpeechLoaderError> {
428 let weight = weight_okc
429 .transpose_axes(&[0, 2, 1])
430 .map_err(|error| SpeechLoaderError::Mlx(error.to_string()))?;
431 let mut output = mlx_rs::ops::conv1d(input_nlc, &weight, stride, padding, None, None)
432 .map_err(|error| SpeechLoaderError::Mlx(error.to_string()))?;
433 if let Some(bias) = bias {
434 output += bias;
435 }
436 Ok(output)
437}
438
439#[cfg(feature = "native-mlx")]
440pub fn layer_norm(
441 input: &mlx_rs::Array,
442 weight: Option<&mlx_rs::Array>,
443 bias: Option<&mlx_rs::Array>,
444 eps: f32,
445) -> Result<mlx_rs::Array, SpeechLoaderError> {
446 mlx_rs::fast::layer_norm(input, weight, bias, eps)
447 .map_err(|error| SpeechLoaderError::Mlx(error.to_string()))
448}
449
450#[cfg(feature = "native-mlx")]
451pub fn gelu(input: &mlx_rs::Array) -> Result<mlx_rs::Array, SpeechLoaderError> {
452 mlx_rs::nn::gelu(input).map_err(|error| SpeechLoaderError::Mlx(error.to_string()))
453}
454
455#[cfg(feature = "native-mlx")]
456pub fn scaled_dot_product_attention(
457 queries: &mlx_rs::Array,
458 keys: &mlx_rs::Array,
459 values: &mlx_rs::Array,
460 scale: f32,
461) -> Result<mlx_rs::Array, SpeechLoaderError> {
462 mlx_rs::fast::scaled_dot_product_attention(queries, keys, values, scale, None)
463 .map_err(|error| SpeechLoaderError::Mlx(error.to_string()))
464}
465
466#[derive(Debug, Clone, Copy, PartialEq)]
467pub struct MelSpectrogramConfig {
468 pub sample_rate_hz: u32,
469 pub fft_size: usize,
470 pub hop_length: usize,
471 pub mel_bins: usize,
472 pub min_frequency_hz: f32,
473 pub max_frequency_hz: f32,
474}
475
476impl MelSpectrogramConfig {
477 pub fn whisper() -> Self {
478 Self {
479 sample_rate_hz: 16_000,
480 fft_size: 400,
481 hop_length: 160,
482 mel_bins: 80,
483 min_frequency_hz: 0.0,
484 max_frequency_hz: 8_000.0,
485 }
486 }
487}
488
489#[derive(Debug, Clone, PartialEq)]
490pub struct MelSpectrogram {
491 pub frames: usize,
492 pub bins: usize,
493 pub values: Vec<f32>,
494}
495
496pub fn log_mel_spectrogram(
497 samples: &[f32],
498 config: MelSpectrogramConfig,
499) -> Result<MelSpectrogram, SpeechLoaderError> {
500 if samples.is_empty() {
501 return Err(SpeechLoaderError::InvalidAudio(
502 "audio sample buffer is empty".to_string(),
503 ));
504 }
505 if config.fft_size == 0 || config.hop_length == 0 || config.mel_bins == 0 {
506 return Err(SpeechLoaderError::InvalidAudio(
507 "mel spectrogram dimensions must be non-zero".to_string(),
508 ));
509 }
510
511 let padded = if samples.len() < config.fft_size {
512 let mut padded = samples.to_vec();
513 padded.resize(config.fft_size, 0.0);
514 padded
515 } else {
516 samples.to_vec()
517 };
518 let frames = 1 + (padded.len() - config.fft_size) / config.hop_length;
519 let window = hann_window(config.fft_size);
520 let filters = mel_filterbank(config);
521 let fft_bins = config.fft_size / 2 + 1;
522 let mut values = Vec::with_capacity(frames * config.mel_bins);
523
524 for frame_index in 0..frames {
525 let offset = frame_index * config.hop_length;
526 let mut power = vec![0.0_f32; fft_bins];
527 for (bin, power_bin) in power.iter_mut().enumerate() {
528 let mut real = 0.0_f32;
529 let mut imag = 0.0_f32;
530 for n in 0..config.fft_size {
531 let angle = -2.0 * PI * bin as f32 * n as f32 / config.fft_size as f32;
532 let sample = padded[offset + n] * window[n];
533 real += sample * angle.cos();
534 imag += sample * angle.sin();
535 }
536 *power_bin = real.mul_add(real, imag * imag);
537 }
538
539 for mel in 0..config.mel_bins {
540 let mut energy = 0.0_f32;
541 for bin in 0..fft_bins {
542 energy += power[bin] * filters[mel * fft_bins + bin];
543 }
544 values.push(energy.max(1.0e-10).log10());
545 }
546 }
547
548 if let Some(max_log_mel) = values
549 .iter()
550 .copied()
551 .max_by(|left, right| left.partial_cmp(right).unwrap_or(std::cmp::Ordering::Equal))
552 {
553 let floor = max_log_mel - 8.0;
554 for value in &mut values {
555 *value = (value.max(floor) + 4.0) / 4.0;
556 }
557 }
558
559 Ok(MelSpectrogram {
560 frames,
561 bins: config.mel_bins,
562 values,
563 })
564}
565
566fn hann_window(size: usize) -> Vec<f32> {
567 (0..size)
568 .map(|index| 0.5 - 0.5 * (2.0 * PI * index as f32 / size as f32).cos())
569 .collect()
570}
571
572fn mel_filterbank(config: MelSpectrogramConfig) -> Vec<f32> {
573 let fft_bins = config.fft_size / 2 + 1;
574 let min_mel = hz_to_mel(config.min_frequency_hz);
575 let max_mel = hz_to_mel(config.max_frequency_hz);
576 let mel_points: Vec<f32> = (0..config.mel_bins + 2)
577 .map(|i| min_mel + (max_mel - min_mel) * i as f32 / (config.mel_bins + 1) as f32)
578 .map(mel_to_hz)
579 .collect();
580
581 let bin_points: Vec<usize> = mel_points
582 .iter()
583 .map(|hz| ((config.fft_size as f32 + 1.0) * hz / config.sample_rate_hz as f32) as usize)
584 .map(|bin| bin.min(fft_bins - 1))
585 .collect();
586
587 let mut filters = vec![0.0_f32; config.mel_bins * fft_bins];
588 for mel in 0..config.mel_bins {
589 let left = bin_points[mel];
590 let center = bin_points[mel + 1].max(left + 1);
591 let right = bin_points[mel + 2].max(center + 1).min(fft_bins - 1);
592
593 for bin in left..center.min(fft_bins) {
594 filters[mel * fft_bins + bin] = (bin - left) as f32 / (center - left) as f32;
595 }
596 for bin in center..=right {
597 filters[mel * fft_bins + bin] = (right - bin) as f32 / (right - center) as f32;
598 }
599 }
600 filters
601}
602
603fn hz_to_mel(hz: f32) -> f32 {
604 2595.0 * (1.0 + hz / 700.0).log10()
605}
606
607fn mel_to_hz(mel: f32) -> f32 {
608 700.0 * (10.0_f32.powf(mel / 2595.0) - 1.0)
609}
610
611#[cfg(test)]
612mod tests {
613 use super::*;
614
615 #[test]
616 fn discovers_single_safetensors_file() {
617 let root =
618 std::env::temp_dir().join(format!("vona-mlx-speech-test-{}", std::process::id()));
619 let _ = std::fs::remove_dir_all(&root);
620 std::fs::create_dir_all(&root).unwrap();
621 std::fs::write(root.join("config.json"), "{}").unwrap();
622 std::fs::write(root.join("model.safetensors"), b"stub").unwrap();
623
624 let files = SpeechModelFiles::discover(&root).unwrap();
625 assert_eq!(
626 files.safetensors_files,
627 vec![root.join("model.safetensors")]
628 );
629 let _ = std::fs::remove_dir_all(&root);
630 }
631
632 #[test]
633 fn computes_log_mel_shape() {
634 let samples = vec![0.0_f32; 480];
635 let spec = log_mel_spectrogram(&samples, MelSpectrogramConfig::whisper()).unwrap();
636 assert_eq!(spec.bins, 80);
637 assert_eq!(spec.frames, 1);
638 assert_eq!(spec.values.len(), 80);
639 }
640}