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