Spaces:
Running on Zero
Running on Zero
Download tests/fixtures/make_asr_ab_results.py from WolfDavid/japanese-learning-avatar: direct link, hf CLI and curl.
- Browser
- Download file 12.8 kB
-
https://huggingface.co/spaces/WolfDavid/japanese-learning-avatar/resolve/fbd985ff7754c8ea5773c2af17e7023a7b8e24fd/tests/fixtures/make_asr_ab_results.py
- Command line
-
hf download hf://spaces/WolfDavid/japanese-learning-avatar@fbd985ff7754c8ea5773c2af17e7023a7b8e24fd/tests/fixtures/make_asr_ab_results.py
-
curl -L -o make_asr_ab_results.py https://huggingface.co/spaces/WolfDavid/japanese-learning-avatar/resolve/fbd985ff7754c8ea5773c2af17e7023a7b8e24fd/tests/fixtures/make_asr_ab_results.py
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()) | |