Skip to content

Commit af8574c

Browse files
authored
Merge pull request #3655 from modelscope/codex/kws-optional-output-20260907
fix(kws): honor optional output directory on every call
2 parents f85d8f5 + beb92d9 commit af8574c

6 files changed

Lines changed: 253 additions & 10 deletions

File tree

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
name: Validate KWS output handling
2+
3+
on:
4+
pull_request:
5+
paths:
6+
- "funasr/models/fsmn_kws/**"
7+
- "funasr/models/fsmn_kws_mt/**"
8+
- "funasr/models/sanm_kws/**"
9+
- "funasr/models/sanm_kws_streaming/**"
10+
- "funasr/utils/datadir_writer.py"
11+
- "funasr/utils/kws_utils.py"
12+
- "tests/test_kws_optional_output.py"
13+
- ".github/workflows/test-kws-output.yml"
14+
push:
15+
branches: [main]
16+
paths:
17+
- "funasr/models/fsmn_kws/**"
18+
- "funasr/models/fsmn_kws_mt/**"
19+
- "funasr/models/sanm_kws/**"
20+
- "funasr/models/sanm_kws_streaming/**"
21+
- "funasr/utils/datadir_writer.py"
22+
- "funasr/utils/kws_utils.py"
23+
- "tests/test_kws_optional_output.py"
24+
- ".github/workflows/test-kws-output.yml"
25+
26+
permissions:
27+
contents: read
28+
29+
jobs:
30+
kws-output:
31+
runs-on: ubuntu-latest
32+
timeout-minutes: 15
33+
steps:
34+
- uses: actions/checkout@v4
35+
- uses: actions/setup-python@v5
36+
with:
37+
python-version: "3.11"
38+
cache: pip
39+
- name: Install CPU test dependencies
40+
run: |
41+
python -m pip install torch --index-url https://download.pytorch.org/whl/cpu
42+
python -m pip install "numpy<2" "kaldiio>=2.17.0" librosa "rapidfuzz>=3.0.0" six
43+
- name: Check KWS results and optional file output without model weights
44+
run: python -m unittest discover -s tests -p test_kws_optional_output.py -v

