#!/usr/bin/env python3 import io import json import re import sys import time import unicodedata from pathlib import Path import jiwer import numpy as np import pyarrow.parquet as pq import soundfile as sf import torch from scipy.signal import resample_poly sys.path.insert(0, "/workspace/golha/vibevoice_repo") from vibevoice.modular.modeling_vibevoice_asr import VibeVoiceASRForConditionalGeneration from vibevoice.processor.vibevoice_asr_processor import VibeVoiceASRProcessor PERSIAN_DIGITS = "۰۱۲۳۴۵۶۷۸۹٠١٢٣٤٥٦٧٨٩" ASCII_DIGITS = "01234567890123456789" DIGIT_MAP = str.maketrans(PERSIAN_DIGITS, ASCII_DIGITS) PUNCT_RE = re.compile(r"[،؛؟,.\!?\(\)\[\]\{\}\"'`«»“”‘’\-_=+/\\:;]") def normalize(text): text = unicodedata.normalize("NFKC", text or "") text = text.replace("‌", " ").replace("​", "") text = text.replace("ي", "ی").replace("ك", "ک").translate(DIGIT_MAP) text = PUNCT_RE.sub(" ", text) return re.sub(r"\s+", " ", text).strip() def read_audio(path_or_bytes, base_dir=None): if isinstance(path_or_bytes, (bytes, bytearray)): wav, sr = sf.read(io.BytesIO(path_or_bytes), dtype="float32") else: path = Path(path_or_bytes) if base_dir is not None and not path.is_absolute(): path = Path(base_dir) / path wav, sr = sf.read(path, dtype="float32") if wav.ndim > 1: wav = wav.mean(axis=-1) return wav, sr def resample_to_24k(wav, sr): if sr == 24000: return wav.astype(np.float32) g = np.gcd(sr, 24000) return resample_poly(wav, 24000 // g, sr // g).astype(np.float32) def load_gold(path="/workspace/golha/gold.jsonl"): base = Path(path).parent rows = [] for line in open(path): item = json.loads(line) ref = (item.get("gold_transcript") or "").strip() if not ref: continue wav, sr = read_audio(item["audio_path"], base) rows.append({"id": str(item.get("programmeitem_id")), "wav": resample_to_24k(wav, sr), "ref": ref}) return rows def load_fleurs(fleurs_dir="/workspace/golha/persian-eval"): files = sorted(Path(fleurs_dir).glob("data/*.parquet")) or sorted(Path(fleurs_dir).glob("*.parquet")) records = pq.read_table([str(p) for p in files]).to_pylist() rows = [] for idx, item in enumerate(records): audio = item["audio"] audio_bytes = audio.get("bytes") if isinstance(audio, dict) else None if audio_bytes is None: continue wav, sr = read_audio(audio_bytes) rows.append( { "id": str(item.get("id", idx)), "wav": resample_to_24k(wav, sr), "ref": item.get("transcription") or item.get("raw_transcription") or "", } ) return rows class VibeVoiceRunner: def __init__(self, model_path="microsoft/VibeVoice-ASR"): self.processor = VibeVoiceASRProcessor.from_pretrained( model_path, language_model_pretrained_name="Qwen/Qwen2.5-7B", trust_remote_code=True, ) self.model = VibeVoiceASRForConditionalGeneration.from_pretrained( model_path, dtype=torch.bfloat16, attn_implementation="sdpa", trust_remote_code=True, ).to("cuda") self.model.eval() def transcribe(self, wav_24k): inputs = self.processor( audio=[wav_24k], sampling_rate=24000, return_tensors="pt", padding=True, add_generation_prompt=True, ) inputs = {k: v.to("cuda") if isinstance(v, torch.Tensor) else v for k, v in inputs.items()} with torch.no_grad(): 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, ) 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 return raw.strip(), raw, segments def eval_set(runner, set_name, clips, pred_dir): refs = [] preds = [] times = [] pred_path = pred_dir / f"microsoft__VibeVoice-ASR__{set_name}_predictions.jsonl" start = time.time() with pred_path.open("w") as f: for i, clip in enumerate(clips, 1): torch.cuda.synchronize() t0 = time.time() hyp, raw, segments = runner.transcribe(clip["wav"]) torch.cuda.synchronize() times.append(time.time() - t0) ref_norm = normalize(clip["ref"]) hyp_norm = normalize(hyp) refs.append(ref_norm) preds.append(hyp_norm) f.write( json.dumps( { "id": clip["id"], "ref": clip["ref"], "hyp": hyp, "raw": raw, "segments": segments, "ref_norm": ref_norm, "hyp_norm": hyp_norm, }, ensure_ascii=False, ) + "\n" ) if i % 25 == 0: print(f"[progress] microsoft/VibeVoice-ASR {set_name} {i}/{len(clips)} mean_ms={np.mean(times)*1000:.1f}", flush=True) return { "model": "microsoft/VibeVoice-ASR", "repo": "microsoft/VibeVoice-ASR", "family": "VibeVoice ASR", "precision": "bf16", "params_b": 8.674, "set": set_name, "n": len(clips), "wer": jiwer.wer(refs, preds), "cer": jiwer.cer(refs, preds), "mean_decode_ms": float(np.mean(times) * 1000.0), "wall_sec": time.time() - start, "predictions": str(pred_path), "status": "complete", } def main(): pred_dir = Path("/workspace/golha/public_asr_double_benchmark_predictions") pred_dir.mkdir(exist_ok=True) out = Path("/workspace/golha/public_asr_vibevoice_benchmark.jsonl") done = set() if out.exists(): for line in out.read_text().splitlines(): try: row = json.loads(line) done.add((row.get("model"), row.get("set"))) except Exception: pass runner = VibeVoiceRunner() evalsets = {"gold69": load_gold(), "fleurs": load_fleurs()} with out.open("a", buffering=1) as f: for set_name, clips in evalsets.items(): key = ("microsoft/VibeVoice-ASR", set_name) if key in done: print("[skip] microsoft/VibeVoice-ASR", set_name, flush=True) continue print("[eval] microsoft/VibeVoice-ASR", set_name, "n=" + str(len(clips)), flush=True) rec = eval_set(runner, set_name, clips, pred_dir) print(f"[result] microsoft/VibeVoice-ASR {set_name} WER={rec['wer']*100:.2f}% CER={rec['cer']*100:.2f}% wall={rec['wall_sec']:.1f}s", flush=True) f.write(json.dumps(rec, ensure_ascii=False) + "\n") if __name__ == "__main__": main()