WeLe Agentic AI: LangGraph multi-agent CRM assistant with voice
This commit is contained in:
@@ -0,0 +1,155 @@
|
||||
"""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)
|
||||
Reference in New Issue
Block a user