128 lines
4.2 KiB
Python
128 lines
4.2 KiB
Python
"""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
|