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
12 changes: 7 additions & 5 deletions funasr/frontends/default.py
Original file line number Diff line number Diff line change
Expand Up @@ -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] == "<AddShift>":
line_item = lines[i + 1].split()
if line_item[0] == "<LearnRateCoef>":
line_item = lines[i + 1].split() if i + 1 < len(lines) else []
if line_item and line_item[0] == "<LearnRateCoef>":
add_shift_line = line_item[3 : (len(line_item) - 1)]
means_list = list(add_shift_line)
continue
elif line_item[0] == "<Rescale>":
line_item = lines[i + 1].split()
if line_item[0] == "<LearnRateCoef>":
line_item = lines[i + 1].split() if i + 1 < len(lines) else []
if line_item and line_item[0] == "<LearnRateCoef>":
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 <AddShift>/<Rescale> 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
12 changes: 7 additions & 5 deletions funasr/frontends/wav_frontend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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] == "<AddShift>":
line_item = lines[i + 1].split()
if line_item[0] == "<LearnRateCoef>":
line_item = lines[i + 1].split() if i + 1 < len(lines) else []
if line_item and line_item[0] == "<LearnRateCoef>":
add_shift_line = line_item[3 : (len(line_item) - 1)]
means_list = list(add_shift_line)
continue
elif line_item[0] == "<Rescale>":
line_item = lines[i + 1].split()
if line_item[0] == "<LearnRateCoef>":
line_item = lines[i + 1].split() if i + 1 < len(lines) else []
if line_item and line_item[0] == "<LearnRateCoef>":
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 <AddShift>/<Rescale> 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])
Expand Down
56 changes: 56 additions & 0 deletions tests/test_numpy_compatibility.py
Original file line number Diff line number Diff line change
Expand Up @@ -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<AddShift>\n \t \n<LearnRateCoef> 0 0 "
+ " ".join(["0.1"] * 8)
+ " </LearnRateCoef>\n</AddShift>\n\n"
+ "<Rescale>\n<LearnRateCoef> 0 0 "
+ " ".join(["1.25"] * 8)
+ " </LearnRateCoef>\n</Rescale>\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(
"<Nnet> <AddShift> <LearnRateCoef> 1 [ -1.0 -2.0 ] </LearnRateCoef> "
"<Rescale> <LearnRateCoef> 1 [ 0.5 0.5 ] </Rescale> </Nnet>\n",
encoding="utf-8",
)
truncated = tmp_path / "truncated.cmvn"
truncated.write_text("<AddShift> 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

Expand Down
Loading