urdu-s2s-mvp / scripts /run_s2s_live_batch.py
sufinity's picture
Deploy Urdu S2S MVP
c759578 verified
Raw
History Blame Contribute Delete
6.97 kB
#!/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())