Skip to content

Commit 33b0403

Browse files
LiangRx0501gemini-code-assist[bot]LauraGPT
authored
fix: disable duplicate dynamic silence logic in DynamicStreamingVAD (#3240)
* fix: forward silence_schedule to streaming FSMN-VAD * fix: forward silence_schedule to streaming FSMN-VAD * style: remove debug comment * fix: disable duplicate dynamic silence logic in DynamicStreamingVAD Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> * fix(vad): initialize dynamic thresholds on first feed Pass the wrapper silence and noise thresholds into the initial FSMN-VAD cache so first-call behavior does not depend on caller chunking. Add a regression that exercises the wrapper with the production cache initializer. Assisted-by: Codex:gpt-5.6 --------- Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> Co-authored-by: zhifu gao <zhifu.gzf@alibaba-inc.com>
1 parent 1d8080a commit 33b0403

3 files changed

Lines changed: 97 additions & 1 deletion

File tree

funasr/models/fsmn_vad_streaming/dynamic_vad.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -130,10 +130,18 @@ def feed(self, audio_chunk: torch.Tensor, is_final: bool = False) -> List[List[i
130130
self.accumulated_since_cut_ms += int(chunk_samples * 1000 / self.sample_rate)
131131

132132
self._apply_dynamic_threshold()
133+
initial_cache_kwargs = {}
134+
if "stats" not in self.cache:
135+
initial_cache_kwargs = {
136+
"max_end_silence_time": self._get_silence_threshold(),
137+
"speech_noise_thres": self.speech_noise_thres,
138+
}
133139

134140
res = self.model.generate(
135141
input=[audio_chunk], cache=self.cache,
136142
is_final=is_final, chunk_size=self.chunk_size_ms,
143+
dynamic_silence=False,
144+
**initial_cache_kwargs,
137145
)
138146

139147
signals = res[0].get("value", [])

funasr/models/fsmn_vad_streaming/model.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -871,7 +871,9 @@ def init_cache(self, cache: dict = None, **kwargs):
871871
sil_pdf_ids=self.vad_opts.sil_pdf_ids,
872872
max_end_sil_frame_cnt_thresh=self.vad_opts.max_end_silence_time
873873
- self.vad_opts.speech_to_sil_time_thres,
874-
speech_noise_thres=self.vad_opts.speech_noise_thres,
874+
speech_noise_thres=kwargs.get(
875+
"speech_noise_thres", self.vad_opts.speech_noise_thres
876+
),
875877
)
876878
cache["windows_detector"] = windows_detector
877879
cache["stats"] = stats
Lines changed: 86 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,86 @@
1+
import unittest
2+
from types import SimpleNamespace
3+
4+
import torch
5+
6+
from funasr.models.fsmn_vad_streaming import model as vad_model
7+
from funasr.models.fsmn_vad_streaming.dynamic_vad import DynamicStreamingVAD
8+
9+
10+
class _ThresholdAwareModel:
11+
"""Small AutoModel stand-in that uses the production cache initializer."""
12+
13+
sample_rate = 16000
14+
15+
def __init__(self):
16+
self.model = vad_model.FsmnVADStreaming.__new__(vad_model.FsmnVADStreaming)
17+
self.model.vad_opts = SimpleNamespace(
18+
window_size_ms=200,
19+
sil_to_speech_time_thres=150,
20+
speech_to_sil_time_thres=150,
21+
frame_in_ms=10,
22+
sil_pdf_ids=[0],
23+
max_end_silence_time=800,
24+
speech_noise_thres=0.5,
25+
)
26+
27+
def generate(self, input, cache, **kwargs):
28+
if not cache:
29+
self.model.init_cache(cache, **kwargs)
30+
31+
audio = torch.cat((cache.get("_test_audio", torch.empty(0)), input[0]))
32+
cache["_test_audio"] = audio
33+
34+
speech_indices = torch.nonzero(audio.abs() > 0.5)
35+
if not len(speech_indices) or cache.get("_test_emitted"):
36+
return [{"value": []}]
37+
38+
last_speech_sample = speech_indices[-1].item()
39+
trailing_silence_ms = int(
40+
(len(audio) - last_speech_sample - 1) * 1000 / self.sample_rate
41+
)
42+
stats = cache["stats"]
43+
threshold_ms = (
44+
stats.max_end_sil_frame_cnt_thresh
45+
+ self.model.vad_opts.speech_to_sil_time_thres
46+
)
47+
if trailing_silence_ms < threshold_ms:
48+
return [{"value": []}]
49+
50+
cache["_test_emitted"] = True
51+
speech_end_ms = int((last_speech_sample + 1) * 1000 / self.sample_rate)
52+
return [{"value": [[0, speech_end_ms]]}]
53+
54+
55+
class TestDynamicStreamingVadFirstCall(unittest.TestCase):
56+
def _new_vad(self):
57+
return DynamicStreamingVAD(
58+
_ThresholdAwareModel(),
59+
silence_schedule=[(float("inf"), 10000)],
60+
speech_noise_thres=0.73,
61+
)
62+
63+
def test_first_feed_initializes_wrapper_thresholds(self):
64+
vad = self._new_vad()
65+
66+
vad.feed(torch.ones(960))
67+
68+
self.assertEqual(vad.cache["stats"].max_end_sil_frame_cnt_thresh, 9850)
69+
self.assertAlmostEqual(vad.cache["stats"].speech_noise_thres, 0.73)
70+
71+
def test_first_feed_is_chunking_invariant(self):
72+
audio = torch.cat((torch.ones(16000), torch.zeros(32000)))
73+
74+
one_chunk_vad = self._new_vad()
75+
one_chunk_segments = one_chunk_vad.feed(audio)
76+
77+
split_vad = self._new_vad()
78+
split_segments = split_vad.feed(audio[:960])
79+
split_segments.extend(split_vad.feed(audio[960:]))
80+
81+
self.assertEqual(one_chunk_segments, split_segments)
82+
self.assertEqual(one_chunk_segments, [])
83+
84+
85+
if __name__ == "__main__":
86+
unittest.main()

0 commit comments

Comments
 (0)