WeLe Agentic AI: LangGraph multi-agent CRM assistant with voice
This commit is contained in:
@@ -0,0 +1,82 @@
|
||||
"""Voice service configuration.
|
||||
|
||||
Deliberately small: this process does one job — turn audio into text and text
|
||||
into audio. Everything about *what to say* lives in the Node agentic service.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def _int(name: str, default: int) -> int:
|
||||
try:
|
||||
return int(os.environ.get(name, default))
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
|
||||
def _flag(name: str, default: bool) -> bool:
|
||||
return os.environ.get(name, str(default)).lower() in {"1", "true", "yes"}
|
||||
|
||||
|
||||
class Settings:
|
||||
host: str = os.environ.get("VOICE_HOST", "127.0.0.1")
|
||||
port: int = _int("VOICE_PORT", 4100)
|
||||
|
||||
# ── Models ───────────────────────────────────────────────────────────────
|
||||
# IndicConformer is a hybrid CTC + RNNT model. CTC decoding is used because
|
||||
# it is a single forward pass — RNNT is more accurate but decodes
|
||||
# autoregressively, and in a voice loop the latency costs more than the
|
||||
# accuracy buys.
|
||||
stt_model: str = os.environ.get("STT_MODEL", "ai4bharat/indic-conformer-600m-multilingual")
|
||||
stt_decoding: str = os.environ.get("STT_DECODING", "ctc") # ctc | rnnt
|
||||
|
||||
# English is not one of IndicConformer's 22 codes, and it cannot identify
|
||||
# languages. Whisper covers both — multilingual, not the .en checkpoint,
|
||||
# because language ID is what makes "auto" work.
|
||||
english_model: str = os.environ.get("ENGLISH_MODEL", "openai/whisper-small")
|
||||
preload_english: bool = _flag("PRELOAD_ENGLISH", True)
|
||||
|
||||
# Used when auto-detection is not confident enough to overrule the user.
|
||||
default_lang: str = os.environ.get("VOICE_DEFAULT_LANG", "ta")
|
||||
|
||||
tts_model: str = os.environ.get("TTS_MODEL", "ai4bharat/indic-parler-tts")
|
||||
|
||||
# ── Device / precision ───────────────────────────────────────────────────
|
||||
# float16 on CUDA: both models together are ~3 GB in half precision, which
|
||||
# fits the 6 GB card with room for activations. float32 would not.
|
||||
device: str = os.environ.get("VOICE_DEVICE", "cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
@property
|
||||
def dtype(self) -> torch.dtype:
|
||||
return torch.float16 if self.device == "cuda" else torch.float32
|
||||
|
||||
# ── Audio ────────────────────────────────────────────────────────────────
|
||||
sample_rate_in: int = 16000 # what the browser worklet sends
|
||||
# Indic Parler-TTS emits 44.1 kHz — verified from model.config.sampling_rate,
|
||||
# not the 24 kHz the upstream Parler-TTS Mini uses. This is only a fallback;
|
||||
# the real rate is read from the loaded model and sent to the browser, which
|
||||
# configures its playback worklet from it.
|
||||
sample_rate_out: int = _int("TTS_SAMPLE_RATE", 44100)
|
||||
|
||||
# ── Endpointing (Silero VAD) ─────────────────────────────────────────────
|
||||
# Silero operates on fixed 512-sample frames at 16 kHz (32 ms).
|
||||
vad_frame: int = 512
|
||||
vad_threshold: float = float(os.environ.get("VAD_THRESHOLD", "0.5"))
|
||||
# How much trailing silence ends a turn. Too short truncates people who
|
||||
# pause mid-sentence; too long makes the assistant feel sluggish.
|
||||
vad_silence_ms: int = _int("VAD_SILENCE_MS", 700)
|
||||
# Ignore blips so a cough or a door does not open a turn.
|
||||
vad_min_speech_ms: int = _int("VAD_MIN_SPEECH_MS", 250)
|
||||
# Audio kept from *before* detected speech, so word onsets are not clipped.
|
||||
vad_prefix_ms: int = _int("VAD_PREFIX_MS", 300)
|
||||
vad_max_utterance_ms: int = _int("VAD_MAX_UTTERANCE_MS", 30000)
|
||||
|
||||
# Warm the models at startup rather than on the first user turn — a cold
|
||||
# CUDA graph on the first utterance costs several seconds.
|
||||
warmup: bool = _flag("VOICE_WARMUP", True)
|
||||
|
||||
|
||||
settings = Settings()
|
||||
@@ -0,0 +1,233 @@
|
||||
"""Voice service — STT and TTS over one WebSocket.
|
||||
|
||||
This process holds the GPU models and nothing else. It has no idea what the
|
||||
CRM is: it receives audio and returns text, receives text and returns audio.
|
||||
All orchestration, auth and business logic stay in the Node service, so voice
|
||||
is just another channel into the same agent rather than a parallel system.
|
||||
|
||||
Protocol (ws /ws/voice), JSON control + binary audio:
|
||||
|
||||
client → server
|
||||
binary 16 kHz mono PCM16 mic frames
|
||||
{"type":"config","lang":"ta"} set the session language
|
||||
{"type":"speak","text":"…","id":"…"} synthesise
|
||||
{"type":"cancel"} stop speaking now (barge-in)
|
||||
{"type":"reset"} clear the endpointer
|
||||
|
||||
server → client
|
||||
{"type":"ready", …}
|
||||
{"type":"speech_start"} VAD opened a turn → caller ducks TTS
|
||||
{"type":"transcript","text":…} a finished utterance
|
||||
{"type":"audio_start","id":…,"sample_rate":24000}
|
||||
binary float32 mono TTS chunks
|
||||
{"type":"audio_end","id":…}
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
||||
|
||||
from .config import settings
|
||||
from .stt import SUPPORTED, Transcriber
|
||||
from .tts import Synthesizer, split_sentences
|
||||
from .vad import Endpointer
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)-5s %(message)s", datefmt="%H:%M:%S")
|
||||
logger = logging.getLogger("voice")
|
||||
|
||||
app = FastAPI(title="WeLe Voice Service")
|
||||
|
||||
stt = Transcriber()
|
||||
tts = Synthesizer()
|
||||
|
||||
|
||||
@app.on_event("startup")
|
||||
async def _startup() -> None:
|
||||
t0 = time.perf_counter()
|
||||
logger.info("loading models on %s…", settings.device)
|
||||
stt.load()
|
||||
tts.load()
|
||||
if settings.warmup:
|
||||
stt.warmup()
|
||||
tts.warmup()
|
||||
logger.info("voice service ready in %.1fs", time.perf_counter() - t0)
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
async def health() -> dict:
|
||||
import torch
|
||||
|
||||
return {
|
||||
"ok": True,
|
||||
"device": settings.device,
|
||||
"stt_model": settings.stt_model,
|
||||
"tts_model": settings.tts_model,
|
||||
"tts_sample_rate": tts.sample_rate,
|
||||
"languages": SUPPORTED,
|
||||
"vram_gb": round(torch.cuda.memory_reserved() / 1e9, 2) if settings.device == "cuda" else None,
|
||||
}
|
||||
|
||||
|
||||
@app.get("/languages")
|
||||
async def languages() -> dict:
|
||||
return {"languages": SUPPORTED}
|
||||
|
||||
|
||||
class Session:
|
||||
"""One browser connection. Owns its endpointer and its speaking state."""
|
||||
|
||||
def __init__(self, ws: WebSocket) -> None:
|
||||
self.ws = ws
|
||||
# "auto" detects per utterance; `prefer` breaks ties when the detector
|
||||
# is unsure, which is common on short code-mixed ("Tanglish") speech.
|
||||
self.lang = "auto"
|
||||
self.prefer = settings.default_lang
|
||||
self.endpointer = Endpointer()
|
||||
self.endpointer.load()
|
||||
self._speak_task: asyncio.Task | None = None
|
||||
self._cancel = asyncio.Event()
|
||||
|
||||
async def send(self, payload: dict) -> None:
|
||||
await self.ws.send_text(json.dumps(payload, ensure_ascii=False))
|
||||
|
||||
# ── microphone ───────────────────────────────────────────────────────────
|
||||
async def on_audio(self, raw: bytes) -> None:
|
||||
pcm = np.frombuffer(raw, dtype=np.int16).astype(np.float32) / 32768.0
|
||||
loop = asyncio.get_running_loop()
|
||||
# VAD is a small torch model but still blocking; keep the event loop free.
|
||||
utterances, started = await loop.run_in_executor(None, self.endpointer.push, pcm)
|
||||
|
||||
if started:
|
||||
# Barge-in: the user talking wins immediately.
|
||||
await self.stop_speaking()
|
||||
await self.send({"type": "speech_start"})
|
||||
|
||||
for utt in utterances:
|
||||
request_lang = f"auto:{self.prefer}" if self.lang == "auto" else self.lang
|
||||
result = await loop.run_in_executor(None, stt.transcribe, utt.audio, request_lang)
|
||||
if result["text"]:
|
||||
await self.send({"type": "transcript", **result, "truncated": utt.truncated})
|
||||
else:
|
||||
await self.send({"type": "transcript_empty", "reason": result.get("note", "no speech")})
|
||||
|
||||
# ── speaking ─────────────────────────────────────────────────────────────
|
||||
async def speak(self, text: str, msg_id: str, lang: str | None = None, voice: str | None = None) -> None:
|
||||
await self.stop_speaking()
|
||||
self._cancel.clear()
|
||||
self._speak_task = asyncio.create_task(self._speak(text, msg_id, lang or self.lang, voice))
|
||||
|
||||
async def _speak(self, text: str, msg_id: str, lang: str, voice: str | None) -> None:
|
||||
loop = asyncio.get_running_loop()
|
||||
try:
|
||||
await self.send({"type": "audio_start", "id": msg_id, "sample_rate": tts.sample_rate})
|
||||
|
||||
# Sentence at a time: shorter prompts reach first audio sooner, and
|
||||
# a boundary is a clean place to stop when interrupted.
|
||||
for sentence in split_sentences(text):
|
||||
if self._cancel.is_set():
|
||||
break
|
||||
queue: asyncio.Queue = asyncio.Queue(maxsize=32)
|
||||
|
||||
def produce() -> None:
|
||||
try:
|
||||
for chunk in tts.stream(sentence, lang, voice):
|
||||
if self._cancel.is_set():
|
||||
break
|
||||
asyncio.run_coroutine_threadsafe(queue.put(chunk), loop).result()
|
||||
finally:
|
||||
asyncio.run_coroutine_threadsafe(queue.put(None), loop).result()
|
||||
|
||||
loop.run_in_executor(None, produce)
|
||||
while True:
|
||||
chunk = await queue.get()
|
||||
if chunk is None:
|
||||
break
|
||||
if self._cancel.is_set():
|
||||
continue # drain, don't send
|
||||
await self.ws.send_bytes(np.asarray(chunk, dtype=np.float32).tobytes())
|
||||
|
||||
await self.send({"type": "audio_end", "id": msg_id, "cancelled": self._cancel.is_set()})
|
||||
except WebSocketDisconnect:
|
||||
pass
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.exception("synthesis failed")
|
||||
try:
|
||||
await self.send({"type": "error", "where": "tts", "message": str(e)[:200]})
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
async def stop_speaking(self) -> None:
|
||||
if self._speak_task and not self._speak_task.done():
|
||||
self._cancel.set()
|
||||
try:
|
||||
await asyncio.wait_for(self._speak_task, timeout=2.0)
|
||||
except (asyncio.TimeoutError, asyncio.CancelledError):
|
||||
self._speak_task.cancel()
|
||||
self._speak_task = None
|
||||
|
||||
|
||||
@app.websocket("/ws/voice")
|
||||
async def voice(ws: WebSocket) -> None:
|
||||
await ws.accept()
|
||||
session = Session(ws)
|
||||
await session.send({
|
||||
"type": "ready",
|
||||
"sample_rate_in": settings.sample_rate_in,
|
||||
"sample_rate_out": tts.sample_rate,
|
||||
"languages": SUPPORTED,
|
||||
})
|
||||
logger.info("voice session opened")
|
||||
|
||||
try:
|
||||
while True:
|
||||
msg = await ws.receive()
|
||||
if msg["type"] == "websocket.disconnect":
|
||||
break
|
||||
|
||||
if (raw := msg.get("bytes")) is not None:
|
||||
await session.on_audio(raw)
|
||||
continue
|
||||
|
||||
if (text := msg.get("text")) is None:
|
||||
continue
|
||||
try:
|
||||
data = json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
kind = data.get("type")
|
||||
if kind == "config":
|
||||
session.lang = data.get("lang", session.lang)
|
||||
if session.lang != "auto":
|
||||
session.prefer = session.lang
|
||||
elif data.get("prefer"):
|
||||
session.prefer = data["prefer"]
|
||||
session.endpointer.reset()
|
||||
await session.send({"type": "config_ok", "lang": session.lang, "prefer": session.prefer})
|
||||
elif kind == "speak":
|
||||
await session.speak(data.get("text", ""), data.get("id", ""), data.get("lang"), data.get("voice"))
|
||||
elif kind == "cancel":
|
||||
await session.stop_speaking()
|
||||
await session.send({"type": "cancelled"})
|
||||
elif kind == "reset":
|
||||
session.endpointer.reset()
|
||||
except WebSocketDisconnect:
|
||||
pass
|
||||
finally:
|
||||
await session.stop_speaking()
|
||||
logger.info("voice session closed")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
import uvicorn
|
||||
|
||||
uvicorn.run(app, host=settings.host, port=settings.port, log_level="info", ws_max_size=16 * 1024 * 1024)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,200 @@
|
||||
"""Speech-to-text — AI4Bharat IndicConformer, with an English path and auto routing.
|
||||
|
||||
Why two models rather than one:
|
||||
|
||||
* **IndicConformer** decodes 22 Indian languages, and decodes *as* the language
|
||||
you name — it does not detect. Handing it "ta" for English speech produces
|
||||
Tamil-script nonsense. English is not one of its codes at all.
|
||||
* **Whisper (multilingual)** covers English well and, usefully, can identify the
|
||||
spoken language in a single decoder step.
|
||||
|
||||
So the default mode is `auto`: Whisper identifies the language from the audio,
|
||||
English is transcribed by Whisper directly (the mel is already computed, so
|
||||
this costs nothing extra), and anything Indic is routed to IndicConformer,
|
||||
which is far stronger on those languages than Whisper is.
|
||||
|
||||
Code-mixed speech ("Tanglish") is the awkward case: language ID can land either
|
||||
side of the fence on a short, mixed utterance. When Whisper is not confident,
|
||||
the caller's preferred language wins rather than a coin toss — a Tamil-speaking
|
||||
office gets Tamil, and the occasional English sentence still routes correctly
|
||||
when it is clearly English.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from .config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# The codes IndicConformer accepts.
|
||||
INDIC_LANGS = {
|
||||
"as", "bn", "brx", "doi", "gu", "hi", "kn", "kok", "ks", "mai", "ml",
|
||||
"mni", "mr", "ne", "or", "pa", "sa", "sat", "sd", "ta", "te", "ur",
|
||||
}
|
||||
|
||||
# Offered by the UI. "auto" first: most WeLe agents switch language mid-shift.
|
||||
SUPPORTED = [
|
||||
{"code": "auto", "label": "Auto-detect", "native": "Auto"},
|
||||
{"code": "ta", "label": "Tamil", "native": "தமிழ்"},
|
||||
{"code": "en", "label": "English", "native": "English"},
|
||||
{"code": "hi", "label": "Hindi", "native": "हिन्दी"},
|
||||
{"code": "te", "label": "Telugu", "native": "తెలుగు"},
|
||||
{"code": "kn", "label": "Kannada", "native": "ಕನ್ನಡ"},
|
||||
{"code": "ml", "label": "Malayalam", "native": "മലയാളം"},
|
||||
{"code": "mr", "label": "Marathi", "native": "मराठी"},
|
||||
{"code": "bn", "label": "Bengali", "native": "বাংলা"},
|
||||
]
|
||||
|
||||
# Below this, trust the user's stated preference over the detector.
|
||||
DETECT_CONFIDENCE = 0.60
|
||||
|
||||
|
||||
class Transcriber:
|
||||
def __init__(self) -> None:
|
||||
self._indic = None
|
||||
self._whisper = None
|
||||
self._whisper_proc = None
|
||||
|
||||
# ── loading ──────────────────────────────────────────────────────────────
|
||||
def load(self) -> None:
|
||||
from transformers import AutoModel
|
||||
|
||||
t0 = time.perf_counter()
|
||||
# float32: the checkpoint ships custom remote code that assumes fp32.
|
||||
# ~2.4 GB at 600M, which still leaves room for Whisper and the TTS model.
|
||||
self._indic = AutoModel.from_pretrained(settings.stt_model, trust_remote_code=True)
|
||||
self._indic.to(settings.device).eval()
|
||||
logger.info("STT loaded (%s) in %.1fs", settings.stt_model, time.perf_counter() - t0)
|
||||
|
||||
if settings.preload_english:
|
||||
self._load_whisper()
|
||||
|
||||
def _load_whisper(self) -> None:
|
||||
"""English + language ID. Multilingual on purpose — the .en checkpoint
|
||||
cannot identify languages, which is the whole point of auto mode."""
|
||||
if self._whisper is not None:
|
||||
return
|
||||
from transformers import WhisperForConditionalGeneration, WhisperProcessor
|
||||
|
||||
t0 = time.perf_counter()
|
||||
self._whisper_proc = WhisperProcessor.from_pretrained(settings.english_model)
|
||||
self._whisper = WhisperForConditionalGeneration.from_pretrained(
|
||||
settings.english_model, torch_dtype=settings.dtype,
|
||||
).to(settings.device).eval()
|
||||
logger.info("English/ID model loaded (%s) in %.1fs", settings.english_model, time.perf_counter() - t0)
|
||||
|
||||
# ── inference ────────────────────────────────────────────────────────────
|
||||
@torch.inference_mode()
|
||||
def transcribe(self, audio: np.ndarray, lang: str = "auto") -> dict:
|
||||
"""audio: float32 mono @16 kHz in [-1, 1].
|
||||
|
||||
`lang` may be an explicit code, or "auto" / "auto:ta" to detect with a
|
||||
fallback preference.
|
||||
"""
|
||||
t0 = time.perf_counter()
|
||||
if audio.size < settings.sample_rate_in // 5: # under 200 ms
|
||||
return {"text": "", "lang": lang, "ms": 0, "note": "too short"}
|
||||
|
||||
detected = None
|
||||
confidence = None
|
||||
|
||||
if lang.startswith("auto"):
|
||||
prefer = lang.split(":", 1)[1] if ":" in lang else settings.default_lang
|
||||
feats = self._features(audio)
|
||||
detected, confidence = self._detect(feats)
|
||||
|
||||
if confidence is not None and confidence < DETECT_CONFIDENCE:
|
||||
logger.info("language ID low confidence (%s @ %.2f) — using preferred %s",
|
||||
detected, confidence, prefer)
|
||||
use = prefer
|
||||
elif detected == "en" or detected in INDIC_LANGS:
|
||||
use = detected
|
||||
else:
|
||||
# Whisper reported something we cannot decode (e.g. "nn" on
|
||||
# noise). Fall back rather than fail.
|
||||
use = prefer
|
||||
|
||||
text = self._english(audio, feats=feats) if use == "en" else self._indic_decode(audio, use)
|
||||
else:
|
||||
use = lang
|
||||
text = self._english(audio) if lang == "en" else self._indic_decode(audio, lang)
|
||||
|
||||
ms = int((time.perf_counter() - t0) * 1000)
|
||||
audio_ms = int(1000 * audio.size / settings.sample_rate_in)
|
||||
logger.info("STT %s%s: %dms audio → %dms → %r",
|
||||
use, f" (detected {detected} {confidence:.2f})" if confidence is not None else "",
|
||||
audio_ms, ms, text[:80])
|
||||
|
||||
return {
|
||||
"text": text.strip(),
|
||||
"lang": use,
|
||||
"detected": detected,
|
||||
"confidence": round(confidence, 3) if confidence is not None else None,
|
||||
"ms": ms,
|
||||
"audio_ms": audio_ms,
|
||||
}
|
||||
|
||||
# ── internals ────────────────────────────────────────────────────────────
|
||||
def _features(self, audio: np.ndarray):
|
||||
self._load_whisper()
|
||||
return self._whisper_proc(
|
||||
audio, sampling_rate=settings.sample_rate_in, return_tensors="pt",
|
||||
).input_features.to(settings.device, settings.dtype)
|
||||
|
||||
def _detect(self, feats) -> tuple[str | None, float | None]:
|
||||
"""One decoder step: read the language-token distribution."""
|
||||
try:
|
||||
tok = self._whisper_proc.tokenizer
|
||||
sot = tok.convert_tokens_to_ids("<|startoftranscript|>")
|
||||
start = torch.tensor([[sot]], device=settings.device)
|
||||
logits = self._whisper(feats, decoder_input_ids=start).logits[:, -1]
|
||||
|
||||
lang_ids, codes = [], []
|
||||
for code in {*INDIC_LANGS, "en"}:
|
||||
tid = tok.convert_tokens_to_ids(f"<|{code}|>")
|
||||
# Unknown languages map to the unk id; skip those.
|
||||
if tid is not None and tid != tok.unk_token_id:
|
||||
lang_ids.append(tid)
|
||||
codes.append(code)
|
||||
if not lang_ids:
|
||||
return None, None
|
||||
|
||||
probs = torch.softmax(logits[0, lang_ids].float(), dim=-1)
|
||||
best = int(probs.argmax())
|
||||
return codes[best], float(probs[best])
|
||||
except Exception as e: # noqa: BLE001 — detection must never break a turn
|
||||
logger.warning("language ID failed (%s) — falling back to preference", e)
|
||||
return None, None
|
||||
|
||||
def _indic_decode(self, audio: np.ndarray, lang: str) -> str:
|
||||
if lang not in INDIC_LANGS:
|
||||
logger.warning("unsupported STT language %r — using %s", lang, settings.default_lang)
|
||||
lang = settings.default_lang
|
||||
wav = torch.from_numpy(audio).unsqueeze(0).to(settings.device) # (1, N)
|
||||
out = self._indic(wav, lang, settings.stt_decoding)
|
||||
if isinstance(out, (list, tuple)):
|
||||
return str(out[0]) if out else ""
|
||||
return str(out)
|
||||
|
||||
def _english(self, audio: np.ndarray, feats=None) -> str:
|
||||
self._load_whisper()
|
||||
if feats is None:
|
||||
feats = self._features(audio)
|
||||
ids = self._whisper.generate(feats, language="en", task="transcribe", max_new_tokens=180)
|
||||
return self._whisper_proc.batch_decode(ids, skip_special_tokens=True)[0]
|
||||
|
||||
def warmup(self) -> None:
|
||||
"""Silent pass so the first real utterance isn't paying for CUDA init."""
|
||||
try:
|
||||
silence = np.zeros(settings.sample_rate_in, dtype=np.float32)
|
||||
self.transcribe(silence, settings.default_lang)
|
||||
if settings.preload_english:
|
||||
self.transcribe(silence, "en")
|
||||
logger.info("STT warm")
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("STT warmup skipped: %s", e)
|
||||
@@ -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)
|
||||
@@ -0,0 +1,127 @@
|
||||
"""Endpointing with Silero VAD.
|
||||
|
||||
Turn boundaries are decided here rather than in the browser for two reasons:
|
||||
the same decision then applies to every future channel (a phone bridge has no
|
||||
AudioWorklet), and barge-in needs the server to know someone started talking
|
||||
while the assistant was still speaking.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections import deque
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from .config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Utterance:
|
||||
audio: np.ndarray # float32 mono @16k, in [-1, 1]
|
||||
duration_ms: int
|
||||
truncated: bool = False # hit the max-length guard rather than silence
|
||||
|
||||
|
||||
@dataclass
|
||||
class VADState:
|
||||
speaking: bool = False
|
||||
speech_ms: int = 0
|
||||
silence_ms: int = 0
|
||||
buffer: list[np.ndarray] = field(default_factory=list)
|
||||
|
||||
|
||||
class Endpointer:
|
||||
"""Streaming VAD that emits one Utterance per detected turn.
|
||||
|
||||
Silero wants exactly 512 samples at 16 kHz, but the browser sends 40 ms
|
||||
(640-sample) chunks. Rather than force the client to match, incoming audio
|
||||
is accumulated and drained in exact frames.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._model = None
|
||||
self._pending = np.zeros(0, dtype=np.float32)
|
||||
self.state = VADState()
|
||||
# Pre-roll: speech is only *detected* a frame or two in, so without a
|
||||
# prefix the first phoneme is already gone by the time we start saving.
|
||||
prefix_frames = max(1, (settings.vad_prefix_ms * settings.sample_rate_in) // (1000 * settings.vad_frame))
|
||||
self._prefix: deque[np.ndarray] = deque(maxlen=prefix_frames)
|
||||
|
||||
def load(self) -> None:
|
||||
from silero_vad import load_silero_vad
|
||||
|
||||
self._model = load_silero_vad()
|
||||
logger.info("silero VAD loaded")
|
||||
|
||||
def reset(self) -> None:
|
||||
self.state = VADState()
|
||||
self._pending = np.zeros(0, dtype=np.float32)
|
||||
self._prefix.clear()
|
||||
if self._model is not None:
|
||||
self._model.reset_states()
|
||||
|
||||
@property
|
||||
def is_speaking(self) -> bool:
|
||||
return self.state.speaking
|
||||
|
||||
def push(self, pcm: np.ndarray) -> tuple[list[Utterance], bool]:
|
||||
"""Feed float32 audio.
|
||||
|
||||
Returns (completed utterances, speech_started_this_call). The second
|
||||
value drives barge-in: the caller cuts TTS playback the moment it flips.
|
||||
"""
|
||||
assert self._model is not None, "call load() first"
|
||||
|
||||
self._pending = np.concatenate([self._pending, pcm]) if self._pending.size else pcm
|
||||
frame = settings.vad_frame
|
||||
frame_ms = int(1000 * frame / settings.sample_rate_in)
|
||||
|
||||
done: list[Utterance] = []
|
||||
started = False
|
||||
|
||||
while self._pending.size >= frame:
|
||||
chunk = self._pending[:frame]
|
||||
self._pending = self._pending[frame:]
|
||||
|
||||
with torch.no_grad():
|
||||
prob = float(self._model(torch.from_numpy(chunk), settings.sample_rate_in).item())
|
||||
|
||||
voiced = prob >= settings.vad_threshold
|
||||
st = self.state
|
||||
|
||||
if not st.speaking:
|
||||
self._prefix.append(chunk)
|
||||
if voiced:
|
||||
st.speech_ms += frame_ms
|
||||
if st.speech_ms >= settings.vad_min_speech_ms:
|
||||
# Commit: open the turn with the pre-roll included.
|
||||
st.speaking = True
|
||||
st.silence_ms = 0
|
||||
st.buffer = list(self._prefix)
|
||||
self._prefix.clear()
|
||||
started = True
|
||||
else:
|
||||
st.speech_ms = 0
|
||||
continue
|
||||
|
||||
# Speaking.
|
||||
st.buffer.append(chunk)
|
||||
if voiced:
|
||||
st.silence_ms = 0
|
||||
else:
|
||||
st.silence_ms += frame_ms
|
||||
|
||||
spoken_ms = len(st.buffer) * frame_ms
|
||||
ended = st.silence_ms >= settings.vad_silence_ms
|
||||
too_long = spoken_ms >= settings.vad_max_utterance_ms
|
||||
|
||||
if ended or too_long:
|
||||
audio = np.concatenate(st.buffer)
|
||||
done.append(Utterance(audio=audio, duration_ms=spoken_ms, truncated=too_long and not ended))
|
||||
self.reset()
|
||||
|
||||
return done, started
|
||||
Reference in New Issue
Block a user