"""Hugging Face Pipeline for Gnani Prisma v2.5 (STT) and Gnani Timbre v2.0 (TTS). Routes ``"automatic-speech-recognition"`` and ``"text-to-speech"`` through Gnani's hosted inference endpoints. No local model weights are used. """ from __future__ import annotations import asyncio import os from pathlib import Path from typing import Any, Dict, List from transformers import Pipeline from configuration_gnani import GnaniVachanaConfig _CREDENTIAL_URL = "https://app.gnani.ai" _DOCS_URL = "https://docs.gnani.ai/api/introduction/introduction" SUPPORTED_STT_LANGUAGES = { "en-IN": "English (India)", "hi-IN": "Hindi", "ta-IN": "Tamil", "te-IN": "Telugu", "kn-IN": "Kannada", "ml-IN": "Malayalam", "mr-IN": "Marathi", "bn-IN": "Bengali", "gu-IN": "Gujarati", "pa-IN": "Punjabi", "en-IN,hi-IN": "English-Hindi (code-mixed)", } SUPPORTED_TTS_VOICES = { "Pranav": "Male — Bold, Trustworthy", "Kaveri": "Female — Confident, Bright", "Shubhra": "Female — Gentle, Expressive", "Deepak": "Male — Grounded, Conversational", } """See https://docs.gnani.ai/api/TTS/tts-sse#available-voices""" def _require_env(name: str) -> str: """Return the value of env var *name* or raise with a helpful message.""" value = os.environ.get(name, "").strip() if not value: raise RuntimeError( f"{name} not set. Get your key at {_CREDENTIAL_URL} " f"or email speechstack@gnani.ai" ) return value class GnaniVachanaPipeline(Pipeline): """Unified Hugging Face Pipeline for Gnani Prisma v2.5 (STT) and Gnani Timbre v2.0 (TTS). This pipeline delegates to Gnani's hosted speech API. It supports two tasks selected via ``pipeline(task, ...)``: * ``"automatic-speech-recognition"`` — Gnani Prisma v2.5 STT * ``"text-to-speech"`` — Gnani Timbre v2.0 TTS Streaming variants (WebSocket / SSE) are available by passing ``use_streaming=True`` or setting it in the model config. **No model weights are downloaded.** Authentication is via environment variables (see README for details). See Also -------- Full API docs: https://docs.gnani.ai/api/introduction/introduction """ def __init__(self, model: Any = None, **kwargs): if model is None: model = GnaniVachanaConfig() if isinstance(model, dict): model = GnaniVachanaConfig(**model) super().__init__(model=model, **kwargs) # ------------------------------------------------------------------ # _sanitize_parameters — split caller kwargs into pre/fwd/post dicts # ------------------------------------------------------------------ def _sanitize_parameters(self, **kwargs) -> tuple[dict, dict, dict]: """Split user-supplied kwargs into preprocess / forward / postprocess groups. Returns ------- tuple[dict, dict, dict] (preprocess_kwargs, forward_kwargs, postprocess_kwargs) """ forward_kwargs: Dict[str, Any] = {} postprocess_kwargs: Dict[str, Any] = {} if self.task == "automatic-speech-recognition": for key in ("language_code", "format", "use_streaming"): if key in kwargs: forward_kwargs[key] = kwargs[key] elif self.task == "text-to-speech": for key in ( "voice", "sample_rate", "container", "encoding", "use_streaming", "tts_mode", ): if key in kwargs: forward_kwargs[key] = kwargs[key] else: raise ValueError( f"Unsupported task '{self.task}'. " "Use 'automatic-speech-recognition' or 'text-to-speech'." ) return {}, forward_kwargs, postprocess_kwargs # ------------------------------------------------------------------ # preprocess # ------------------------------------------------------------------ def preprocess(self, inputs: Any, **kwargs) -> Dict[str, Any]: """Pass inputs through unchanged — the API handles encoding. Parameters ---------- inputs File path (str/Path) for STT, or text string for TTS. Returns ------- dict ``{"inputs": inputs}`` """ return {"inputs": inputs} # ------------------------------------------------------------------ # _forward # ------------------------------------------------------------------ def _forward(self, model_inputs: Dict[str, Any], **kwargs) -> Dict[str, Any]: """Dispatch to the correct Gnani client based on task and streaming flag. Parameters ---------- model_inputs : dict Must contain ``"inputs"`` — a file path (STT) or text (TTS). **kwargs Task-specific overrides (language_code, voice, etc.). Returns ------- dict Raw result dict forwarded to :meth:`postprocess`. """ raw_input = model_inputs["inputs"] use_streaming = kwargs.pop( "use_streaming", getattr(self.model, "use_streaming", False) ) if self.task == "automatic-speech-recognition": return self._forward_stt(raw_input, use_streaming, **kwargs) return self._forward_tts(raw_input, use_streaming, **kwargs) # --- STT helpers (Gnani Prisma v2.5) --------------------------------- def _forward_stt( self, audio_input: Any, use_streaming: bool, **kwargs ) -> Dict[str, Any]: """Run speech-to-text via REST or WebSocket.""" language_code = kwargs.get( "language_code", getattr(self.model, "default_language_code", "hi-IN"), ) if use_streaming: return self._forward_stt_stream(audio_input, language_code) return self._forward_stt_rest(audio_input, language_code) def _forward_stt_rest( self, audio_input: Any, language_code: str ) -> Dict[str, Any]: """Transcribe an audio file using the REST STT client.""" from gnani.stt import GnaniSTTClient api_key = _require_env("GNANI_API_KEY") client = GnaniSTTClient(api_key=api_key) if isinstance(audio_input, bytes): response = client.transcribe_bytes( audio_input, language_code=language_code ) else: response = client.transcribe( str(audio_input), language_code=language_code ) return {"transcript": response.get("transcript", "")} @staticmethod def _wav_to_pcm16k(wav_path: Path) -> bytes: """Read a WAV file and return raw PCM resampled to 16 kHz mono 16-bit.""" import audioop import wave with wave.open(str(wav_path), "rb") as wf: n_ch = wf.getnchannels() sw = wf.getsampwidth() fr = wf.getframerate() raw = wf.readframes(wf.getnframes()) if n_ch > 1: raw = audioop.tomono(raw, sw, 1.0, 0.0) if sw != 2: raw = audioop.lin2lin(raw, sw, 2) if fr != 16000: raw, _ = audioop.ratecv(raw, 2, 1, fr, 16000, None) return raw def _forward_stt_stream( self, audio_input: Any, language_code: str ) -> Dict[str, Any]: """Transcribe audio via the realtime WebSocket STT client. Converts WAV files to raw PCM (16 kHz, mono, 16-bit) before streaming. """ from gnani.stt import GnaniSTTStreamClient api_key = _require_env("GNANI_API_KEY") if isinstance(audio_input, (str, Path)): pcm_bytes = self._wav_to_pcm16k(Path(audio_input)) else: pcm_bytes = audio_input trailing_silence = b"\x00" * 64000 # 2s silence to trigger VAD endpoint async def _run() -> List[str]: client = GnaniSTTStreamClient( api_key=api_key, language_code=language_code, sample_rate=16000 ) await client.connect() stream_bytes = pcm_bytes + trailing_silence chunk_size = 1024 for i in range(0, len(stream_bytes), chunk_size): await client.send_audio(stream_bytes[i : i + chunk_size]) await asyncio.sleep(0.032) await asyncio.sleep(8.0) transcripts = await client.close() return [t.text for t in transcripts] texts = self._run_async(_run()) return {"transcript": " ".join(texts)} # --- TTS helpers (Gnani Timbre v2.0) --------------------------------- def _forward_tts( self, text: str, use_streaming: bool, **kwargs ) -> Dict[str, Any]: """Run Gnani Timbre v2.0 text-to-speech via REST, SSE, or WebSocket. The transport is selected via ``tts_mode``: * ``"rest"`` -- single synchronous HTTP call (default) * ``"sse"`` -- Server-Sent Events streaming (lower latency) * ``"realtime"`` -- WebSocket streaming (lowest latency) For backward compatibility, ``use_streaming=True`` without an explicit ``tts_mode`` maps to ``"realtime"``. """ tts_mode = kwargs.pop("tts_mode", None) if tts_mode is None: tts_mode = "realtime" if use_streaming else "rest" valid_modes = ("rest", "sse", "realtime") if tts_mode not in valid_modes: raise ValueError( f"Invalid tts_mode '{tts_mode}'. Choose from {valid_modes}." ) voice = kwargs.get( "voice", getattr(self.model, "default_voice", "Pranav") ) sample_rate = kwargs.get( "sample_rate", getattr(self.model, "default_sample_rate", 22050) ) container = kwargs.get( "container", getattr(self.model, "default_container", "wav") ) encoding = kwargs.get( "encoding", getattr(self.model, "default_encoding", "linear_pcm") ) tts_args = (text, voice, sample_rate, container, encoding) if tts_mode == "sse": return self._forward_tts_sse(*tts_args) if tts_mode == "realtime": return self._forward_tts_realtime(*tts_args) return self._forward_tts_rest(*tts_args) def _forward_tts_rest( self, text: str, voice: str, sample_rate: int, container: str, encoding: str, ) -> Dict[str, Any]: """Synthesise speech using the REST TTS client.""" from gnani.tts import AudioConfig, GnaniTTSClient api_key = _require_env("GNANI_API_KEY") client = GnaniTTSClient(api_key=api_key) audio_config = AudioConfig( sample_rate=sample_rate, encoding=encoding, container=container ) audio_bytes = client.synthesize( text, voice=voice, audio_config=audio_config ) return {"audio": audio_bytes, "sampling_rate": sample_rate} def _forward_tts_sse( self, text: str, voice: str, sample_rate: int, container: str, encoding: str, ) -> Dict[str, Any]: """Synthesise speech via SSE streaming (lower latency than REST).""" from gnani.tts import AudioConfig, GnaniTTSStreamClient api_key = _require_env("GNANI_API_KEY") client = GnaniTTSStreamClient(api_key=api_key) audio_config = AudioConfig( sample_rate=sample_rate, encoding=encoding, container=container ) audio_bytes = client.synthesize( text, voice=voice, audio_config=audio_config ) return {"audio": audio_bytes, "sampling_rate": sample_rate} def _forward_tts_realtime( self, text: str, voice: str, sample_rate: int, container: str, encoding: str, ) -> Dict[str, Any]: """Synthesise speech via the WebSocket TTS client.""" from gnani.tts import AudioConfig, GnaniTTSRealtimeClient api_key = _require_env("GNANI_API_KEY") audio_config = AudioConfig( sample_rate=sample_rate, encoding=encoding, container=container ) async def _run() -> bytes: client = GnaniTTSRealtimeClient(api_key=api_key) return await client.synthesize_and_collect( text, voice=voice, audio_config=audio_config ) audio_bytes = self._run_async(_run()) return {"audio": audio_bytes, "sampling_rate": sample_rate} # ------------------------------------------------------------------ # postprocess # ------------------------------------------------------------------ def postprocess( self, model_outputs: Dict[str, Any], **kwargs ) -> Dict[str, Any]: """Format the raw API response for the caller. Returns ------- dict For STT: ``{"text": str}`` For TTS: ``{"audio": bytes, "sampling_rate": int}`` """ if self.task == "automatic-speech-recognition": return {"text": model_outputs.get("transcript", "")} return { "audio": model_outputs.get("audio", b""), "sampling_rate": model_outputs.get("sampling_rate", 22050), } # ------------------------------------------------------------------ # helpers # ------------------------------------------------------------------ @staticmethod def _run_async(coro): """Run an async coroutine from sync code, handling nested event loops.""" try: loop = asyncio.get_running_loop() except RuntimeError: loop = None if loop and loop.is_running(): import concurrent.futures with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool: return pool.submit(asyncio.run, coro).result() return asyncio.run(coro)