summaryrefslogtreecommitdiff
path: root/tools/whistle-eval/evaluate.py
diff options
context:
space:
mode:
authorTom Cooks <tommasogagliardi+github@gmail.com>2026-10-10 15:42:04 -0400
committerTom Cooks <tommasogagliardi+github@gmail.com>2026-10-10 15:42:04 -0400
commit901bf58c5cf9151942f8c4cd4e1427a02dec4b52 (patch)
treee9e312fb9893881f142a6f79e30417cd0c3f2be8 /tools/whistle-eval/evaluate.py
parent6db1714d0354d1e1105e330ce84c46a54487c3bb (diff)
downloadreccoon-901bf58c5cf9151942f8c4cd4e1427a02dec4b52.tar.gz
RECCoon 1.0.0: rebrand to Tom Cooks, Parakeet-only, GPLv3, F-Droid-ready
Diffstat (limited to 'tools/whistle-eval/evaluate.py')
-rw-r--r--tools/whistle-eval/evaluate.py141
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())