#!/usr/bin/env python3 """Run a small live Urdu S2S batch from a benchmark manifest.""" from __future__ import annotations import argparse import csv import json from pathlib import Path import sys import traceback 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 urdu_s2s.asr_providers import FasterWhisperASRProvider # noqa: E402 from urdu_s2s.live_providers import ( # noqa: E402 OpenAIBridgeProvider, OpenAICompatibleChatClient, OpenAIReplyProvider, ) from urdu_s2s.pipeline import SpeechToSpeechPipeline # noqa: E402 from urdu_s2s.schemas import SpeechToSpeechRequest # noqa: E402 from urdu_s2s.tts_providers import ChatterboxPraxyTTSProvider # noqa: E402 from run_s2s_live import DEFAULT_PRAXY_ANCHOR, result_to_payload, write_json_result # noqa: E402 DEFAULT_MANIFEST = ROOT / "artifacts/live_s2s_smoke10_manifest.csv" def read_manifest(path: Path) -> list[dict[str, str]]: with path.open(newline="", encoding="utf-8") as handle: return list(csv.DictReader(handle)) 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 resolve_repo_path(path: Path) -> Path: return path if path.is_absolute() else ROOT / path def write_summary_csv(path: Path, rows: list[dict[str, object]]) -> None: path.parent.mkdir(parents=True, exist_ok=True) fieldnames = [ "id", "status", "audio_path", "prompt_roman_urdu", "asr_transcript", "assistant_reply_urdu", "devanagari_tts_text", "tts_audio_path", "json_path", "error", ] 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 parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--manifest", type=Path, default=DEFAULT_MANIFEST) parser.add_argument("--ids", default="", help="Comma-separated IDs. Defaults to every row.") parser.add_argument("--output-dir", type=Path, default=ROOT / "reports/evals/s2s_live_smoke10") parser.add_argument("--summary-csv", type=Path, default=ROOT / "reports/evals/s2s_live_smoke10_summary.csv") parser.add_argument("--model", default="") parser.add_argument("--base-url", default="") parser.add_argument("--whisper-model", default="large-v3") parser.add_argument("--whisper-language", default="ur") parser.add_argument("--whisper-device", default="cuda") parser.add_argument("--whisper-compute-type", default="float16") parser.add_argument("--voice-prompt-audio-path", type=Path, default=DEFAULT_PRAXY_ANCHOR) parser.add_argument("--chatterbox-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() manifest_path = resolve_repo_path(args.manifest) output_dir = resolve_repo_path(args.output_dir) voice_prompt_audio_path = resolve_repo_path(args.voice_prompt_audio_path) selected_ids = parse_ids(args.ids) rows = read_manifest(manifest_path) if selected_ids is not None: rows = [row for row in rows if row.get("id") in selected_ids] if not rows: raise ValueError(f"No manifest rows selected from {manifest_path}") chat_client = OpenAICompatibleChatClient( base_url=args.base_url or None, model=args.model or None, ) asr_provider = FasterWhisperASRProvider( model_name=args.whisper_model, language=args.whisper_language, device=args.whisper_device, compute_type=args.whisper_compute_type, ) tts_provider = ChatterboxPraxyTTSProvider( output_audio_path=output_dir / "placeholder.wav", voice_prompt_audio_path=voice_prompt_audio_path, device=args.chatterbox_device, t3_model=args.chatterbox_t3_model, ) pipeline = SpeechToSpeechPipeline( asr_provider=asr_provider, reply_provider=OpenAIReplyProvider(chat_client=chat_client), bridge_provider=OpenAIBridgeProvider(chat_client=chat_client), tts_provider=tts_provider, ) summary_rows: list[dict[str, object]] = [] output_dir.mkdir(parents=True, exist_ok=True) for index, row in enumerate(rows, start=1): bench_id = row["id"] audio_path = resolve_repo_path(Path(row["audio_path"])) wav_path = output_dir / f"{bench_id}_praxy.wav" json_path = output_dir / f"{bench_id}.json" print(f"[{index}/{len(rows)}] {bench_id} -> {wav_path}", flush=True) try: tts_provider.output_audio_path = wav_path request = SpeechToSpeechRequest( request_id=f"{bench_id}_live_v2", audio_path=audio_path, metadata={"prompt_roman_urdu": row.get("prompt_roman_urdu", "")}, ) payload = result_to_payload(pipeline.run(request)) write_json_result(payload, json_path) summary_rows.append( { "id": bench_id, "status": "ok", "audio_path": str(audio_path), "prompt_roman_urdu": row.get("prompt_roman_urdu", ""), "asr_transcript": payload["asr_transcript"], "assistant_reply_urdu": payload["assistant_reply_urdu"], "devanagari_tts_text": payload["devanagari_tts_text"], "tts_audio_path": payload["tts_audio_path"], "json_path": str(json_path), "error": "", } ) except Exception as exc: # noqa: BLE001 - batch runner should report per-item 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", "audio_path": str(audio_path), "prompt_roman_urdu": row.get("prompt_roman_urdu", ""), "error": error, } ) if args.fail_fast: break write_summary_csv(resolve_repo_path(args.summary_csv), summary_rows) print(f"summary={resolve_repo_path(args.summary_csv)}") print(f"ok={sum(1 for row in summary_rows if row['status'] == 'ok')} total={len(summary_rows)}") return 0 if all(row["status"] == "ok" for row in summary_rows) else 1 if __name__ == "__main__": raise SystemExit(main())