Add fftfreq, rfftfreq and scalar axes for fftshift/ifftshift (#3298)

This commit is contained in:
declanhealy2
2026-03-31 18:29:16 -07:00
committed by GitHub
parent 1944cf67a2
commit 2105df91da
5 changed files with 140 additions and 14 deletions
+55 -14
View File
@@ -11,6 +11,7 @@
#include "mlx/fft.h"
#include "mlx/ops.h"
#include "python/src/small_vector.h"
#include "python/src/utils.h"
namespace mx = mlx::core;
namespace nb = nanobind;
@@ -541,15 +542,55 @@ void init_fft(nb::module_& parent_module) {
Returns:
array: The real array containing the inverse of :func:`rfftn`.
)pbdoc");
m.def(
"fftfreq",
[](int n, double d, mx::StreamOrDevice s) {
return mx::fft::fftfreq(n, d, s);
},
"n"_a,
"d"_a = 1.0,
"stream"_a = nb::none(),
R"pbdoc(
Return the discrete Fourier Transform sample frequencies.
Args:
n (int): Window length.
d (float, optional): Sample spacing. The default is ``1.0``.
Returns:
array: The sample frequencies as a one-dimensional array of type ``float32``.
)pbdoc");
m.def(
"rfftfreq",
[](int n, double d, mx::StreamOrDevice s) {
return mx::fft::rfftfreq(n, d, s);
},
"n"_a,
"d"_a = 1.0,
"stream"_a = nb::none(),
R"pbdoc(
Return the discrete Fourier Transform sample frequencies
for use with :func:`rfft` and :func:`irfft`.
The returned array contains the non-negative frequency terms
in the range ``[0, floor(n/2)]``.
Args:
n (int): Window length.
d (float, optional): Sample spacing. The default is ``1.0``.
Returns:
array: The sample frequencies as a one-dimensional array of type ``float32``.
)pbdoc");
m.def(
"fftshift",
[](const mx::array& a,
const std::optional<std::vector<int>>& axes,
mx::StreamOrDevice s) {
if (axes.has_value()) {
return mx::fft::fftshift(a, axes.value(), s);
} else {
[](const mx::array& a, const IntOrVec& axes, mx::StreamOrDevice s) {
if (std::holds_alternative<std::monostate>(axes)) {
return mx::fft::fftshift(a, s);
} else if (auto pv = std::get_if<int>(&axes); pv) {
return mx::fft::fftshift(a, {*pv}, s);
} else {
return mx::fft::fftshift(a, std::get<std::vector<int>>(axes), s);
}
},
"a"_a,
@@ -560,7 +601,7 @@ void init_fft(nb::module_& parent_module) {
Args:
a (array): The input array.
axes (list(int), optional): Axes over which to perform the shift.
axes (int or list(int), optional): Axis or axes over which to perform the shift.
If ``None``, shift all axes.
Returns:
@@ -568,13 +609,13 @@ void init_fft(nb::module_& parent_module) {
)pbdoc");
m.def(
"ifftshift",
[](const mx::array& a,
const std::optional<std::vector<int>>& axes,
mx::StreamOrDevice s) {
if (axes.has_value()) {
return mx::fft::ifftshift(a, axes.value(), s);
} else {
[](const mx::array& a, const IntOrVec& axes, mx::StreamOrDevice s) {
if (std::holds_alternative<std::monostate>(axes)) {
return mx::fft::ifftshift(a, s);
} else if (auto pv = std::get_if<int>(&axes); pv) {
return mx::fft::ifftshift(a, {*pv}, s);
} else {
return mx::fft::ifftshift(a, std::get<std::vector<int>>(axes), s);
}
},
"a"_a,
@@ -586,7 +627,7 @@ void init_fft(nb::module_& parent_module) {
Args:
a (array): The input array.
axes (list(int), optional): Axes over which to perform the inverse shift.
axes (int or list(int), optional): Axis or axes over which to perform the inverse shift.
If ``None``, shift all axes.
Returns:
+48
View File
@@ -246,6 +246,48 @@ class TestFFT(mlx_tests.MLXTestCase):
with self.assertRaises(ValueError):
mx.fft.fft(mx.array([1.0]), norm="invalid")
def test_fftfreq(self):
for n, d in [(1, 1.0), (4, 0.5), (5, 0.25), (8, -0.5), (6, 1.0)]:
out = mx.fft.fftfreq(n, d=d)
expected = np.fft.fftfreq(n, d=d).astype(np.float32)
self.assertEqual(out.dtype, mx.float32)
np.testing.assert_allclose(out, expected, atol=0.0, rtol=0.0)
with self.assertRaises(ValueError):
mx.fft.fftfreq(0)
with self.assertRaises(ValueError):
mx.fft.fftfreq(-1)
# Test default d=1.0
out = mx.fft.fftfreq(8)
expected = np.fft.fftfreq(8).astype(np.float32)
np.testing.assert_allclose(out, expected, atol=0.0, rtol=0.0)
with self.assertRaises(ValueError):
mx.fft.fftfreq(4, d=0.0)
def test_rfftfreq(self):
for n, d in [(1, 1.0), (4, 0.5), (5, 0.25), (8, -0.5), (6, 1.0)]:
out = mx.fft.rfftfreq(n, d=d)
expected = np.fft.rfftfreq(n, d=d).astype(np.float32)
self.assertEqual(out.dtype, mx.float32)
np.testing.assert_allclose(out, expected, atol=0.0, rtol=0.0)
# Test default d=1.0
out = mx.fft.rfftfreq(8)
expected = np.fft.rfftfreq(8).astype(np.float32)
np.testing.assert_allclose(out, expected, atol=0.0, rtol=0.0)
with self.assertRaises(ValueError):
mx.fft.rfftfreq(0)
with self.assertRaises(ValueError):
mx.fft.rfftfreq(-1)
with self.assertRaises(ValueError):
mx.fft.rfftfreq(4, d=0.0)
def test_fftshift(self):
# Test 1D arrays
r = np.random.rand(100).astype(np.float32)
@@ -253,11 +295,14 @@ class TestFFT(mlx_tests.MLXTestCase):
# Test with specific axis
r = np.random.rand(4, 6).astype(np.float32)
self.check_mx_np(mx.fft.fftshift, np.fft.fftshift, r, axes=0)
self.check_mx_np(mx.fft.fftshift, np.fft.fftshift, r, axes=1)
self.check_mx_np(mx.fft.fftshift, np.fft.fftshift, r, axes=[0])
self.check_mx_np(mx.fft.fftshift, np.fft.fftshift, r, axes=[1])
self.check_mx_np(mx.fft.fftshift, np.fft.fftshift, r, axes=[0, 1])
# Test with negative axes
self.check_mx_np(mx.fft.fftshift, np.fft.fftshift, r, axes=-1)
self.check_mx_np(mx.fft.fftshift, np.fft.fftshift, r, axes=[-1])
# Test with odd lengths
@@ -278,11 +323,14 @@ class TestFFT(mlx_tests.MLXTestCase):
# Test with specific axis
r = np.random.rand(4, 6).astype(np.float32)
self.check_mx_np(mx.fft.ifftshift, np.fft.ifftshift, r, axes=0)
self.check_mx_np(mx.fft.ifftshift, np.fft.ifftshift, r, axes=1)
self.check_mx_np(mx.fft.ifftshift, np.fft.ifftshift, r, axes=[0])
self.check_mx_np(mx.fft.ifftshift, np.fft.ifftshift, r, axes=[1])
self.check_mx_np(mx.fft.ifftshift, np.fft.ifftshift, r, axes=[0, 1])
# Test with negative axes
self.check_mx_np(mx.fft.ifftshift, np.fft.ifftshift, r, axes=-1)
self.check_mx_np(mx.fft.ifftshift, np.fft.ifftshift, r, axes=[-1])
# Test with odd lengths