"""Validate the Russian CosyVoice sample pack and optionally run GigaAM ASR.""" from __future__ import annotations import argparse import csv from datetime import datetime, timezone import json from pathlib import Path import re from statistics import mean from typing import Any import numpy as np import soundfile as sf try: from glados_ru.align_pitch import median_f0 except ModuleNotFoundError: # Support ``python glados_ru/qa_cosyvoice_batch.py``. from align_pitch import median_f0 ROOT = Path(__file__).parent DEFAULT_REFERENCES = ROOT / "references" DEFAULT_OUTPUTS = ROOT / "generated_cosyvoice_vc_batch_gpu_pitch" DEFAULT_REPORT = ROOT / "qa_report.json" def transcript_words(text: str) -> list[str]: """Return comparable Russian words; ASR commonly folds ``ё`` into ``е``.""" return re.findall(r"[а-я0-9]+", text.lower().replace("ё", "е")) def word_error_rate(expected: str, actual: str) -> float: """Compute Levenshtein distance over normalized words.""" left = transcript_words(expected) right = transcript_words(actual) distance = list(range(len(right) + 1)) for left_index, left_word in enumerate(left, 1): updated = [left_index] for right_index, right_word in enumerate(right, 1): updated.append( min( updated[-1] + 1, distance[right_index] + 1, distance[right_index - 1] + (left_word != right_word), ) ) distance = updated return distance[-1] / max(1, len(left)) def wav_ids(directory: Path) -> set[str]: return {path.stem for path in directory.glob("*.wav")} def require_matching_ids(directory: Path, expected: set[str]) -> None: actual = wav_ids(directory) if actual != expected: missing = ",".join(sorted(expected - actual)) or "-" extra = ",".join(sorted(actual - expected)) or "-" raise ValueError(f"{directory}: missing={missing}; extra={extra}") def load_translations(root: Path = ROOT) -> dict[str, str]: translations: dict[str, str] = {} for path in sorted(root.glob("translations_*.tsv")): with path.open(encoding="utf-8", newline="") as handle: for row in csv.DictReader(handle, delimiter="\t"): translations[row["id"].zfill(4)] = row["ru"] return translations def _load_asr(enabled: bool) -> Any | None: if not enabled: return None try: import onnx_asr except ImportError as exc: raise RuntimeError("Install onnx-asr to use --asr") from exc return onnx_asr.load_model( "gigaam-multilingual-ctc", quantization="int8", providers=["CPUExecutionProvider"], ) def audit( *, references: Path, outputs: Path, translations: dict[str, str], asr: bool, max_duration_delta: float, max_f0_error: float, max_peak: float, max_asr_wer: float, ) -> dict[str, Any]: expected_ids = set(translations) require_matching_ids(references, expected_ids) require_matching_ids(outputs, expected_ids) recognizer = _load_asr(asr) samples: list[dict[str, Any]] = [] for item_id in sorted(expected_ids): reference = references / f"{item_id}.wav" output = outputs / f"{item_id}.wav" reference_info = sf.info(reference) output_info = sf.info(output) audio, _ = sf.read(output, always_2d=False) peak = float(np.max(np.abs(audio))) duration_delta = abs(output_info.duration - reference_info.duration) reference_f0 = median_f0(reference) output_f0 = median_f0(output) f0_error = abs(output_f0 - reference_f0) / reference_f0 checks = { "sample_rate": output_info.samplerate == 24_000, "mono": output_info.channels == 1, "pcm16": output_info.subtype == "PCM_16", "finite": bool(np.isfinite(audio).all()), "duration": duration_delta <= max_duration_delta, "peak": peak <= max_peak, "median_f0": f0_error <= max_f0_error, } transcript = None wer = None if recognizer is not None: transcript = recognizer.recognize(str(output)) wer = word_error_rate(translations[item_id], transcript) checks["asr_wer"] = wer <= max_asr_wer samples.append( { "id": item_id, "passed": all(checks.values()), "checks": checks, "sample_rate": output_info.samplerate, "channels": output_info.channels, "subtype": output_info.subtype, "duration_seconds": round(output_info.duration, 6), "reference_duration_seconds": round(reference_info.duration, 6), "duration_delta_seconds": round(duration_delta, 6), "peak": round(peak, 6), "median_f0_hz": round(output_f0, 3), "reference_median_f0_hz": round(reference_f0, 3), "median_f0_relative_error": round(f0_error, 6), "expected_text": translations[item_id], "asr_transcript": transcript, "asr_wer": round(wer, 6) if wer is not None else None, } ) wers = [sample["asr_wer"] for sample in samples if sample["asr_wer"] is not None] return { "generated_at": datetime.now(timezone.utc).isoformat(), "recognizer": "onnx-asr/gigaam-multilingual-ctc-int8" if asr else None, "thresholds": { "max_duration_delta_seconds": max_duration_delta, "max_median_f0_relative_error": max_f0_error, "max_peak": max_peak, "max_asr_wer": max_asr_wer if asr else None, }, "summary": { "sample_count": len(samples), "passed_count": sum(sample["passed"] for sample in samples), "exact_asr_count": sum(wer == 0 for wer in wers) if wers else None, "mean_asr_wer": round(mean(wers), 6) if wers else None, "max_asr_wer": round(max(wers), 6) if wers else None, "passed": all(sample["passed"] for sample in samples), }, "samples": samples, } def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--references", type=Path, default=DEFAULT_REFERENCES) parser.add_argument("--outputs", type=Path, default=DEFAULT_OUTPUTS) parser.add_argument("--report", type=Path, default=DEFAULT_REPORT) parser.add_argument("--asr", action="store_true") parser.add_argument("--max-duration-delta", type=float, default=0.08) parser.add_argument("--max-f0-error", type=float, default=0.025) parser.add_argument("--max-peak", type=float, default=0.951) parser.add_argument("--max-asr-wer", type=float, default=0.10) args = parser.parse_args() report = audit( references=args.references, outputs=args.outputs, translations=load_translations(), asr=args.asr, max_duration_delta=args.max_duration_delta, max_f0_error=args.max_f0_error, max_peak=args.max_peak, max_asr_wer=args.max_asr_wer, ) args.report.parent.mkdir(parents=True, exist_ok=True) args.report.write_text( json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8", ) summary = report["summary"] print(json.dumps(summary, ensure_ascii=False)) if not summary["passed"]: raise SystemExit(1) if __name__ == "__main__": main()