japanese-learning-avatar / tests /fixtures /make_asr_ab_results.py
WolfDavid's picture
feat(01-07): add tiered browser ASR and wire push-to-talk into the shared turn loop
b4ca58d
Raw History Blame
12.8 kB
"""Produce tests/fixtures/asr_ab_results.json: the measured ASR model A/B.
Run it, do not hand-edit its output::
uv run python tests/fixtures/make_asr_ab_results.py
What it does, and why each part is measured rather than looked up:
* Drives ``avatar/asr-harness.html`` over a loopback static server with real Playwright
Chromium, so every number comes from the same code path the Space will run.
* Gives each candidate a **fresh browser profile** for its cold load, then reopens that
profile for the warm load. Deleting the runtime's CacheStorage instead would be
cheaper, but ``caches.delete()`` followed by ``caches.open()`` throws "Unexpected
internal error" in Chromium 151 while the runtime still holds handles into it.
* Takes first-load size from the ``Content-Length`` of the exact ONNX files the runtime
fetches. Cross-origin resource timing reports ``transferSize: 0`` without
``Timing-Allow-Origin``, so the browser cannot honestly report this to itself.
* Includes the ``q8`` candidates **on purpose, expecting them to fail on WASM**. The
failure is the single most consequential finding of this exercise and it belongs in the
committed record, not in a commit message.
WebGPU note: headless Chromium exposes ``navigator.gpu`` but returns a null adapter, so
this script runs HEADED. A headless run will silently record every WebGPU candidate as
having fallen back to WASM, which is true but useless as an A/B.
"""
from __future__ import annotations
import argparse
import datetime as dt
import functools
import http.server
import json
import shutil
import socket
import sys
import tempfile
import threading
import time
import urllib.request
from pathlib import Path
HERE = Path(__file__).resolve().parent
REPO_ROOT = HERE.parent.parent
OUT = HERE / "asr_ab_results.json"
# The ground truth is exact because we generated the audio: plan 01-04 synthesised these
# three clips from these three strings with VOICEVOX. No human transcribed anything.
CLIPS = [
{"wav": "speech_ja.wav", "reference": "こんにちは"},
{
"wav": "speech_ja_long.wav",
"reference": "今日はいい天気ですから、公園を散歩してから、買い物に行きました。",
},
{
"wav": "speech_ja_slow.wav",
"reference": "今日はいい天気ですから、公園を散歩してから、買い物に行きました。",
},
]
GATE_FIXTURES = ["silence_30s.wav", "cafe_noise_30s.wav", *[c["wav"] for c in CLIPS]]
CANDIDATES = [
{"model": "onnx-community/whisper-base", "dtype": "q4", "device": "wasm"},
{"model": "onnx-community/whisper-base", "dtype": "q4", "device": "webgpu"},
{"model": "onnx-community/whisper-small", "dtype": "q4", "device": "wasm"},
{"model": "onnx-community/whisper-small", "dtype": "q4", "device": "webgpu"},
{"model": "onnx-community/whisper-large-v3-turbo", "dtype": "q4f16", "device": "webgpu"},
{"model": "onnx-community/whisper-large-v3-turbo", "dtype": "q4", "device": "wasm"},
# The q8 control group. 01-RESEARCH.md recommends exactly this configuration.
{"model": "onnx-community/whisper-base", "dtype": "q8", "device": "wasm"},
{"model": "onnx-community/whisper-base", "dtype": "q8", "device": "webgpu"},
{"model": "onnx-community/whisper-small", "dtype": "q8", "device": "wasm"},
]
SIZE_URL = "https://huggingface.co/{model}/resolve/main/onnx/{file}"
# Punctuation is a rendering choice, not a transcription error, so CER is reported both
# ways. A tutor cares about the second number.
PUNCTUATION = "、。,.,.!?!?  \n\t"
def levenshtein(a: str, b: str) -> int:
if a == b:
return 0
if not a:
return len(b)
if not b:
return len(a)
previous = list(range(len(b) + 1))
for i, ca in enumerate(a, start=1):
current = [i]
for j, cb in enumerate(b, start=1):
current.append(min(previous[j] + 1, current[j - 1] + 1, previous[j - 1] + (ca != cb)))
previous = current
return previous[-1]
def cer(hypothesis: str | None, reference: str) -> float | None:
if hypothesis is None:
return None
if not reference:
return None
return round(levenshtein(hypothesis, reference) / len(reference), 4)
def strip_punctuation(text: str) -> str:
return "".join(c for c in text if c not in PUNCTUATION)
def content_length(url: str) -> int | None:
request = urllib.request.Request(url, method="HEAD") # noqa: S310
request.add_header("User-Agent", "japanese-learning-avatar/asr-ab")
try:
with urllib.request.urlopen(request, timeout=60) as response: # noqa: S310
# LFS-backed files answer with x-linked-size; Content-Length is the pointer.
raw = response.headers.get("x-linked-size") or response.headers.get("Content-Length")
return int(raw) if raw else None
except OSError:
return None
# transformers.js does not name the file after the dtype: `q8` resolves to the
# `_quantized` suffix, not `_q8`. Getting this wrong silently reports "size unknown" for
# exactly the candidate the whole q8 control group exists to size.
DTYPE_SUFFIX = {
"fp32": "",
"fp16": "_fp16",
"q8": "_quantized",
"int8": "_int8",
"uint8": "_uint8",
"q4": "_q4",
"q4f16": "_q4f16",
"bnb4": "_bnb4",
}
def first_load_bytes(model: str, dtype: str) -> dict:
suffix = DTYPE_SUFFIX.get(dtype, f"_{dtype}")
files = [f"encoder_model{suffix}.onnx", f"decoder_model_merged{suffix}.onnx"]
sizes = {f: content_length(SIZE_URL.format(model=model, file=f)) for f in files}
known = [v for v in sizes.values() if isinstance(v, int)]
total = sum(known)
return {
"files": sizes,
"totalBytes": total if len(known) == len(files) else None,
"totalMB": round(total / 1024 / 1024, 1) if len(known) == len(files) else None,
}
def start_server(port: int):
handler = functools.partial(http.server.SimpleHTTPRequestHandler, directory=str(REPO_ROOT))
try:
server = http.server.ThreadingHTTPServer(("127.0.0.1", port), handler)
except OSError:
with socket.socket() as probe:
probe.bind(("127.0.0.1", 0))
port = probe.getsockname()[1]
server = http.server.ThreadingHTTPServer(("127.0.0.1", port), handler)
server.daemon_threads = True
threading.Thread(target=server.serve_forever, daemon=True).start()
return server, f"http://127.0.0.1:{server.server_port}"
WARM_LOAD = """
async (candidate) => {
const started = performance.now();
const asr = window.__newAsr(candidate);
await asr.init();
return { warmLoadMs: Math.round(performance.now() - started), tier: asr.getTier() };
}
"""
ENVIRONMENT = """
async () => ({
userAgent: navigator.userAgent,
platform: navigator.platform,
webgpuAvailable: !!navigator.gpu,
adapter: await (async () => {
if (!navigator.gpu) return null;
const a = await navigator.gpu.requestAdapter();
if (!a) return 'adapter-null';
const i = a.info || {};
const parts = [i.vendor, i.architecture, i.device, i.description].filter(Boolean);
return parts.join(' ') || 'adapter-ok';
})(),
})
"""
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--port", type=int, default=8477)
parser.add_argument("--headless", action="store_true", help="records WASM-only numbers")
parser.add_argument("--out", type=Path, default=OUT)
args = parser.parse_args()
try:
from playwright.sync_api import sync_playwright
except ImportError:
print("playwright is required: uv sync --extra dev && playwright install chromium")
return 2
server, base = start_server(args.port)
clips = [
{"url": f"{base}/tests/fixtures/{c['wav']}", "reference": c["reference"]} for c in CLIPS
]
profiles = Path(tempfile.mkdtemp(prefix="jla-asr-ab-"))
environment: dict = {}
gate_rows: list[dict] = []
records: list[dict] = []
def open_page(pw, profile: Path):
context = pw.chromium.launch_persistent_context(
user_data_dir=str(profile),
headless=args.headless,
args=["--autoplay-policy=no-user-gesture-required"],
)
page = context.new_page()
page.set_default_timeout(0)
page.goto(f"{base}/avatar/asr-harness.html")
page.wait_for_function("() => window.__harnessReady === true", timeout=120_000)
return context, page
try:
with sync_playwright() as pw:
for index, candidate in enumerate(CANDIDATES):
profile = profiles / f"p{index}"
profile.mkdir(parents=True, exist_ok=True)
label = f"{candidate['model']} {candidate['dtype']} {candidate['device']}"
print(f"--- cold {label}", flush=True)
started = time.time()
context, page = open_page(pw, profile)
if not environment:
environment = page.evaluate(ENVIRONMENT)
print(f" {environment}", flush=True)
# The gate table: measured off the WAV files themselves, so it is
# independent of any microphone, any fake device and any model.
for wav in GATE_FIXTURES:
gate_rows.append(
page.evaluate(
"(u) => window.__measureClip(u)",
f"{base}/tests/fixtures/{wav}",
)
| {"fixture": wav}
)
try:
record = page.evaluate(
"async ([c, clips]) => (await window.__runAB([c], clips))[0]",
[candidate, clips],
)
except Exception as err: # noqa: BLE001 - a driver failure is data too
record = {**candidate, "tier": None, "error": f"driver: {err}", "clips": []}
context.close()
record["requestedDevice"] = candidate["device"]
record.pop("cache", None)
print(f" tier={record.get('tier')} {time.time() - started:.1f}s", flush=True)
if record.get("tier"):
context, page = open_page(pw, profile)
try:
record["warmLoadMs"] = page.evaluate(WARM_LOAD, candidate)["warmLoadMs"]
except Exception as err: # noqa: BLE001
record["warmLoadError"] = str(err)[:300]
context.close()
print(f" warm={record.get('warmLoadMs')} ms", flush=True)
shutil.rmtree(profile, ignore_errors=True)
record["firstLoad"] = first_load_bytes(candidate["model"], candidate["dtype"])
for clip in record.get("clips", []):
transcript = clip.get("transcript")
reference = clip["reference"]
clip["cer"] = cer(transcript, reference)
clip["cerNoPunctuation"] = (
cer(strip_punctuation(transcript), strip_punctuation(reference))
if transcript is not None
else None
)
scored = [c["cerNoPunctuation"] for c in record.get("clips", [])]
scored = [v for v in scored if v is not None]
record["meanCerNoPunctuation"] = (
round(sum(scored) / len(scored), 4) if scored else None
)
records.append(record)
finally:
server.shutdown()
server.server_close()
shutil.rmtree(profiles, ignore_errors=True)
payload = {
"generated": dt.date.today().isoformat(),
"harness": "avatar/asr-harness.html driven by Playwright Chromium",
"runtime": "https://esm.sh/@huggingface/transformers@4.2.0",
"headless": args.headless,
"environment": environment,
"referenceIsGroundTruth": (
"The reference strings are the exact inputs plan 01-04 handed to VOICEVOX, so "
"CER is measured against known text, not against a human transcription."
),
"caveat": (
"These clips are VOICEVOX-synthesised speech: cleaner and more regular than a "
"learner speaking into a laptop microphone. Every CER here is a BEST CASE."
),
"gate": gate_rows,
"clips": CLIPS,
"results": records,
}
args.out.write_text(json.dumps(payload, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
print(f"wrote {args.out}", flush=True)
return 0
if __name__ == "__main__":
sys.exit(main())