WeLe Agentic AI: LangGraph multi-agent CRM assistant with voice

This commit is contained in:
2026-08-28 02:16:03 +05:30
commit 105e58e02a
69 changed files with 11501 additions and 0 deletions
+82
View File
@@ -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()
+233
View File
@@ -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()
+200
View File
@@ -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)
+155
View File
@@ -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)
+127
View File
@@ -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