VAD fallback (#97)

* Silero VAD preferred with webrtc fallback

* webrtc VAD neds a different sample size

* fixup

* fixup
This commit is contained in:
chadbailey59
2024-04-04 13:31:07 -05:00
committed by GitHub
parent 385b51ac83
commit 03ea208361
2 changed files with 59 additions and 21 deletions

View File

@@ -5,6 +5,7 @@ import signal
import threading import threading
import types import types
from enum import Enum
from functools import partial from functools import partial
from typing import Any from typing import Any
@@ -33,6 +34,11 @@ except ModuleNotFoundError as e:
from dailyai.transports.threaded_transport import ThreadedTransport from dailyai.transports.threaded_transport import ThreadedTransport
NUM_CHANNELS = 1
SPEECH_THRESHOLD = 0.90
VAD_RESET_PERIOD_MS = 2000
class DailyTransport(ThreadedTransport, EventHandler): class DailyTransport(ThreadedTransport, EventHandler):
_daily_initialized = False _daily_initialized = False
@@ -55,6 +61,7 @@ class DailyTransport(ThreadedTransport, EventHandler):
start_transcription: bool = False, start_transcription: bool = False,
**kwargs, **kwargs,
): ):
kwargs['has_webrtc_vad'] = True
# This will call ThreadedTransport.__init__ method, not EventHandler # This will call ThreadedTransport.__init__ method, not EventHandler
super().__init__(**kwargs) super().__init__(**kwargs)
@@ -86,6 +93,12 @@ class DailyTransport(ThreadedTransport, EventHandler):
self._event_handlers = {} self._event_handlers = {}
self.webrtc_vad = Daily.create_native_vad(
reset_period_ms=VAD_RESET_PERIOD_MS,
sample_rate=self._speaker_sample_rate,
channels=NUM_CHANNELS
)
def _patch_method(self, event_name, *args, **kwargs): def _patch_method(self, event_name, *args, **kwargs):
try: try:
for handler in self._event_handlers[event_name]: for handler in self._event_handlers[event_name]:
@@ -106,6 +119,18 @@ class DailyTransport(ThreadedTransport, EventHandler):
self._logger.error(f"Exception in event handler {event_name}: {e}") self._logger.error(f"Exception in event handler {event_name}: {e}")
raise e raise e
def _webrtc_vad_analyze(self):
buffer = self.read_audio_frames(
int(self._vad_samples))
if len(buffer) > 0:
confidence = self.webrtc_vad.analyze_frames(buffer)
# yeses = int(confidence * 20.0)
# nos = 20 - yeses
# out = "!" * yeses + "." * nos
# print(f"!!! confidence: {out} {confidence}")
talking = confidence > SPEECH_THRESHOLD
return talking
def add_event_handler(self, event_name: str, handler): def add_event_handler(self, event_name: str, handler):
if not event_name.startswith("on_"): if not event_name.startswith("on_"):
raise Exception( raise Exception(

View File

@@ -40,9 +40,6 @@ def int2float(sound):
return sound return sound
SAMPLE_RATE = 16000
class VADState(Enum): class VADState(Enum):
QUIET = 1 QUIET = 1
STARTING = 2 STARTING = 2
@@ -61,11 +58,12 @@ class ThreadedTransport(AbstractTransport):
self._vad_stop_s = kwargs.get("vad_stop_s") or 0.8 self._vad_stop_s = kwargs.get("vad_stop_s") or 0.8
self._context = kwargs.get("context") or [] self._context = kwargs.get("context") or []
self._vad_enabled = kwargs.get("vad_enabled") or False self._vad_enabled = kwargs.get("vad_enabled") or False
self._has_webrtc_vad = kwargs.get("has_webrtc_vad") or False
if self._vad_enabled and self._speaker_enabled: if self._vad_enabled and self._speaker_enabled:
raise Exception( raise Exception(
"Sorry, you can't use speaker_enabled and vad_enabled at the same time. Please set one to False." "Sorry, you can't use speaker_enabled and vad_enabled at the same time. Please set one to False."
) )
self._vad_samples = 1536
if self._vad_enabled: if self._vad_enabled:
try: try:
@@ -79,14 +77,19 @@ class ThreadedTransport(AbstractTransport):
(self.model, self.utils) = torch.hub.load( (self.model, self.utils) = torch.hub.load(
repo_or_dir="snakers4/silero-vad", model="silero_vad", force_reload=False repo_or_dir="snakers4/silero-vad", model="silero_vad", force_reload=False
) )
self._logger.debug("Loaded Silero VAD")
except ModuleNotFoundError as e: except ModuleNotFoundError as e:
print(f"Exception: {e}") if self._has_webrtc_vad:
print("In order to use VAD, you'll need to install the `torch` and `torchaudio` modules.") self._logger.debug(f"Couldn't load torch; using webrtc VAD")
raise Exception(f"Missing module(s): {e}") self._vad_samples = int(self._speaker_sample_rate / 100.0)
else:
self._logger.error(f"Exception: {e}")
self._logger.error(
"In order to use VAD, you'll need to install the `torch` and `torchaudio` modules.")
raise Exception(f"Missing module(s): {e}")
self._vad_samples = 1536 vad_frame_s = self._vad_samples / self._speaker_sample_rate
vad_frame_s = self._vad_samples / SAMPLE_RATE
self._vad_start_frames = round(self._vad_start_s / vad_frame_s) self._vad_start_frames = round(self._vad_start_s / vad_frame_s)
self._vad_stop_frames = round(self._vad_stop_s / vad_frame_s) self._vad_stop_frames = round(self._vad_stop_s / vad_frame_s)
self._vad_starting_count = 0 self._vad_starting_count = 0
@@ -262,19 +265,28 @@ class ThreadedTransport(AbstractTransport):
def _prerun(self): def _prerun(self):
pass pass
def _vad(self): def _silero_vad_analyze(self):
# CB: Starting silero VAD stuff audio_chunk = self.read_audio_frames(self._vad_samples)
# TODO-CB: Probably need to force virtual speaker creation if we're audio_int16 = np.frombuffer(audio_chunk, np.int16)
# going to build this in? audio_float32 = int2float(audio_int16)
# TODO-CB: pyaudio installation new_confidence = self.model(
while not self._stop_threads.is_set(): torch.from_numpy(audio_float32), 16000).item()
audio_chunk = self.read_audio_frames(self._vad_samples) # yeses = int(new_confidence * 20.0)
audio_int16 = np.frombuffer(audio_chunk, np.int16) # nos = 20 - yeses
audio_float32 = int2float(audio_int16) # out = "!" * yeses + "." * nos
new_confidence = self.model( # print(f"!!! confidence: {out}")
torch.from_numpy(audio_float32), 16000).item() speaking = new_confidence > 0.5
speaking = new_confidence > 0.5 return speaking
def _vad(self):
while not self._stop_threads.is_set():
if hasattr(self, 'model'): # we can use Silero
speaking = self._silero_vad_analyze()
elif self._has_webrtc_vad:
speaking = self._webrtc_vad_analyze()
else:
raise Exception("VAD is running with no VAD service available")
if speaking: if speaking:
match self._vad_state: match self._vad_state:
case VADState.QUIET: case VADState.QUIET:
@@ -311,6 +323,7 @@ class ThreadedTransport(AbstractTransport):
self._vad_state == VADState.STOPPING self._vad_state == VADState.STOPPING
and self._vad_stopping_count >= self._vad_stop_frames and self._vad_stopping_count >= self._vad_stop_frames
): ):
if self._loop: if self._loop:
asyncio.run_coroutine_threadsafe( asyncio.run_coroutine_threadsafe(
self.receive_queue.put( self.receive_queue.put(