Download Models/Russian_CosyVoice3/qa_cosyvoice_batch.py from Random118/GLaDOS_TTS: direct link, hf CLI and curl.
- Browser
- Download file 7.61 kB
-
https://huggingface.co/Random118/GLaDOS_TTS/resolve/main/Models/Russian_CosyVoice3/qa_cosyvoice_batch.py
- Command line
-
hf download hf://Random118/GLaDOS_TTS/Models/Russian_CosyVoice3/qa_cosyvoice_batch.py
-
curl -L -o qa_cosyvoice_batch.py https://huggingface.co/Random118/GLaDOS_TTS/resolve/main/Models/Russian_CosyVoice3/qa_cosyvoice_batch.py
7.61 kB
| """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() | |