#!/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())