From d0b9cc97faaf7d98b823d7cce39a871a57424b4e Mon Sep 17 00:00:00 2001 From: Lesereingrape <90967079+Lesereingrape@users.noreply.github.com> Date: Sun, 4 Oct 2026 02:27:30 +0800 Subject: [PATCH] fix(frontend): read cmvn statistics without crashing on blank lines Both package cmvn readers index line_item[0] and lines[i + 1] without checking, so a blank or whitespace-only line anywhere in the file (trailing newline at EOF, a separator between sections) aborts frontend construction with IndexError, and a file whose tags the walker does not recognize returns empty statistics that only fail much later at inference. Skip blank lines, bound the lookahead, and raise a ValueError naming the file when neither nor parsed. --- funasr/frontends/default.py | 12 ++++--- funasr/frontends/wav_frontend.py | 12 ++++--- tests/test_numpy_compatibility.py | 56 +++++++++++++++++++++++++++++++ 3 files changed, 70 insertions(+), 10 deletions(-) diff --git a/funasr/frontends/default.py b/funasr/frontends/default.py index 67958cafe..324ecea2d 100644 --- a/funasr/frontends/default.py +++ b/funasr/frontends/default.py @@ -394,23 +394,25 @@ def _load_cmvn(self, cmvn_file): cmvn_file: TODO. """ with open(cmvn_file, "r", encoding="utf-8") as f: - lines = f.readlines() + lines = [line for line in f.readlines() if line.split()] means_list = [] vars_list = [] for i in range(len(lines)): line_item = lines[i].split() if line_item[0] == "": - line_item = lines[i + 1].split() - if line_item[0] == "": + line_item = lines[i + 1].split() if i + 1 < len(lines) else [] + if line_item and line_item[0] == "": add_shift_line = line_item[3 : (len(line_item) - 1)] means_list = list(add_shift_line) continue elif line_item[0] == "": - line_item = lines[i + 1].split() - if line_item[0] == "": + line_item = lines[i + 1].split() if i + 1 < len(lines) else [] + if line_item and line_item[0] == "": rescale_line = line_item[3 : (len(line_item) - 1)] vars_list = list(rescale_line) continue + if not means_list or not vars_list: + raise ValueError(f"No / statistics found in cmvn file: {cmvn_file}") means = np.array(means_list).astype(np.float64) vars = np.array(vars_list).astype(np.float64) return means, vars diff --git a/funasr/frontends/wav_frontend.py b/funasr/frontends/wav_frontend.py index 45f34095e..593cc1446 100644 --- a/funasr/frontends/wav_frontend.py +++ b/funasr/frontends/wav_frontend.py @@ -19,23 +19,25 @@ def load_cmvn(cmvn_file): cmvn_file: TODO. """ with open(cmvn_file, "r", encoding="utf-8") as f: - lines = f.readlines() + lines = [line for line in f.readlines() if line.split()] means_list = [] vars_list = [] for i in range(len(lines)): line_item = lines[i].split() if line_item[0] == "": - line_item = lines[i + 1].split() - if line_item[0] == "": + line_item = lines[i + 1].split() if i + 1 < len(lines) else [] + if line_item and line_item[0] == "": add_shift_line = line_item[3 : (len(line_item) - 1)] means_list = list(add_shift_line) continue elif line_item[0] == "": - line_item = lines[i + 1].split() - if line_item[0] == "": + line_item = lines[i + 1].split() if i + 1 < len(lines) else [] + if line_item and line_item[0] == "": rescale_line = line_item[3 : (len(line_item) - 1)] vars_list = list(rescale_line) continue + if not means_list or not vars_list: + raise ValueError(f"No / statistics found in cmvn file: {cmvn_file}") means = np.array(means_list).astype(np.float32) vars = np.array(vars_list).astype(np.float32) cmvn = np.array([means, vars]) diff --git a/tests/test_numpy_compatibility.py b/tests/test_numpy_compatibility.py index ff194b745..9e8dad057 100644 --- a/tests/test_numpy_compatibility.py +++ b/tests/test_numpy_compatibility.py @@ -37,6 +37,62 @@ def test_cmvn_load_returns_float64_arrays(cmvn_path): np.testing.assert_array_equal(scales, np.full(8, 1.25, dtype=np.float64)) +def test_cmvn_readers_skip_blank_lines(tmp_path): + from funasr.frontends.default import MultiChannelFrontend + from funasr.frontends.wav_frontend import load_cmvn + + path = tmp_path / "blank_lines.cmvn" + path.write_text( + "\n\n \t \n 0 0 " + + " ".join(["0.1"] * 8) + + " \n\n\n" + + "\n 0 0 " + + " ".join(["1.25"] * 8) + + " \n\n\n", + encoding="utf-8", + ) + means, scales = MultiChannelFrontend._load_cmvn(None, path) + np.testing.assert_array_equal(means, np.full(8, 0.1, dtype=np.float64)) + np.testing.assert_array_equal(scales, np.full(8, 1.25, dtype=np.float64)) + cmvn = load_cmvn(str(path)) + assert cmvn.shape == (2, 8) + torch.testing.assert_close(cmvn[0], torch.full((8,), 0.1)) + torch.testing.assert_close(cmvn[1], torch.full((8,), 1.25)) + + +def test_cmvn_readers_reject_files_without_statistics(tmp_path): + from funasr.frontends.default import MultiChannelFrontend + from funasr.frontends.wav_frontend import load_cmvn + + # Kaldi also emits the whole nnet on a single line; the readers walk lines, + # so nothing matches and empty statistics used to be returned silently. + compact = tmp_path / "compact.cmvn" + compact.write_text( + " 1 [ -1.0 -2.0 ] " + " 1 [ 0.5 0.5 ] \n", + encoding="utf-8", + ) + truncated = tmp_path / "truncated.cmvn" + truncated.write_text(" 2 2\n", encoding="utf-8") + for path in (compact, truncated): + with pytest.raises(ValueError) as excinfo: + MultiChannelFrontend._load_cmvn(None, path) + assert str(path) in str(excinfo.value) + with pytest.raises(ValueError) as excinfo: + load_cmvn(str(path)) + assert str(path) in str(excinfo.value) + + +def test_cmvn_reader_loads_the_am_mvn_shipped_by_the_runtime(): + from funasr.frontends.wav_frontend import load_cmvn + + path = Path(__file__).resolve().parents[1] / ( + "runtime/triton_gpu/model_repo_sense_voice_small/feature_extractor/am.mvn" + ) + cmvn = load_cmvn(str(path)) + assert cmvn.shape == (2, 560) and torch.isfinite(cmvn).all() + + def test_frontend_applies_cmvn_and_preserves_padding(cmvn_path): from funasr.frontends.default import MultiChannelFrontend