#!/usr/bin/env python3 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()