Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 25 additions & 3 deletions funasr/models/fsmn_vad_streaming/dynamic_vad.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 中查询静音阈值。"""
Expand Down Expand Up @@ -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 = {}
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
161 changes: 161 additions & 0 deletions tests/test_dynamic_streaming_vad.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Loading