PersianASR-TrippleThreat / benchmark_gold69_v2_vosk.py
Reza2kn's picture
Update NVIDIA Gold69 v2 standardized benchmark
8414ba9 verified
Raw
History Blame
3.67 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 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()