GLaDOS_TTS / Models /Russian_CosyVoice3 /qa_cosyvoice_batch.py
Random118's picture
Add Russian CosyVoice 3 dubbing pack and audio gallery
8c2913a verified
Raw History Blame Contribute Delete
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()