Add Benchmark tab (WAV + reference -> fastest & most accurate); fix STT Test
New benchmark.py: word-error-rate accuracy + a run() that times each STT engine on a reference clip. Settings gains a Benchmark tab: pick a .wav and a matching .txt, run all STT engines, see Time + Accuracy per engine and a summary of the fastest and most accurate. Add presets for each model you want compared. Fix "Local engine selected but the model isn't loaded." on Test: a cached _transcriber_for() loads the local model on demand (the daemon only preloads it when a local engine is active). Tested: local small 2.17s/100% vs remote :8010. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
parent
db387cab78
commit
1b0b4bbc1f
82
linux/blitztext/benchmark.py
Normal file
82
linux/blitztext/benchmark.py
Normal file
@ -0,0 +1,82 @@
|
||||
"""Benchmark STT engines against a reference clip.
|
||||
|
||||
Given a WAV and its reference transcript, transcribe with each engine, measure
|
||||
the time, and score accuracy as 1 − WER (word error rate). Used by the Settings
|
||||
Benchmark tab to find the fastest and most accurate engine/model.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
from . import stt
|
||||
from .routing import normalize
|
||||
|
||||
|
||||
@dataclass
|
||||
class BenchRow:
|
||||
engine: str
|
||||
model: str
|
||||
ok: bool
|
||||
seconds: float
|
||||
wer: float
|
||||
accuracy: float # percent, max(0, 1-wer)*100
|
||||
text: str
|
||||
error: str = ""
|
||||
|
||||
|
||||
def _edit_distance(a: list[str], b: list[str]) -> int:
|
||||
m, n = len(a), len(b)
|
||||
dp = list(range(n + 1))
|
||||
for i in range(1, m + 1):
|
||||
prev = dp[0]
|
||||
dp[0] = i
|
||||
for j in range(1, n + 1):
|
||||
cur = dp[j]
|
||||
dp[j] = min(dp[j] + 1, dp[j - 1] + 1, prev + (a[i - 1] != b[j - 1]))
|
||||
prev = cur
|
||||
return dp[n]
|
||||
|
||||
|
||||
def wer(reference: str, hypothesis: str) -> float:
|
||||
"""Word error rate (0 = perfect). Text is normalized (case/punct-insensitive)."""
|
||||
ref = normalize(reference)
|
||||
hyp = normalize(hypothesis)
|
||||
if not ref:
|
||||
return 0.0 if not hyp else 1.0
|
||||
return _edit_distance(ref, hyp) / len(ref)
|
||||
|
||||
|
||||
def run(engines, wav_path: Path, reference: str, *, language: str = "",
|
||||
get_local_transcriber=None, progress=None) -> list[BenchRow]:
|
||||
"""Benchmark each engine; calls progress(row) as each finishes."""
|
||||
rows: list[BenchRow] = []
|
||||
for e in engines:
|
||||
tr = get_local_transcriber(e) if (e.is_local and get_local_transcriber) else None
|
||||
res = stt.benchmark(e, wav_path, language=language, local_transcriber=tr)
|
||||
w = wer(reference, res.text) if res.ok else 1.0
|
||||
row = BenchRow(
|
||||
engine=e.name,
|
||||
model=e.model or ("local" if e.is_local else "(default)"),
|
||||
ok=res.ok,
|
||||
seconds=res.seconds,
|
||||
wer=w,
|
||||
accuracy=max(0.0, 1.0 - w) * 100.0,
|
||||
text=res.text,
|
||||
error=res.error,
|
||||
)
|
||||
rows.append(row)
|
||||
if progress:
|
||||
progress(row)
|
||||
return rows
|
||||
|
||||
|
||||
def best(rows: list[BenchRow]) -> tuple[BenchRow | None, BenchRow | None]:
|
||||
"""Return (fastest, most_accurate) among successful rows."""
|
||||
ok = [r for r in rows if r.ok]
|
||||
if not ok:
|
||||
return None, None
|
||||
fastest = min(ok, key=lambda r: r.seconds)
|
||||
most_accurate = max(ok, key=lambda r: (r.accuracy, -r.seconds))
|
||||
return fastest, most_accurate
|
||||
@ -19,9 +19,9 @@ import gi
|
||||
|
||||
gi.require_version("Gtk", "3.0")
|
||||
gi.require_version("Gdk", "3.0")
|
||||
from gi.repository import Gdk, GLib, Gtk # noqa: E402
|
||||
from gi.repository import Gdk, GLib, Gtk, Pango # noqa: E402
|
||||
|
||||
from . import audio, autostart, llm, logbuffer, stt # noqa: E402
|
||||
from . import audio, autostart, benchmark, llm, logbuffer, stt # noqa: E402
|
||||
from .config import Config, save # noqa: E402
|
||||
from .llm import LLMEngine # noqa: E402
|
||||
from .stt import STTEngine # noqa: E402
|
||||
@ -195,6 +195,7 @@ class SettingsDialog:
|
||||
self.cfg = cfg
|
||||
self.daemon = daemon
|
||||
self._meter = None
|
||||
self._tr_cache: dict = {}
|
||||
self._wf_idx = self._stt_idx = self._llm_idx = 0
|
||||
|
||||
self.dlg = Gtk.Dialog(title="Blitztext — Settings", transient_for=parent, modal=True)
|
||||
@ -209,6 +210,7 @@ class SettingsDialog:
|
||||
self._build_engines(_page(nb, "Engines"))
|
||||
self._build_input(_page(nb, "Input"))
|
||||
self._build_general(_page(nb, "General"))
|
||||
self._build_benchmark(_page(nb, "Benchmark"))
|
||||
self._build_log(_page(nb, "Log"))
|
||||
|
||||
self._bind_entry = None
|
||||
@ -442,13 +444,30 @@ class SettingsDialog:
|
||||
self.stt_result.set_markup("<i>Recording 4s — speak now…</i>")
|
||||
threading.Thread(target=self._run_stt_test, args=(e,), daemon=True).start()
|
||||
|
||||
def _transcriber_for(self, engine):
|
||||
"""A local Transcriber for the engine, loading on demand (cached)."""
|
||||
if not engine.is_local:
|
||||
return None
|
||||
model = engine.model or self.cfg.model
|
||||
active = self.cfg.active_stt
|
||||
if (self.daemon and self.daemon.transcriber and active.is_local
|
||||
and (active.model or self.cfg.model) == model):
|
||||
return self.daemon.transcriber
|
||||
key = (model, self.cfg.device, self.cfg.compute_type)
|
||||
if key not in self._tr_cache:
|
||||
from .transcribe import Transcriber
|
||||
logbuffer.log(f"Loading local model '{model}' ({self.cfg.device}) for test/benchmark…")
|
||||
self._tr_cache[key] = Transcriber(model, self.cfg.device, self.cfg.compute_type, self.cfg.beam_size)
|
||||
return self._tr_cache[key]
|
||||
|
||||
def _run_stt_test(self, engine):
|
||||
from .recorder import Recording, detect_recorder
|
||||
try:
|
||||
rec = Recording(detect_recorder(self.cfg.recorder), self.cfg.mic)
|
||||
time.sleep(4.0)
|
||||
wav = rec.stop()
|
||||
tr = self.daemon.transcriber if (self.daemon and engine.is_local) else None
|
||||
GLib.idle_add(self.stt_result.set_markup, "<i>Transcribing…</i>")
|
||||
tr = self._transcriber_for(engine)
|
||||
res = stt.benchmark(engine, wav, language=self.cfg.language, local_transcriber=tr)
|
||||
wav.unlink(missing_ok=True)
|
||||
if res.ok:
|
||||
@ -648,6 +667,76 @@ class SettingsDialog:
|
||||
if self._meter is not None:
|
||||
self._meter.stop(); self._meter = None
|
||||
|
||||
# ===== Benchmark ========================================================
|
||||
def _build_benchmark(self, page: Gtk.Box) -> None:
|
||||
page.pack_start(Gtk.Label(
|
||||
label="Benchmark all your STT engines against a reference clip "
|
||||
"(add presets for each model you want compared).",
|
||||
xalign=0.0, wrap=True), False, False, 0)
|
||||
|
||||
wavf = Gtk.FileChooserButton(title="WAV file", action=Gtk.FileChooserAction.OPEN)
|
||||
fa = Gtk.FileFilter(); fa.set_name("Audio (.wav)"); fa.add_pattern("*.wav"); wavf.add_filter(fa)
|
||||
self.bench_wav = _labeled(page, "Audio (.wav)", wavf)
|
||||
reff = Gtk.FileChooserButton(title="Reference transcript", action=Gtk.FileChooserAction.OPEN)
|
||||
ft = Gtk.FileFilter(); ft.set_name("Text (.txt)"); ft.add_pattern("*.txt"); reff.add_filter(ft)
|
||||
self.bench_ref = _labeled(page, "Reference (.txt)", reff)
|
||||
|
||||
run = Gtk.Button(label="Run benchmark"); run.connect("clicked", self._run_bench)
|
||||
run.set_halign(Gtk.Align.START)
|
||||
page.pack_start(run, False, False, 6)
|
||||
|
||||
self.bench_store = Gtk.ListStore(str, str, str, str, str)
|
||||
tree = Gtk.TreeView(model=self.bench_store)
|
||||
for title, i, expand in [("Engine", 0, False), ("Model", 1, False),
|
||||
("Time (s)", 2, False), ("Accuracy", 3, False), ("Output", 4, True)]:
|
||||
r = Gtk.CellRendererText()
|
||||
if i == 4:
|
||||
r.set_property("ellipsize", Pango.EllipsizeMode.END)
|
||||
col = Gtk.TreeViewColumn(title, r, text=i); col.set_resizable(True)
|
||||
col.set_expand(expand)
|
||||
tree.append_column(col)
|
||||
sw = Gtk.ScrolledWindow(); sw.set_policy(Gtk.PolicyType.AUTOMATIC, Gtk.PolicyType.AUTOMATIC)
|
||||
sw.add(tree); page.pack_start(sw, True, True, 4)
|
||||
|
||||
self.bench_summary = Gtk.Label(xalign=0.0); self.bench_summary.set_line_wrap(True)
|
||||
page.pack_start(self.bench_summary, False, False, 4)
|
||||
|
||||
def _run_bench(self, _b) -> None:
|
||||
wav = self.bench_wav.get_filename()
|
||||
refp = self.bench_ref.get_filename()
|
||||
if not wav or not refp:
|
||||
self._error("Pick a .wav file and a matching reference .txt file.")
|
||||
return
|
||||
reference = Path(refp).read_text(errors="replace")
|
||||
self.bench_store.clear()
|
||||
self.bench_summary.set_markup("<i>Running… (local models load on first use)</i>")
|
||||
self._stt_commit()
|
||||
engines = list(self.cfg.stt_engines)
|
||||
|
||||
def work():
|
||||
def prog(row):
|
||||
GLib.idle_add(self._bench_add_row, row)
|
||||
rows = benchmark.run(engines, Path(wav), reference, language=self.cfg.language,
|
||||
get_local_transcriber=self._transcriber_for, progress=prog)
|
||||
GLib.idle_add(self._bench_done, rows)
|
||||
threading.Thread(target=work, daemon=True).start()
|
||||
|
||||
def _bench_add_row(self, row) -> bool:
|
||||
acc = f"{row.accuracy:.1f}%" if row.ok else "—"
|
||||
out = row.text if row.ok else f"⚠ {row.error}"
|
||||
self.bench_store.append([row.engine, row.model, f"{row.seconds:.2f}", acc, out])
|
||||
return False
|
||||
|
||||
def _bench_done(self, rows) -> bool:
|
||||
fastest, acc = benchmark.best(rows)
|
||||
if not fastest:
|
||||
self.bench_summary.set_markup('<span foreground="#ff3b30">All engines failed — check the Log tab.</span>')
|
||||
return False
|
||||
self.bench_summary.set_markup(
|
||||
f"<b>Fastest:</b> {GLib.markup_escape_text(fastest.engine)} ({fastest.seconds:.2f}s)"
|
||||
f" · <b>Most accurate:</b> {GLib.markup_escape_text(acc.engine)} ({acc.accuracy:.1f}%)")
|
||||
return False
|
||||
|
||||
# ===== Log ==============================================================
|
||||
def _build_log(self, page: Gtk.Box) -> None:
|
||||
sw = Gtk.ScrolledWindow(); sw.set_policy(Gtk.PolicyType.AUTOMATIC, Gtk.PolicyType.AUTOMATIC)
|
||||
|
||||
Loading…
Reference in New Issue
Block a user