Skip to main content

vona_mlx_speech/
lib.rs

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}