|
| 1 | +"""Adapter that makes Silero VAD return FunASR-compatible millisecond segments.""" |
| 2 | + |
| 3 | +import time |
| 4 | + |
| 5 | +import torch |
| 6 | + |
| 7 | +from funasr.register import tables |
| 8 | +from funasr.utils.load_utils import load_audio_text_image_video |
| 9 | + |
| 10 | + |
| 11 | +@tables.register("model_classes", "SileroVad") |
| 12 | +class SileroVad(torch.nn.Module): |
| 13 | + """Offline Silero VAD adapter used by ``AutoModel(vad_model='silero-vad')``. |
| 14 | +
|
| 15 | + Requires the official ``silero-vad`` Python package. |
| 16 | + """ |
| 17 | + |
| 18 | + def __init__(self, **kwargs): |
| 19 | + super().__init__() |
| 20 | + self.anchor = torch.nn.Parameter(torch.empty(0), requires_grad=False) |
| 21 | + try: |
| 22 | + from silero_vad import get_speech_timestamps, load_silero_vad |
| 23 | + except ImportError as error: |
| 24 | + raise ImportError( |
| 25 | + "Silero VAD requires the optional dependency. Install it with " |
| 26 | + '`python -m pip install "funasr[silero]"` or ' |
| 27 | + "`python -m pip install silero-vad`." |
| 28 | + ) from error |
| 29 | + self.model = load_silero_vad(onnx=kwargs.get("silero_onnx", False)) |
| 30 | + self.get_speech_timestamps = get_speech_timestamps |
| 31 | + |
| 32 | + @staticmethod |
| 33 | + def _split_long_segments(segments, max_single_segment_time): |
| 34 | + if not max_single_segment_time: |
| 35 | + return segments |
| 36 | + limit_ms = int(max_single_segment_time) |
| 37 | + split = [] |
| 38 | + for start, end in segments: |
| 39 | + while end - start > limit_ms: |
| 40 | + split.append([start, start + limit_ms]) |
| 41 | + start += limit_ms |
| 42 | + split.append([start, end]) |
| 43 | + return split |
| 44 | + |
| 45 | + def inference(self, data_in, key=None, **kwargs): |
| 46 | + sample_rate = int(kwargs.get("silero_sampling_rate", 16000)) |
| 47 | + if sample_rate not in (8000, 16000): |
| 48 | + raise ValueError("Silero VAD supports silero_sampling_rate=8000 or 16000") |
| 49 | + audio_list = load_audio_text_image_video( |
| 50 | + data_in, |
| 51 | + fs=sample_rate, |
| 52 | + audio_fs=kwargs.get("fs", sample_rate), |
| 53 | + data_type=kwargs.get("data_type", "sound"), |
| 54 | + ) |
| 55 | + if not isinstance(audio_list, list): |
| 56 | + audio_list = [audio_list] |
| 57 | + |
| 58 | + started = time.perf_counter() |
| 59 | + results = [] |
| 60 | + for index, audio in enumerate(audio_list): |
| 61 | + waveform = torch.as_tensor(audio, dtype=torch.float32).flatten().cpu() |
| 62 | + timestamps = self.get_speech_timestamps( |
| 63 | + waveform, |
| 64 | + self.model, |
| 65 | + sampling_rate=sample_rate, |
| 66 | + threshold=kwargs.get("silero_threshold", 0.5), |
| 67 | + min_speech_duration_ms=kwargs.get("silero_min_speech_duration_ms", 250), |
| 68 | + min_silence_duration_ms=kwargs.get("silero_min_silence_duration_ms", 100), |
| 69 | + speech_pad_ms=kwargs.get("silero_speech_pad_ms", 30), |
| 70 | + ) |
| 71 | + segments = [ |
| 72 | + [int(item["start"] * 1000 / sample_rate), int(item["end"] * 1000 / sample_rate)] |
| 73 | + for item in timestamps |
| 74 | + ] |
| 75 | + segments = self._split_long_segments( |
| 76 | + segments, kwargs.get("max_single_segment_time") |
| 77 | + ) |
| 78 | + results.append({"key": key[index] if key else str(index), "value": segments}) |
| 79 | + elapsed = time.perf_counter() - started |
| 80 | + total_samples = sum(len(torch.as_tensor(audio)) for audio in audio_list) |
| 81 | + return results, {"batch_data_time": total_samples / sample_rate, "forward": elapsed} |
0 commit comments