Add fftfreq, rfftfreq and scalar axes for fftshift/ifftshift (#3298)
This commit is contained in:
+55
-14
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user