|
| 1 | +import unittest |
| 2 | +from types import SimpleNamespace |
| 3 | +from unittest.mock import patch |
| 4 | + |
| 5 | +import torch |
| 6 | + |
| 7 | +from funasr.models.fsmn_vad_streaming import model as vad_model |
| 8 | + |
| 9 | + |
| 10 | +class TestFsmnVadDynamicSilence(unittest.TestCase): |
| 11 | + def _run_inference(self, **kwargs): |
| 12 | + vad = vad_model.FsmnVADStreaming.__new__(vad_model.FsmnVADStreaming) |
| 13 | + vad.vad_opts = SimpleNamespace(speech_to_sil_time_thres=100) |
| 14 | + vad.forward = lambda **batch: [] |
| 15 | + |
| 16 | + cache = { |
| 17 | + "frontend": {}, |
| 18 | + "prev_samples": torch.empty(0), |
| 19 | + "encoder": {}, |
| 20 | + "stats": SimpleNamespace( |
| 21 | + vad_state_machine=vad_model.VadStateMachine.kVadInStateInSpeechSegment, |
| 22 | + max_end_sil_frame_cnt_thresh=200, |
| 23 | + speech_noise_thres=0.6, |
| 24 | + ), |
| 25 | + } |
| 26 | + frontend = SimpleNamespace(fs=16000, frame_shift=10, lfr_n=1) |
| 27 | + |
| 28 | + def fake_extract_fbank(*args, **kwargs): |
| 29 | + cache["frontend"]["waveforms"] = torch.zeros(1, 16000) |
| 30 | + return torch.zeros(1, 1, 80), torch.tensor([100]) |
| 31 | + |
| 32 | + with ( |
| 33 | + patch.object( |
| 34 | + vad_model, |
| 35 | + "load_audio_text_image_video", |
| 36 | + return_value=[torch.zeros(16000)], |
| 37 | + ), |
| 38 | + patch.object(vad_model, "extract_fbank", side_effect=fake_extract_fbank), |
| 39 | + ): |
| 40 | + vad_model.FsmnVADStreaming.inference( |
| 41 | + vad, |
| 42 | + torch.zeros(16000), |
| 43 | + frontend=frontend, |
| 44 | + cache=cache, |
| 45 | + key=["utt"], |
| 46 | + chunk_size=1000, |
| 47 | + is_final=False, |
| 48 | + device="cpu", |
| 49 | + **kwargs, |
| 50 | + ) |
| 51 | + |
| 52 | + return cache |
| 53 | + |
| 54 | + def test_explicit_max_end_silence_time_keeps_fixed_threshold_by_default(self): |
| 55 | + cache = self._run_inference(max_end_silence_time=300) |
| 56 | + |
| 57 | + self.assertEqual(cache["stats"].max_end_sil_frame_cnt_thresh, 200) |
| 58 | + self.assertNotIn("_dynamic_accumulated_ms", cache) |
| 59 | + |
| 60 | + def test_explicit_dynamic_silence_still_enables_schedule(self): |
| 61 | + cache = self._run_inference(max_end_silence_time=300, dynamic_silence=True) |
| 62 | + |
| 63 | + self.assertEqual(cache["stats"].max_end_sil_frame_cnt_thresh, 1900) |
| 64 | + self.assertEqual(cache["_dynamic_accumulated_ms"], 1000) |
| 65 | + |
| 66 | + |
| 67 | +if __name__ == "__main__": |
| 68 | + unittest.main() |
0 commit comments