diff --git a/src/torchmetrics/__init__.py b/src/torchmetrics/__init__.py index d660d3354b9..2df716f3f6d 100644 --- a/src/torchmetrics/__init__.py +++ b/src/torchmetrics/__init__.py @@ -27,13 +27,6 @@ if not hasattr(PIL, "PILLOW_VERSION"): PIL.PILLOW_VERSION = PIL.__version__ -if package_available("scipy"): - import scipy.signal - - # back compatibility patch due to SMRMpy using scipy.signal.hamming - if not hasattr(scipy.signal, "hamming"): - scipy.signal.hamming = scipy.signal.windows.hamming - from torchmetrics import functional # noqa: E402 from torchmetrics.aggregation import ( # noqa: E402 CatMetric, diff --git a/src/torchmetrics/audio/__init__.py b/src/torchmetrics/audio/__init__.py index da5271a4cda..ad6206fed64 100644 --- a/src/torchmetrics/audio/__init__.py +++ b/src/torchmetrics/audio/__init__.py @@ -29,17 +29,9 @@ _PESQ_AVAILABLE, _PYSTOI_AVAILABLE, _REQUESTS_AVAILABLE, - _SCIPI_AVAILABLE, _TORCHAUDIO_AVAILABLE, ) -if _SCIPI_AVAILABLE: - import scipy.signal - - # back compatibility patch due to SMRMpy using scipy.signal.hamming - if not hasattr(scipy.signal, "hamming"): - scipy.signal.hamming = scipy.signal.windows.hamming - __all__ = [ "ComplexScaleInvariantSignalNoiseRatio", "PermutationInvariantTraining", diff --git a/src/torchmetrics/functional/audio/__init__.py b/src/torchmetrics/functional/audio/__init__.py index 09faa97334a..2f21bf01c85 100644 --- a/src/torchmetrics/functional/audio/__init__.py +++ b/src/torchmetrics/functional/audio/__init__.py @@ -29,17 +29,9 @@ _PESQ_AVAILABLE, _PYSTOI_AVAILABLE, _REQUESTS_AVAILABLE, - _SCIPI_AVAILABLE, _TORCHAUDIO_AVAILABLE, ) -if _SCIPI_AVAILABLE: - import scipy.signal - - # back compatibility patch due to SMRMpy using scipy.signal.hamming - if not hasattr(scipy.signal, "hamming"): - scipy.signal.hamming = scipy.signal.windows.hamming - __all__ = [ "complex_scale_invariant_signal_noise_ratio", "permutation_invariant_training", diff --git a/tests/unittests/audio/__init__.py b/tests/unittests/audio/__init__.py index dd8a3c0014f..83e7a59fe15 100644 --- a/tests/unittests/audio/__init__.py +++ b/tests/unittests/audio/__init__.py @@ -5,6 +5,17 @@ from unittests import _PATH_ALL_TESTS +# SRMRpy calls `scipy.signal.hamming`, which SciPy moved to `scipy.signal.windows.hamming`. +# Only `test_srmr.py` imports SRMRpy, and this package is imported before it, so the shim +# lives here rather than in `torchmetrics`, which never imports SRMRpy at all. +try: + import scipy.signal + + if not hasattr(scipy.signal, "hamming"): + scipy.signal.hamming = scipy.signal.windows.hamming +except ImportError: + pass + _SAMPLE_AUDIO_SPEECH = os.path.join(_PATH_ALL_TESTS, "_data", "audio", "audio_speech.wav") _SAMPLE_AUDIO_SPEECH_BAB_DB = os.path.join(_PATH_ALL_TESTS, "_data", "audio", "audio_speech_bab_0dB.wav") _SAMPLE_NUMPY_ISSUE_895 = os.path.join(_PATH_ALL_TESTS, "_data", "audio", "issue_895.npz")