#!/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()