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("Gtk", "3.0")
|
||||||
gi.require_version("Gdk", "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 .config import Config, save # noqa: E402
|
||||||
from .llm import LLMEngine # noqa: E402
|
from .llm import LLMEngine # noqa: E402
|
||||||
from .stt import STTEngine # noqa: E402
|
from .stt import STTEngine # noqa: E402
|
||||||
@ -195,6 +195,7 @@ class SettingsDialog:
|
|||||||
self.cfg = cfg
|
self.cfg = cfg
|
||||||
self.daemon = daemon
|
self.daemon = daemon
|
||||||
self._meter = None
|
self._meter = None
|
||||||
|
self._tr_cache: dict = {}
|
||||||
self._wf_idx = self._stt_idx = self._llm_idx = 0
|
self._wf_idx = self._stt_idx = self._llm_idx = 0
|
||||||
|
|
||||||
self.dlg = Gtk.Dialog(title="Blitztext — Settings", transient_for=parent, modal=True)
|
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_engines(_page(nb, "Engines"))
|
||||||
self._build_input(_page(nb, "Input"))
|
self._build_input(_page(nb, "Input"))
|
||||||
self._build_general(_page(nb, "General"))
|
self._build_general(_page(nb, "General"))
|
||||||
|
self._build_benchmark(_page(nb, "Benchmark"))
|
||||||
self._build_log(_page(nb, "Log"))
|
self._build_log(_page(nb, "Log"))
|
||||||
|
|
||||||
self._bind_entry = None
|
self._bind_entry = None
|
||||||
@ -442,13 +444,30 @@ class SettingsDialog:
|
|||||||
self.stt_result.set_markup("<i>Recording 4s — speak now…</i>")
|
self.stt_result.set_markup("<i>Recording 4s — speak now…</i>")
|
||||||
threading.Thread(target=self._run_stt_test, args=(e,), daemon=True).start()
|
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):
|
def _run_stt_test(self, engine):
|
||||||
from .recorder import Recording, detect_recorder
|
from .recorder import Recording, detect_recorder
|
||||||
try:
|
try:
|
||||||
rec = Recording(detect_recorder(self.cfg.recorder), self.cfg.mic)
|
rec = Recording(detect_recorder(self.cfg.recorder), self.cfg.mic)
|
||||||
time.sleep(4.0)
|
time.sleep(4.0)
|
||||||
wav = rec.stop()
|
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)
|
res = stt.benchmark(engine, wav, language=self.cfg.language, local_transcriber=tr)
|
||||||
wav.unlink(missing_ok=True)
|
wav.unlink(missing_ok=True)
|
||||||
if res.ok:
|
if res.ok:
|
||||||
@ -648,6 +667,76 @@ class SettingsDialog:
|
|||||||
if self._meter is not None:
|
if self._meter is not None:
|
||||||
self._meter.stop(); self._meter = 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 ==============================================================
|
# ===== Log ==============================================================
|
||||||
def _build_log(self, page: Gtk.Box) -> None:
|
def _build_log(self, page: Gtk.Box) -> None:
|
||||||
sw = Gtk.ScrolledWindow(); sw.set_policy(Gtk.PolicyType.AUTOMATIC, Gtk.PolicyType.AUTOMATIC)
|
sw = Gtk.ScrolledWindow(); sw.set_policy(Gtk.PolicyType.AUTOMATIC, Gtk.PolicyType.AUTOMATIC)
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user