diff --git a/mkl_fft/tests/test_fft1d.py b/mkl_fft/tests/test_fft1d.py index 16464e75..6cd0e2dc 100644 --- a/mkl_fft/tests/test_fft1d.py +++ b/mkl_fft/tests/test_fft1d.py @@ -357,14 +357,18 @@ def test_array6(self): def test_fft_out_strided(axis, func): shape = (20, 33, 54) x = rnd.random(shape) + 1j * rnd.random(shape) - out = np.empty(shape, dtype=x.dtype) + base = np.full(shape, -1 - 1j) x = x[::2, ::3, ::4] - out = np.empty(x.shape, dtype=x.dtype) + out = base[::2, ::3, ::4] result = getattr(mkl_fft, func)(x, axis=axis, out=out) - expected = getattr(np.fft, func)(x, axis=axis, out=out) + expected = getattr(np.fft, func)(x, axis=axis) + assert result is out assert_allclose(result, expected) + # nothing outside of out was written to + out[...] = -1 - 1j + assert np.all(base == -1 - 1j) @requires_numpy_2 @@ -378,14 +382,18 @@ def test_rfft_out_strided(axis): out_sh = (20, 18, 54) else: # axis == 2 out_sh = (20, 33, 32) - out = np.empty(out_sh, dtype=np.complex128) + base = np.full(out_sh, -1 - 1j) x = x[::2, ::3, ::4] - out = out[::2, ::3, ::4] + out = base[::2, ::3, ::4] result = mkl_fft.rfft(x, axis=axis, out=out) - expected = np.fft.rfft(x, axis=axis, out=out) + expected = np.fft.rfft(x, axis=axis) + assert result is out assert_allclose(result, expected) + # nothing outside of out was written to + out[...] = -1 - 1j + assert np.all(base == -1 - 1j) @requires_numpy_2 @@ -399,14 +407,18 @@ def test_irfft_out_strided(axis): out_sh = (20, 60, 54) else: # axis == 2 out_sh = (20, 33, 104) - out = np.empty(out_sh, dtype=np.float64) + base = np.full(out_sh, -1.0) x = x[::2, ::3, ::4] - out = out[::2, ::3, ::4] + out = base[::2, ::3, ::4] result = mkl_fft.irfft(x, axis=axis, out=out) - expected = np.fft.irfft(x, axis=axis, out=out) + expected = np.fft.irfft(x, axis=axis) + assert result is out assert_allclose(result, expected) + # nothing outside of out was written to + out[...] = -1.0 + assert np.all(base == -1.0) @requires_numpy_2 diff --git a/mkl_fft/tests/test_fftnd.py b/mkl_fft/tests/test_fftnd.py index 4aa6f159..923f8a4c 100644 --- a/mkl_fft/tests/test_fftnd.py +++ b/mkl_fft/tests/test_fftnd.py @@ -299,14 +299,18 @@ def test_repeated_axes(dtype, axes, func): def test_out_strided(axes, func): shape = (20, 30, 40) x = rnd.random(shape) + 1j * rnd.random(shape) - out = np.empty(shape, dtype=x.dtype) + base = np.full(shape, -1 - 1j) x = x[::2, ::3, ::4] - out = out[::2, ::3, ::4] + out = base[::2, ::3, ::4] result = getattr(mkl_fft, func)(x, axes=axes, out=out) - expected = getattr(np.fft, func)(x, axes=axes, out=out) + expected = getattr(np.fft, func)(x, axes=axes) + assert result is out assert_allclose(result, expected, strict=True) + # nothing outside of out was written to + out[...] = -1 - 1j + assert np.all(base == -1 - 1j) @pytest.mark.parametrize(