Skip to content
Open
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
2 changes: 1 addition & 1 deletion funasr/auto/auto_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -1124,7 +1124,7 @@ def inference_with_vad(self, input, input_len=None, **cfg):
input_len=None,
model=self.spk_model,
kwargs=self.spk_kwargs,
**cfg,
**{**cfg, "fs": fs},
)
spk_embs = torch.cat([r["spk_embedding"] for r in spk_res], dim=0)
results[_b]["spk_embedding"] = spk_embs
Expand Down
134 changes: 134 additions & 0 deletions tests/test_pcm_input_format.py
Original file line number Diff line number Diff line change
Expand Up @@ -235,3 +235,137 @@ def test_vad_source_rates_are_not_reapplied_to_resampled_segments(
)
assert len(loaded) == 3200
np.testing.assert_allclose(loaded, expected, atol=1e-6, rtol=0)


class _RecordingSpeaker(torch.nn.Module):
"""Use CAMPPlus preprocessing with a stand-in embedding network."""

def __init__(self, **kwargs):
super().__init__()
self.anchor = torch.nn.Parameter(torch.zeros(1))
self.calls = []

def inference(self, data_in, key, **kwargs):
from funasr.models.campplus.model import CAMPPlus

self.calls.append(kwargs)
return CAMPPlus.inference(self, data_in, key=key, **kwargs)

def forward(self, features):
return torch.ones(features.shape[0], 4)


@pytest.mark.parametrize(
"base_rate,call_rate,speaker_rate,source_rate",
[
(None, None, None, 16000),
(8000, None, None, 8000),
(None, 8000, None, 8000),
(None, 48000, None, 48000),
(8000, 48000, 8000, 48000),
(48000, 8000, 48000, 8000),
(16000, None, 48000, 16000),
],
)
@pytest.mark.parametrize("batch", [False, True])
@pytest.mark.parametrize("input_kind", ["array", "pcm"])
def test_vad_speaker_uses_resampled_rate(
base_rate, call_rate, speaker_rate, source_rate, batch, input_kind, monkeypatch
):
from funasr.models.campplus import model as campplus_model

feature_inputs = []

def capture_features(audio):
feature_inputs.extend(sample.numpy().copy() for sample in audio)
return torch.zeros(len(audio), 20, 80), [20] * len(audio), [len(a) for a in audio]

monkeypatch.setattr(campplus_model, "extract_feature", capture_features)
wrapper = _wrapper(vad=True)
wrapper.vad_model.segments = [[0, 200]]
wrapper.spk_model = _RecordingSpeaker()
wrapper.spk_kwargs = {"device": "cpu"}
wrapper.kwargs["return_spk_res"] = False
if base_rate is not None:
wrapper.kwargs["fs"] = base_rate
if speaker_rate is not None:
wrapper.spk_kwargs["fs"] = speaker_rate
wrapper._store_base_configs()

pcm = np.round(
np.sin(2 * np.pi * 440 * np.arange(source_rate // 5) / source_rate) * 12000
).astype("<i2").tobytes()
waveform = _expected(pcm)
resampled = load_utils.load_audio_text_image_video(
waveform, fs=16000, audio_fs=source_rate
)
# sv_chunk pads the already-resampled 200 ms segment to 1.5 seconds.
expected = np.pad(resampled, (0, 24000 - len(resampled)))
options = {} if call_rate is None else {"fs": call_rate}
if input_kind == "pcm":
options["input_format"] = "pcm_s16le"
value = pcm if input_kind == "pcm" else waveform
results = wrapper.generate(input=[value, value] if batch else value, **options)
assert len(results) == (2 if batch else 1)
assert len(feature_inputs) == len(results)
for audio in feature_inputs:
assert len(audio) == 24000
np.testing.assert_allclose(audio, expected, atol=1e-6, rtol=0)
for _, vad_options in wrapper.vad_model.calls:
assert vad_options["fs"] == source_rate
for _, asr_options in wrapper.model.calls:
assert asr_options["fs"] == 16000
for speaker_options in wrapper.spk_model.calls:
assert speaker_options.get("fs", 16000) == 16000
assert wrapper._base_kwargs_map["kwargs"].get("fs") == base_rate
assert wrapper._base_kwargs_map["spk_kwargs"].get("fs") == speaker_rate

# A subsequent call without an override must use the constructor source rate.
feature_inputs.clear()
next_rate = base_rate or 16000
wrapper.generate(input=np.zeros(next_rate // 5, dtype=np.float32))
assert wrapper.vad_model.calls[-1][1]["fs"] == next_rate
assert wrapper.model.calls[-1][1]["fs"] == 16000
assert wrapper.spk_model.calls[-1].get("fs", 16000) == 16000
assert len(feature_inputs[0]) == 24000


def test_speaker_rate_override_preserves_caller_config(monkeypatch):
import copy
import importlib

from funasr.register import tables

auto_module = importlib.import_module("funasr.auto.auto_model")

class ConfiguredModel(_RecordingModel):
def __init__(self, model, **kwargs):
super().__init__(vad=model == "rate-test-vad")
self.segments = [[0, 200]]

monkeypatch.setitem(tables.model_classes, "rate-test-asr", ConfiguredModel)
monkeypatch.setitem(tables.model_classes, "rate-test-vad", ConfiguredModel)
monkeypatch.setitem(tables.model_classes, "rate-test-speaker", _RecordingSpeaker)
monkeypatch.setitem(
tables.frontend_classes, "rate-test-frontend", lambda **kw: SimpleNamespace(fs=16000)
)
monkeypatch.setattr(
auto_module, "ClusterBackend", lambda **kw: torch.nn.Identity(), raising=False
)
monkeypatch.setattr(auto_module, "download_model", _forbid_decode)
caller_config = {"model_conf": {}, "device": "cpu", "fs": 48000}
original = copy.deepcopy(caller_config)
wrapper = AutoModel(
model="rate-test-asr", model_conf={}, frontend="rate-test-frontend",
vad_model="rate-test-vad", vad_kwargs={"model_conf": {}},
spk_model="rate-test-speaker", spk_kwargs=caller_config,
fs=48000, device="cpu", ncpu=torch.get_num_threads(),
disable_update=True, disable_pbar=True, return_spk_res=False,
)
for rate in (8000, 48000):
options = {"fs": rate} if rate == 8000 else {}
wrapper.generate(input=np.zeros(rate // 5, dtype=np.float32), **options)
assert wrapper.vad_model.calls[-1][1]["fs"] == rate
assert wrapper.spk_model.calls[-1]["fs"] == 16000
assert caller_config == original
assert wrapper._base_kwargs_map["spk_kwargs"]["fs"] == 48000