| |
| from __future__ import annotations |
|
|
| import argparse |
| import io |
| import json |
| import time |
| from pathlib import Path |
|
|
| import numpy as np |
| import pandas as pd |
| import soundfile as sf |
| from scipy.signal import resample_poly |
| from vosk import KaldiRecognizer, Model |
|
|
| from rescore_gold69_v2_fair import fair_text |
|
|
|
|
| def load_gold(path: str) -> list[dict]: |
| df = pd.read_parquet(path) |
| rows = [] |
| for i, row in df.iterrows(): |
| audio = row["audio"] |
| raw = audio.get("bytes") if isinstance(audio, dict) else audio |
| wav, sr = sf.read(io.BytesIO(raw), dtype="float32") |
| if wav.ndim > 1: |
| wav = wav.mean(axis=-1) |
| rows.append( |
| { |
| "id": str(row.get("programmeitem_id") or row.get("id") or i), |
| "wav": wav, |
| "sr": int(sr), |
| "ref": str(row.get("transcript") or row.get("gold_transcript") or row.get("text") or ""), |
| } |
| ) |
| return rows |
|
|
|
|
| def transcribe_one(model: Model, wav: np.ndarray, sr: int) -> str: |
| if sr != 16000: |
| g = np.gcd(sr, 16000) |
| wav = resample_poly(wav, 16000 // g, sr // g).astype("float32", copy=False) |
| sr = 16000 |
| rec = KaldiRecognizer(model, sr) |
| rec.SetWords(False) |
| pcm16 = (np.clip(wav, -1, 1) * 32767).astype("int16").tobytes() |
| chunk = sr * 2 |
| pieces = [] |
| for i in range(0, len(pcm16), chunk): |
| if rec.AcceptWaveform(pcm16[i : i + chunk]): |
| part = json.loads(rec.Result()).get("text") or "" |
| if part: |
| pieces.append(part) |
| final = json.loads(rec.FinalResult()).get("text") or "" |
| if final: |
| pieces.append(final) |
| return " ".join(pieces).strip() |
|
|
|
|
| def main() -> None: |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--gold-parquet", required=True) |
| ap.add_argument("--model-dir", required=True) |
| ap.add_argument("--out-dir", required=True) |
| args = ap.parse_args() |
|
|
| out_dir = Path(args.out_dir) |
| out_dir.mkdir(parents=True, exist_ok=True) |
| clips = load_gold(args.gold_parquet) |
| model = Model(args.model_dir) |
| pred_path = out_dir / "Vosk__gold69_v2_predictions.jsonl" |
| refs, hyps, times = [], [], [] |
| with pred_path.open("w", encoding="utf-8") as f: |
| for i, clip in enumerate(clips, 1): |
| t0 = time.time() |
| hyp = transcribe_one(model, clip["wav"], clip["sr"]) |
| times.append(time.time() - t0) |
| refs.append(clip["ref"]) |
| hyps.append(hyp) |
| f.write(json.dumps({"id": clip["id"], "ref": clip["ref"], "hyp": hyp}, ensure_ascii=False) + "\n") |
| print(f"[pred] Vosk {i}/{len(clips)}", flush=True) |
| import jiwer |
|
|
| scores = { |
| "gold69_v2_fair_wer": jiwer.wer( |
| [fair_text(r, keep_space=True) for r in refs], |
| [fair_text(h, keep_space=True) for h in hyps], |
| ), |
| "gold69_v2_fair_cer": jiwer.cer( |
| [fair_text(r, keep_space=False) for r in refs], |
| [fair_text(h, keep_space=False) for h in hyps], |
| ), |
| } |
| rec = { |
| "model": "Vosk", |
| "n": len(clips), |
| **scores, |
| "mean_decode_ms": float(np.mean(times) * 1000.0), |
| "params_b": 0.494933674, |
| "predictions": str(pred_path), |
| "status": "complete", |
| } |
| out = out_dir / "public_gold69_v2_results.json" |
| rows = json.loads(out.read_text()) if out.exists() else [] |
| rows = [r for r in rows if r.get("model") != rec["model"]] + [rec] |
| out.write_text(json.dumps(rows, ensure_ascii=False, indent=2) + "\n") |
| print(json.dumps(rec, ensure_ascii=False), flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|