322 lines
12 KiB
Python
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
|