PersianASR-TrippleThreat / benchmark_gold69_v2_ctc.py
Reza2kn's picture
Update leaderboard with corrected Gold69 v2 fair scoring
23073ab verified
Raw
History Blame
3.93 kB
#!/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 transformers import AutoModelForCTC, AutoProcessor
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
@torch.no_grad()
def run_model(model_id: str, rows, device: str, batch_size: int):
processor = AutoProcessor.from_pretrained(model_id)
model = AutoModelForCTC.from_pretrained(model_id).to(device).eval()
dtype = next(model.parameters()).dtype
params = sum(p.numel() for p in model.parameters())
preds = []
times = []
for start in range(0, len(rows), batch_size):
batch = rows[start:start + batch_size]
inputs = processor([r["wav"] for r in batch], sampling_rate=16000, return_tensors="pt", padding=True)
inputs = {k: (v.to(device).to(dtype) if v.is_floating_point() else v.to(device)) for k, v in inputs.items()}
if device.startswith("cuda"):
torch.cuda.synchronize()
t0 = time.time()
logits = model(**inputs).logits
if device.startswith("cuda"):
torch.cuda.synchronize()
times.append((time.time() - t0) / len(batch))
pred_ids = torch.argmax(logits, dim=-1)
preds.extend(processor.batch_decode(pred_ids))
print(f"[pred] {model_id} {start + len(batch)}/{len(rows)}", flush=True)
del model
if device.startswith("cuda"):
torch.cuda.empty_cache()
return preds, times, params
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)
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("--batch-size", type=int, default=8)
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)
results = []
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.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)
print(json.dumps(rec, ensure_ascii=False), flush=True)
(out_dir / "ctc_gold69_v2_results.json").write_text(json.dumps(results, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()