diff --git a/CHANGELOG.md b/CHANGELOG.md index 7b4d5485..047f0319 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -25,6 +25,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 * Fixed possible memory leaks when `PyMem_Malloc` fails, and raise `MemoryError` [gh-373](https://github.com/IntelPython/mkl_fft/pull/373) * Fixed possible memory leaks when multi-iterator constructors fail [gh-373](https://github.com/IntelPython/mkl_fft/pull/373) * Fixed N-D transforms returning a success status when a scratch allocation fails [gh-373](https://github.com/IntelPython/mkl_fft/pull/373) +* Fixed `fft`, `ifft`, and the `fftn`/`fft2` family returning `out` unmodified when `out` is given and the input dtype is not `float32`, `float64`, `complex64` or `complex128` (e.g. integer, `bool`, `float16`) [gh-387](https://github.com/IntelPython/mkl_fft/pull/387) +* Fixed wrong results and writes outside of `out` in `fft`/`ifft` when the input is cast or zero-padded into a contiguous copy while `out` has the same non-contiguous layout as the original input [gh-387](https://github.com/IntelPython/mkl_fft/pull/387) ## [2.3.2] - 2026-08-04 diff --git a/mkl_fft/_pydfti.pyx b/mkl_fft/_pydfti.pyx index 94fc82f1..e53eb94e 100644 --- a/mkl_fft/_pydfti.pyx +++ b/mkl_fft/_pydfti.pyx @@ -426,9 +426,7 @@ def _c2c_fft1d_impl(x, n=None, axis=-1, direction=+1, double fsc=1.0, out=None): x_arr = _process_arguments(x, n, axis, &axis_, &n_, &in_place, &xnd, 0) x_type = cnp.PyArray_TYPE(x_arr) - if out is not None: - in_place = 0 - elif x_type is cnp.NPY_CFLOAT or x_type is cnp.NPY_CDOUBLE: + if x_type is cnp.NPY_CFLOAT or x_type is cnp.NPY_CDOUBLE: # we can operate in place if requested. if in_place: if not cnp.PyArray_ISONESEGMENT(x_arr): @@ -453,6 +451,10 @@ def _c2c_fft1d_impl(x, n=None, axis=-1, direction=+1, double fsc=1.0, out=None): x_type = cnp.PyArray_TYPE(x_arr) in_place = 1 + # checked only after the cast above, which is needed even if out is given + if out is not None: + in_place = 0 + if in_place: _cache_capsule = _tls_dfti_cache_capsule() _cache = cpython.pycapsule.PyCapsule_GetPointer( @@ -502,8 +504,9 @@ def _c2c_fft1d_impl(x, n=None, axis=-1, direction=+1, double fsc=1.0, out=None): _validate_out_array(out, x, out_dtype, axis=axis_, n=n_) # out array that is used in OneMKL c2c FFT must have the exact same # stride as input array. If not, we need to allocate a new array. + # Compare with x_arr, which may be a cast or padded copy of x. # TODO: check to see if this condition can be relaxed - if _get_element_strides(x) == _get_element_strides(out): + if _get_element_strides(x_arr) == _get_element_strides(out): f_arr = out else: f_arr = _allocate_result(x_arr, n_, axis_, f_type) @@ -544,6 +547,10 @@ def _c2c_fft1d_impl(x, n=None, axis=-1, direction=+1, double fsc=1.0, out=None): status = cdouble_cdouble_mkl_fft1d_out( x_arr, n_, axis_, f_arr, fsc, _cache ) + else: + raise ValueError( + "An input argument x is not of a supported type" + ) else: if x_type is cnp.NPY_FLOAT: if direction < 0: diff --git a/mkl_fft/tests/test_fft1d.py b/mkl_fft/tests/test_fft1d.py index 16464e75..ccf4332c 100644 --- a/mkl_fft/tests/test_fft1d.py +++ b/mkl_fft/tests/test_fft1d.py @@ -367,6 +367,73 @@ def test_fft_out_strided(axis, func): assert_allclose(result, expected) +@pytest.mark.parametrize("n", [None, 8, 24]) +@pytest.mark.parametrize("dt", ["?", "i1", "u1", "i4", "i8", "f2"]) +@pytest.mark.parametrize("func", ["fft", "ifft"]) +def test_fft_out_cast_input(func, dt, n): + # dtypes other than f4/f8/c8/c16 are cast to complex128 + x = np.arange(1, 17).astype(dt) + # sentinel: catches out being returned untouched + out = np.full(16 if n is None else n, -1 - 1j) + result = getattr(mkl_fft, func)(x, n=n, out=out) + expected = getattr(np.fft, func)(x.astype(np.float64), n=n) + + assert result is out + assert_allclose(result, expected, atol=1e-12) + + +@pytest.mark.parametrize("axis", [0, 1, 2]) +@pytest.mark.parametrize("dt", ["i8", "c16"]) +@pytest.mark.parametrize("func", ["fft", "ifft"]) +def test_fft_out_strided_input_copied(func, dt, axis): + # x and out have the same strides, but x is cast (i8) or padded (c16) + # into a contiguous copy, so out cannot be handed to MKL as is + shape = (20, 33, 54) + base = np.full(shape, -1 - 1j) + out = base[::2, ::3, ::4] + x = rnd.randint(-50, 50, size=shape).astype(dt)[::2, ::3, ::4] + n = None + if dt == "c16": + ind = [slice(None)] * x.ndim + ind[axis] = slice(0, x.shape[axis] - 3) + x = x[tuple(ind)] + n = out.shape[axis] + + result = getattr(mkl_fft, func)(x, n=n, axis=axis, out=out) + expected = getattr(np.fft, func)(x.astype(np.complex128), n=n, axis=axis) + + assert result is out + assert_allclose(result, expected, atol=1e-10) + # nothing outside of out was written to + out[...] = -1 - 1j + assert np.all(base == -1 - 1j) + + +@pytest.mark.parametrize("axis", [0, 1, 2]) +@pytest.mark.parametrize("func", ["fft", "ifft"]) +def test_fft_out_fortran_cast_input(func, axis): + x = np.asfortranarray(rnd.randint(-50, 50, size=(4, 5, 6))) + out = np.full(x.shape, -1 - 1j, order="F") + result = getattr(mkl_fft, func)(x, axis=axis, out=out) + expected = getattr(np.fft, func)(x.astype(np.float64), axis=axis) + + assert result is out + assert_allclose(result, expected, atol=1e-10) + + +@pytest.mark.skipif( + np.can_cast(np.longdouble, np.complex128), + reason="long double is the same as double", +) +@pytest.mark.parametrize("dt", [np.longdouble, np.clongdouble]) +@pytest.mark.parametrize("use_out", [False, True]) +def test_fft_longdouble_unsupported(dt, use_out): + x = np.ones(8, dtype=dt) + out = np.empty(8, dtype=np.complex128) if use_out else None + with pytest.raises(ValueError, match="single or double precision"): + mkl_fft.fft(x, out=out) + + @requires_numpy_2 @pytest.mark.parametrize("axis", [0, 1, 2]) def test_rfft_out_strided(axis): diff --git a/mkl_fft/tests/test_fftnd.py b/mkl_fft/tests/test_fftnd.py index 4aa6f159..0b2f9b0c 100644 --- a/mkl_fft/tests/test_fftnd.py +++ b/mkl_fft/tests/test_fftnd.py @@ -309,6 +309,28 @@ def test_out_strided(axes, func): assert_allclose(result, expected, strict=True) +@pytest.mark.parametrize("strided", [False, True]) +@pytest.mark.parametrize("axes", [None, (0, 1), (0, 2), (1, 2)]) +@pytest.mark.parametrize("func", ["fftn", "ifftn"]) +def test_out_int_input(func, axes, strided): + # integer input is cast to complex128 before being transformed into out + shape = (20, 30, 40) + x = rnd.randint(-50, 50, size=shape) + base = np.full(shape, -1 - 1j) + out = base + if strided: + x = x[::2, ::3, ::4] + out = base[::2, ::3, ::4] + result = getattr(mkl_fft, func)(x, axes=axes, out=out) + expected = getattr(np.fft, func)(x.astype(np.float64), axes=axes) + + assert result is out + assert_allclose(result, expected, rtol=reps_64, atol=1e-9) + # nothing outside of out was written to + out[...] = -1 - 1j + assert np.all(base == -1 - 1j) + + @pytest.mark.parametrize( "dtype", [np.float32, np.float64, np.complex64, np.complex128] )