From 2df180c636b1e543ee9d87c178f66828679d7292 Mon Sep 17 00:00:00 2001 From: zhifu gao Date: Tue, 8 Sep 2026 20:15:17 +0000 Subject: [PATCH] fix(vad): exclude idle time from dynamic silence schedule Signed-off-by: zhifu gao --- .../models/fsmn_vad_streaming/dynamic_vad.py | 28 ++- tests/test_dynamic_streaming_vad.py | 161 ++++++++++++++++++ 2 files changed, 186 insertions(+), 3 deletions(-) diff --git a/funasr/models/fsmn_vad_streaming/dynamic_vad.py b/funasr/models/fsmn_vad_streaming/dynamic_vad.py index f880fd2f82..944b60b9ff 100644 --- a/funasr/models/fsmn_vad_streaming/dynamic_vad.py +++ b/funasr/models/fsmn_vad_streaming/dynamic_vad.py @@ -94,6 +94,18 @@ def __init__( self.confirmed_segments: List[List[int]] = [] self.current_speech_start: Optional[int] = None self.accumulated_since_cut_ms: int = 0 + self._total_audio_samples: int = 0 + self._needs_stream_reset: bool = False + + def _update_speech_duration(self): + if self.current_speech_start is None: + self.accumulated_since_cut_ms = 0 + else: + # VAD start signals use the session clock, including preceding silence. + audio_end_ms = self._total_audio_samples * 1000 // self.sample_rate + self.accumulated_since_cut_ms = max( + 0, audio_end_ms - int(self.current_speech_start) + ) def _get_silence_threshold(self) -> int: """根据当前累积时长,从 schedule 中查询静音阈值。""" @@ -123,11 +135,14 @@ def feed(self, audio_chunk: torch.Tensor, is_final: bool = False) -> List[List[i 新确认的语音段列表,每段为 [start_ms, end_ms]。 仅在检测到语音结束时返回非空列表。 """ + if self._needs_stream_reset: + self.reset() if audio_chunk.dim() > 1: audio_chunk = audio_chunk.squeeze() chunk_samples = len(audio_chunk) - self.accumulated_since_cut_ms += int(chunk_samples * 1000 / self.sample_rate) + self._total_audio_samples += chunk_samples + self._update_speech_duration() self._apply_dynamic_threshold() initial_cache_kwargs = {} @@ -163,17 +178,22 @@ def feed(self, audio_chunk: torch.Tensor, is_final: bool = False) -> List[List[i self.current_speech_start = None self.accumulated_since_cut_ms = 0 + self._update_speech_duration() + # Preserve final state for callers' tail fallback; reset before the next stream. + self._needs_stream_reset = is_final return new_confirmed def finalize(self) -> List[List[int]]: """结束流式处理,返回最后可能未结束的语音段。 - 调用此方法后,VAD 状态会被重置。 - 如果当前有正在进行的语音段,会被强制结束。 + 本次状态保留到下一次 feed,便于调用方补齐未返回结束信号的尾段。 + 下一次 feed 会重置状态,开始新的流。 Returns: 最后确认的语音段列表。 """ + if self._needs_stream_reset: + return [] # Feed empty with is_final=True to flush empty = torch.zeros(int(self.sample_rate * 0.01), dtype=torch.float32) return self.feed(empty, is_final=True) @@ -228,3 +248,5 @@ def reset(self): self.confirmed_segments = [] self.current_speech_start = None self.accumulated_since_cut_ms = 0 + self._total_audio_samples = 0 + self._needs_stream_reset = False diff --git a/tests/test_dynamic_streaming_vad.py b/tests/test_dynamic_streaming_vad.py index 00f256f273..c4ae399957 100644 --- a/tests/test_dynamic_streaming_vad.py +++ b/tests/test_dynamic_streaming_vad.py @@ -82,5 +82,166 @@ def test_first_feed_is_chunking_invariant(self): self.assertEqual(one_chunk_segments, []) +class _SignalModel(_ThresholdAwareModel): + def __init__(self, signals): + super().__init__() + self.signals = iter(signals) + self.thresholds = [] + + def generate(self, input, cache, **kwargs): + if not cache: + self.model.init_cache(cache, **kwargs) + self.thresholds.append( + cache["stats"].max_end_sil_frame_cnt_thresh + + self.model.vad_opts.speech_to_sil_time_thres + ) + return [{"value": next(self.signals, [])}] + + +class TestDynamicStreamingVadSpeechDuration(unittest.TestCase): + def test_initial_silence_does_not_advance_speech_schedule(self): + vad = DynamicStreamingVAD(_SignalModel([[]])) + + vad.feed(torch.zeros(60 * 16000)) + + self.assertEqual(vad.current_duration_ms, 0) + self.assertEqual(vad.current_threshold_ms, 2000) + + def test_speech_after_initial_silence_uses_short_utterance_threshold(self): + model = _SignalModel([[], [[60000, -1]]]) + vad = DynamicStreamingVAD(model) + vad.feed(torch.zeros(60 * 16000)) + + vad.feed(torch.ones(16000)) + + self.assertEqual(model.thresholds, [2000, 2000]) + self.assertEqual(vad.current_duration_ms, 1000) + self.assertEqual(vad.current_threshold_ms, 2000) + + def test_inter_utterance_silence_does_not_age_next_utterance(self): + model = _SignalModel([[[0, -1]], [[-1, 1000]], [], [[22000, -1]]]) + vad = DynamicStreamingVAD(model) + vad.feed(torch.ones(16000)) + vad.feed(torch.zeros(2 * 16000)) + vad.feed(torch.zeros(19 * 16000)) + + vad.feed(torch.ones(16000)) + + self.assertEqual(model.thresholds[-1], 2000) + self.assertEqual(vad.current_duration_ms, 1000) + + def test_duration_uses_detected_start_and_keeps_long_speech_schedule(self): + vad = DynamicStreamingVAD(_SignalModel([[], [[6000, -1]], []])) + vad.feed(torch.zeros(5 * 16000)) + vad.feed(torch.ones(2 * 16000)) + self.assertEqual(vad.current_duration_ms, 1000) + + vad.feed(torch.ones(6 * 16000)) + + self.assertEqual(vad.current_duration_ms, 7000) + self.assertEqual(vad.current_threshold_ms, 1500) + + def test_submillisecond_packets_do_not_lose_elapsed_samples(self): + vad = DynamicStreamingVAD(_SignalModel([[[0, -1]], [], []])) + for count in (1, 15, 160): + vad.feed(torch.ones(count)) + + self.assertEqual(vad.current_duration_ms, 11) + + def test_last_open_signal_in_packet_defines_current_duration(self): + vad = DynamicStreamingVAD(_SignalModel([[[0, 1000], [2000, -1]]])) + + self.assertEqual(vad.feed(torch.ones(10 * 16000)), [[0, 1000]]) + + self.assertEqual(vad.current_duration_ms, 8000) + self.assertEqual(vad.current_threshold_ms, 1500) + + def test_reset_starts_a_new_sample_clock(self): + vad = DynamicStreamingVAD(_SignalModel([[], [[0, -1]]])) + vad.feed(torch.zeros(60 * 16000)) + vad.reset() + + vad.feed(torch.ones(16000)) + + self.assertEqual(vad.current_duration_ms, 1000) + self.assertEqual(vad.current_threshold_ms, 2000) + + def test_backdated_start_in_previous_packet_counts_entire_segment(self): + vad = DynamicStreamingVAD(_SignalModel([[], [[4800, -1]]])) + vad.feed(torch.zeros(5 * 16000)) + + vad.feed(torch.ones(16000)) + + self.assertEqual(vad.current_duration_ms, 1200) + + def test_start_end_start_in_one_packet_uses_last_start(self): + vad = DynamicStreamingVAD( + _SignalModel([[[1000, -1], [-1, 2000], [3000, -1]]]) + ) + + self.assertEqual(vad.feed(torch.ones(10 * 16000)), [[1000, 2000]]) + self.assertEqual(vad.current_speech_start, 3000) + self.assertEqual(vad.current_duration_ms, 7000) + + def test_finalize_without_end_preserves_start_for_service_fallback(self): + vad = DynamicStreamingVAD(_SignalModel([[[500, -1]], []])) + vad.feed(torch.ones(16000)) + + self.assertEqual(vad.finalize(), []) + + self.assertEqual(vad.current_speech_start, 500) + self.assertTrue(vad.is_speaking) + + def test_feed_after_finalize_starts_a_new_stream(self): + model = _SignalModel([[], [], [[0, -1]]]) + vad = DynamicStreamingVAD(model) + vad.feed(torch.zeros(60 * 16000)) + vad.finalize() + + vad.feed(torch.ones(16000)) + + self.assertEqual(vad.current_duration_ms, 1000) + self.assertEqual(model.thresholds[-1], 2000) + + def test_feed_after_direct_final_clears_previous_stream(self): + model = _SignalModel([[[0, 1000], [2000, -1]], [[0, -1]]]) + vad = DynamicStreamingVAD(model) + self.assertEqual(vad.feed(torch.ones(60 * 16000), is_final=True), [[0, 1000]]) + self.assertEqual(vad.current_speech_start, 2000) + old_cache = vad.cache + + vad.feed(torch.ones(16000)) + + self.assertEqual(vad.current_duration_ms, 1000) + self.assertEqual(vad.confirmed_segments, []) + self.assertIsNot(vad.cache, old_cache) + self.assertEqual(model.thresholds[-1], 2000) + + def test_repeated_finalize_preserves_final_state_without_feeding_again(self): + model = _SignalModel([[[0, 100], [500, -1]], []]) + vad = DynamicStreamingVAD(model) + vad.feed(torch.ones(16000)) + vad.finalize() + final_cache = vad.cache + + self.assertEqual(vad.finalize(), []) + + self.assertEqual(vad.current_speech_start, 500) + self.assertEqual(vad.confirmed_segments, [[0, 100]]) + self.assertIs(vad.cache, final_cache) + self.assertEqual(len(model.thresholds), 2) + + def test_finalize_after_process_preserves_results_and_fallback_state(self): + model = _SignalModel([[[0, 30]], [[80, -1]]]) + vad = DynamicStreamingVAD(model) + self.assertEqual(vad.process(torch.ones(1920)), [[0, 30]]) + + self.assertEqual(vad.finalize(), []) + + self.assertEqual(vad.confirmed_segments, [[0, 30]]) + self.assertEqual(vad.current_speech_start, 80) + self.assertEqual(len(model.thresholds), 2) + + if __name__ == "__main__": unittest.main()