urdu-s2s-mvp / scripts /run_s2s_live.py
sufinity's picture
Deploy Urdu S2S MVP
c759578 verified
Raw
History Blame Contribute Delete
8.74 kB
#!/usr/bin/env python3
"""Run the reusable Urdu S2S service with live reply and speech-text providers."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import sys
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 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 ( # noqa: E402
ASRResult,
BridgeResult,
ReplyResult,
SpeechToSpeechRequest,
SpeechToSpeechResult,
TTSResult,
)
from urdu_s2s.tts_providers import ChatterboxPraxyTTSProvider # noqa: E402
from urdu_s2s.tracing import to_jsonable # noqa: E402
DEFAULT_PRAXY_ANCHOR = ROOT / "data/processed/voice_anchors/chatterbox_praxy_v1/bench_025.wav"
class TranscriptASRProvider:
"""Temporary ASR adapter for live S2S testing from a known transcript."""
def __init__(self, transcript: str) -> None:
self.transcript = transcript
def transcribe(self, request: SpeechToSpeechRequest) -> ASRResult:
return ASRResult(
text=self.transcript,
provider="provided_transcript",
model="manual_asr_transcript",
language="ur",
)
class PlaceholderTTSProvider:
"""Temporary TTS adapter until the live Chatterbox/Praxy provider is wired."""
def __init__(self, audio_path: Path) -> None:
self.audio_path = audio_path
def synthesize(
self,
request: SpeechToSpeechRequest,
reply: ReplyResult,
bridge: BridgeResult,
) -> TTSResult:
return TTSResult(
audio_path=self.audio_path,
provider="placeholder_tts",
model="not_synthesized_yet",
)
def build_text_live_pipeline(
*,
audio_path: Path,
request_id: str,
asr_transcript: str,
prompt_roman_urdu: str,
tts_audio_path: Path,
chat_client: OpenAICompatibleChatClient,
) -> tuple[SpeechToSpeechPipeline, SpeechToSpeechRequest]:
return build_live_pipeline(
audio_path=audio_path,
request_id=request_id,
asr_provider_name="transcript",
asr_transcript=asr_transcript,
prompt_roman_urdu=prompt_roman_urdu,
tts_provider_name="placeholder",
tts_audio_path=tts_audio_path,
chat_client=chat_client,
)
def build_live_pipeline(
*,
audio_path: Path,
request_id: str,
asr_provider_name: str,
asr_transcript: str,
prompt_roman_urdu: str,
tts_provider_name: str = "placeholder",
tts_audio_path: Path,
voice_prompt_audio_path: Path = DEFAULT_PRAXY_ANCHOR,
chat_client: OpenAICompatibleChatClient,
whisper_model_factory: Any | None = None,
whisper_model: str = "large-v3",
whisper_language: str = "ur",
whisper_device: str = "cpu",
whisper_compute_type: str = "int8",
chatterbox_model_loader: Any | None = None,
chatterbox_wav_writer: Any | None = None,
chatterbox_duration_reader: Any | None = None,
chatterbox_device: str = "cuda",
chatterbox_t3_model: str = "v3",
) -> tuple[SpeechToSpeechPipeline, SpeechToSpeechRequest]:
if asr_provider_name == "transcript":
if not asr_transcript.strip():
raise ValueError("--asr-transcript is required when --asr-provider transcript")
asr_provider = TranscriptASRProvider(asr_transcript)
elif asr_provider_name == "faster_whisper":
asr_provider = FasterWhisperASRProvider(
model_name=whisper_model,
language=whisper_language,
device=whisper_device,
compute_type=whisper_compute_type,
model_factory=whisper_model_factory,
)
else:
raise ValueError(f"Unknown ASR provider: {asr_provider_name}")
if tts_provider_name == "placeholder":
tts_provider = PlaceholderTTSProvider(tts_audio_path)
elif tts_provider_name == "chatterbox_praxy":
tts_provider = ChatterboxPraxyTTSProvider(
output_audio_path=tts_audio_path,
voice_prompt_audio_path=voice_prompt_audio_path,
device=chatterbox_device,
t3_model=chatterbox_t3_model,
model_loader=chatterbox_model_loader,
wav_writer=chatterbox_wav_writer,
duration_reader=chatterbox_duration_reader,
)
else:
raise ValueError(f"Unknown TTS provider: {tts_provider_name}")
pipeline = SpeechToSpeechPipeline(
asr_provider=asr_provider,
reply_provider=OpenAIReplyProvider(chat_client=chat_client),
bridge_provider=OpenAIBridgeProvider(chat_client=chat_client),
tts_provider=tts_provider,
)
request = SpeechToSpeechRequest(
request_id=request_id,
audio_path=audio_path,
metadata={"prompt_roman_urdu": prompt_roman_urdu},
)
return pipeline, request
def result_to_payload(result: SpeechToSpeechResult) -> dict[str, object]:
return {
"request_id": result.request.request_id,
"input_audio_path": str(result.request.audio_path),
"asr_transcript": result.asr.text,
"assistant_reply_urdu": result.reply.text_urdu,
"devanagari_tts_text": result.bridge.text_devanagari,
"tts_audio_path": str(result.tts.audio_path),
"trace": to_jsonable(result.trace),
}
def write_json_result(payload: dict[str, object], output_path: Path) -> None:
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(
json.dumps(payload, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--audio-path", required=True, type=Path)
parser.add_argument(
"--asr-provider",
choices=["transcript", "faster_whisper"],
default="transcript",
)
parser.add_argument(
"--asr-transcript",
default="",
help="Temporary transcript input until a live ASR provider is wired.",
)
parser.add_argument("--whisper-model", default="large-v3")
parser.add_argument("--whisper-language", default="ur")
parser.add_argument("--whisper-device", default="cpu")
parser.add_argument("--whisper-compute-type", default="int8")
parser.add_argument("--prompt-roman-urdu", default="")
parser.add_argument("--request-id", default="")
parser.add_argument("--model", default="")
parser.add_argument("--base-url", default="")
parser.add_argument("--output-json", type=Path, default=ROOT / "reports/evals/s2s_live_result.json")
parser.add_argument(
"--tts-provider",
choices=["placeholder", "chatterbox_praxy"],
default="placeholder",
)
parser.add_argument(
"--tts-audio-path",
type=Path,
default=ROOT / "reports/evals/s2s_live_placeholder_tts.wav",
help="Response WAV path for Chatterbox, or placeholder path in placeholder mode.",
)
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")
return parser.parse_args()
def main() -> int:
args = parse_args()
request_id = args.request_id or args.audio_path.stem
chat_client = OpenAICompatibleChatClient(
base_url=args.base_url or None,
model=args.model or None,
)
pipeline, request = build_live_pipeline(
audio_path=args.audio_path,
request_id=request_id,
asr_provider_name=args.asr_provider,
asr_transcript=args.asr_transcript,
prompt_roman_urdu=args.prompt_roman_urdu,
tts_provider_name=args.tts_provider,
tts_audio_path=args.tts_audio_path,
voice_prompt_audio_path=args.voice_prompt_audio_path,
chat_client=chat_client,
whisper_model=args.whisper_model,
whisper_language=args.whisper_language,
whisper_device=args.whisper_device,
whisper_compute_type=args.whisper_compute_type,
chatterbox_device=args.chatterbox_device,
chatterbox_t3_model=args.chatterbox_t3_model,
)
payload = result_to_payload(pipeline.run(request))
write_json_result(payload, args.output_json)
print(json.dumps(payload, ensure_ascii=False, indent=2))
print(f"wrote={args.output_json}")
return 0
if __name__ == "__main__":
raise SystemExit(main())