419 lines
15 KiB
Rust
419 lines
15 KiB
Rust
//! ASR bench harness — the Parakeet-vs-Whisper bake-off (test_parakeet branch).
|
|
//!
|
|
//! Drives the *production* transcribers (`WhisperTranscriber`, `OnnxTranscriber`)
|
|
//! plus a sherpa-onnx `TransducerRecognizer` (Parakeet TDT) over a shared corpus,
|
|
//! and reports WER / RTF / process-CPU-seconds per contender. One contender per
|
|
//! process invocation, deliberately: (1) process CPU time then measures exactly
|
|
//! one engine, and (2) the ort crate (OpenVINO runtime) and sherpa's bundled
|
|
//! onnxruntime never coexist in one address space.
|
|
//!
|
|
//! Usage:
|
|
//! asr_bench <contender> [--corpus DIR] [--out CSV]
|
|
//! asr_bench <contender> --live WAV # sequential 10 s window latency
|
|
//! asr_bench --self-test # WER scorer assertions
|
|
//!
|
|
//! Contenders: w1-cpu-base | w2-vulkan-base | w3-vulkan-medium | w4-npu-onnx |
|
|
//! p1-parakeet-t<N> (N = sherpa num_threads)
|
|
//!
|
|
//! Corpus layout: {name}.wav (16 kHz mono) + {name}.txt reference per utterance
|
|
//! (staged by scripts/bench/fetch-corpus.ps1).
|
|
|
|
use std::io::Write as _;
|
|
use std::path::{Path, PathBuf};
|
|
use std::time::Instant;
|
|
|
|
use whispassist_lib::audio::read_wav_mono_16k;
|
|
use whispassist_lib::models::BackendId;
|
|
use whispassist_lib::paths;
|
|
use whispassist_lib::transcription::{
|
|
onnx_models, AudioWindow, OnnxTranscriber, Transcriber, WhisperTranscriber,
|
|
};
|
|
|
|
const WINDOW_SECS: usize = 10;
|
|
const SAMPLE_RATE: usize = 16_000;
|
|
|
|
// ---- process CPU time (kernel + user) via GetProcessTimes — no crate
|
|
// features needed, kernel32 is always linked on Windows. ----
|
|
#[cfg(windows)]
|
|
mod cputime {
|
|
#[repr(C)]
|
|
#[derive(Default, Clone, Copy)]
|
|
struct Filetime {
|
|
lo: u32,
|
|
hi: u32,
|
|
}
|
|
extern "system" {
|
|
fn GetCurrentProcess() -> isize;
|
|
fn GetProcessTimes(
|
|
h: isize,
|
|
creation: *mut Filetime,
|
|
exit: *mut Filetime,
|
|
kernel: *mut Filetime,
|
|
user: *mut Filetime,
|
|
) -> i32;
|
|
}
|
|
/// Cumulative process CPU time in milliseconds.
|
|
pub fn process_cpu_ms() -> u64 {
|
|
let (mut c, mut e, mut k, mut u) = (
|
|
Filetime::default(),
|
|
Filetime::default(),
|
|
Filetime::default(),
|
|
Filetime::default(),
|
|
);
|
|
// SAFETY: pseudo-handle + four valid out-pointers, per the API contract.
|
|
let ok = unsafe { GetProcessTimes(GetCurrentProcess(), &mut c, &mut e, &mut k, &mut u) };
|
|
if ok == 0 {
|
|
return 0;
|
|
}
|
|
let ms = |f: Filetime| (((f.hi as u64) << 32) | f.lo as u64) / 10_000; // 100 ns units
|
|
ms(k) + ms(u)
|
|
}
|
|
}
|
|
#[cfg(not(windows))]
|
|
mod cputime {
|
|
pub fn process_cpu_ms() -> u64 {
|
|
0
|
|
}
|
|
}
|
|
|
|
// ---- WER ----
|
|
|
|
/// Lowercase, strip everything but alphanumerics/apostrophes, split to words —
|
|
/// LibriSpeech references are uppercase without punctuation; hypotheses carry
|
|
/// casing + punctuation, so both sides normalize through here.
|
|
fn normalize(s: &str) -> Vec<String> {
|
|
s.to_lowercase()
|
|
.chars()
|
|
.map(|c| {
|
|
if c.is_alphanumeric() || c == '\'' {
|
|
c
|
|
} else {
|
|
' '
|
|
}
|
|
})
|
|
.collect::<String>()
|
|
.split_whitespace()
|
|
.map(str::to_string)
|
|
.collect()
|
|
}
|
|
|
|
/// Word-level Levenshtein distance (two-row DP).
|
|
fn edit_distance(a: &[String], b: &[String]) -> usize {
|
|
let mut prev: Vec<usize> = (0..=b.len()).collect();
|
|
let mut cur = vec![0usize; b.len() + 1];
|
|
for (i, wa) in a.iter().enumerate() {
|
|
cur[0] = i + 1;
|
|
for (j, wb) in b.iter().enumerate() {
|
|
let sub = prev[j] + usize::from(wa != wb);
|
|
cur[j + 1] = sub.min(prev[j + 1] + 1).min(cur[j] + 1);
|
|
}
|
|
std::mem::swap(&mut prev, &mut cur);
|
|
}
|
|
prev[b.len()]
|
|
}
|
|
|
|
fn self_test() {
|
|
assert_eq!(edit_distance(&normalize("a b c"), &normalize("a b c")), 0);
|
|
assert_eq!(edit_distance(&normalize("a b c"), &normalize("a x c")), 1); // sub
|
|
assert_eq!(edit_distance(&normalize("a b c"), &normalize("a c")), 1); // del
|
|
assert_eq!(edit_distance(&normalize("a c"), &normalize("a b c")), 1); // ins
|
|
assert_eq!(
|
|
normalize("The QUICK, brown-fox's."),
|
|
vec!["the", "quick", "brown", "fox's"]
|
|
);
|
|
assert!(edit_distance(&normalize("hello world"), &normalize("")) == 2);
|
|
println!("self-test OK");
|
|
}
|
|
|
|
// ---- engines ----
|
|
|
|
// ponytail: one Engine per process for its whole lifetime — variant size skew is irrelevant.
|
|
#[allow(clippy::large_enum_variant)]
|
|
enum Engine {
|
|
Whisper(WhisperTranscriber),
|
|
Onnx(OnnxTranscriber),
|
|
Parakeet(std::sync::Mutex<sherpa_rs::transducer::TransducerRecognizer>),
|
|
}
|
|
|
|
impl Engine {
|
|
fn transcribe_wav(&self, wav: &Path) -> Result<String, String> {
|
|
match self {
|
|
Engine::Whisper(t) => Ok(join_segments(
|
|
t.transcribe_file(wav).map_err(|e| e.to_string())?,
|
|
)),
|
|
Engine::Onnx(t) => Ok(join_segments(
|
|
t.transcribe_file(wav).map_err(|e| e.to_string())?,
|
|
)),
|
|
Engine::Parakeet(r) => {
|
|
let samples = read_wav_mono_16k(wav).map_err(|e| e.to_string())?;
|
|
Ok(r.lock()
|
|
.map_err(|_| "poisoned".to_string())?
|
|
.transcribe(SAMPLE_RATE as u32, &samples))
|
|
}
|
|
}
|
|
}
|
|
|
|
/// One live window; returns the wall time only (text is discarded — this
|
|
/// mode measures latency, corpus mode measures accuracy).
|
|
fn time_window(&self, samples: Vec<f32>, offset_ms: u64) -> Result<u128, String> {
|
|
let t0 = Instant::now();
|
|
match self {
|
|
Engine::Whisper(t) => {
|
|
let (tx, rx) = std::sync::mpsc::channel();
|
|
t.transcribe_stream(AudioWindow { samples, offset_ms }, tx)
|
|
.map_err(|e| e.to_string())?;
|
|
drop(rx);
|
|
}
|
|
Engine::Onnx(t) => {
|
|
let (tx, rx) = std::sync::mpsc::channel();
|
|
t.transcribe_stream(AudioWindow { samples, offset_ms }, tx)
|
|
.map_err(|e| e.to_string())?;
|
|
drop(rx);
|
|
}
|
|
Engine::Parakeet(r) => {
|
|
let _ = r
|
|
.lock()
|
|
.map_err(|_| "poisoned".to_string())?
|
|
.transcribe(SAMPLE_RATE as u32, &samples);
|
|
}
|
|
}
|
|
Ok(t0.elapsed().as_millis())
|
|
}
|
|
}
|
|
|
|
fn join_segments(segs: Vec<whispassist_lib::models::TranscriptSegment>) -> String {
|
|
segs.iter()
|
|
.map(|s| s.text.trim())
|
|
.collect::<Vec<_>>()
|
|
.join(" ")
|
|
}
|
|
|
|
/// First file in `dir` whose name starts with `prefix` and ends with `.onnx`.
|
|
fn find_onnx(dir: &Path, prefix: &str) -> Result<PathBuf, String> {
|
|
std::fs::read_dir(dir)
|
|
.map_err(|e| format!("{}: {e}", dir.display()))?
|
|
.filter_map(|e| e.ok())
|
|
.map(|e| e.path())
|
|
.find(|p| {
|
|
p.file_name()
|
|
.and_then(|n| n.to_str())
|
|
.is_some_and(|n| n.starts_with(prefix) && n.ends_with(".onnx"))
|
|
})
|
|
.ok_or_else(|| format!("no {prefix}*.onnx under {}", dir.display()))
|
|
}
|
|
|
|
fn load_engine(contender: &str) -> Result<Engine, String> {
|
|
let whisper = |model: &str, backend: BackendId| -> Result<Engine, String> {
|
|
let path = paths::whisper_model_file(model);
|
|
WhisperTranscriber::load(&path, backend, Some("en"))
|
|
.map(Engine::Whisper)
|
|
.map_err(|e| e.to_string())
|
|
};
|
|
match contender {
|
|
"w1-cpu-base" => whisper("base.en-q5_1", BackendId::Cpu),
|
|
"w2-vulkan-base" => whisper("base.en-q5_1", BackendId::Intel),
|
|
"w3-vulkan-medium" => whisper("medium.en-q5_0", BackendId::Intel),
|
|
"w4-npu-onnx" => {
|
|
let dir = onnx_models::model_dir(onnx_models::DEFAULT_ONNX_MODEL);
|
|
OnnxTranscriber::load(&dir, BackendId::Npu, None)
|
|
.map(Engine::Onnx)
|
|
.map_err(|e| e.to_string())
|
|
}
|
|
p if p.starts_with("p1-parakeet-t") || p.starts_with("p2-parakeet-dml-t") => {
|
|
let dml = p.starts_with("p2-");
|
|
let threads: i32 = p
|
|
.rsplit_once('t')
|
|
.and_then(|(_, n)| n.parse().ok())
|
|
.ok_or_else(|| format!("bad thread count in '{p}'"))?;
|
|
let dir = paths::models_dir().join("bench-parakeet");
|
|
let config = sherpa_rs::transducer::TransducerConfig {
|
|
encoder: find_onnx(&dir, "encoder")?.to_string_lossy().into_owned(),
|
|
decoder: find_onnx(&dir, "decoder")?.to_string_lossy().into_owned(),
|
|
joiner: find_onnx(&dir, "joiner")?.to_string_lossy().into_owned(),
|
|
tokens: dir.join("tokens.txt").to_string_lossy().into_owned(),
|
|
model_type: "nemo_transducer".to_string(),
|
|
decoding_method: "greedy_search".to_string(),
|
|
sample_rate: SAMPLE_RATE as i32,
|
|
feature_dim: 80,
|
|
num_threads: threads,
|
|
// "directml" needs the bench-directml cargo feature (DirectML
|
|
// sherpa-onnx binaries); with plain binaries sherpa falls back
|
|
// noisily and the run is invalid — P2 rows only count from a
|
|
// bench-directml build.
|
|
provider: Some(if dml { "directml" } else { "cpu" }.to_string()),
|
|
..Default::default()
|
|
};
|
|
sherpa_rs::transducer::TransducerRecognizer::new(config)
|
|
.map(|r| Engine::Parakeet(std::sync::Mutex::new(r)))
|
|
.map_err(|e| e.to_string())
|
|
}
|
|
other => Err(format!(
|
|
"unknown contender '{other}' (w1-cpu-base | w2-vulkan-base | w3-vulkan-medium | w4-npu-onnx | p1-parakeet-t<N> | p2-parakeet-dml-t<N>)"
|
|
)),
|
|
}
|
|
}
|
|
|
|
// ---- modes ----
|
|
|
|
fn csv_escape(s: &str) -> String {
|
|
format!("\"{}\"", s.replace('"', "\"\""))
|
|
}
|
|
|
|
fn run_corpus(contender: &str, engine: &Engine, load_ms: u128, corpus: &Path, out: &Path) {
|
|
let mut wavs: Vec<PathBuf> = std::fs::read_dir(corpus)
|
|
.unwrap_or_else(|e| panic!("corpus dir {}: {e}", corpus.display()))
|
|
.filter_map(|e| e.ok())
|
|
.map(|e| e.path())
|
|
.filter(|p| p.extension().is_some_and(|x| x == "wav"))
|
|
.collect();
|
|
wavs.sort();
|
|
assert!(!wavs.is_empty(), "no wavs in {}", corpus.display());
|
|
|
|
if let Some(parent) = out.parent() {
|
|
std::fs::create_dir_all(parent).expect("create out dir");
|
|
}
|
|
let new_file = !out.exists();
|
|
let mut csv = std::fs::OpenOptions::new()
|
|
.create(true)
|
|
.append(true)
|
|
.open(out)
|
|
.expect("open csv");
|
|
if new_file {
|
|
writeln!(
|
|
csv,
|
|
"contender,file,audio_secs,load_ms,wall_ms,cpu_ms,wer_errors,ref_words,hypothesis"
|
|
)
|
|
.unwrap();
|
|
}
|
|
|
|
let (mut errors, mut words, mut wall_total, mut cpu_total, mut audio_total) =
|
|
(0usize, 0usize, 0u128, 0u64, 0f64);
|
|
for wav in &wavs {
|
|
let stem = wav.file_stem().unwrap().to_string_lossy().into_owned();
|
|
let reference = std::fs::read_to_string(wav.with_extension("txt")).unwrap_or_else(|e| {
|
|
panic!(
|
|
"missing reference {}: {e}",
|
|
wav.with_extension("txt").display()
|
|
)
|
|
});
|
|
let audio_secs = read_wav_mono_16k(wav)
|
|
.map(|s| s.len() as f64 / SAMPLE_RATE as f64)
|
|
.unwrap_or(0.0);
|
|
|
|
let cpu0 = cputime::process_cpu_ms();
|
|
let t0 = Instant::now();
|
|
let hypothesis = match engine.transcribe_wav(wav) {
|
|
Ok(t) => t,
|
|
Err(e) => {
|
|
eprintln!("FAILED {stem}: {e}");
|
|
continue;
|
|
}
|
|
};
|
|
let wall = t0.elapsed().as_millis();
|
|
let cpu = cputime::process_cpu_ms() - cpu0;
|
|
|
|
let r = normalize(&reference);
|
|
let h = normalize(&hypothesis);
|
|
let err = edit_distance(&r, &h);
|
|
errors += err;
|
|
words += r.len();
|
|
wall_total += wall;
|
|
cpu_total += cpu;
|
|
audio_total += audio_secs;
|
|
|
|
writeln!(
|
|
csv,
|
|
"{contender},{stem},{audio_secs:.2},{load_ms},{wall},{cpu},{err},{},{}",
|
|
r.len(),
|
|
csv_escape(&hypothesis)
|
|
)
|
|
.unwrap();
|
|
println!(
|
|
"{stem}: {audio_secs:.1}s wall={wall}ms cpu={cpu}ms wer={err}/{}",
|
|
r.len()
|
|
);
|
|
}
|
|
|
|
println!("---- {contender} summary ----");
|
|
println!(
|
|
"files: {} audio: {audio_total:.1}s load: {load_ms}ms",
|
|
wavs.len()
|
|
);
|
|
println!(
|
|
"WER: {:.2}% ({errors}/{words})",
|
|
100.0 * errors as f64 / words.max(1) as f64
|
|
);
|
|
println!(
|
|
"RTF (wall): {:.3}",
|
|
wall_total as f64 / 1000.0 / audio_total.max(0.001)
|
|
);
|
|
println!(
|
|
"CPU-sec per audio-sec: {:.3}",
|
|
cpu_total as f64 / 1000.0 / audio_total.max(0.001)
|
|
);
|
|
}
|
|
|
|
fn run_live(contender: &str, engine: &Engine, wav: &Path) {
|
|
let samples = read_wav_mono_16k(wav).expect("read live wav");
|
|
let chunk = WINDOW_SECS * SAMPLE_RATE;
|
|
let mut latencies: Vec<u128> = Vec::new();
|
|
for (i, part) in samples.chunks(chunk).enumerate() {
|
|
if part.len() < SAMPLE_RATE {
|
|
continue; // sub-second tail: skip, same spirit as MIN_DIARIZE_SAMPLES
|
|
}
|
|
let ms = engine
|
|
.time_window(part.to_vec(), (i * WINDOW_SECS * 1000) as u64)
|
|
.expect("window decode");
|
|
println!("window {i}: {ms}ms");
|
|
latencies.push(ms);
|
|
}
|
|
latencies.sort_unstable();
|
|
let pct = |p: f64| latencies[(((latencies.len() - 1) as f64) * p) as usize];
|
|
println!(
|
|
"---- {contender} live ({WINDOW_SECS}s windows, n={}) ----",
|
|
latencies.len()
|
|
);
|
|
println!(
|
|
"p50={}ms p95={}ms max={}ms (budget: {}ms)",
|
|
pct(0.50),
|
|
pct(0.95),
|
|
latencies.last().unwrap(),
|
|
WINDOW_SECS * 1000
|
|
);
|
|
}
|
|
|
|
fn main() {
|
|
let args: Vec<String> = std::env::args().skip(1).collect();
|
|
if args.iter().any(|a| a == "--self-test") {
|
|
self_test();
|
|
return;
|
|
}
|
|
let contender = args
|
|
.first()
|
|
.expect("usage: asr_bench <contender> [--corpus DIR|--live WAV]");
|
|
let flag = |name: &str| {
|
|
args.iter()
|
|
.position(|a| a == name)
|
|
.and_then(|i| args.get(i + 1))
|
|
.map(PathBuf::from)
|
|
};
|
|
let corpus = flag("--corpus").unwrap_or_else(|| PathBuf::from("bench-corpus"));
|
|
let out = flag("--out").unwrap_or_else(|| PathBuf::from("bench-results/raw.csv"));
|
|
|
|
let t0 = Instant::now();
|
|
let engine = match load_engine(contender) {
|
|
Ok(e) => e,
|
|
Err(e) => {
|
|
eprintln!("load failed: {e}");
|
|
std::process::exit(1);
|
|
}
|
|
};
|
|
let load_ms = t0.elapsed().as_millis();
|
|
println!("{contender}: loaded in {load_ms}ms");
|
|
|
|
match flag("--live") {
|
|
Some(wav) => run_live(contender, &engine, &wav),
|
|
None => run_corpus(contender, &engine, load_ms, &corpus, &out),
|
|
}
|
|
}
|