use std::path::Path;
use crate::traits::AsrError;
#[derive(Debug, Clone)]
pub struct Cmvn {
pub means: Vec<f32>,
pub vars: Vec<f32>,
}
impl Cmvn {
pub fn dim(&self) -> usize {
self.means.len()
}
}
pub fn apply_lfr(
feats: &[f32],
t_in: usize,
feat_dim: usize,
lfr_m: usize,
lfr_n: usize,
) -> (Vec<f32>, usize) {
assert!(lfr_m > 0 && lfr_n > 0, "lfr_m/lfr_n must be positive");
if t_in == 0 {
return (Vec::new(), 0);
}
let t_out = t_in.div_ceil(lfr_n);
let out_dim = feat_dim * lfr_m;
let mut out = Vec::with_capacity(t_out * out_dim);
let left_pad = (lfr_m - 1) / 2;
let frame = |raw: i64| -> &[f32] {
let clamped = raw.clamp(0, (t_in - 1) as i64) as usize;
&feats[clamped * feat_dim..(clamped + 1) * feat_dim]
};
for i in 0..t_out {
let centre = (i * lfr_n) as i64;
for k in 0..lfr_m {
let idx = centre + k as i64 - left_pad as i64;
out.extend_from_slice(frame(idx));
}
}
debug_assert_eq!(out.len(), t_out * out_dim);
(out, t_out)
}
pub fn apply_cmvn(feats: &mut [f32], feat_dim: usize, cmvn: &Cmvn) {
assert_eq!(
cmvn.means.len(),
feat_dim,
"cmvn dim {} != feat dim {}",
cmvn.means.len(),
feat_dim
);
assert_eq!(cmvn.vars.len(), feat_dim);
for chunk in feats.chunks_exact_mut(feat_dim) {
for ((x, m), v) in chunk
.iter_mut()
.zip(cmvn.means.iter())
.zip(cmvn.vars.iter())
{
*x = (*x + *m) * *v;
}
}
}
pub fn load_cmvn(path: &Path) -> Result<Cmvn, AsrError> {
let text = std::fs::read_to_string(path).map_err(|e| AsrError::ModelLoad(format!("read am.mvn {}: {e}", path.display())))?;
let means = extract_block(&text, "<AddShift>").ok_or_else(|| AsrError::ModelLoad(format!("am.mvn {}: missing <AddShift> block", path.display())))?;
let vars = extract_block(&text, "<Rescale>").ok_or_else(|| AsrError::ModelLoad(format!("am.mvn {}: missing <Rescale> block", path.display())))?;
if means.len() != vars.len() {
return Err(AsrError::ModelLoad(format!(
"am.mvn {}: mean dim {} != var dim {}",
path.display(),
means.len(),
vars.len()
)));
}
if means.is_empty() {
return Err(AsrError::ModelLoad(format!("am.mvn {}: empty CMVN vectors", path.display())));
}
Ok(Cmvn { means, vars })
}
fn extract_block(text: &str, tag: &str) -> Option<Vec<f32>> {
let tag_pos = text.find(tag)?;
let after_tag = &text[tag_pos + tag.len()..];
let open = after_tag.find('[')?;
let close = after_tag[open..].find(']')?;
let body = &after_tag[open + 1..open + close];
let vals: Vec<f32> = body
.split_whitespace()
.filter_map(|t| t.parse::<f32>().ok())
.collect();
if vals.is_empty() {
None
} else {
Some(vals)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn flatten(rows: Vec<Vec<f32>>) -> (Vec<f32>, usize, usize) {
let t = rows.len();
let d = rows.first().map_or(0, |r| r.len());
let mut out = Vec::with_capacity(t * d);
for r in rows {
out.extend(r);
}
(out, t, d)
}
#[test]
fn apply_lfr_paraformer_dimensions() {
let (feats, t, d) = flatten(vec![vec![0.0_f32; 80]; 16]);
let (out, t_out) = apply_lfr(&feats, t, d, 7, 6);
assert_eq!(t_out, 3);
assert_eq!(out.len(), 3 * 7 * 80);
}
#[test]
fn apply_lfr_short_input_pads_with_replication() {
let row: Vec<f32> = (0..4).map(|x| x as f32).collect();
let (feats, t, d) = flatten(vec![row.clone()]);
let (out, t_out) = apply_lfr(&feats, t, d, 7, 6);
assert_eq!(t_out, 1);
assert_eq!(out.len(), 7 * 4);
for chunk in out.chunks_exact(4) {
assert_eq!(chunk, row.as_slice());
}
}
#[test]
fn apply_lfr_centres_window_with_left_padding() {
let mut rows = Vec::new();
for i in 0..7 {
rows.push(vec![i as f32; 2]);
}
let (feats, t, d) = flatten(rows);
let (out, t_out) = apply_lfr(&feats, t, d, 7, 6);
assert_eq!(t_out, 2);
let first: Vec<f32> = out[..7 * 2].to_vec();
let expected_first: Vec<f32> = [0.0, 0.0, 0.0, 0.0, 1.0, 2.0, 3.0]
.iter()
.flat_map(|v| [*v, *v])
.collect();
assert_eq!(first, expected_first);
}
#[test]
fn apply_cmvn_matches_reference_formula() {
let mut feats = vec![1.0_f32, 2.0, 3.0, 4.0]; let cmvn = Cmvn {
means: vec![-1.0, -2.0],
vars: vec![2.0, 0.5],
};
apply_cmvn(&mut feats, 2, &cmvn);
assert_eq!(feats, vec![0.0, 0.0, 4.0, 1.0]);
}
#[test]
fn load_cmvn_parses_addshift_and_rescale() {
let dir =
std::env::temp_dir().join(format!("euhadra_paraformer_mvn_{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("am.mvn");
let body = "<Nnet>\n\
<AddShift> <LearnRateCoef> 0 [ -8.31 -8.42 -8.55 ]\n\
<Rescale> <LearnRateCoef> 0 [ 0.139 0.142 0.145 ]\n\
</Nnet>\n";
std::fs::write(&path, body).unwrap();
let cmvn = load_cmvn(&path).unwrap();
std::fs::remove_dir_all(&dir).ok();
assert_eq!(cmvn.dim(), 3);
assert!((cmvn.means[0] - -8.31).abs() < 1e-5);
assert!((cmvn.vars[2] - 0.145).abs() < 1e-5);
}
#[test]
fn load_cmvn_rejects_dim_mismatch() {
let dir =
std::env::temp_dir().join(format!("euhadra_paraformer_mvn_bad_{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("am.mvn");
std::fs::write(&path, "<AddShift> [ -1 -2 -3 ]\n<Rescale> [ 1 2 ]\n").unwrap();
let res = load_cmvn(&path);
std::fs::remove_dir_all(&dir).ok();
assert!(res.is_err());
}
#[test]
fn load_cmvn_handles_multiline_brackets() {
let dir = std::env::temp_dir().join(format!(
"euhadra_paraformer_mvn_multi_{}",
std::process::id()
));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("am.mvn");
let body = "<Nnet>\n\
<AddShift> <LearnRateCoef> 0 [\n -1.0 -2.0\n -3.0 ]\n\
<Rescale> <LearnRateCoef> 0 [ 1.0\n 2.0\n 3.0 ]\n";
std::fs::write(&path, body).unwrap();
let cmvn = load_cmvn(&path).unwrap();
std::fs::remove_dir_all(&dir).ok();
assert_eq!(cmvn.dim(), 3);
assert!((cmvn.means[1] - -2.0).abs() < 1e-5);
assert!((cmvn.vars[1] - 2.0).abs() < 1e-5);
}
#[test]
fn load_cmvn_missing_file_errors() {
let res = load_cmvn(Path::new("/nonexistent/euhadra/am.mvn"));
assert!(res.is_err());
}
}