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
14 changes: 11 additions & 3 deletions funasr/utils/load_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,]
Expand Down Expand Up @@ -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

Expand All @@ -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.
Expand All @@ -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"
]

Expand Down
114 changes: 114 additions & 0 deletions tests/test_load_audio_bytes.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,8 @@
import wave

import numpy as np
import pytest
import soundfile as sf


ROOT = Path(__file__).resolve().parents[1]
Expand Down Expand Up @@ -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"
Expand Down Expand Up @@ -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("<i2").tobytes())
monkeypatch.setattr(LOAD_UTILS.torchaudio, "load", _consume_then_fail)
monkeypatch.setattr(sf, "read", _consume_then_fail)

actual = LOAD_UTILS.load_audio_text_image_video(
str(path), fs=target_rate, audio_fs=source_rate
)

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)


@pytest.mark.skipif(not LOAD_UTILS.is_ffmpeg_installed(), reason="ffmpeg is required")
@pytest.mark.parametrize("sample_rate", [8000, 16000, 48000])
def test_direct_ffmpeg_raw_pcm_defaults_to_requested_rate(sample_rate, tmp_path):
samples = _sine_pcm(sample_rate, duration=1.0)
path = tmp_path / "audio.pcm"
path.write_bytes(samples.astype("<i2").tobytes())

actual = LOAD_UTILS._load_audio_ffmpeg(str(path), sr=sample_rate)

np.testing.assert_array_equal(actual, samples.astype(np.float32) / 32768.0)


if __name__ == "__main__":
unittest.main()
35 changes: 35 additions & 0 deletions tests/test_pcm_input_format.py
Original file line number Diff line number Diff line change
Expand Up @@ -256,6 +256,41 @@ def test_vad_source_rates_are_not_reapplied_to_resampled_segments(
np.testing.assert_allclose(loaded, expected, atol=1e-6, rtol=0)


@pytest.mark.skipif(not load_utils.is_ffmpeg_installed(), reason="ffmpeg is required")
@pytest.mark.parametrize("source_rate", [8000, 48000])
@pytest.mark.parametrize("batch", [False, True])
def test_vad_ffmpeg_fallback_preserves_resampled_audio(
source_rate, batch, tmp_path, monkeypatch
):
import soundfile as sf

samples = np.round(
np.sin(2 * np.pi * 440 * np.arange(source_rate) / source_rate) * 12000
).astype("<i2")
path = tmp_path / "audio.wav"
sf.write(path, samples, source_rate, subtype="PCM_16")
expected = load_utils._load_audio_ffmpeg(str(path), sr=16000)

def decoder_unavailable(*args, **kwargs):
raise RuntimeError("decoder unavailable")

monkeypatch.setattr(load_utils.torchaudio, "load", decoder_unavailable)
monkeypatch.setattr(sf, "read", decoder_unavailable)
wrapper = _wrapper(vad=True)
wrapper.vad_model.segments = [[0, 1000]]
value = [str(path), str(path)] if batch else str(path)

results = wrapper.generate(input=value, fs=source_rate)

assert len(results) == (2 if batch else 1)
for _, vad_options in wrapper.vad_model.calls:
assert vad_options["fs"] == source_rate
for segments, asr_options in wrapper.model.calls:
assert asr_options["fs"] == 16000
assert len(segments[0]) == 16000
np.testing.assert_array_equal(segments[0], expected)


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

Expand Down
Loading