Skip to content
Merged
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
42 changes: 41 additions & 1 deletion dpnp/tests/third_party/cupy/fft_tests/test_fft.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,14 @@
from __future__ import annotations

import functools
import math
import warnings

import numpy as np
import pytest

import dpnp as cupy
from dpnp.tests.helper import has_support_aspect64
from dpnp.tests.helper import get_float_dtypes, has_support_aspect64

# from cupy.fft import config
# from cupy.fft._fft import (
Expand Down Expand Up @@ -1288,6 +1289,45 @@ def test_irfftn(self, xp, dtype, order, enable_nd):
return xp.fft.irfftn(a, s=self.s, axes=self.axes, norm=self.norm)


@pytest.mark.parametrize("dtype", get_float_dtypes())
@pytest.mark.parametrize("offset", [0, 1])
@pytest.mark.parametrize("ndim", [1, 3])
@pytest.mark.parametrize("layout", ["view", "strided", "pointer"])
def test_rfft_input_alignment(
dtype: type[np.float32 | np.float64], offset: int, ndim: int, layout: str
) -> None:
shape: tuple[int, ...] = (35,) * ndim
backing: cupy.ndarray
x: cupy.ndarray
if layout == "view":
backing = cupy.ones(shape=(2,) + shape, dtype=dtype)
x = backing[offset]
elif layout == "strided":
backing = cupy.ones(shape=shape[:-1] + (2 * shape[-1],), dtype=dtype)
x = backing[..., offset::2]
else:
backing = cupy.ones(shape=math.prod(shape) + 1, dtype=dtype)
x = cupy.ndarray(shape, dtype=dtype, buffer=backing, offset=offset)
assert x.data.ptr % (2 * x.itemsize) == offset * x.itemsize

out: cupy.ndarray = cupy.fft.rfft(x) if ndim == 1 else cupy.fft.rfftn(x)
expected: np.ndarray = np.fft.rfftn(np.ones(shape=shape, dtype=dtype))
# oneMKL f32 err in zero bins scales with DC value (prod(shape))
testing.assert_allclose(
actual=out,
desired=expected,
rtol=1e-5 if dtype is np.float32 else 1e-12,
atol=(
np.finfo(dtype).eps * math.prod(shape)
if dtype is np.float32
else 1e-9
),
)
testing.assert_array_equal(
actual=backing, desired=np.ones(shape=backing.shape, dtype=dtype)
)


# Only those tests in which a legit plan can be obtained are kept
@testing.with_requires("numpy>=2.0")
@pytest.mark.usefixtures("skip_forward_backward")
Expand Down
Loading