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