156 lines
5.9 KiB
Python
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)
|