Initial commit
This commit is contained in:
@@ -0,0 +1,4 @@
|
||||
from realtime_voice_service.core.session import CallSession, SessionState
|
||||
from realtime_voice_service.core.vad import BaseVAD, SileroVADDetector, VADFrameResult
|
||||
|
||||
__all__ = ["BaseVAD", "CallSession", "SessionState", "SileroVADDetector", "VADFrameResult"]
|
||||
+389
@@ -0,0 +1,389 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Iterable
|
||||
from enum import Enum
|
||||
|
||||
from realtime_voice_service.core.vad import BaseVAD, SileroVADDetector
|
||||
from realtime_voice_service.providers.base import BaseLLM, BaseSTT, BaseTTS, MockLLM, MockSTT, MockTTS
|
||||
from realtime_voice_service.transports.base import BaseMediaTransport
|
||||
|
||||
|
||||
LOGGER = logging.getLogger("uvicorn.error")
|
||||
|
||||
|
||||
class SessionState(str, Enum):
|
||||
LISTENING = "LISTENING"
|
||||
USER_SPEAKING = "USER_SPEAKING"
|
||||
ASSISTANT_THINKING = "ASSISTANT_THINKING"
|
||||
ASSISTANT_SPEAKING = "ASSISTANT_SPEAKING"
|
||||
|
||||
|
||||
class GenerationInterrupted(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class TextChunker:
|
||||
_BOUNDARY_CHARS = ".!?\n"
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._buffer = ""
|
||||
|
||||
def feed(self, fragment: str) -> list[str]:
|
||||
if not fragment:
|
||||
return []
|
||||
self._buffer += fragment
|
||||
return self._extract_ready_chunks()
|
||||
|
||||
def flush(self) -> list[str]:
|
||||
chunks = self._extract_ready_chunks()
|
||||
tail = self._buffer.strip()
|
||||
self._buffer = ""
|
||||
if tail:
|
||||
chunks.append(tail)
|
||||
return chunks
|
||||
|
||||
def reset(self) -> None:
|
||||
self._buffer = ""
|
||||
|
||||
def _extract_ready_chunks(self) -> list[str]:
|
||||
chunks: list[str] = []
|
||||
while True:
|
||||
next_chunk = self._pop_next_chunk()
|
||||
if next_chunk is None:
|
||||
break
|
||||
chunks.append(next_chunk)
|
||||
return chunks
|
||||
|
||||
def _pop_next_chunk(self) -> str | None:
|
||||
for index, char in enumerate(self._buffer):
|
||||
if char not in self._BOUNDARY_CHARS:
|
||||
continue
|
||||
end_index = index + 1
|
||||
while end_index < len(self._buffer) and self._buffer[end_index] in self._BOUNDARY_CHARS:
|
||||
end_index += 1
|
||||
while end_index < len(self._buffer) and self._buffer[end_index].isspace():
|
||||
end_index += 1
|
||||
chunk = self._buffer[:end_index].strip()
|
||||
self._buffer = self._buffer[end_index:]
|
||||
if chunk:
|
||||
return chunk
|
||||
return None
|
||||
|
||||
|
||||
class CallSession:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
session_id: str,
|
||||
transport: BaseMediaTransport,
|
||||
vad: BaseVAD | None = None,
|
||||
stt: BaseSTT | None = None,
|
||||
llm: BaseLLM | None = None,
|
||||
tts: BaseTTS | None = None,
|
||||
) -> None:
|
||||
self.session_id = session_id
|
||||
self.transport = transport
|
||||
self.state = SessionState.LISTENING
|
||||
self.generation_epoch = 0
|
||||
self.interruptions: list[str] = []
|
||||
self.last_latency_ms: dict[str, int] = {}
|
||||
self._vad = vad or SileroVADDetector(sample_rate_hz=transport.sample_rate_hz)
|
||||
self._stt = stt or MockSTT(sample_rate_hz=transport.sample_rate_hz)
|
||||
self._llm = llm or MockLLM()
|
||||
self._tts = tts or MockTTS(sample_rate_hz=transport.sample_rate_hz)
|
||||
self._conversation: list[tuple[str, str]] = []
|
||||
self._assistant_task: asyncio.Task[None] | None = None
|
||||
self._sentence_queue: asyncio.Queue[str | None] | None = None
|
||||
self._closed = False
|
||||
|
||||
@property
|
||||
def conversation(self) -> tuple[tuple[str, str], ...]:
|
||||
return tuple(self._conversation)
|
||||
|
||||
async def run(self) -> None:
|
||||
if self._assistant_task is not None:
|
||||
raise RuntimeError("CallSession.run() can only be called once per session")
|
||||
try:
|
||||
await self.media_loop()
|
||||
finally:
|
||||
await self.stop()
|
||||
|
||||
async def media_loop(self) -> None:
|
||||
while not self._closed:
|
||||
audio_chunk = await self.transport.receive_audio()
|
||||
if audio_chunk is None:
|
||||
LOGGER.info("realtime session %s transport closed", self.session_id)
|
||||
break
|
||||
if not audio_chunk:
|
||||
continue
|
||||
|
||||
vad_result = self._vad.feed(audio_chunk)
|
||||
|
||||
if vad_result.is_speech and self.state != SessionState.USER_SPEAKING:
|
||||
if self.state in {SessionState.ASSISTANT_THINKING, SessionState.ASSISTANT_SPEAKING}:
|
||||
self.interrupt("barge-in")
|
||||
self._set_state(
|
||||
SessionState.USER_SPEAKING,
|
||||
reason=f"speech detected prob={vad_result.speech_probability:.3f}",
|
||||
)
|
||||
|
||||
if vad_result.speech_ended and vad_result.utterance_audio:
|
||||
speech_end_monotonic = time.perf_counter()
|
||||
self._set_state(SessionState.ASSISTANT_THINKING, reason="speech end detected")
|
||||
self._start_assistant_turn(
|
||||
epoch=self.generation_epoch,
|
||||
utterance_audio=vad_result.utterance_audio,
|
||||
speech_end_monotonic=speech_end_monotonic,
|
||||
)
|
||||
|
||||
def interrupt(self, reason: str = "interrupt") -> int:
|
||||
self.generation_epoch += 1
|
||||
self.interruptions.append(reason)
|
||||
self._clear_sentence_queue()
|
||||
if self._assistant_task is not None and not self._assistant_task.done():
|
||||
self._assistant_task.cancel()
|
||||
LOGGER.info(
|
||||
"realtime session %s interrupted: epoch=%s reason=%s",
|
||||
self.session_id,
|
||||
self.generation_epoch,
|
||||
reason,
|
||||
)
|
||||
return self.generation_epoch
|
||||
|
||||
async def stop(self) -> None:
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
self._vad.reset()
|
||||
self._clear_sentence_queue()
|
||||
if self._assistant_task is not None and not self._assistant_task.done():
|
||||
self._assistant_task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError):
|
||||
await self._assistant_task
|
||||
await self.transport.close()
|
||||
|
||||
def _start_assistant_turn(
|
||||
self,
|
||||
*,
|
||||
epoch: int,
|
||||
utterance_audio: bytes,
|
||||
speech_end_monotonic: float,
|
||||
) -> None:
|
||||
if self._assistant_task is not None and not self._assistant_task.done():
|
||||
self._assistant_task.cancel()
|
||||
self._assistant_task = asyncio.create_task(
|
||||
self._run_assistant_turn(
|
||||
epoch=epoch,
|
||||
utterance_audio=utterance_audio,
|
||||
speech_end_monotonic=speech_end_monotonic,
|
||||
),
|
||||
name=f"{self.session_id}-assistant-{epoch}",
|
||||
)
|
||||
|
||||
async def _run_assistant_turn(
|
||||
self,
|
||||
*,
|
||||
epoch: int,
|
||||
utterance_audio: bytes,
|
||||
speech_end_monotonic: float,
|
||||
) -> None:
|
||||
sentence_queue: asyncio.Queue[str | None] | None = None
|
||||
playback_task: asyncio.Task[None] | None = None
|
||||
assistant_fragments: list[str] = []
|
||||
try:
|
||||
transcript = await self._stt.transcribe(utterance_audio)
|
||||
self._ensure_generation(epoch)
|
||||
self._log_latency(
|
||||
"stt_latency",
|
||||
speech_end_monotonic,
|
||||
epoch=epoch,
|
||||
message="speech end -> transcript ready",
|
||||
)
|
||||
transcript = transcript.strip()
|
||||
if not transcript:
|
||||
LOGGER.info("realtime session %s produced empty transcript", self.session_id)
|
||||
self._set_state(SessionState.LISTENING, reason="empty transcript")
|
||||
return
|
||||
|
||||
self._conversation.append(("user", transcript))
|
||||
sentence_queue = asyncio.Queue()
|
||||
self._sentence_queue = sentence_queue
|
||||
playback_task = asyncio.create_task(
|
||||
self._stream_tts_pipeline(epoch=epoch, sentence_queue=sentence_queue),
|
||||
name=f"{self.session_id}-playback-{epoch}",
|
||||
)
|
||||
llm_started_monotonic = time.perf_counter()
|
||||
chunker = TextChunker()
|
||||
first_token_seen = False
|
||||
async for token in self._llm.generate_stream(transcript, list(self._conversation)):
|
||||
self._ensure_generation(epoch)
|
||||
if not token:
|
||||
continue
|
||||
if not first_token_seen:
|
||||
first_token_seen = True
|
||||
self._log_latency(
|
||||
"ttft",
|
||||
llm_started_monotonic,
|
||||
epoch=epoch,
|
||||
message="llm request -> first token",
|
||||
)
|
||||
assistant_fragments.append(token)
|
||||
await self._enqueue_chunks(
|
||||
epoch=epoch,
|
||||
sentence_queue=sentence_queue,
|
||||
chunks=chunker.feed(token),
|
||||
)
|
||||
|
||||
assistant_text = "".join(assistant_fragments).strip()
|
||||
if not assistant_text:
|
||||
await self._finish_sentence_queue(sentence_queue)
|
||||
LOGGER.warning("realtime session %s llm produced no text", self.session_id)
|
||||
self._set_state(SessionState.LISTENING, reason="empty llm response")
|
||||
return
|
||||
|
||||
await self._enqueue_chunks(
|
||||
epoch=epoch,
|
||||
sentence_queue=sentence_queue,
|
||||
chunks=chunker.flush(),
|
||||
)
|
||||
await self._finish_sentence_queue(sentence_queue)
|
||||
if playback_task is not None:
|
||||
await playback_task
|
||||
self._ensure_generation(epoch)
|
||||
self._conversation.append(("assistant", assistant_text))
|
||||
self._set_state(SessionState.LISTENING, reason="assistant turn completed")
|
||||
except GenerationInterrupted:
|
||||
LOGGER.info(
|
||||
"realtime session %s ignored stale generation %s",
|
||||
self.session_id,
|
||||
epoch,
|
||||
)
|
||||
except asyncio.CancelledError:
|
||||
LOGGER.info(
|
||||
"realtime session %s cancelled generation %s",
|
||||
self.session_id,
|
||||
epoch,
|
||||
)
|
||||
raise
|
||||
except Exception:
|
||||
LOGGER.exception(
|
||||
"realtime session %s generation %s failed",
|
||||
self.session_id,
|
||||
epoch,
|
||||
)
|
||||
if epoch == self.generation_epoch and not self._closed:
|
||||
self._set_state(SessionState.LISTENING, reason="assistant generation failed")
|
||||
finally:
|
||||
self._clear_sentence_queue(sentence_queue)
|
||||
if playback_task is not None and not playback_task.done():
|
||||
playback_task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError):
|
||||
await playback_task
|
||||
current_task = asyncio.current_task()
|
||||
if self._assistant_task is current_task:
|
||||
self._assistant_task = None
|
||||
|
||||
def _ensure_generation(self, epoch: int) -> None:
|
||||
if self._closed or epoch != self.generation_epoch:
|
||||
raise GenerationInterrupted(f"stale generation {epoch}")
|
||||
|
||||
def _log_latency(
|
||||
self,
|
||||
metric_name: str,
|
||||
started_monotonic: float,
|
||||
*,
|
||||
epoch: int,
|
||||
message: str,
|
||||
) -> None:
|
||||
latency_ms = max(int((time.perf_counter() - started_monotonic) * 1000.0), 0)
|
||||
self.last_latency_ms[metric_name] = latency_ms
|
||||
LOGGER.info(
|
||||
"realtime session %s %s=%sms epoch=%s (%s)",
|
||||
self.session_id,
|
||||
metric_name,
|
||||
latency_ms,
|
||||
epoch,
|
||||
message,
|
||||
)
|
||||
|
||||
async def _enqueue_chunks(
|
||||
self,
|
||||
*,
|
||||
epoch: int,
|
||||
sentence_queue: asyncio.Queue[str | None],
|
||||
chunks: Iterable[str],
|
||||
) -> None:
|
||||
for chunk in chunks:
|
||||
self._ensure_generation(epoch)
|
||||
normalized = chunk.strip()
|
||||
if not normalized:
|
||||
continue
|
||||
await sentence_queue.put(normalized)
|
||||
|
||||
async def _finish_sentence_queue(self, sentence_queue: asyncio.Queue[str | None]) -> None:
|
||||
await sentence_queue.put(None)
|
||||
|
||||
async def _stream_tts_pipeline(
|
||||
self,
|
||||
*,
|
||||
epoch: int,
|
||||
sentence_queue: asyncio.Queue[str | None],
|
||||
) -> None:
|
||||
first_audio_seen = False
|
||||
first_tts_started_monotonic: float | None = None
|
||||
while True:
|
||||
self._ensure_generation(epoch)
|
||||
sentence = await sentence_queue.get()
|
||||
if sentence is None:
|
||||
return
|
||||
normalized = sentence.strip()
|
||||
if not normalized:
|
||||
continue
|
||||
if first_tts_started_monotonic is None:
|
||||
first_tts_started_monotonic = time.perf_counter()
|
||||
async for audio_chunk in self._tts.synthesize_stream(normalized):
|
||||
self._ensure_generation(epoch)
|
||||
if not audio_chunk:
|
||||
continue
|
||||
if not first_audio_seen:
|
||||
first_audio_seen = True
|
||||
self._log_latency(
|
||||
"ttfa",
|
||||
first_tts_started_monotonic or time.perf_counter(),
|
||||
epoch=epoch,
|
||||
message="tts request -> first audio",
|
||||
)
|
||||
self._set_state(SessionState.ASSISTANT_SPEAKING, reason="assistant playback started")
|
||||
await self.transport.send_audio(audio_chunk)
|
||||
|
||||
def _clear_sentence_queue(self, sentence_queue: asyncio.Queue[str | None] | None = None) -> None:
|
||||
queue = sentence_queue if sentence_queue is not None else self._sentence_queue
|
||||
if queue is None:
|
||||
return
|
||||
while True:
|
||||
try:
|
||||
queue.get_nowait()
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
with contextlib.suppress(asyncio.QueueFull):
|
||||
queue.put_nowait(None)
|
||||
if queue is self._sentence_queue:
|
||||
self._sentence_queue = None
|
||||
|
||||
def _set_state(self, state: SessionState, *, reason: str | None = None) -> None:
|
||||
if self.state == state:
|
||||
return
|
||||
previous_state = self.state
|
||||
self.state = state
|
||||
suffix = f" ({reason})" if reason else ""
|
||||
LOGGER.info(
|
||||
"realtime session %s state %s -> %s%s",
|
||||
self.session_id,
|
||||
previous_state.value,
|
||||
state.value,
|
||||
suffix,
|
||||
)
|
||||
+221
@@ -0,0 +1,221 @@
|
||||
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")
|
||||
|
||||
|
||||
@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 flush(self) -> bytes | None:
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def reset(self) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class SileroVADDetector(BaseVAD):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
sample_rate_hz: int = 8000,
|
||||
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.speech_end_silence_ms = max(speech_end_silence_ms, 32)
|
||||
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_end_silence_samples = int(self.sample_rate_hz * self.speech_end_silence_ms / 1000.0)
|
||||
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.reset()
|
||||
|
||||
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
|
||||
if self._model is not None and hasattr(self._model, "reset_states"):
|
||||
self._model.reset_states()
|
||||
|
||||
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
|
||||
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 = bytes(self._utterance_audio)
|
||||
speech_samples = self._speech_samples
|
||||
self._reset_segment()
|
||||
result.is_speech = False
|
||||
if speech_samples >= self._min_speech_samples:
|
||||
result.speech_ended = True
|
||||
result.utterance_audio = utterance_audio
|
||||
|
||||
return result
|
||||
|
||||
def flush(self) -> bytes | None:
|
||||
if not self._triggered or not self._utterance_audio:
|
||||
return None
|
||||
utterance_audio = bytes(self._utterance_audio)
|
||||
speech_samples = self._speech_samples
|
||||
self._reset_segment()
|
||||
if speech_samples < self._min_speech_samples:
|
||||
return None
|
||||
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 _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
|
||||
Reference in New Issue
Block a user