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
2 changes: 2 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
15 changes: 11 additions & 4 deletions mkl_fft/_pydfti.pyx
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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 = <DftiCache *>cpython.pycapsule.PyCapsule_GetPointer(
Expand Down Expand Up @@ -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 = <cnp.ndarray> out
else:
f_arr = _allocate_result(x_arr, n_, axis_, f_type)
Expand Down Expand Up @@ -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_, <int> 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:
Expand Down
67 changes: 67 additions & 0 deletions mkl_fft/tests/test_fft1d.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
22 changes: 22 additions & 0 deletions mkl_fft/tests/test_fftnd.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
)
Expand Down
Loading