| |
| 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() |
|
|