#!/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 pyarrow.parquet as pq import soundfile as sf import torch from rescore_gold69_v2_fair import fair_text def read_audio(cell): wav, sr = sf.read(io.BytesIO(cell["bytes"]), dtype="float32") if wav.ndim > 1: wav = wav.mean(axis=-1) return wav, sr def load_gold(parquet_path: Path): table = pq.read_table(parquet_path) rows = [] for row in table.to_pylist(): wav, sr = read_audio(row["audio"]) rows.append({ "id": str(row["programmeitem_id"]), "wav": wav, "sr": sr, "ref": row["gold_transcript"] or "", }) return rows def score(rows, preds): import jiwer refs_words = [fair_text(r["ref"], keep_space=True) for r in rows] hyps_words = [fair_text(h, keep_space=True) for h in preds] refs_chars = [fair_text(r["ref"], keep_space=False) for r in rows] hyps_chars = [fair_text(h, keep_space=False) for h in preds] return jiwer.wer(refs_words, hyps_words), jiwer.cer(refs_chars, hyps_chars) @torch.no_grad() def run_model(model_id: str, rows, device: str, dtype_name: str, batch_size: int): from qwen_asr import Qwen3ASRModel dtype_map = { "bf16": torch.bfloat16, "bfloat16": torch.bfloat16, "fp16": torch.float16, "float16": torch.float16, "fp32": torch.float32, "float32": torch.float32, } dtype = dtype_map[dtype_name] model = Qwen3ASRModel.from_pretrained( model_id, dtype=dtype, device_map=device, max_inference_batch_size=batch_size, max_new_tokens=512, ) params = sum(p.numel() for p in model.model.parameters()) preds = [] times = [] for i, row in enumerate(rows, 1): if device.startswith("cuda"): torch.cuda.synchronize() t0 = time.time() result = model.transcribe((row["wav"], row["sr"]), language="Persian")[0] if device.startswith("cuda"): torch.cuda.synchronize() times.append(time.time() - t0) preds.append(result.text) print(f"[pred] {model_id} {i}/{len(rows)}", flush=True) del model if device.startswith("cuda"): torch.cuda.empty_cache() return preds, times, params def main(): ap = argparse.ArgumentParser() ap.add_argument("--gold-parquet", required=True) ap.add_argument("--out-dir", required=True) ap.add_argument("--device", default="cuda:0") ap.add_argument("--dtype", default="bf16") ap.add_argument("--batch-size", type=int, default=1) ap.add_argument("models", nargs="+") args = ap.parse_args() rows = load_gold(Path(args.gold_parquet)) out_dir = Path(args.out_dir) out_dir.mkdir(parents=True, exist_ok=True) existing = [] out_path = out_dir / "public_gold69_v2_results.json" if out_path.exists(): existing = json.loads(out_path.read_text()) results = [r for r in existing if r.get("model") not in set(args.models)] for model_id in args.models: safe = model_id.replace("/", "__") pred_path = out_dir / f"{safe}__gold69_v2_predictions.jsonl" preds, times, params = run_model(model_id, rows, args.device, args.dtype, args.batch_size) with pred_path.open("w") as f: for row, hyp in zip(rows, preds): f.write(json.dumps({"id": row["id"], "ref": row["ref"], "hyp": hyp}, ensure_ascii=False) + "\n") wer, cer = score(rows, preds) rec = { "model": model_id, "n": len(rows), "gold69_v2_fair_wer": wer, "gold69_v2_fair_cer": cer, "mean_decode_ms": float(np.mean(times) * 1000.0), "params_b": params / 1e9, "predictions": str(pred_path), "status": "complete", } results.append(rec) out_path.write_text(json.dumps(results, ensure_ascii=False, indent=2)) print(json.dumps(rec, ensure_ascii=False), flush=True) if __name__ == "__main__": main()