summaryrefslogtreecommitdiff
path: root/src/mic_clipper/vad.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/mic_clipper/vad.py')
-rw-r--r--src/mic_clipper/vad.py59
1 files changed, 59 insertions, 0 deletions
diff --git a/src/mic_clipper/vad.py b/src/mic_clipper/vad.py
new file mode 100644
index 0000000..705dfd8
--- /dev/null
+++ b/src/mic_clipper/vad.py
@@ -0,0 +1,59 @@
+"""Stateful, local Silero VAD ONNX inference."""
+
+from __future__ import annotations
+
+from importlib import resources
+from pathlib import Path
+
+import numpy as np
+import onnxruntime
+
+
+FRAME_SAMPLES = 512
+CONTEXT_SAMPLES = 64
+
+
+class SileroVad:
+ """Run the bundled Silero v6 model on exact 32 ms PCM frames."""
+
+ def __init__(self, model_path: Path | None = None) -> None:
+ self.model_path = Path(model_path) if model_path else self._default_model_path()
+ options = onnxruntime.SessionOptions()
+ options.inter_op_num_threads = 1
+ options.intra_op_num_threads = 1
+ options.enable_cpu_mem_arena = False
+ options.log_severity_level = 4
+ self.session = onnxruntime.InferenceSession(
+ self.model_path,
+ providers=["CPUExecutionProvider"],
+ sess_options=options,
+ )
+ self._context = np.zeros(CONTEXT_SAMPLES, dtype=np.float32)
+ self._h = np.zeros((1, 1, 128), dtype=np.float32)
+ self._c = np.zeros((1, 1, 128), dtype=np.float32)
+
+ def probability(self, samples: np.ndarray) -> float:
+ samples = np.asarray(samples, dtype=np.float32)
+ if samples.ndim != 1 or len(samples) != FRAME_SAMPLES:
+ raise ValueError("Silero VAD requires exactly 512 mono samples at 16 kHz")
+ model_input = np.concatenate((self._context, samples)).reshape(1, -1)
+ output, self._h, self._c = self.session.run(
+ None,
+ {"input": model_input, "h": self._h, "c": self._c},
+ )
+ self._context = samples[-CONTEXT_SAMPLES:].copy()
+ return float(output[0])
+
+ @staticmethod
+ def _default_model_path() -> Path:
+ bundled = resources.files("mic_clipper.assets").joinpath("silero_vad_v6.onnx")
+ if bundled.is_file():
+ return Path(bundled)
+ try:
+ import faster_whisper
+ except ImportError as error:
+ raise RuntimeError(
+ "Silero VAD model is missing; bundle silero_vad_v6.onnx or install the "
+ "pre-provisioned faster-whisper runtime asset"
+ ) from error
+ return Path(faster_whisper.__file__).parent / "assets" / "silero_vad_v6.onnx"