Files
Agentic-AI/voice-service/app/tts.py
T

156 lines
5.9 KiB
Python

"""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)