Skip to content
Merged
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
44 changes: 37 additions & 7 deletions src/pyrecest/_backend/pytorch/fft.py
Original file line number Diff line number Diff line change
Expand Up @@ -124,6 +124,26 @@ def _normalize_fft_shape_args(args, kwargs):
return args, kwargs


def _normalize_fft_dim_args(args, kwargs, position, normalizer):
"""Normalize a positional or keyword ``dim`` FFT argument."""
if position is not None and len(args) > position:
args = list(args)
args[position] = normalizer(args[position])
return tuple(args), kwargs
if "dim" not in kwargs:
return args, kwargs
kwargs = dict(kwargs)
kwargs["dim"] = normalizer(kwargs["dim"])
return args, kwargs


def _get_fft_dim_arg(args, kwargs, position):
"""Return the effective positional or keyword ``dim`` FFT argument."""
if position is not None and len(args) > position:
return args[position]
return kwargs.get("dim")


def _with_dim_alias(kwargs, alias, func_name, *, none_alias_means_default=True):
if alias not in kwargs:
return kwargs
Expand Down Expand Up @@ -153,6 +173,7 @@ def _wrap_arraylike_fft(
*,
func_name,
dim_alias=None,
dim_arg_position=None,
empty_dim_is_noop=False,
normalize_scalar_dim=False,
normalize_dim_sequence=False,
Expand All @@ -171,18 +192,21 @@ def fft_func(value, *args, **kwargs):
)
if validate_real_length:
_validate_real_fft_length_args(args, kwargs)
if normalize_scalar_dim and "dim" in kwargs:
kwargs = dict(kwargs)
kwargs["dim"] = _normalize_single_fft_dim(kwargs["dim"])
if normalize_dim_sequence and "dim" in kwargs:
kwargs = dict(kwargs)
kwargs["dim"] = _normalize_fft_dim_sequence(kwargs["dim"])
if normalize_scalar_dim:
args, kwargs = _normalize_fft_dim_args(
args, kwargs, dim_arg_position, _normalize_single_fft_dim
)
if normalize_dim_sequence:
args, kwargs = _normalize_fft_dim_args(
args, kwargs, dim_arg_position, _normalize_fft_dim_sequence
)
if normalize_shape_sequence:
args, kwargs = _normalize_fft_shape_args(args, kwargs)
value = _as_fft_tensor(value)
dim = _get_fft_dim_arg(args, kwargs, dim_arg_position)
if (
empty_dim_is_noop
and _is_empty_dim(kwargs.get("dim"))
and _is_empty_dim(dim)
and _empty_dim_noop_is_valid(args, kwargs)
):
return value
Expand All @@ -195,6 +219,7 @@ def fft_func(value, *args, **kwargs):
_torch.fft.rfft,
func_name="rfft",
dim_alias="axis",
dim_arg_position=1,
normalize_scalar_dim=True,
validate_real_length=True,
none_alias_means_default=False,
Expand All @@ -203,6 +228,7 @@ def fft_func(value, *args, **kwargs):
_torch.fft.irfft,
func_name="irfft",
dim_alias="axis",
dim_arg_position=1,
normalize_scalar_dim=True,
validate_real_length=True,
none_alias_means_default=False,
Expand All @@ -211,20 +237,23 @@ def fft_func(value, *args, **kwargs):
_torch.fft.fftshift,
func_name="fftshift",
dim_alias="axes",
dim_arg_position=0,
empty_dim_is_noop=True,
normalize_dim_sequence=True,
)
ifftshift = _wrap_arraylike_fft(
_torch.fft.ifftshift,
func_name="ifftshift",
dim_alias="axes",
dim_arg_position=0,
empty_dim_is_noop=True,
normalize_dim_sequence=True,
)
fftn = _wrap_arraylike_fft(
_torch.fft.fftn,
func_name="fftn",
dim_alias="axes",
dim_arg_position=1,
empty_dim_is_noop=True,
normalize_dim_sequence=True,
normalize_shape_sequence=True,
Expand All @@ -233,6 +262,7 @@ def fft_func(value, *args, **kwargs):
_torch.fft.ifftn,
func_name="ifftn",
dim_alias="axes",
dim_arg_position=1,
empty_dim_is_noop=True,
normalize_dim_sequence=True,
normalize_shape_sequence=True,
Expand Down
83 changes: 83 additions & 0 deletions tests/backend_support/test_pytorch_fft_positional_axes_contract.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,83 @@
import importlib.util
import os
import subprocess
import sys

import pytest


@pytest.mark.backend_portable
def test_pytorch_fft_positional_axes_accept_numpy_integer_arrays():
if importlib.util.find_spec("torch") is None:
pytest.skip("torch is not installed")

env = os.environ.copy()
env["PYRECEST_BACKEND"] = "pytorch"
src_path = os.path.abspath("src")
env["PYTHONPATH"] = (
src_path
if not env.get("PYTHONPATH")
else os.pathsep.join([src_path, env["PYTHONPATH"]])
)

code = """
import numpy as np
import numpy.testing as npt
import pyrecest.backend as backend

values_np = np.array([[0.0, 1.0], [2.0, 3.0]])
values = backend.array(values_np)

for axis in (np.array(0), np.int64(0)):
npt.assert_allclose(
backend.to_numpy(backend.fft.rfft(values, None, axis)),
np.fft.rfft(values_np, None, axis),
)
npt.assert_allclose(
backend.to_numpy(backend.fft.irfft(values, None, axis)),
np.fft.irfft(values_np, None, axis),
)

axes_cases = (
np.array([0]),
np.array([0, 1]),
[np.array(0)],
(np.array(1),),
)
for axes in axes_cases:
npt.assert_array_equal(
backend.to_numpy(backend.fft.fftshift(values, axes)),
np.fft.fftshift(values_np, axes),
)
npt.assert_array_equal(
backend.to_numpy(backend.fft.ifftshift(values, axes)),
np.fft.ifftshift(values_np, axes),
)
npt.assert_allclose(
backend.to_numpy(backend.fft.fftn(values, None, axes)),
np.fft.fftn(values_np, None, axes),
)
npt.assert_allclose(
backend.to_numpy(backend.fft.ifftn(values, None, axes)),
np.fft.ifftn(values_np, None, axes),
)

for axes in ((), []):
npt.assert_array_equal(
backend.to_numpy(backend.fft.fftshift(values, axes)),
values_np,
)
npt.assert_array_equal(
backend.to_numpy(backend.fft.ifftshift(values, axes)),
values_np,
)
npt.assert_array_equal(
backend.to_numpy(backend.fft.fftn(values, None, axes)),
values_np,
)
npt.assert_array_equal(
backend.to_numpy(backend.fft.ifftn(values, None, axes)),
values_np,
)
"""
subprocess.run([sys.executable, "-c", code], check=True, env=env)
Loading