From 581406e5abc5ef325c8eaeb9331ae061b54febd6 Mon Sep 17 00:00:00 2001 From: makaveli10 Date: Wed, 27 Sep 2023 15:43:28 +0800 Subject: [PATCH 1/4] add vad model to suppress warnings --- whisper_live/server.py | 7 +-- whisper_live/vad.py | 112 +++++++++++++++++++++++++++++++++++++++++ 2 files changed, 114 insertions(+), 5 deletions(-) create mode 100644 whisper_live/vad.py diff --git a/whisper_live/server.py b/whisper_live/server.py index aa4439a..a81f4f3 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -17,6 +17,7 @@ import torch import numpy as np import time from whisper_live.transcriber import WhisperModel +from whisper_live.vad import VoiceActivityDetection class TranscriptionServer: @@ -29,11 +30,7 @@ class TranscriptionServer: RATE = 16000 def __init__(self): # voice activity detection model - self.vad_model, _ = torch.hub.load(repo_or_dir='snakers4/silero-vad', - model='silero_vad', - force_reload=True, - onnx=True - ) + self.vad_model = VoiceActivityDetection() self.vad_threshold = 0.4 self.clients = {} self.websockets = {} diff --git a/whisper_live/vad.py b/whisper_live/vad.py new file mode 100644 index 0000000..e3a6c88 --- /dev/null +++ b/whisper_live/vad.py @@ -0,0 +1,112 @@ +import os +import subprocess +import torch +import numpy as np +import onnxruntime + + +class VoiceActivityDetection(): + + def __init__(self, force_onnx_cpu=True): + path = self.download() + opts = onnxruntime.SessionOptions() + opts.log_severity_level = 3 + + opts.inter_op_num_threads = 1 + opts.intra_op_num_threads = 1 + + if force_onnx_cpu and 'CPUExecutionProvider' in onnxruntime.get_available_providers(): + self.session = onnxruntime.InferenceSession(path, providers=['CPUExecutionProvider'], sess_options=opts) + else: + self.session = onnxruntime.InferenceSession(path, providers=['CUDAExecutionProvider'], sess_options=opts) + + self.reset_states() + self.sample_rates = [8000, 16000] + + def _validate_input(self, x, sr: int): + if x.dim() == 1: + x = x.unsqueeze(0) + if x.dim() > 2: + raise ValueError(f"Too many dimensions for input audio chunk {x.dim()}") + + if sr != 16000 and (sr % 16000 == 0): + step = sr // 16000 + x = x[:,::step] + sr = 16000 + + if sr not in self.sample_rates: + raise ValueError(f"Supported sampling rates: {self.sample_rates} (or multiply of 16000)") + + if sr / x.shape[1] > 31.25: + raise ValueError("Input audio chunk is too short") + + return x, sr + + def reset_states(self, batch_size=1): + self._h = np.zeros((2, batch_size, 64)).astype('float32') + self._c = np.zeros((2, batch_size, 64)).astype('float32') + self._last_sr = 0 + self._last_batch_size = 0 + + def __call__(self, x, sr: int): + + x, sr = self._validate_input(x, sr) + batch_size = x.shape[0] + + if not self._last_batch_size: + self.reset_states(batch_size) + if (self._last_sr) and (self._last_sr != sr): + self.reset_states(batch_size) + if (self._last_batch_size) and (self._last_batch_size != batch_size): + self.reset_states(batch_size) + + if sr in [8000, 16000]: + ort_inputs = {'input': x.numpy(), 'h': self._h, 'c': self._c, 'sr': np.array(sr, dtype='int64')} + ort_outs = self.session.run(None, ort_inputs) + out, self._h, self._c = ort_outs + else: + raise ValueError() + + self._last_sr = sr + self._last_batch_size = batch_size + + out = torch.tensor(out) + return out + + def audio_forward(self, x, sr: int, num_samples: int = 512): + outs = [] + x, sr = self._validate_input(x, sr) + + if x.shape[1] % num_samples: + pad_num = num_samples - (x.shape[1] % num_samples) + x = torch.nn.functional.pad(x, (0, pad_num), 'constant', value=0.0) + + self.reset_states(x.shape[0]) + for i in range(0, x.shape[1], num_samples): + wavs_batch = x[:, i:i+num_samples] + out_chunk = self.__call__(wavs_batch, sr) + outs.append(out_chunk) + + stacked = torch.cat(outs, dim=1) + return stacked.cpu() + + @staticmethod + def download(model_url="https://github.com/snakers4/silero-vad/raw/master/files/silero_vad.onnx"): + target_dir = os.path.expanduser("~/.cache/whisper-live/") + + # Ensure the target directory exists + os.makedirs(target_dir, exist_ok=True) + + # Define the target file path + model_filename = os.path.join(target_dir, "silero_vad.onnx") + + # Check if the model file already exists + if not os.path.exists(model_filename): + # If it doesn't exist, download the model using wget + print("Downloading VAD ONNX model...") + try: + subprocess.run(["wget", "-O", model_filename, model_url], check=True) + except subprocess.CalledProcessError: + print("Failed to download the model using wget.") + return model_filename + From 26b58577471bfd55a2d9120a6b59f6f694cfd471 Mon Sep 17 00:00:00 2001 From: makaveli10 Date: Wed, 27 Sep 2023 15:44:48 +0800 Subject: [PATCH 2/4] add vad model link to original script --- whisper_live/vad.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/whisper_live/vad.py b/whisper_live/vad.py index e3a6c88..c84e5dc 100644 --- a/whisper_live/vad.py +++ b/whisper_live/vad.py @@ -1,3 +1,5 @@ +# original: https://github.com/snakers4/silero-vad/blob/master/utils_vad.py + import os import subprocess import torch From 0444cea0bac56a5bc6cf2de99df748facd64be92 Mon Sep 17 00:00:00 2001 From: makaveli10 Date: Wed, 27 Sep 2023 16:16:23 +0800 Subject: [PATCH 3/4] fix onnxruntime version 1.16 --- requirements/server.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements/server.txt b/requirements/server.txt index 360e392..1c98750 100644 --- a/requirements/server.txt +++ b/requirements/server.txt @@ -4,4 +4,4 @@ faster-whisper==0.6.0 torch==1.10.1 torchaudio==0.10.1 websockets -onnxruntime \ No newline at end of file +onnxruntime==1.16.0 \ No newline at end of file From 5fb6be41bb0196866de988cd2a924533118afe4d Mon Sep 17 00:00:00 2001 From: makaveli10 Date: Wed, 27 Sep 2023 16:17:20 +0800 Subject: [PATCH 4/4] allow only CPU Execution --- whisper_live/vad.py | 7 ++----- 1 file changed, 2 insertions(+), 5 deletions(-) diff --git a/whisper_live/vad.py b/whisper_live/vad.py index c84e5dc..50a406c 100644 --- a/whisper_live/vad.py +++ b/whisper_live/vad.py @@ -9,7 +9,7 @@ import onnxruntime class VoiceActivityDetection(): - def __init__(self, force_onnx_cpu=True): + def __init__(self): path = self.download() opts = onnxruntime.SessionOptions() opts.log_severity_level = 3 @@ -17,10 +17,7 @@ class VoiceActivityDetection(): opts.inter_op_num_threads = 1 opts.intra_op_num_threads = 1 - if force_onnx_cpu and 'CPUExecutionProvider' in onnxruntime.get_available_providers(): - self.session = onnxruntime.InferenceSession(path, providers=['CPUExecutionProvider'], sess_options=opts) - else: - self.session = onnxruntime.InferenceSession(path, providers=['CUDAExecutionProvider'], sess_options=opts) + self.session = onnxruntime.InferenceSession(path, providers=['CPUExecutionProvider'], sess_options=opts) self.reset_states() self.sample_rates = [8000, 16000]