use std::collections::BTreeMap;
use std::path::PathBuf;
use std::time::Instant;
use clap::{Parser, ValueEnum};
use euhadra::eval::annotations::load_jsonl as load_annotations;
use euhadra::eval::baseline::{LanguageLayerBaseline, LatencyMicrosRecord};
use euhadra::eval::f1::{aggregate, iou_f1, strict_f1, F1Stats, Span};
use euhadra::eval::fixtures::{load_jsonl as load_fixtures, Fixture};
use euhadra::eval::latency::Samples;
use euhadra::eval::metrics::{cer, wer};
use euhadra::prelude::*;
#[derive(Parser, Debug)]
#[command(about = "L3: direct F1 (self-correction / filler) + ablation on natural speech")]
struct Cli {
#[arg(long, value_enum)]
task: Task,
#[arg(long, default_value = "ja")]
lang: String,
#[arg(long)]
input: PathBuf,
#[arg(long)]
dict: Option<PathBuf>,
#[arg(long)]
base_dict: Option<PathBuf>,
#[arg(long, default_value_t = 0.5)]
iou_threshold: f64,
#[arg(long)]
output: Option<PathBuf>,
#[arg(long)]
verbose: bool,
#[arg(long)]
embedder_dir: Option<PathBuf>,
#[arg(long, default_value_t = 1.0)]
alpha: f32,
#[arg(long, default_value_t = 0.85)]
threshold: f32,
#[arg(long, default_value_t = 0.65)]
composite_threshold: f32,
#[arg(long)]
min_f1: Option<f64>,
}
fn enforce_min_f1(min_f1: Option<f64>, measured: f64, metric: &str) -> Result<(), String> {
use std::cmp::Ordering;
match min_f1 {
Some(min)
if matches!(
measured.partial_cmp(&min),
Some(Ordering::Less) | None
) =>
{
Err(format!(
"{metric} F1 {measured:.4} is below the required minimum {min:.4}"
))
}
_ => Ok(()),
}
}
#[cfg(test)]
mod tests {
use super::enforce_min_f1;
#[test]
fn no_minimum_never_fails() {
assert!(enforce_min_f1(None, 0.0, "x").is_ok());
}
#[test]
fn meeting_the_minimum_passes() {
assert!(enforce_min_f1(Some(0.9), 0.9, "x").is_ok());
assert!(enforce_min_f1(Some(0.9), 1.0, "x").is_ok());
}
#[test]
fn falling_short_fails_with_both_numbers() {
let err = enforce_min_f1(Some(0.9), 0.88, "correction-pair").unwrap_err();
assert!(err.contains("0.8800"), "{err}");
assert!(err.contains("0.9000"), "{err}");
assert!(err.contains("correction-pair"), "{err}");
}
#[test]
fn nan_f1_fails_rather_than_silently_passing() {
assert!(enforce_min_f1(Some(0.5), f64::NAN, "x").is_err());
}
}
#[derive(Clone, Copy, Debug, ValueEnum)]
enum Task {
SelfCorrection,
Ablation,
Filler,
PhonemeCorrection,
}
#[tokio::main(flavor = "current_thread")]
async fn main() {
if let Err(e) = run().await {
eprintln!("error: {e}");
std::process::exit(2);
}
}
async fn run() -> Result<(), String> {
let cli = Cli::parse();
match cli.task {
Task::SelfCorrection => run_self_correction(&cli).await,
Task::Ablation => run_ablation(&cli).await,
Task::Filler => run_filler(&cli).await,
Task::PhonemeCorrection => run_phoneme_correction(&cli).await,
}
}
async fn run_self_correction(cli: &Cli) -> Result<(), String> {
let annotations = load_annotations(&cli.input)
.map_err(|e| format!("loading {}: {e}", cli.input.display()))?;
if annotations.is_empty() {
return Err(format!("annotation file {} is empty", cli.input.display()));
}
let detector = SelfCorrectionDetector::new();
let ctx = ContextSnapshot::default();
let mut utt_tp = 0usize;
let mut utt_fp = 0usize;
let mut utt_fn = 0usize;
let mut utt_tn = 0usize;
let mut span_stats: Vec<F1Stats> = Vec::new();
let mut strict_stats: Vec<F1Stats> = Vec::new();
let cues = cue_set_for(&cli.lang)?;
for anno in &annotations {
let result = detector
.process(&anno.text, &ctx)
.await
.map_err(|e| format!("detector on {}: {e}", anno.utterance_id))?;
let predicted: Vec<Span> = if result.corrections.is_empty() {
Vec::new()
} else {
match diff_removed_span(&anno.text, &result.text) {
Some(raw) => vec![trim_trailing_cue(&anno.text, raw, &cues)],
None => Vec::new(),
}
};
let gold: Vec<Span> = anno.repairs.iter().map(|r| r.reparandum.span()).collect();
match (predicted.is_empty(), gold.is_empty()) {
(false, false) => utt_tp += 1,
(false, true) => utt_fp += 1,
(true, false) => utt_fn += 1,
(true, true) => utt_tn += 1,
}
if !predicted.is_empty() || !gold.is_empty() {
let strict = strict_f1(&predicted, &gold);
span_stats.push(iou_f1(&predicted, &gold, cli.iou_threshold));
strict_stats.push(strict);
if cli.verbose && (strict.fp > 0 || strict.fn_ > 0) {
let chars: Vec<char> = anno.text.chars().collect();
let span_text = |s: &Span| -> String {
chars
.get(s.start..s.end.min(chars.len()))
.map(|s| s.iter().collect::<String>())
.unwrap_or_default()
};
let pred_str: Vec<String> = predicted
.iter()
.map(|s| format!("{:?}={:?}", (s.start, s.end), span_text(s)))
.collect();
let gold_str: Vec<String> = gold
.iter()
.map(|s| format!("{:?}={:?}", (s.start, s.end), span_text(s)))
.collect();
if predicted != gold {
println!(
" [diff] {} text={:?}\n predicted={:?}\n gold={:?}",
anno.utterance_id, anno.text, pred_str, gold_str,
);
}
}
}
}
let utt_f1 = F1Stats::from_counts(utt_tp, utt_fp, utt_fn);
let span_iou_agg = aggregate(&span_stats);
let span_strict_agg = aggregate(&strict_stats);
println!("=== L3 self-correction direct F1 ({}) ===", cli.lang);
println!("annotations: {}", annotations.len());
println!(
"utterance-level tp={} fp={} fn={} tn={}",
utt_tp, utt_fp, utt_fn, utt_tn
);
println!(
" precision={} recall={} F1={}",
fmt_pct(utt_f1.precision),
fmt_pct(utt_f1.recall),
fmt_pct(utt_f1.f1),
);
println!(
"span-level (IoU≥{:.2}) tp={} fp={} fn={} precision={} recall={} F1={}",
cli.iou_threshold,
span_iou_agg.tp,
span_iou_agg.fp,
span_iou_agg.fn_,
fmt_pct(span_iou_agg.precision),
fmt_pct(span_iou_agg.recall),
fmt_pct(span_iou_agg.f1),
);
println!(
"span-level (strict) tp={} fp={} fn={} precision={} recall={} F1={}",
span_strict_agg.tp,
span_strict_agg.fp,
span_strict_agg.fn_,
fmt_pct(span_strict_agg.precision),
fmt_pct(span_strict_agg.recall),
fmt_pct(span_strict_agg.f1),
);
if let Some(out) = &cli.output {
let report = serde_json::json!({
"task": "self-correction",
"lang": cli.lang,
"annotations": annotations.len(),
"utterance_level": {
"tp": utt_tp, "fp": utt_fp, "fn": utt_fn, "tn": utt_tn,
"precision": utt_f1.precision, "recall": utt_f1.recall, "f1": utt_f1.f1,
},
"span_level_iou": {
"iou_threshold": cli.iou_threshold,
"tp": span_iou_agg.tp, "fp": span_iou_agg.fp, "fn": span_iou_agg.fn_,
"precision": span_iou_agg.precision, "recall": span_iou_agg.recall, "f1": span_iou_agg.f1,
},
"span_level_strict": {
"tp": span_strict_agg.tp, "fp": span_strict_agg.fp, "fn": span_strict_agg.fn_,
"precision": span_strict_agg.precision, "recall": span_strict_agg.recall, "f1": span_strict_agg.f1,
},
});
if let Some(parent) = out.parent() {
std::fs::create_dir_all(parent).ok();
}
std::fs::write(out, serde_json::to_string_pretty(&report).unwrap())
.map_err(|e| format!("write {}: {e}", out.display()))?;
eprintln!("report written to {}", out.display());
}
enforce_min_f1(cli.min_f1, span_strict_agg.f1, "self-correction span-level")
}
fn cue_set_for(lang: &str) -> Result<Vec<&'static str>, String> {
match lang {
"en" | "english" => Ok(en_cue_set()),
"ja" | "japanese" => Ok(ja_cue_set()),
"zh" | "chinese" => Ok(zh_cue_set()),
"es" | "spanish" => Ok(es_cue_set()),
"ko" | "korean" => Ok(ko_cue_set()),
other => Err(format!(
"self-correction task: --lang {other} not wired \
(expected one of: en, ja, zh, es, ko)"
)),
}
}
fn en_cue_set() -> Vec<&'static str> {
vec![
"no wait",
"or rather",
"i mean",
"actually",
"rather",
"sorry",
"wait",
"no",
]
}
fn ja_cue_set() -> Vec<&'static str> {
vec![
"いや",
"じゃなくて",
"じゃなく",
"ではなく",
"ていうか",
"っていうか",
"じゃない",
]
}
fn zh_cue_set() -> Vec<&'static str> {
vec![
"我的意思是",
"确切地说",
"应该说",
"我是说",
"不对",
"不是",
"算了",
]
}
fn es_cue_set() -> Vec<&'static str> {
vec![
"mejor dicho",
"quiero decir",
"o sea",
"perdón",
"mejor",
"digo",
"no es",
"no",
]
}
fn ko_cue_set() -> Vec<&'static str> {
vec![
"그게 아니라",
"그게 아니고",
"잘못 말했다",
"잘못 말했네",
"아 잠깐",
"잠깐만",
"아니에요",
"아니라",
"아니야",
"아니",
]
}
fn diff_removed_span(input: &str, output: &str) -> Option<Span> {
let in_chars: Vec<char> = input.chars().collect();
let out_chars: Vec<char> = output.chars().collect();
let mut suffix_len = 0;
while suffix_len < in_chars.len()
&& suffix_len < out_chars.len()
&& in_chars[in_chars.len() - 1 - suffix_len] == out_chars[out_chars.len() - 1 - suffix_len]
{
suffix_len += 1;
}
let in_suffix_start = in_chars.len() - suffix_len;
let out_suffix_start = out_chars.len() - suffix_len;
let mut prefix_len = 0;
while prefix_len < in_suffix_start
&& prefix_len < out_suffix_start
&& in_chars[prefix_len] == out_chars[prefix_len]
{
prefix_len += 1;
}
if prefix_len < in_suffix_start {
Some(Span {
start: prefix_len,
end: in_suffix_start,
})
} else {
None
}
}
fn trim_trailing_cue(input: &str, raw: Span, cues: &[&str]) -> Span {
let chars: Vec<char> = input.chars().collect();
if raw.start >= raw.end || raw.end > chars.len() {
return raw;
}
let is_sep = |c: char| {
matches!(c, '、' | ',' | '。' | '?' | '!')
|| c.is_whitespace()
|| matches!(c, ',' | '.' | ';' | ':' | '!' | '?')
};
let mut end = raw.end;
while end > raw.start && is_sep(chars[end - 1]) {
end -= 1;
}
let span_text: String = chars[raw.start..end].iter().collect();
let mut sorted_cues: Vec<&&str> = cues.iter().collect();
sorted_cues.sort_by_key(|c| std::cmp::Reverse(c.chars().count()));
for cue in sorted_cues {
if span_text.ends_with(*cue) {
let cue_chars = cue.chars().count();
end -= cue_chars;
while end > raw.start && is_sep(chars[end - 1]) {
end -= 1;
}
break;
}
}
Span {
start: raw.start,
end,
}
}
fn fmt_pct(x: f64) -> String {
if x.is_nan() {
"n/a".to_string()
} else {
format!("{:.3}", x)
}
}
type RuleDetector = Box<dyn Fn(&str) -> Vec<Span>>;
async fn run_filler(cli: &Cli) -> Result<(), String> {
let lang = cli.lang.as_str();
let detect_spans: RuleDetector = match lang {
"en" | "english" => {
let filter = SimpleFillerFilter::english();
Box::new(move |t| filter.detect_spans(t))
}
"ja" | "japanese" => {
let filter = JapaneseFillerFilter::new();
Box::new(move |t| filter.detect_spans(t))
}
"zh" | "chinese" => {
let filter = ChineseFillerFilter::new();
Box::new(move |t| filter.detect_spans(t))
}
"es" | "spanish" => {
let filter = SpanishFillerFilter::new();
Box::new(move |t| filter.detect_spans(t))
}
"ko" | "korean" => {
let filter = SimpleFillerFilter::korean();
Box::new(move |t| filter.detect_spans(t))
}
other => {
return Err(format!(
"filler task: --lang {other} not wired \
(expected one of: en, ja, zh, es, ko)"
));
}
};
score_filler(cli, detect_spans).await
}
async fn score_filler(cli: &Cli, detect_spans: RuleDetector) -> Result<(), String> {
let annotations = load_annotations(&cli.input)
.map_err(|e| format!("loading {}: {e}", cli.input.display()))?;
if annotations.is_empty() {
return Err(format!("annotation file {} is empty", cli.input.display()));
}
let mut utt_tp = 0usize;
let mut utt_fp = 0usize;
let mut utt_fn = 0usize;
let mut utt_tn = 0usize;
let mut span_stats: Vec<F1Stats> = Vec::new();
for anno in &annotations {
let predicted = detect_spans(&anno.text);
let gold: Vec<Span> = anno.fillers.iter().map(|f| f.span()).collect();
match (predicted.is_empty(), gold.is_empty()) {
(false, false) => utt_tp += 1,
(false, true) => utt_fp += 1,
(true, false) => utt_fn += 1,
(true, true) => utt_tn += 1,
}
if !predicted.is_empty() || !gold.is_empty() {
let stats = strict_f1(&predicted, &gold);
span_stats.push(stats);
if cli.verbose && (stats.fp > 0 || stats.fn_ > 0) {
let chars: Vec<char> = anno.text.chars().collect();
let span_text = |s: &Span| -> String {
chars
.get(s.start..s.end.min(chars.len()))
.map(|s| s.iter().collect::<String>())
.unwrap_or_default()
};
let pred_str: Vec<String> = predicted
.iter()
.map(|s| format!("{:?}={:?}", (s.start, s.end), span_text(s)))
.collect();
let gold_str: Vec<String> = gold
.iter()
.map(|s| format!("{:?}={:?}", (s.start, s.end), span_text(s)))
.collect();
println!(
" [diff] {} text={:?}\n predicted={:?}\n gold={:?}",
anno.utterance_id, anno.text, pred_str, gold_str,
);
}
}
}
let utt_f1 = F1Stats::from_counts(utt_tp, utt_fp, utt_fn);
let span_agg = aggregate(&span_stats);
println!("=== L3 filler direct F1 ({}) ===", cli.lang);
println!("annotations: {}", annotations.len());
println!(
"utterance-level tp={} fp={} fn={} tn={}",
utt_tp, utt_fp, utt_fn, utt_tn
);
println!(
" precision={} recall={} F1={}",
fmt_pct(utt_f1.precision),
fmt_pct(utt_f1.recall),
fmt_pct(utt_f1.f1),
);
println!(
"span-level (strict) tp={} fp={} fn={} precision={} recall={} F1={}",
span_agg.tp,
span_agg.fp,
span_agg.fn_,
fmt_pct(span_agg.precision),
fmt_pct(span_agg.recall),
fmt_pct(span_agg.f1),
);
if let Some(out) = &cli.output {
let report = serde_json::json!({
"task": "filler",
"lang": cli.lang,
"annotations": annotations.len(),
"utterance_level": {
"tp": utt_tp, "fp": utt_fp, "fn": utt_fn, "tn": utt_tn,
"precision": utt_f1.precision, "recall": utt_f1.recall, "f1": utt_f1.f1,
},
"span_level_strict": {
"tp": span_agg.tp, "fp": span_agg.fp, "fn": span_agg.fn_,
"precision": span_agg.precision, "recall": span_agg.recall, "f1": span_agg.f1,
},
});
if let Some(parent) = out.parent() {
std::fs::create_dir_all(parent).ok();
}
std::fs::write(out, serde_json::to_string_pretty(&report).unwrap())
.map_err(|e| format!("write {}: {e}", out.display()))?;
eprintln!("report written to {}", out.display());
}
enforce_min_f1(cli.min_f1, span_agg.f1, "span-level")
}
async fn run_phoneme_correction(cli: &Cli) -> Result<(), String> {
let lang = cli.lang.as_str();
if !matches!(lang, "en" | "english") {
return Err(format!(
"phoneme-correction task: --lang {lang} not wired \
(en only in v1)"
));
}
let dict_path = cli
.dict
.as_ref()
.ok_or_else(|| "phoneme-correction task requires --dict <path>".to_string())?;
let dict_raw = std::fs::read_to_string(dict_path)
.map_err(|e| format!("loading dict {}: {e}", dict_path.display()))?;
let dict_map: BTreeMap<String, String> = serde_json::from_str(&dict_raw)
.map_err(|e| format!("parsing dict {}: {e}", dict_path.display()))?;
if dict_map.is_empty() {
return Err(format!("dict {} is empty", dict_path.display()));
}
let custom_entries: Vec<euhadra::phoneme::CustomEntry> = dict_map
.iter()
.map(|(word, phonemes)| euhadra::phoneme::CustomEntry {
word: word.clone(),
phonemes: phonemes.clone(),
embedding: None,
})
.collect();
let base_dict = match cli.base_dict.as_ref() {
Some(path) => euhadra::phoneme::IpaDictionary::load(path)
.map_err(|e| format!("loading base dict {}: {}", path.display(), e))?,
None => euhadra::phoneme::IpaDictionary::empty(),
};
#[allow(unused_mut)]
let mut corrector = euhadra::phoneme::PhonemeCorrector::new(base_dict, custom_entries)
.with_threshold(cli.threshold)
.with_composite_threshold(cli.composite_threshold);
#[cfg(feature = "onnx")]
if let Some(dir) = &cli.embedder_dir {
let embedder = euhadra::phoneme::OnnxTextEmbedder::load(dir)
.map_err(|e| format!("loading embedder {}: {}", dir.display(), e))?;
corrector = corrector.with_embedder(embedder, cli.alpha);
eprintln!(
"[phoneme] embedder={} alpha={:.2} composite_threshold={:.2}",
dir.display(),
cli.alpha,
cli.composite_threshold
);
}
#[cfg(not(feature = "onnx"))]
if cli.embedder_dir.is_some() {
return Err("--embedder-dir requires --features onnx".to_string());
}
let corrector = corrector;
let ctx = ContextSnapshot::default();
let annotations = load_annotations(&cli.input)
.map_err(|e| format!("loading {}: {e}", cli.input.display()))?;
if annotations.is_empty() {
return Err(format!("annotation file {} is empty", cli.input.display()));
}
let mut utt_tp = 0usize; let mut utt_fp = 0usize; let mut utt_fn = 0usize; let mut utt_tn = 0usize;
let mut pair_tp = 0usize;
let mut pair_fp = 0usize;
let mut pair_fn = 0usize;
for anno in &annotations {
let result = corrector
.process(&anno.text, &ctx)
.await
.map_err(|e| format!("corrector on {}: {e}", anno.utterance_id))?;
let expected_text = anno.expected_text.as_deref().unwrap_or(&anno.text);
let output_text = result.text.as_str();
let output_norm = output_text.trim();
let expected_norm = expected_text.trim();
match (!anno.corrections.is_empty(), output_norm == expected_norm) {
(true, true) => utt_tp += 1,
(true, false) => utt_fn += 1,
(false, true) => utt_tn += 1,
(false, false) => utt_fp += 1,
}
let mut predicted_pairs: Vec<(String, String)> = result
.corrections
.iter()
.filter(|c| matches!(c.kind, euhadra::processor::CorrectionKind::DictionaryMatch))
.map(|c| (c.original.clone(), c.replacement.clone()))
.collect();
predicted_pairs.sort();
let mut gold_pairs: Vec<(String, String)> = anno
.corrections
.iter()
.map(|c| (c.original.clone(), c.replacement.clone()))
.collect();
gold_pairs.sort();
let (mut pi, mut gi) = (0usize, 0usize);
while pi < predicted_pairs.len() && gi < gold_pairs.len() {
match predicted_pairs[pi].cmp(&gold_pairs[gi]) {
std::cmp::Ordering::Equal => {
pair_tp += 1;
pi += 1;
gi += 1;
}
std::cmp::Ordering::Less => {
pair_fp += 1;
pi += 1;
}
std::cmp::Ordering::Greater => {
pair_fn += 1;
gi += 1;
}
}
}
pair_fp += predicted_pairs.len() - pi;
pair_fn += gold_pairs.len() - gi;
if cli.verbose && (output_norm != expected_norm || predicted_pairs != gold_pairs) {
println!(
" [diff] {} text={:?}\n predicted_text={:?}\n expected_text={:?}\n predicted_pairs={:?}\n gold_pairs={:?}",
anno.utterance_id,
anno.text,
output_norm,
expected_norm,
predicted_pairs,
gold_pairs,
);
}
}
let utt_f1 = F1Stats::from_counts(utt_tp, utt_fp, utt_fn);
let pair_f1 = F1Stats::from_counts(pair_tp, pair_fp, pair_fn);
println!("=== L3 phoneme-correction direct F1 ({}) ===", cli.lang);
println!("annotations: {}", annotations.len());
println!("dict words: {}", dict_map.len());
println!(
"utterance-level tp={} fp={} fn={} tn={}",
utt_tp, utt_fp, utt_fn, utt_tn
);
println!(
" precision={} recall={} F1={}",
fmt_pct(utt_f1.precision),
fmt_pct(utt_f1.recall),
fmt_pct(utt_f1.f1),
);
println!(
"correction-pair tp={} fp={} fn={}",
pair_tp, pair_fp, pair_fn
);
println!(
" precision={} recall={} F1={}",
fmt_pct(pair_f1.precision),
fmt_pct(pair_f1.recall),
fmt_pct(pair_f1.f1),
);
if let Some(out) = &cli.output {
let report = serde_json::json!({
"task": "phoneme-correction",
"lang": cli.lang,
"annotations": annotations.len(),
"dict_words": dict_map.len(),
"utterance": {
"tp": utt_tp, "fp": utt_fp, "fn": utt_fn, "tn": utt_tn,
"precision": utt_f1.precision,
"recall": utt_f1.recall,
"f1": utt_f1.f1,
},
"correction_pair": {
"tp": pair_tp, "fp": pair_fp, "fn": pair_fn,
"precision": pair_f1.precision,
"recall": pair_f1.recall,
"f1": pair_f1.f1,
},
});
std::fs::write(out, serde_json::to_string_pretty(&report).unwrap())
.map_err(|e| format!("write {}: {e}", out.display()))?;
eprintln!("report written to {}", out.display());
}
enforce_min_f1(cli.min_f1, pair_f1.f1, "correction-pair")
}
async fn run_ablation(cli: &Cli) -> Result<(), String> {
let fixtures = load_fixtures(&cli.input)
.map_err(|e| format!("loading fixtures {}: {e}", cli.input.display()))?;
if fixtures.is_empty() {
return Err(format!("fixture file {} is empty", cli.input.display()));
}
let result = evaluate_ablation_for_lang(&cli.lang, &fixtures).await?;
println!("=== L3 ablation ({}) ===", cli.lang);
println!("fixtures: {}", result.fixtures);
let primary = match cli.lang.as_str() {
"en" | "es" => "WER",
_ => "CER",
};
for (cfg, er) in &result.ablation {
println!(" ablation/{cfg:30} {primary}={er:.4}");
}
for (layer, lat) in &result.layer_latency_us {
println!(
" latency/{layer:30} p50={:.1}μs p95={:.1}μs",
lat.p50, lat.p95
);
}
if let Some(out) = &cli.output {
let json = serde_json::to_string_pretty(&result).map_err(|e| format!("json: {e}"))?;
if let Some(parent) = out.parent() {
std::fs::create_dir_all(parent).ok();
}
std::fs::write(out, json).map_err(|e| format!("write {}: {e}", out.display()))?;
eprintln!("report written to {}", out.display());
}
Ok(())
}
#[derive(Debug, Clone, Copy)]
struct LayerConfig {
name: &'static str,
filter: bool,
self_correction: bool,
punctuation: bool,
}
const FULL: LayerConfig = LayerConfig {
name: "full",
filter: true,
self_correction: true,
punctuation: true,
};
const WITHOUT_FILLER: LayerConfig = LayerConfig {
name: "without_filler",
filter: false,
self_correction: true,
punctuation: true,
};
const WITHOUT_SC: LayerConfig = LayerConfig {
name: "without_self_correction",
filter: true,
self_correction: false,
punctuation: true,
};
const WITHOUT_PUNCT: LayerConfig = LayerConfig {
name: "without_punctuation",
filter: true,
self_correction: true,
punctuation: false,
};
async fn evaluate_ablation_for_lang(
lang: &str,
fixtures: &[Fixture],
) -> Result<LanguageLayerBaseline, String> {
let configs: Vec<LayerConfig> = match lang {
"en" | "ja" | "zh" | "es" | "ko" => {
vec![FULL, WITHOUT_FILLER, WITHOUT_SC, WITHOUT_PUNCT]
}
other => return Err(format!("unsupported lang {other}")),
};
let mut ablation = BTreeMap::new();
for cfg in &configs {
let er = mean_error_rate(lang, fixtures, cfg).await?;
ablation.insert(cfg.name.to_string(), round4(er));
}
let layer_latency = bench_layer_latency(lang, fixtures, 10, 100).await;
Ok(LanguageLayerBaseline {
fixtures: fixtures.len(),
ablation,
layer_latency_us: layer_latency,
})
}
async fn mean_error_rate(
lang: &str,
fixtures: &[Fixture],
cfg: &LayerConfig,
) -> Result<f64, String> {
let mut sum = 0.0;
let mut counted = 0;
for fix in fixtures {
let pipeline = build_pipeline(lang, cfg, &fix.asr_hypothesis)?;
let audio = vec![AudioChunk {
samples: vec![0.0; 160],
sample_rate: 16000,
channels: 1,
}];
let result = pipeline.transcribe(&audio).await
.map_err(|e| format!("pipeline: {e}"))?;
let RefinementOutput::TextInsertion { text, .. } = &result.output else {
return Err("expected TextInsertion".into());
};
let er = match lang {
"en" | "es" => wer(&fix.reference, text),
_ => cer(&fix.reference, text),
};
if !er.is_nan() {
sum += er;
counted += 1;
}
}
if counted == 0 {
return Err("no scorable fixtures".into());
}
Ok(sum / counted as f64)
}
fn build_pipeline(lang: &str, cfg: &LayerConfig, hypothesis: &str) -> Result<Pipeline, String> {
let mut builder = Pipeline::builder()
.asr(MockAsr::new(hypothesis))
.refiner(MockRefiner::passthrough())
.context(MockContextProvider::new())
.emitter(MockEmitter::new());
if cfg.filter {
builder = match lang {
"en" => builder.filter(SimpleFillerFilter::english()),
"ja" => builder.filter(JapaneseFillerFilter::new()),
"zh" => builder.filter(ChineseFillerFilter::new()),
"es" => builder.filter(SpanishFillerFilter::new()),
"ko" => builder.filter(SimpleFillerFilter::korean()),
other => return Err(format!("unsupported lang {other}")),
};
}
if cfg.self_correction {
builder = builder.processor(SelfCorrectionDetector::new());
}
if cfg.punctuation {
builder = builder.processor(BasicPunctuationRestorer);
}
builder.build().map_err(|e| format!("build pipeline: {e}"))
}
async fn bench_layer_latency(
lang: &str,
fixtures: &[Fixture],
warmup: usize,
iters: usize,
) -> BTreeMap<String, LatencyMicrosRecord> {
let mut out = BTreeMap::new();
match lang {
"en" => {
let f = SimpleFillerFilter::english();
out.insert(
"filler".to_string(),
bench_filter(&f, fixtures, warmup, iters).await,
);
}
"ja" => {
let f = JapaneseFillerFilter::new();
out.insert(
"filler".to_string(),
bench_filter(&f, fixtures, warmup, iters).await,
);
}
"zh" => {
let f = ChineseFillerFilter::new();
out.insert(
"filler".to_string(),
bench_filter(&f, fixtures, warmup, iters).await,
);
}
"es" => {
let f = SpanishFillerFilter::new();
out.insert(
"filler".to_string(),
bench_filter(&f, fixtures, warmup, iters).await,
);
}
"ko" => {
let f = SimpleFillerFilter::korean();
out.insert(
"filler".to_string(),
bench_filter(&f, fixtures, warmup, iters).await,
);
}
_ => {}
}
let sc = SelfCorrectionDetector::new();
out.insert(
"self_correction".to_string(),
bench_processor(&sc, fixtures, warmup, iters).await,
);
let punct = BasicPunctuationRestorer;
out.insert(
"punctuation".to_string(),
bench_processor(&punct, fixtures, warmup, iters).await,
);
out
}
async fn bench_filter<F: TextFilter>(
layer: &F,
fixtures: &[Fixture],
warmup: usize,
iters: usize,
) -> LatencyMicrosRecord {
for fix in fixtures.iter().take(warmup.min(fixtures.len())) {
let _ = layer.filter(&fix.asr_hypothesis).await;
}
let mut samples = Samples::new();
for _ in 0..iters {
for fix in fixtures {
let start = Instant::now();
let _ = layer.filter(&fix.asr_hypothesis).await;
samples.record(start.elapsed());
}
}
samples.summary().expect("non-empty fixtures").into()
}
async fn bench_processor<P: TextProcessor>(
layer: &P,
fixtures: &[Fixture],
warmup: usize,
iters: usize,
) -> LatencyMicrosRecord {
let ctx = ContextSnapshot::default();
for fix in fixtures.iter().take(warmup.min(fixtures.len())) {
let _ = layer.process(&fix.asr_hypothesis, &ctx).await;
}
let mut samples = Samples::new();
for _ in 0..iters {
for fix in fixtures {
let start = Instant::now();
let _ = layer.process(&fix.asr_hypothesis, &ctx).await;
samples.record(start.elapsed());
}
}
samples.summary().expect("non-empty fixtures").into()
}
fn round4(x: f64) -> f64 {
(x * 10_000.0).round() / 10_000.0
}