diff --git a/funasr/utils/load_utils.py b/funasr/utils/load_utils.py index b689d65ec..7c93dd34b 100644 --- a/funasr/utils/load_utils.py +++ b/funasr/utils/load_utils.py @@ -140,7 +140,11 @@ def load_audio_text_image_video( except: if hasattr(data_or_path_or_list, "seek"): data_or_path_or_list.seek(0) - data_or_path_or_list = _load_audio_ffmpeg(data_or_path_or_list, sr=fs) + data_or_path_or_list = _load_audio_ffmpeg( + data_or_path_or_list, sr=fs, input_sr=audio_fs + ) + # FFmpeg has already resampled its output to the target rate. + audio_fs = fs data_or_path_or_list = torch.from_numpy( data_or_path_or_list ).squeeze() # [n_samples,] @@ -438,7 +442,7 @@ def extract_fbank(data, data_len=None, data_type: str = "sound", frontend=None, return data.to(torch.float32), data_len.to(torch.int32) -def _load_audio_ffmpeg(file, sr: int = 16000): +def _load_audio_ffmpeg(file, sr: int = 16000, input_sr=None): """ Open an audio file and read as mono waveform, resampling as necessary @@ -450,6 +454,10 @@ def _load_audio_ffmpeg(file, sr: int = 16000): sr: int The sample rate to resample the audio if necessary + input_sr: int or None + The source sample rate for raw PCM files. Defaults to sr; ignored for + containers, which provide their own sample rate. + Returns ------- A NumPy array containing the audio waveform, in float32 dtype. @@ -470,7 +478,7 @@ def _load_audio_ffmpeg(file, sr: int = 16000): if isinstance(input_source, str) and input_source.lower().endswith('.pcm'): pcm_params = [ "-f", "s16le", - "-ar", str(sr), + "-ar", str(sr if input_sr is None else input_sr), "-ac", "1" ] diff --git a/tests/test_load_audio_bytes.py b/tests/test_load_audio_bytes.py index d23ae04a7..6cbeb6eca 100644 --- a/tests/test_load_audio_bytes.py +++ b/tests/test_load_audio_bytes.py @@ -11,6 +11,8 @@ import wave import numpy as np +import pytest +import soundfile as sf ROOT = Path(__file__).resolve().parents[1] @@ -47,6 +49,13 @@ def _load_utils_module(): LOAD_UTILS = _load_utils_module() + +def _consume_then_fail(value, *args, **kwargs): + if hasattr(value, "read"): + value.read() + raise RuntimeError("decoder unavailable") + + # 80 ms, 16 kHz mono, generated without ID3 or Xing metadata by ffmpeg 6.1.1. NO_ID3_MP3 = base64.b64decode( "//NIxAAa0LJcBVlIABsLTlpy05ZMsmWnLloB0Uy2BhEGUUDiEXzRhOvE88TrdOJY1lDSQMoYBCIAFBGIO4CYJgHAGBsEw2TtwUQIECBBBwnB8H8oGJTlA/wcOYgB/WDhzIA/wI7n+jhgHz+BDnfg+BAQBDB//B8H1QJghAAQSMFwa3j/mFYgGAQU1hUCfLMm//NIxBcg+gJkDZ2gADIa5vcxhhUcxwIo5huQh2CQ4GXugbI6BkioGWkgiPkQAyWIBoMXQsSBs0DYYRz/hlkMijHCghC3/kNFyi5SaHOHO//KJFSKmReJox//yKkVMi8XjEul1L//8mi8Yl0upF42EoS//Pf///XYuxb//41DPjII3hgHzBwACWYI5Npk2mcG//NIxBYiUm5MAZ6oABXAdm1MCaSE6GJWZaBm8zAWVAg8DDsXAAA46gGGABlEABoH+F1A6xxid//IwiAuAskT//HGWCcIIdJ///J84VCcOm5P///nDQuHTcvqNC5///+s3N1Ghos3N1Ghz//WD4gCIPiAIg////gwEQuQYFyiD4Eqijt1s/+gX9S2f8CFjqrn//NIxA8hCpaUAZqYAPlgf8AdQXUCtxZpIEW4sIuMgA7jJMuo+OAmyIkTLikklo/KxOF8qk+YIqSUtH8rE4ZlUvoGy6SqK/5cTPF9R8uLPUlUV0lf8vqPmizyaj6Cz1FdJVFdJX/6c+hPJz5wBpK2LsX/hgBmwwCYHDAJgcMKtVaq1VVMQU1FMy4xMDBVVVVV//NIxA0AAANIAcAAAFVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVVV" @@ -231,5 +240,110 @@ def test_corrupt_container_error_is_actionable(self): LOAD_UTILS.load_bytes(corrupt_wav) +@pytest.mark.skipif(not LOAD_UTILS.is_ffmpeg_installed(), reason="ffmpeg is required") +@pytest.mark.parametrize("source_rate", [8000, 16000, 48000]) +@pytest.mark.parametrize("target_rate", [16000, 24000]) +@pytest.mark.parametrize("rate_hint", [8000, 16000, 48000]) +@pytest.mark.parametrize("file_like", [False, True]) +def test_ffmpeg_fallback_uses_decoded_rate( + source_rate, target_rate, rate_hint, file_like, tmp_path, monkeypatch +): + samples = _sine_pcm(source_rate, duration=1.0) + wav_data = _wav_bytes(samples, source_rate) + path = tmp_path / "audio.wav" + path.write_bytes(wav_data) + source = io.BytesIO(wav_data) if file_like else str(path) + expected = LOAD_UTILS._load_audio_ffmpeg(str(path), sr=target_rate) + monkeypatch.setattr(LOAD_UTILS.torchaudio, "load", _consume_then_fail) + monkeypatch.setattr(sf, "read", _consume_then_fail) + + actual = LOAD_UTILS.load_audio_text_image_video( + source, fs=target_rate, audio_fs=rate_hint + ) + + assert len(actual) == target_rate + np.testing.assert_array_equal(actual, expected) + + +@pytest.mark.skipif(not LOAD_UTILS.is_ffmpeg_installed(), reason="ffmpeg is required") +@pytest.mark.parametrize("source_rate", [8000, 48000]) +def test_ffmpeg_fallback_without_torchaudio(source_rate, monkeypatch): + source = io.BytesIO(_wav_bytes(_sine_pcm(source_rate, 1.0), source_rate)) + expected = LOAD_UTILS._load_audio_ffmpeg(source, sr=16000) + monkeypatch.setattr(LOAD_UTILS, "torchaudio", None) + monkeypatch.setattr(sf, "read", _consume_then_fail) + + actual = LOAD_UTILS.load_audio_text_image_video( + source, fs=16000, audio_fs=source_rate + ) + + assert len(actual) == 16000 + np.testing.assert_array_equal(actual, expected) + + +@pytest.mark.parametrize("source_rate", [8000, 16000, 48000]) +@pytest.mark.parametrize("target_rate", [16000, 24000]) +@pytest.mark.parametrize("decoder", ["torchaudio", "soundfile"]) +def test_normal_decoders_use_container_rate( + source_rate, target_rate, decoder, tmp_path, monkeypatch +): + samples = _sine_pcm(source_rate, duration=1.0) + path = tmp_path / "audio.wav" + path.write_bytes(_wav_bytes(samples, source_rate)) + if decoder == "soundfile": + monkeypatch.setattr(LOAD_UTILS.torchaudio, "load", _consume_then_fail) + else: + # torchaudio 2.9+ requires the optional torchcodec package to load audio. + try: + LOAD_UTILS.torchaudio.load(str(path)) + except (ImportError, RuntimeError) as error: + pytest.skip(f"torchaudio decoder unavailable: {error}") + monkeypatch.setattr(sf, "read", mock.Mock(side_effect=AssertionError)) + fallback = mock.Mock(side_effect=AssertionError("unexpected ffmpeg fallback")) + monkeypatch.setattr(LOAD_UTILS, "_load_audio_ffmpeg", fallback) + + actual = LOAD_UTILS.load_audio_text_image_video( + str(path), fs=target_rate, audio_fs=12345 + ) + + assert len(actual) == target_rate + expected = _sine_pcm(target_rate, 1.0).astype(np.float32) / 32768.0 + np.testing.assert_allclose(actual[100:-100], expected[100:-100], atol=0.001) + fallback.assert_not_called() + + +@pytest.mark.skipif(not LOAD_UTILS.is_ffmpeg_installed(), reason="ffmpeg is required") +@pytest.mark.parametrize("source_rate", [8000, 16000, 48000]) +@pytest.mark.parametrize("target_rate", [16000, 24000]) +def test_ffmpeg_raw_pcm_preserves_source_rate( + source_rate, target_rate, tmp_path, monkeypatch +): + samples = _sine_pcm(source_rate, duration=1.0) + path = tmp_path / "audio.pcm" + path.write_bytes(samples.astype("