‎funasr/models/fsmn_kws/model.py‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -263,12 +263,13 @@ def inference(
263263
is_deted, det_keyword, det_score = detect_result[0], detect_result[1], detect_result[2]
264264

265265
if is_deted:
266-
self.writer["detect"][key[i]] = "detected " + det_keyword + " " + str(det_score)
267266
det_info = "detected " + det_keyword + " " + str(det_score)
268267
else:
269-
self.writer["detect"][key[i]] = "rejected"
270268
det_info = "rejected"
271269

270+
if kwargs.get("output_dir") is not None:
271+
self.writer["detect"][key[i]] = det_info
272+
272273
result_i = {"key": key[i], "text": det_info}
273274
results.append(result_i)
274275

‎funasr/models/fsmn_kws_mt/model.py‎

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -325,23 +325,25 @@ def inference(
325325
is_deted, det_keyword, det_score = detect_result[0], detect_result[1], detect_result[2]
326326

327327
if is_deted:
328-
self.writer["detect"][key[i]] = "detected " + det_keyword + " " + str(det_score)
329328
det_info = "detected " + det_keyword + " " + str(det_score)
330329
else:
331-
self.writer["detect"][key[i]] = "rejected"
332330
det_info = "rejected"
333331

332+
if kwargs.get("output_dir") is not None:
333+
self.writer["detect"][key[i]] = det_info
334+
334335
x2 = encoder_out2[i, :encoder_out_lens[i], :]
335336
detect_result2 = self.kws_decoder2.decode(x2)
336337
is_deted2, det_keyword2, det_score2 = detect_result2[0], detect_result2[1], detect_result2[2]
337338

338339
if is_deted2:
339-
self.writer["detect2"][key[i]] = "detected " + det_keyword2 + " " + str(det_score2)
340340
det_info2 = "detected " + det_keyword2 + " " + str(det_score2)
341341
else:
342-
self.writer["detect2"][key[i]] = "rejected"
343342
det_info2 = "rejected"
344343

344+
if kwargs.get("output_dir") is not None:
345+
self.writer["detect2"][key[i]] = det_info2
346+
345347
result_i = {"key": key[i], "text": det_info, "text2": det_info2}
346348
results.append(result_i)
347349

‎funasr/models/sanm_kws/model.py‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -286,12 +286,13 @@ def inference(
286286
is_deted, det_keyword, det_score = detect_result[0], detect_result[1], detect_result[2]
287287

288288
if is_deted:
289-
self.writer["detect"][key[i]] = "detected " + det_keyword + " " + str(det_score)
290289
det_info = "detected " + det_keyword + " " + str(det_score)
291290
else:
292-
self.writer["detect"][key[i]] = "rejected"
293291
det_info = "rejected"
294292

293+
if kwargs.get("output_dir") is not None:
294+
self.writer["detect"][key[i]] = det_info
295+
295296
result_i = {"key": key[i], "text": det_info}
296297
results.append(result_i)
297298

‎funasr/models/sanm_kws_streaming/model.py‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -271,12 +271,13 @@ def generate_chunk(
271271
is_deted, det_keyword, det_score = detect_result[0], detect_result[1], detect_result[2]
272272

273273
if is_deted:
274-
self.writer["detect"][key[i]] = "detected " + det_keyword + " " + str(det_score)
275274
det_info = "detected " + det_keyword + " " + str(det_score)
276275
else:
277-
self.writer["detect"][key[i]] = "rejected"
278276
det_info = "rejected"
279277

278+
if kwargs.get("output_dir") is not None:
279+
self.writer["detect"][key[i]] = det_info
280+
280281
result_i = {"key": key[i], "text": det_info}
281282
results.append(result_i)
282283

‎tests/test_kws_optional_output.py‎

Lines changed: 194 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,194 @@
1+
"""Exercise production KWS output handling without loading acoustic models."""
2+
3+
import itertools
4+
import tempfile
5+
import unittest
6+
from pathlib import Path
7+
from types import SimpleNamespace
8+
from unittest.mock import patch
9+
10+
import torch
11+
12+
from funasr.models.fsmn_kws.model import FsmnKWS
13+
from funasr.models.fsmn_kws_mt.model import FsmnKWSMT
14+
from funasr.models.sanm_kws.model import SanmKWS
15+
from funasr.models.sanm_kws_streaming.model import SanmKWSStreaming
16+
17+
18+
class FixedDecoder:
19+
def __init__(self, ctc, keywords=None, token_list=None, seg_dict=None):
20+
self.result = ctc
21+
22+
def decode(self, encoder_out):
23+
return self.result
24+
25+
26+
VARIANTS = (FsmnKWS, FsmnKWSMT, SanmKWS, SanmKWSStreaming)
27+
OMITTED = object()
28+
29+
30+
class KwsOptionalOutputTest(unittest.TestCase):
31+
def setUp(self):
32+
self.directory = tempfile.TemporaryDirectory()
33+
self.addCleanup(self.directory.cleanup)
34+
self.patch_decoder = patch(
35+
"funasr.utils.kws_utils.KwsCtcPrefixDecoder", FixedDecoder
36+
)
37+
self.patch_decoder.start()
38+
self.addCleanup(self.patch_decoder.stop)
39+
40+
def model(self, variant, detected=True, detected2=False):
41+
result = (detected, "wake", 0.9)
42+
model = SimpleNamespace(ctc=result, ctc2=(detected2, "hello", 0.8))
43+
model.encode = lambda speech, lengths: (speech, lengths)
44+
if variant is FsmnKWSMT:
45+
model.encode = lambda speech, lengths: (speech, speech, lengths)
46+
model.encode_chunk = lambda speech, lengths, **kwargs: (speech, lengths)
47+
model.kws_decoder = FixedDecoder(result)
48+
self.addCleanup(
49+
lambda: model.writer.close() if hasattr(model, "writer") else None
50+
)
51+
return model
52+
53+
def call(
54+
self, variant, model, output_dir=OMITTED, key="sample", final=True, cache=None
55+
):
56+
kwargs = {"device": "cpu", "data_type": "fbank", "keywords": "wake"}
57+
if output_dir is not OMITTED:
58+
kwargs["output_dir"] = output_dir
59+
speech = torch.zeros(1, 3, 1)
60+
lengths = torch.tensor([3])
61+
tokenizer = SimpleNamespace(token_list=["wake"], seg_dict={})
62+
if variant is FsmnKWSMT:
63+
tokenizer = [tokenizer, tokenizer]
64+
if variant is SanmKWSStreaming:
65+
if cache is None:
66+
cache = {
67+
"encoder": {
68+
"chunk_size": [0, 3, 0],
69+
"encoder_out": None,
70+
"encoder_out_lens": None,
71+
}
72+
}
73+
return variant.generate_chunk(
74+
model,
75+
speech,
76+
lengths,
77+
key=[key],
78+
tokenizer=tokenizer,
79+
cache=cache,
80+
is_final=final,
81+
**kwargs,
82+
)
83+
results, _ = variant.inference(
84+
model,
85+
speech,
86+
data_lengths=lengths[:, None],
87+
key=[key],
88+
tokenizer=tokenizer,
89+
**kwargs,
90+
)
91+
return results
92+
93+
def expected(self, variant, detected=True, detected2=False, key="sample"):
94+
result = {"key": key, "text": "detected wake 0.9" if detected else "rejected"}
95+
if variant is FsmnKWSMT:
96+
result["text2"] = "detected hello 0.8" if detected2 else "rejected"
97+
return [result]
98+
99+
def test_omitted_and_none_return_results_without_creating_writer(self):
100+
for variant, detected, detected2, output_dir in itertools.product(
101+
VARIANTS, (False, True), (False, True), (OMITTED, None)
102+
):
103+
with self.subTest(
104+
variant=variant.__name__,
105+
detected=detected,
106+
detected2=detected2,
107+
output_dir=output_dir,
108+
):
109+
model = self.model(variant, detected, detected2)
110+
self.assertEqual(
111+
self.call(variant, model, output_dir),
112+
self.expected(variant, detected, detected2),
113+
)
114+
self.assertFalse(hasattr(model, "writer"))
115+
116+
def test_enabled_output_preserves_result_and_file_format(self):
117+
for variant, detected, detected2 in itertools.product(
118+
VARIANTS, (False, True), (False, True)
119+
):
120+
with self.subTest(
121+
variant=variant.__name__, detected=detected, detected2=detected2
122+
):
123+
path = (
124+
Path(self.directory.name)
125+
/ f"{variant.__name__}-{detected}-{detected2}"
126+
)
127+
model = self.model(variant, detected, detected2)
128+
expected = self.expected(variant, detected, detected2)
129+
self.assertEqual(self.call(variant, model, str(path)), expected)
130+
self.assertEqual(
131+
(path / "detect").read_text(), f"sample {expected[0]['text']}\n"
132+
)
133+
if variant is FsmnKWSMT:
134+
self.assertEqual(
135+
(path / "detect2").read_text(),
136+
f"sample {expected[0]['text2']}\n",
137+
)
138+
139+
def test_disabled_output_does_not_reuse_cached_writer(self):
140+
for variant, output_dir in itertools.product(VARIANTS, (OMITTED, None)):
141+
with self.subTest(variant=variant.__name__, output_dir=output_dir):
142+
path = (
143+
Path(self.directory.name)
144+
/ f"{variant.__name__}-{output_dir is None}"
145+
)
146+
model = self.model(variant)
147+
self.call(variant, model, str(path), key="first")
148+
before = {p.name: p.read_bytes() for p in path.iterdir()}
149+
self.assertEqual(
150+
self.call(variant, model, output_dir, key="second"),
151+
self.expected(variant, key="second"),
152+
)
153+
self.assertEqual(
154+
{p.name: p.read_bytes() for p in path.iterdir()}, before
155+
)
156+
157+
def test_consecutive_enabled_calls_append(self):
158+
for variant in VARIANTS:
159+
with self.subTest(variant=variant.__name__):
160+
path = Path(self.directory.name) / variant.__name__
161+
model = self.model(variant)
162+
for key in ("first", "second"):
163+
self.call(variant, model, str(path), key=key)
164+
self.assertEqual(
165+
(path / "detect").read_text(),
166+
"first detected wake 0.9\nsecond detected wake 0.9\n",
167+
)
168+
if variant is FsmnKWSMT:
169+
self.assertEqual(
170+
(path / "detect2").read_text(),
171+
"first rejected\nsecond rejected\n",
172+
)
173+
174+
def test_streaming_nonfinal_accumulates_and_final_returns_without_output(self):
175+
model = self.model(SanmKWSStreaming)
176+
cache = {
177+
"encoder": {
178+
"chunk_size": [0, 3, 0],
179+
"encoder_out": None,
180+
"encoder_out_lens": None,
181+
}
182+
}
183+
self.assertIsNone(self.call(SanmKWSStreaming, model, final=False, cache=cache))
184+
self.assertFalse(hasattr(model, "writer"))
185+
self.assertEqual(
186+
self.call(SanmKWSStreaming, model, final=True, cache=cache),
187+
self.expected(SanmKWSStreaming),
188+
)
189+
self.assertEqual(cache["encoder"]["encoder_out"].shape[1], 6)
190+
self.assertFalse(hasattr(model, "writer"))
191+
192+
193+
if __name__ == "__main__":
194+
unittest.main()

0 commit comments

Comments
 (0)