Files
2026-05-01 17:50:41 +05:00

322 lines
12 KiB
Python

from __future__ import annotations
import logging
from abc import ABC, abstractmethod
from dataclasses import dataclass
from typing import Any, Callable
LOGGER = logging.getLogger("uvicorn.error")
def _duration_ms_from_samples(samples: int, *, sample_rate_hz: int) -> int:
if samples <= 0 or sample_rate_hz <= 0:
return 0
return int((samples / float(sample_rate_hz)) * 1000.0)
@dataclass
class VADFrameResult:
is_speech: bool = False
speech_started: bool = False
speech_ended: bool = False
speech_probability: float = 0.0
utterance_audio: bytes | None = None
class BaseVAD(ABC):
@abstractmethod
def feed(self, audio_chunk: bytes) -> VADFrameResult:
raise NotImplementedError
@abstractmethod
def set_speech_end_silence_ms(self, value: int) -> None:
raise NotImplementedError
@abstractmethod
def restore_default_speech_end_silence_ms(self) -> None:
raise NotImplementedError
@abstractmethod
def current_utterance_audio(self) -> bytes:
raise NotImplementedError
@abstractmethod
def flush(self) -> bytes | None:
raise NotImplementedError
@abstractmethod
def reset(self) -> None:
raise NotImplementedError
class SileroVADDetector(BaseVAD):
def __init__(
self,
*,
sample_rate_hz: int = 16000,
threshold: float = 0.5,
negative_threshold: float | None = None,
speech_end_silence_ms: int = 1600,
speech_pad_ms: int = 64,
min_speech_duration_ms: int = 0,
model: Any | None = None,
prediction_fn: Callable[[bytes], float] | None = None,
use_onnx: bool = False,
) -> None:
if sample_rate_hz not in {8000, 16000}:
raise ValueError("Silero VAD supports only 8000 Hz and 16000 Hz sample rates")
self.sample_rate_hz = sample_rate_hz
self.threshold = max(min(threshold, 1.0), 0.0)
self.negative_threshold = (
max(min(negative_threshold, 1.0), 0.0)
if negative_threshold is not None
else max(self.threshold - 0.15, 0.01)
)
self._default_speech_end_silence_ms = max(speech_end_silence_ms, 32)
self.speech_end_silence_ms = self._default_speech_end_silence_ms
self.speech_pad_ms = max(speech_pad_ms, 0)
self.min_speech_duration_ms = max(min_speech_duration_ms, 0)
self.window_samples = 512 if self.sample_rate_hz == 16000 else 256
self.window_bytes = self.window_samples * 2
self._speech_pad_bytes = int(self.sample_rate_hz * self.speech_pad_ms / 1000.0) * 2
self._min_speech_samples = int(self.sample_rate_hz * self.min_speech_duration_ms / 1000.0)
self._model = model
self._prediction_fn = prediction_fn
self._use_onnx = use_onnx
self._set_speech_end_silence_ms(self.speech_end_silence_ms)
self.reset()
LOGGER.info(
"Silero VAD config: sample_rate=%s threshold=%.3f negative_threshold=%.3f "
"silence_timeout_ms=%s speech_pad_ms=%s min_speech_ms=%s window_bytes=%s onnx=%s",
self.sample_rate_hz,
self.threshold,
self.negative_threshold,
self.speech_end_silence_ms,
self.speech_pad_ms,
self.min_speech_duration_ms,
self.window_bytes,
self._use_onnx,
)
@property
def default_speech_end_silence_ms(self) -> int:
return self._default_speech_end_silence_ms
@property
def current_speech_end_silence_ms(self) -> int:
return self.speech_end_silence_ms
def reset(self) -> None:
self._window_buffer = bytearray()
self._pre_speech_audio = bytearray()
self._utterance_audio = bytearray()
self._triggered = False
self._silence_samples = 0
self._speech_samples = 0
self.restore_default_speech_end_silence_ms()
if self._model is not None and hasattr(self._model, "reset_states"):
self._model.reset_states()
LOGGER.info("Silero VAD reset: sample_rate=%s silence_timeout_ms=%s", self.sample_rate_hz, self.speech_end_silence_ms)
def set_speech_end_silence_ms(self, value: int) -> None:
self._set_speech_end_silence_ms(value)
def restore_default_speech_end_silence_ms(self) -> None:
self._set_speech_end_silence_ms(self._default_speech_end_silence_ms)
def current_utterance_audio(self) -> bytes:
return bytes(self._utterance_audio)
def feed(self, audio_chunk: bytes) -> VADFrameResult:
result = VADFrameResult()
if not audio_chunk:
return result
self._window_buffer.extend(audio_chunk)
while len(self._window_buffer) >= self.window_bytes:
window = bytes(self._window_buffer[: self.window_bytes])
del self._window_buffer[: self.window_bytes]
speech_probability = self._predict_speech_probability(window)
result.speech_probability = speech_probability
result.is_speech = result.is_speech or self._triggered or speech_probability >= self.threshold
if not self._triggered:
if speech_probability >= self.threshold:
self._triggered = True
self._speech_samples = self.window_samples
self._silence_samples = 0
self._utterance_audio = bytearray(self._pre_speech_audio)
self._utterance_audio.extend(window)
self._pre_speech_audio.clear()
result.speech_started = True
result.is_speech = True
LOGGER.info(
"Silero VAD speech_started: probability=%.3f pre_speech_bytes=%s window_bytes=%s",
speech_probability,
len(self._utterance_audio) - len(window),
len(window),
)
else:
self._append_pre_speech_window(window)
continue
self._utterance_audio.extend(window)
result.is_speech = True
if speech_probability >= self.threshold:
self._silence_samples = 0
self._speech_samples += self.window_samples
continue
if speech_probability >= self.negative_threshold:
self._silence_samples = 0
continue
self._silence_samples += self.window_samples
if self._silence_samples < self._speech_end_silence_samples:
continue
utterance_audio = self._trim_trailing_silence(
bytes(self._utterance_audio),
trailing_silence_samples=self._silence_samples,
)
speech_samples = self._speech_samples
self._reset_segment()
result.is_speech = False
if speech_samples >= self._min_speech_samples and utterance_audio:
result.speech_ended = True
result.utterance_audio = utterance_audio
LOGGER.info(
"Silero VAD speech_ended: speech_ms=%s trailing_silence_ms=%s utterance_bytes=%s",
_duration_ms_from_samples(speech_samples, sample_rate_hz=self.sample_rate_hz),
_duration_ms_from_samples(self._speech_end_silence_samples, sample_rate_hz=self.sample_rate_hz),
len(utterance_audio),
)
return result
def flush(self) -> bytes | None:
if not self._triggered or not self._utterance_audio:
return None
utterance_audio = self._trim_trailing_silence(
bytes(self._utterance_audio),
trailing_silence_samples=self._silence_samples,
)
speech_samples = self._speech_samples
self._reset_segment()
if speech_samples < self._min_speech_samples or not utterance_audio:
return None
LOGGER.info(
"Silero VAD flush utterance: speech_ms=%s utterance_bytes=%s",
_duration_ms_from_samples(speech_samples, sample_rate_hz=self.sample_rate_hz),
len(utterance_audio),
)
return utterance_audio
def _reset_segment(self) -> None:
self._triggered = False
self._silence_samples = 0
self._speech_samples = 0
self._utterance_audio = bytearray()
self._pre_speech_audio.clear()
def _append_pre_speech_window(self, window: bytes) -> None:
self._pre_speech_audio.extend(window)
if self._speech_pad_bytes <= 0:
self._pre_speech_audio.clear()
return
overflow = len(self._pre_speech_audio) - self._speech_pad_bytes
if overflow > 0:
del self._pre_speech_audio[:overflow]
def _trim_trailing_silence(self, utterance_audio: bytes, *, trailing_silence_samples: int) -> bytes:
if not utterance_audio or trailing_silence_samples <= 0:
return utterance_audio
trim_bytes = min(trailing_silence_samples * 2, len(utterance_audio))
if trim_bytes <= 0:
return utterance_audio
LOGGER.info(
"Silero VAD trim trailing silence: trim_bytes=%s trim_ms=%s before_bytes=%s after_bytes=%s",
trim_bytes,
_duration_ms_from_samples(trailing_silence_samples, sample_rate_hz=self.sample_rate_hz),
len(utterance_audio),
len(utterance_audio) - trim_bytes,
)
return utterance_audio[:-trim_bytes]
def _set_speech_end_silence_ms(self, value: int) -> None:
previous = getattr(self, "speech_end_silence_ms", None)
self.speech_end_silence_ms = max(int(value), 32)
self._speech_end_silence_samples = int(self.sample_rate_hz * self.speech_end_silence_ms / 1000.0)
if previous != self.speech_end_silence_ms:
LOGGER.info(
"Silero VAD silence timeout changed: previous_ms=%s current_ms=%s samples=%s",
previous,
self.speech_end_silence_ms,
self._speech_end_silence_samples,
)
def _predict_speech_probability(self, window: bytes) -> float:
if self._prediction_fn is not None:
return max(0.0, min(float(self._prediction_fn(window)), 1.0))
model = self._load_model()
try:
import numpy as np
import torch
except Exception as exc: # noqa: BLE001
raise RuntimeError(
"Silero VAD requires numpy and torch at runtime; add them to requirements.txt"
) from exc
pcm = np.frombuffer(window, dtype=np.int16).astype(np.float32) / 32768.0
tensor = torch.from_numpy(pcm)
probability = model(tensor, self.sample_rate_hz).item()
return max(0.0, min(float(probability), 1.0))
def _load_model(self) -> Any:
if self._model is not None:
return self._model
try:
import torch
except Exception as exc: # noqa: BLE001
raise RuntimeError(
"Silero VAD requires torch; install torch/torchaudio or use the ONNX runtime path"
) from exc
torch.set_num_threads(1)
try:
from silero_vad import load_silero_vad
except Exception: # noqa: BLE001
load_silero_vad = None
if load_silero_vad is not None:
self._model = load_silero_vad(onnx=self._use_onnx)
LOGGER.info(
"realtime voice VAD loaded via silero_vad package sample_rate=%s onnx=%s",
self.sample_rate_hz,
self._use_onnx,
)
return self._model
try:
self._model, _ = torch.hub.load(
repo_or_dir="snakers4/silero-vad",
model="silero_vad",
onnx=self._use_onnx,
force_reload=False,
)
except TypeError:
self._model, _ = torch.hub.load(
repo_or_dir="snakers4/silero-vad",
model="silero_vad",
force_reload=False,
)
LOGGER.info(
"realtime voice VAD loaded via torch.hub sample_rate=%s onnx=%s",
self.sample_rate_hz,
self._use_onnx,
)
return self._model