Spaces:
Running on Zero
Running on Zero
| #!/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()) | |