PersianASR-TrippleThreat / benchmark_gold69_v2_vibevoice.py
Reza2kn's picture
Update NVIDIA Gold69 v2 standardized benchmark
8414ba9 verified
Raw
History Blame
5.94 kB
#!/usr/bin/env python3
from __future__ import annotations
import argparse
import io
import json
import sys
import time
from pathlib import Path
import numpy as np
import pyarrow.parquet as pq
import soundfile as sf
import torch
from scipy.signal import resample_poly
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)
if sr != 24000:
g = np.gcd(sr, 24000)
wav = resample_poly(wav, 24000 // g, sr // g).astype("float32", copy=False)
sr = 24000
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)
class Runner:
def __init__(self, model_id: str, device: str, device_map: str | None):
from vibevoice.modular.modeling_vibevoice_asr import VibeVoiceASRForConditionalGeneration
from vibevoice.processor.vibevoice_asr_processor import VibeVoiceASRProcessor
self.processor = VibeVoiceASRProcessor.from_pretrained(
model_id,
language_model_pretrained_name="Qwen/Qwen2.5-7B",
trust_remote_code=True,
)
kwargs = {
"dtype": torch.bfloat16,
"attn_implementation": "sdpa",
"trust_remote_code": True,
}
if device_map:
kwargs["device_map"] = device_map
self.model = VibeVoiceASRForConditionalGeneration.from_pretrained(model_id, **kwargs)
if not device_map:
self.model = self.model.to(device)
self.device = next(self.model.parameters()).device
self.model.eval()
self.params = sum(p.numel() for p in self.model.parameters())
@torch.no_grad()
def transcribe(self, wav):
inputs = self.processor(
audio=[wav],
sampling_rate=24000,
return_tensors="pt",
padding=True,
add_generation_prompt=True,
)
inputs = {k: v.to(self.device) if isinstance(v, torch.Tensor) else v for k, v in inputs.items()}
if self.device.type == "cuda":
torch.cuda.synchronize()
t0 = time.time()
output_ids = self.model.generate(
**inputs,
max_new_tokens=512,
pad_token_id=self.processor.pad_id,
eos_token_id=self.processor.tokenizer.eos_token_id,
do_sample=False,
num_beams=1,
)
if self.device.type == "cuda":
torch.cuda.synchronize()
input_len = inputs["input_ids"].shape[1]
generated_ids = output_ids[0, input_len:]
eos = (generated_ids == self.processor.tokenizer.eos_token_id).nonzero(as_tuple=True)[0]
if len(eos) > 0:
generated_ids = generated_ids[: eos[0] + 1]
raw = self.processor.decode(generated_ids, skip_special_tokens=True)
try:
segments = self.processor.post_process_transcription(raw)
except Exception:
segments = []
if segments:
text = " ".join((seg.get("text") or seg.get("Content") or "").strip() for seg in segments)
if text.strip():
return text.strip(), raw, segments, time.time() - t0
return raw.strip(), raw, segments, time.time() - t0
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("--device-map", default=None)
ap.add_argument("--vibevoice-repo", default="")
ap.add_argument("--model", default="microsoft/VibeVoice-ASR")
args = ap.parse_args()
if args.vibevoice_repo:
sys.path.insert(0, args.vibevoice_repo)
rows = load_gold(Path(args.gold_parquet))
out_dir = Path(args.out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
runner = Runner(args.model, args.device, args.device_map)
pred_path = out_dir / "microsoft__VibeVoice-ASR__gold69_v2_predictions.jsonl"
preds = []
times = []
with pred_path.open("w") as f:
for i, row in enumerate(rows, 1):
hyp, raw, segments, elapsed = runner.transcribe(row["wav"])
preds.append(hyp)
times.append(elapsed)
f.write(json.dumps({"id": row["id"], "ref": row["ref"], "hyp": hyp, "raw": raw, "segments": segments}, ensure_ascii=False) + "\n")
print(f"[pred] microsoft/VibeVoice-ASR {i}/{len(rows)}", flush=True)
wer, cer = score(rows, preds)
rec = {
"model": "microsoft/VibeVoice-ASR",
"n": len(rows),
"gold69_v2_fair_wer": wer,
"gold69_v2_fair_cer": cer,
"mean_decode_ms": float(np.mean(times) * 1000.0),
"params_b": runner.params / 1e9,
"predictions": str(pred_path),
"status": "complete",
}
out_path = out_dir / "public_gold69_v2_results.json"
existing = json.loads(out_path.read_text()) if out_path.exists() else []
existing = [r for r in existing if r.get("model") != rec["model"]] + [rec]
out_path.write_text(json.dumps(existing, ensure_ascii=False, indent=2) + "\n")
print(json.dumps(rec, ensure_ascii=False), flush=True)
if __name__ == "__main__":
main()