diff --git a/funasr/auto/auto_model.py b/funasr/auto/auto_model.py index e814cba84..4ff7f13a1 100644 --- a/funasr/auto/auto_model.py +++ b/funasr/auto/auto_model.py @@ -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 diff --git a/tests/test_pcm_input_format.py b/tests/test_pcm_input_format.py index 94e32a192..979b9bff2 100644 --- a/tests/test_pcm_input_format.py +++ b/tests/test_pcm_input_format.py @@ -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("