Spaces:
Running on Zero
Running on Zero
| #!/usr/bin/env python3 | |
| """Synthesize Praxy/Chatterbox audio from cleaned S2S text CSV rows.""" | |
| from __future__ import annotations | |
| import argparse | |
| import csv | |
| from pathlib import Path | |
| import sys | |
| import traceback | |
| from typing import Any | |
| ROOT = Path(__file__).resolve().parents[1] | |
| for path in (ROOT / "src", ROOT / "scripts"): | |
| if str(path) not in sys.path: | |
| sys.path.insert(0, str(path)) | |
| from run_s2s_live import DEFAULT_PRAXY_ANCHOR # noqa: E402 | |
| from urdu_s2s.schemas import BridgeResult, ReplyResult, SpeechToSpeechRequest # noqa: E402 | |
| from urdu_s2s.tts_providers import ChatterboxPraxyTTSProvider # noqa: E402 | |
| DEFAULT_INPUT_CSV = ROOT / "reports/evals/s2s_live_smoke10_text_v2.csv" | |
| DEFAULT_OUTPUT_DIR = ROOT / "reports/evals/s2s_live_smoke10_text_v2_audio" | |
| DEFAULT_SUMMARY_CSV = ROOT / "reports/evals/s2s_live_smoke10_text_v2_audio_summary.csv" | |
| FIELDNAMES = [ | |
| "id", | |
| "status", | |
| "assistant_reply_urdu", | |
| "devanagari_tts_text", | |
| "tts_audio_path", | |
| "duration_seconds", | |
| "error", | |
| ] | |
| def resolve_repo_path(path: Path) -> Path: | |
| return path if path.is_absolute() else ROOT / path | |
| def parse_ids(raw_ids: str) -> set[str] | None: | |
| ids = {part.strip() for part in raw_ids.split(",") if part.strip()} | |
| return ids or None | |
| def read_rows(path: Path, selected_ids: set[str] | None = None) -> list[dict[str, str]]: | |
| with path.open(newline="", encoding="utf-8") as handle: | |
| rows = list(csv.DictReader(handle)) | |
| if selected_ids is not None: | |
| rows = [row for row in rows if row.get("id") in selected_ids] | |
| return rows | |
| def write_summary(path: Path, rows: list[dict[str, object]]) -> None: | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| with path.open("w", newline="", encoding="utf-8") as handle: | |
| writer = csv.DictWriter(handle, fieldnames=FIELDNAMES) | |
| writer.writeheader() | |
| for row in rows: | |
| writer.writerow({field: row.get(field, "") for field in FIELDNAMES}) | |
| def run_tts_rows( | |
| *, | |
| input_csv: Path, | |
| output_dir: Path, | |
| summary_csv: Path, | |
| voice_prompt_audio_path: Path, | |
| selected_ids: set[str] | None = None, | |
| device: str = "cuda", | |
| t3_model: str = "v3", | |
| model_loader: Any | None = None, | |
| wav_writer: Any | None = None, | |
| duration_reader: Any | None = None, | |
| fail_fast: bool = False, | |
| ) -> list[dict[str, object]]: | |
| rows = read_rows(input_csv, selected_ids) | |
| if not rows: | |
| raise ValueError(f"No rows selected from {input_csv}") | |
| output_dir.mkdir(parents=True, exist_ok=True) | |
| provider = ChatterboxPraxyTTSProvider( | |
| output_audio_path=output_dir / "placeholder.wav", | |
| voice_prompt_audio_path=voice_prompt_audio_path, | |
| device=device, | |
| t3_model=t3_model, | |
| model_loader=model_loader, | |
| wav_writer=wav_writer, | |
| duration_reader=duration_reader, | |
| ) | |
| summary_rows: list[dict[str, object]] = [] | |
| for index, row in enumerate(rows, start=1): | |
| bench_id = row["id"] | |
| wav_path = output_dir / f"{bench_id}_praxy_text_v2.wav" | |
| print(f"[{index}/{len(rows)}] {bench_id} -> {wav_path}", flush=True) | |
| try: | |
| provider.output_audio_path = wav_path | |
| result = provider.synthesize( | |
| SpeechToSpeechRequest( | |
| request_id=f"{bench_id}_tts_text_v2", | |
| audio_path=Path(row.get("audio_path", "")), | |
| ), | |
| ReplyResult( | |
| text_urdu=row.get("assistant_reply_urdu", ""), | |
| provider="text_csv", | |
| model="s2s_live_text_v2", | |
| ), | |
| BridgeResult( | |
| text_devanagari=row.get("devanagari_tts_text", ""), | |
| provider="text_csv", | |
| model="s2s_live_text_v2", | |
| ), | |
| ) | |
| summary_rows.append( | |
| { | |
| "id": bench_id, | |
| "status": "ok", | |
| "assistant_reply_urdu": row.get("assistant_reply_urdu", ""), | |
| "devanagari_tts_text": row.get("devanagari_tts_text", ""), | |
| "tts_audio_path": str(result.audio_path), | |
| "duration_seconds": result.duration_seconds, | |
| "error": "", | |
| } | |
| ) | |
| except Exception as exc: # noqa: BLE001 - batch runner should report row failures. | |
| error = "".join(traceback.format_exception_only(type(exc), exc)).strip() | |
| print(f"[{bench_id}] ERROR: {error}", flush=True) | |
| summary_rows.append( | |
| { | |
| "id": bench_id, | |
| "status": "error", | |
| "assistant_reply_urdu": row.get("assistant_reply_urdu", ""), | |
| "devanagari_tts_text": row.get("devanagari_tts_text", ""), | |
| "tts_audio_path": str(wav_path), | |
| "error": error, | |
| } | |
| ) | |
| if fail_fast: | |
| break | |
| write_summary(summary_csv, summary_rows) | |
| return summary_rows | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--input-csv", type=Path, default=DEFAULT_INPUT_CSV) | |
| parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR) | |
| parser.add_argument("--summary-csv", type=Path, default=DEFAULT_SUMMARY_CSV) | |
| parser.add_argument("--ids", default="", help="Comma-separated IDs. Defaults to all rows.") | |
| parser.add_argument("--voice-prompt-audio-path", type=Path, default=DEFAULT_PRAXY_ANCHOR) | |
| parser.add_argument("--device", default="cuda") | |
| parser.add_argument("--chatterbox-t3-model", default="v3") | |
| parser.add_argument("--fail-fast", action="store_true") | |
| return parser.parse_args() | |
| def main() -> int: | |
| args = parse_args() | |
| summary_rows = run_tts_rows( | |
| input_csv=resolve_repo_path(args.input_csv), | |
| output_dir=resolve_repo_path(args.output_dir), | |
| summary_csv=resolve_repo_path(args.summary_csv), | |
| voice_prompt_audio_path=resolve_repo_path(args.voice_prompt_audio_path), | |
| selected_ids=parse_ids(args.ids), | |
| device=args.device, | |
| t3_model=args.chatterbox_t3_model, | |
| fail_fast=args.fail_fast, | |
| ) | |
| ok_count = sum(1 for row in summary_rows if row["status"] == "ok") | |
| print(f"summary={resolve_repo_path(args.summary_csv)}") | |
| print(f"ok={ok_count} total={len(summary_rows)}") | |
| return 0 if ok_count == len(summary_rows) else 1 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |