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