"""Text-to-speech — AI4Bharat Indic Parler-TTS. Parler is prompted with *two* texts: the words to say, and a natural-language description of how to say them (speaker, pace, room tone). The description is what selects a voice — there is no speaker-id argument. Latency shape: Parler is autoregressive, so a whole paragraph costs whole- paragraph time before the first sample exists. Two things fix that here: 1. `ParlerTTSStreamer` yields audio while generation continues, so playback starts after roughly the first `play_steps` frames rather than at the end. 2. The caller sends one *sentence* at a time. Short prompts reach their first chunk sooner, and a sentence boundary is a natural place for the assistant to be interrupted. """ from __future__ import annotations import logging import re import time from threading import Thread from typing import Iterator import numpy as np import torch from .config import settings logger = logging.getLogger(__name__) # Voices recommended on the model card, per language. VOICES = { "ta": "Jaya", "hi": "Rohit", "te": "Prakash", "kn": "Suresh", "ml": "Anjali", "mr": "Sanjay", "bn": "Arjun", "en": "Mary", } _DESCRIPTION = ( "{speaker} speaks in a warm, clear, professional tone at a natural pace. " "The recording is very high quality with no background noise." ) # Split on sentence enders including the Devanagari danda, keeping it simple — # this only needs to find safe places to cut, not parse language. _SENTENCE_RX = re.compile(r"(?<=[.!?।॥])\s+") def split_sentences(text: str, max_chars: int = 220) -> list[str]: """Break text into TTS-sized pieces at sentence boundaries where possible.""" out: list[str] = [] for part in _SENTENCE_RX.split(text.strip()): part = part.strip() if not part: continue while len(part) > max_chars: cut = part.rfind(" ", 0, max_chars) if cut <= 0: cut = max_chars out.append(part[:cut].strip()) part = part[cut:].strip() if part: out.append(part) return out class Synthesizer: def __init__(self) -> None: self._model = None self._tok = None self._desc_tok = None self.sample_rate = settings.sample_rate_out def load(self) -> None: from parler_tts import ParlerTTSForConditionalGeneration from transformers import AutoTokenizer t0 = time.perf_counter() self._model = ParlerTTSForConditionalGeneration.from_pretrained( settings.tts_model, torch_dtype=settings.dtype, ).to(settings.device).eval() self._tok = AutoTokenizer.from_pretrained(settings.tts_model) self._desc_tok = AutoTokenizer.from_pretrained(self._model.config.text_encoder._name_or_path) self.sample_rate = int(self._model.config.sampling_rate) logger.info( "TTS loaded (%s) in %.1fs @ %d Hz", settings.tts_model, time.perf_counter() - t0, self.sample_rate, ) def _describe(self, lang: str, voice: str | None) -> str: return _DESCRIPTION.format(speaker=voice or VOICES.get(lang, "Jaya")) @torch.inference_mode() def stream(self, text: str, lang: str = "ta", voice: str | None = None) -> Iterator[np.ndarray]: """Yield float32 mono chunks at `self.sample_rate` as they are generated.""" text = (text or "").strip() if not text: return desc = self._desc_tok(self._describe(lang, voice), return_tensors="pt").to(settings.device) prompt = self._tok(text, return_tensors="pt").to(settings.device) kwargs = dict( input_ids=desc.input_ids, attention_mask=desc.attention_mask, prompt_input_ids=prompt.input_ids, prompt_attention_mask=prompt.attention_mask, ) streamer = self._make_streamer() if streamer is None: yield self._generate_blocking(kwargs, text) return t0 = time.perf_counter() # generate() blocks, so it runs on its own thread and the streamer is # drained here as frames become available. thread = Thread(target=self._model.generate, kwargs={**kwargs, "streamer": streamer}, daemon=True) thread.start() first = True for chunk in streamer: if chunk is None or len(chunk) == 0: continue audio = chunk.astype(np.float32) if isinstance(chunk, np.ndarray) else chunk.cpu().numpy().astype(np.float32) if first: logger.info("TTS first chunk in %dms (%d chars)", int((time.perf_counter() - t0) * 1000), len(text)) first = False yield audio thread.join(timeout=1.0) def _make_streamer(self): try: from parler_tts import ParlerTTSStreamer except ImportError: logger.warning("ParlerTTSStreamer unavailable — falling back to blocking synthesis") return None # play_steps trades first-chunk latency against per-chunk overhead; # ~0.5 s of audio keeps playback continuous without stalling generation. frame_rate = getattr(self._model.audio_encoder.config, "frame_rate", 86) return ParlerTTSStreamer(self._model, device=settings.device, play_steps=int(frame_rate / 2)) @torch.inference_mode() def _generate_blocking(self, kwargs: dict, text: str) -> np.ndarray: t0 = time.perf_counter() gen = self._model.generate(**kwargs) audio = gen.cpu().numpy().squeeze().astype(np.float32) logger.info("TTS (blocking) %d chars in %dms", len(text), int((time.perf_counter() - t0) * 1000)) return audio def warmup(self) -> None: try: for _ in self.stream("வணக்கம்", "ta"): break logger.info("TTS warm") except Exception as e: # noqa: BLE001 logger.warning("TTS warmup skipped: %s", e)