diff options
Diffstat (limited to 'tools/whistle-eval/evaluate.py')
| -rw-r--r-- | tools/whistle-eval/evaluate.py | 141 |
1 files changed, 141 insertions, 0 deletions
diff --git a/tools/whistle-eval/evaluate.py b/tools/whistle-eval/evaluate.py new file mode 100644 index 0000000..f4437ea --- /dev/null +++ b/tools/whistle-eval/evaluate.py @@ -0,0 +1,141 @@ +#!/usr/bin/env python3 +"""Compare Cactus Whistle with the app's Whisper models on the same clips. + +Usage: + python evaluate.py samples/jfk.wav en + python evaluate.py samples/*.wav # language guessed from a *.it.wav suffix + +For every clip this runs: + * Whisper tiny (int8, multilingual) through sherpa-onnx + * Whisper base (int8, multilingual) through sherpa-onnx + * Cactus Whistle through its reference Node pipeline +and prints the transcript plus the word error rate (WER) against +`samples/<name>.txt` when that reference file exists. + +Setup (see README.md): + npm install + node build-data.mjs + python3 -m venv .venv && .venv/bin/pip install sherpa-onnx numpy +""" +import argparse +import glob +import os +import re +import subprocess +import sys +import unicodedata +import wave + +HERE = os.path.dirname(os.path.abspath(__file__)) +VENV = os.environ.get("WHISTLE_EVAL_PYTHON", sys.executable) + + +def normalize(text: str) -> list: + text = unicodedata.normalize("NFKD", text.lower()) + text = "".join(c for c in text if not unicodedata.combining(c)) + text = re.sub(r"[^a-z0-9]+", " ", text) + return text.split() + + +def wer(reference: list, hypothesis: list) -> float: + if not reference: + return 0.0 if not hypothesis else 1.0 + previous = list(range(len(hypothesis) + 1)) + for i, ref in enumerate(reference, start=1): + current = [i] + for j, hyp in enumerate(hypothesis, start=1): + cost = 0 if ref == hyp else 1 + current.append(min(previous[j] + 1, current[j - 1] + 1, previous[j - 1] + cost)) + previous = current + return previous[-1] / len(reference) + + +def read_reference(audio: str): + base = os.path.splitext(audio)[0] + for candidate in (base + ".txt", audio + ".txt"): + if os.path.isfile(candidate): + with open(candidate, encoding="utf-8") as handle: + return handle.read().strip() + return None + + +def run_whistle(audio: str) -> str: + out = subprocess.run( + ["node", os.path.join(HERE, "vendor", "js", "example.mjs"), audio], + cwd=HERE, capture_output=True, text=True, timeout=600) + if out.returncode != 0: + return f"<whistle failed: {out.stderr.strip().splitlines()[-1] if out.stderr else '?'}>" + first = out.stdout.strip().splitlines()[0] if out.stdout.strip() else "" + return re.sub(r"^\[[a-z]{2}\]\s*", "", first) + + +def parakeet_available() -> bool: + return os.path.isfile(os.path.join( + HERE, "models", "parakeet-v3-int8", "encoder.int8.onnx")) + + +def run_parakeet(audio: str) -> str: + script = os.path.join(HERE, "parakeet_baseline.py") + out = subprocess.run( + [VENV, script, audio], capture_output=True, text=True, timeout=1200) + if out.returncode != 0: + tail = out.stderr.strip().splitlines() + return f"<parakeet failed: {tail[-1] if tail else '?'}>" + return out.stdout.strip() + + +def run_whisper(audio: str, size: str, language: str) -> str: + script = os.path.join(HERE, "whisper_baseline.py") + out = subprocess.run( + [VENV, script, audio, size, language], + capture_output=True, text=True, timeout=900) + if out.returncode != 0: + tail = out.stderr.strip().splitlines() + return f"<whisper {size} failed: {tail[-1] if tail else '?'}>" + return out.stdout.strip() + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("clips", nargs="+") + parser.add_argument("language", nargs="?", default=None) + args = parser.parse_args() + + clips = [] + for pattern in args.clips: + clips.extend(sorted(glob.glob(pattern)) or [pattern]) + + print(f"{'clip':<16} {'model':<14} {'WER':>7} transcript") + print("-" * 100) + for audio in clips: + if not os.path.isfile(audio): + print(f"{audio}: not found") + continue + language = args.language + name = os.path.basename(audio).lower() + if language is None: + language = "it" if (name.startswith("it") or ".it." in name) else "en" + reference = read_reference(audio) + ref_words = normalize(reference) if reference else None + + results = [ + ("whistle", run_whistle(audio)), + ("whisper tiny", run_whisper(audio, "tiny", language)), + ("whisper base", run_whisper(audio, "base", language)), + ] + if parakeet_available(): + results.append(("parakeet v3", run_parakeet(audio))) + for model, text in results: + if ref_words is None: + score = " n/a" + else: + score = f"{wer(ref_words, normalize(text)) * 100:6.1f}%" + print(f"{os.path.basename(audio):<16} {model:<14} {score} {text}") + if reference: + print(f"{'':<16} {'reference':<14} {'':>7} {reference}") + print() + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) |
