Files
WhispAssist/src-tauri/examples/asr_bench.rs
T

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),
}
}