Add norm parameter to FFT transforms (#3287)
Co-authored-by: Cheng <[email protected]>
This commit is contained in:
co-authored by
Cheng
parent
f8eda2c61b
commit
57c813f042
+113
-32
@@ -2,9 +2,11 @@
|
||||
|
||||
#include <nanobind/nanobind.h>
|
||||
#include <nanobind/stl/optional.h>
|
||||
#include <nanobind/stl/string.h>
|
||||
#include <nanobind/stl/variant.h>
|
||||
#include <nanobind/stl/vector.h>
|
||||
#include <numeric>
|
||||
#include <string_view>
|
||||
|
||||
#include "mlx/fft.h"
|
||||
#include "mlx/ops.h"
|
||||
@@ -14,6 +16,25 @@ namespace mx = mlx::core;
|
||||
namespace nb = nanobind;
|
||||
using namespace nb::literals;
|
||||
|
||||
namespace {
|
||||
|
||||
mx::fft::FFTNorm parse_norm(std::string_view norm, std::string_view op) {
|
||||
if (norm == "backward") {
|
||||
return mx::fft::FFTNorm::Backward;
|
||||
}
|
||||
if (norm == "ortho") {
|
||||
return mx::fft::FFTNorm::Ortho;
|
||||
}
|
||||
if (norm == "forward") {
|
||||
return mx::fft::FFTNorm::Forward;
|
||||
}
|
||||
throw std::invalid_argument(
|
||||
std::string("[") + std::string(op) +
|
||||
"] Invalid norm. Expected one of {'backward', 'ortho', 'forward'}.");
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
void init_fft(nb::module_& parent_module) {
|
||||
auto m = parent_module.def_submodule(
|
||||
"fft", "mlx.core.fft: Fast Fourier Transforms.");
|
||||
@@ -22,16 +43,19 @@ void init_fft(nb::module_& parent_module) {
|
||||
[](const mx::array& a,
|
||||
const std::optional<int>& n,
|
||||
int axis,
|
||||
const std::string& norm,
|
||||
mx::StreamOrDevice s) {
|
||||
auto fft_norm = parse_norm(norm, "fft");
|
||||
if (n.has_value()) {
|
||||
return mx::fft::fft(a, n.value(), axis, s);
|
||||
return mx::fft::fft(a, n.value(), axis, fft_norm, s);
|
||||
} else {
|
||||
return mx::fft::fft(a, axis, s);
|
||||
return mx::fft::fft(a, axis, fft_norm, s);
|
||||
}
|
||||
},
|
||||
"a"_a,
|
||||
"n"_a = nb::none(),
|
||||
"axis"_a = -1,
|
||||
"norm"_a = "backward",
|
||||
"stream"_a = nb::none(),
|
||||
R"pbdoc(
|
||||
One dimensional discrete Fourier Transform.
|
||||
@@ -43,6 +67,8 @@ void init_fft(nb::module_& parent_module) {
|
||||
zeros to match ``n``. The default value is ``a.shape[axis]``.
|
||||
axis (int, optional): Axis along which to perform the FFT. The
|
||||
default is ``-1``.
|
||||
norm (str, optional): One of ``"backward"``, ``"ortho"``, or
|
||||
``"forward"``. Default is ``"backward"``.
|
||||
|
||||
Returns:
|
||||
array: The DFT of the input along the given axis.
|
||||
@@ -52,16 +78,19 @@ void init_fft(nb::module_& parent_module) {
|
||||
[](const mx::array& a,
|
||||
const std::optional<int>& n,
|
||||
int axis,
|
||||
const std::string& norm,
|
||||
mx::StreamOrDevice s) {
|
||||
auto fft_norm = parse_norm(norm, "ifft");
|
||||
if (n.has_value()) {
|
||||
return mx::fft::ifft(a, n.value(), axis, s);
|
||||
return mx::fft::ifft(a, n.value(), axis, fft_norm, s);
|
||||
} else {
|
||||
return mx::fft::ifft(a, axis, s);
|
||||
return mx::fft::ifft(a, axis, fft_norm, s);
|
||||
}
|
||||
},
|
||||
"a"_a,
|
||||
"n"_a = nb::none(),
|
||||
"axis"_a = -1,
|
||||
"norm"_a = "backward",
|
||||
"stream"_a = nb::none(),
|
||||
R"pbdoc(
|
||||
One dimensional inverse discrete Fourier Transform.
|
||||
@@ -73,6 +102,8 @@ void init_fft(nb::module_& parent_module) {
|
||||
zeros to match ``n``. The default value is ``a.shape[axis]``.
|
||||
axis (int, optional): Axis along which to perform the FFT. The
|
||||
default is ``-1``.
|
||||
norm (str, optional): One of ``"backward"``, ``"ortho"``, or
|
||||
``"forward"``. Default is ``"backward"``.
|
||||
|
||||
Returns:
|
||||
array: The inverse DFT of the input along the given axis.
|
||||
@@ -82,21 +113,24 @@ void init_fft(nb::module_& parent_module) {
|
||||
[](const mx::array& a,
|
||||
const std::optional<mx::Shape>& n,
|
||||
const std::optional<std::vector<int>>& axes,
|
||||
const std::string& norm,
|
||||
mx::StreamOrDevice s) {
|
||||
auto fft_norm = parse_norm(norm, "fft2");
|
||||
if (axes.has_value() && n.has_value()) {
|
||||
return mx::fft::fftn(a, n.value(), axes.value(), s);
|
||||
return mx::fft::fftn(a, n.value(), axes.value(), fft_norm, s);
|
||||
} else if (axes.has_value()) {
|
||||
return mx::fft::fftn(a, axes.value(), s);
|
||||
return mx::fft::fftn(a, axes.value(), fft_norm, s);
|
||||
} else if (n.has_value()) {
|
||||
throw std::invalid_argument(
|
||||
"[fft2] `axes` should not be `None` if `s` is not `None`.");
|
||||
} else {
|
||||
return mx::fft::fftn(a, s);
|
||||
return mx::fft::fftn(a, fft_norm, s);
|
||||
}
|
||||
},
|
||||
"a"_a,
|
||||
"s"_a = nb::none(),
|
||||
"axes"_a.none() = std::vector<int>{-2, -1},
|
||||
"norm"_a = "backward",
|
||||
"stream"_a = nb::none(),
|
||||
R"pbdoc(
|
||||
Two dimensional discrete Fourier Transform.
|
||||
@@ -109,6 +143,8 @@ void init_fft(nb::module_& parent_module) {
|
||||
sizes of ``a`` along ``axes``.
|
||||
axes (list(int), optional): Axes along which to perform the FFT.
|
||||
The default is ``[-2, -1]``.
|
||||
norm (str, optional): One of ``"backward"``, ``"ortho"``, or
|
||||
``"forward"``. Default is ``"backward"``.
|
||||
|
||||
Returns:
|
||||
array: The DFT of the input along the given axes.
|
||||
@@ -118,21 +154,24 @@ void init_fft(nb::module_& parent_module) {
|
||||
[](const mx::array& a,
|
||||
const std::optional<mx::Shape>& n,
|
||||
const std::optional<std::vector<int>>& axes,
|
||||
const std::string& norm,
|
||||
mx::StreamOrDevice s) {
|
||||
auto fft_norm = parse_norm(norm, "ifft2");
|
||||
if (axes.has_value() && n.has_value()) {
|
||||
return mx::fft::ifftn(a, n.value(), axes.value(), s);
|
||||
return mx::fft::ifftn(a, n.value(), axes.value(), fft_norm, s);
|
||||
} else if (axes.has_value()) {
|
||||
return mx::fft::ifftn(a, axes.value(), s);
|
||||
return mx::fft::ifftn(a, axes.value(), fft_norm, s);
|
||||
} else if (n.has_value()) {
|
||||
throw std::invalid_argument(
|
||||
"[ifft2] `axes` should not be `None` if `s` is not `None`.");
|
||||
} else {
|
||||
return mx::fft::ifftn(a, s);
|
||||
return mx::fft::ifftn(a, fft_norm, s);
|
||||
}
|
||||
},
|
||||
"a"_a,
|
||||
"s"_a = nb::none(),
|
||||
"axes"_a.none() = std::vector<int>{-2, -1},
|
||||
"norm"_a = "backward",
|
||||
"stream"_a = nb::none(),
|
||||
R"pbdoc(
|
||||
Two dimensional inverse discrete Fourier Transform.
|
||||
@@ -145,6 +184,8 @@ void init_fft(nb::module_& parent_module) {
|
||||
sizes of ``a`` along ``axes``.
|
||||
axes (list(int), optional): Axes along which to perform the FFT.
|
||||
The default is ``[-2, -1]``.
|
||||
norm (str, optional): One of ``"backward"``, ``"ortho"``, or
|
||||
``"forward"``. Default is ``"backward"``.
|
||||
|
||||
Returns:
|
||||
array: The inverse DFT of the input along the given axes.
|
||||
@@ -154,21 +195,24 @@ void init_fft(nb::module_& parent_module) {
|
||||
[](const mx::array& a,
|
||||
const std::optional<mx::Shape>& n,
|
||||
const std::optional<std::vector<int>>& axes,
|
||||
const std::string& norm,
|
||||
mx::StreamOrDevice s) {
|
||||
auto fft_norm = parse_norm(norm, "fftn");
|
||||
if (axes.has_value() && n.has_value()) {
|
||||
return mx::fft::fftn(a, n.value(), axes.value(), s);
|
||||
return mx::fft::fftn(a, n.value(), axes.value(), fft_norm, s);
|
||||
} else if (axes.has_value()) {
|
||||
return mx::fft::fftn(a, axes.value(), s);
|
||||
return mx::fft::fftn(a, axes.value(), fft_norm, s);
|
||||
} else if (n.has_value()) {
|
||||
throw std::invalid_argument(
|
||||
"[fftn] `axes` should not be `None` if `s` is not `None`.");
|
||||
} else {
|
||||
return mx::fft::fftn(a, s);
|
||||
return mx::fft::fftn(a, fft_norm, s);
|
||||
}
|
||||
},
|
||||
"a"_a,
|
||||
"s"_a = nb::none(),
|
||||
"axes"_a = nb::none(),
|
||||
"norm"_a = "backward",
|
||||
"stream"_a = nb::none(),
|
||||
R"pbdoc(
|
||||
n-dimensional discrete Fourier Transform.
|
||||
@@ -182,6 +226,8 @@ void init_fft(nb::module_& parent_module) {
|
||||
axes (list(int), optional): Axes along which to perform the FFT.
|
||||
The default is ``None`` in which case the FFT is over the last
|
||||
``len(s)`` axes are or all axes if ``s`` is also ``None``.
|
||||
norm (str, optional): One of ``"backward"``, ``"ortho"``, or
|
||||
``"forward"``. Default is ``"backward"``.
|
||||
|
||||
Returns:
|
||||
array: The DFT of the input along the given axes.
|
||||
@@ -191,21 +237,24 @@ void init_fft(nb::module_& parent_module) {
|
||||
[](const mx::array& a,
|
||||
const std::optional<mx::Shape>& n,
|
||||
const std::optional<std::vector<int>>& axes,
|
||||
const std::string& norm,
|
||||
mx::StreamOrDevice s) {
|
||||
auto fft_norm = parse_norm(norm, "ifftn");
|
||||
if (axes.has_value() && n.has_value()) {
|
||||
return mx::fft::ifftn(a, n.value(), axes.value(), s);
|
||||
return mx::fft::ifftn(a, n.value(), axes.value(), fft_norm, s);
|
||||
} else if (axes.has_value()) {
|
||||
return mx::fft::ifftn(a, axes.value(), s);
|
||||
return mx::fft::ifftn(a, axes.value(), fft_norm, s);
|
||||
} else if (n.has_value()) {
|
||||
throw std::invalid_argument(
|
||||
"[ifftn] `axes` should not be `None` if `s` is not `None`.");
|
||||
} else {
|
||||
return mx::fft::ifftn(a, s);
|
||||
return mx::fft::ifftn(a, fft_norm, s);
|
||||
}
|
||||
},
|
||||
"a"_a,
|
||||
"s"_a = nb::none(),
|
||||
"axes"_a = nb::none(),
|
||||
"norm"_a = "backward",
|
||||
"stream"_a = nb::none(),
|
||||
R"pbdoc(
|
||||
n-dimensional inverse discrete Fourier Transform.
|
||||
@@ -219,6 +268,8 @@ void init_fft(nb::module_& parent_module) {
|
||||
axes (list(int), optional): Axes along which to perform the FFT.
|
||||
The default is ``None`` in which case the FFT is over the last
|
||||
``len(s)`` axes or all axes if ``s`` is also ``None``.
|
||||
norm (str, optional): One of ``"backward"``, ``"ortho"``, or
|
||||
``"forward"``. Default is ``"backward"``.
|
||||
|
||||
Returns:
|
||||
array: The inverse DFT of the input along the given axes.
|
||||
@@ -228,16 +279,19 @@ void init_fft(nb::module_& parent_module) {
|
||||
[](const mx::array& a,
|
||||
const std::optional<int>& n,
|
||||
int axis,
|
||||
const std::string& norm,
|
||||
mx::StreamOrDevice s) {
|
||||
auto fft_norm = parse_norm(norm, "rfft");
|
||||
if (n.has_value()) {
|
||||
return mx::fft::rfft(a, n.value(), axis, s);
|
||||
return mx::fft::rfft(a, n.value(), axis, fft_norm, s);
|
||||
} else {
|
||||
return mx::fft::rfft(a, axis, s);
|
||||
return mx::fft::rfft(a, axis, fft_norm, s);
|
||||
}
|
||||
},
|
||||
"a"_a,
|
||||
"n"_a = nb::none(),
|
||||
"axis"_a = -1,
|
||||
"norm"_a = "backward",
|
||||
"stream"_a = nb::none(),
|
||||
R"pbdoc(
|
||||
One dimensional discrete Fourier Transform on a real input.
|
||||
@@ -253,6 +307,8 @@ void init_fft(nb::module_& parent_module) {
|
||||
zeros to match ``n``. The default value is ``a.shape[axis]``.
|
||||
axis (int, optional): Axis along which to perform the FFT. The
|
||||
default is ``-1``.
|
||||
norm (str, optional): One of ``"backward"``, ``"ortho"``, or
|
||||
``"forward"``. Default is ``"backward"``.
|
||||
|
||||
Returns:
|
||||
array: The DFT of the input along the given axis. The output
|
||||
@@ -263,16 +319,19 @@ void init_fft(nb::module_& parent_module) {
|
||||
[](const mx::array& a,
|
||||
const std::optional<int>& n,
|
||||
int axis,
|
||||
const std::string& norm,
|
||||
mx::StreamOrDevice s) {
|
||||
auto fft_norm = parse_norm(norm, "irfft");
|
||||
if (n.has_value()) {
|
||||
return mx::fft::irfft(a, n.value(), axis, s);
|
||||
return mx::fft::irfft(a, n.value(), axis, fft_norm, s);
|
||||
} else {
|
||||
return mx::fft::irfft(a, axis, s);
|
||||
return mx::fft::irfft(a, axis, fft_norm, s);
|
||||
}
|
||||
},
|
||||
"a"_a,
|
||||
"n"_a = nb::none(),
|
||||
"axis"_a = -1,
|
||||
"norm"_a = "backward",
|
||||
"stream"_a = nb::none(),
|
||||
R"pbdoc(
|
||||
The inverse of :func:`rfft`.
|
||||
@@ -288,6 +347,8 @@ void init_fft(nb::module_& parent_module) {
|
||||
``a.shape[axis] // 2 + 1``.
|
||||
axis (int, optional): Axis along which to perform the FFT. The
|
||||
default is ``-1``.
|
||||
norm (str, optional): One of ``"backward"``, ``"ortho"``, or
|
||||
``"forward"``. Default is ``"backward"``.
|
||||
|
||||
Returns:
|
||||
array: The real array containing the inverse of :func:`rfft`.
|
||||
@@ -297,21 +358,24 @@ void init_fft(nb::module_& parent_module) {
|
||||
[](const mx::array& a,
|
||||
const std::optional<mx::Shape>& n,
|
||||
const std::optional<std::vector<int>>& axes,
|
||||
const std::string& norm,
|
||||
mx::StreamOrDevice s) {
|
||||
auto fft_norm = parse_norm(norm, "rfft2");
|
||||
if (axes.has_value() && n.has_value()) {
|
||||
return mx::fft::rfftn(a, n.value(), axes.value(), s);
|
||||
return mx::fft::rfftn(a, n.value(), axes.value(), fft_norm, s);
|
||||
} else if (axes.has_value()) {
|
||||
return mx::fft::rfftn(a, axes.value(), s);
|
||||
return mx::fft::rfftn(a, axes.value(), fft_norm, s);
|
||||
} else if (n.has_value()) {
|
||||
throw std::invalid_argument(
|
||||
"[rfft2] `axes` should not be `None` if `s` is not `None`.");
|
||||
} else {
|
||||
return mx::fft::rfftn(a, s);
|
||||
return mx::fft::rfftn(a, fft_norm, s);
|
||||
}
|
||||
},
|
||||
"a"_a,
|
||||
"s"_a = nb::none(),
|
||||
"axes"_a.none() = std::vector<int>{-2, -1},
|
||||
"norm"_a = "backward",
|
||||
"stream"_a = nb::none(),
|
||||
R"pbdoc(
|
||||
Two dimensional real discrete Fourier Transform.
|
||||
@@ -329,6 +393,8 @@ void init_fft(nb::module_& parent_module) {
|
||||
sizes of ``a`` along ``axes``.
|
||||
axes (list(int), optional): Axes along which to perform the FFT.
|
||||
The default is ``[-2, -1]``.
|
||||
norm (str, optional): One of ``"backward"``, ``"ortho"``, or
|
||||
``"forward"``. Default is ``"backward"``.
|
||||
|
||||
Returns:
|
||||
array: The real DFT of the input along the given axes. The output
|
||||
@@ -339,21 +405,24 @@ void init_fft(nb::module_& parent_module) {
|
||||
[](const mx::array& a,
|
||||
const std::optional<mx::Shape>& n,
|
||||
const std::optional<std::vector<int>>& axes,
|
||||
const std::string& norm,
|
||||
mx::StreamOrDevice s) {
|
||||
auto fft_norm = parse_norm(norm, "irfft2");
|
||||
if (axes.has_value() && n.has_value()) {
|
||||
return mx::fft::irfftn(a, n.value(), axes.value(), s);
|
||||
return mx::fft::irfftn(a, n.value(), axes.value(), fft_norm, s);
|
||||
} else if (axes.has_value()) {
|
||||
return mx::fft::irfftn(a, axes.value(), s);
|
||||
return mx::fft::irfftn(a, axes.value(), fft_norm, s);
|
||||
} else if (n.has_value()) {
|
||||
throw std::invalid_argument(
|
||||
"[irfft2] `axes` should not be `None` if `s` is not `None`.");
|
||||
} else {
|
||||
return mx::fft::irfftn(a, s);
|
||||
return mx::fft::irfftn(a, fft_norm, s);
|
||||
}
|
||||
},
|
||||
"a"_a,
|
||||
"s"_a = nb::none(),
|
||||
"axes"_a.none() = std::vector<int>{-2, -1},
|
||||
"norm"_a = "backward",
|
||||
"stream"_a = nb::none(),
|
||||
R"pbdoc(
|
||||
The inverse of :func:`rfft2`.
|
||||
@@ -372,6 +441,8 @@ void init_fft(nb::module_& parent_module) {
|
||||
sizes of ``a`` along ``axes``.
|
||||
axes (list(int), optional): Axes along which to perform the FFT.
|
||||
The default is ``[-2, -1]``.
|
||||
norm (str, optional): One of ``"backward"``, ``"ortho"``, or
|
||||
``"forward"``. Default is ``"backward"``.
|
||||
|
||||
Returns:
|
||||
array: The real array containing the inverse of :func:`rfft2`.
|
||||
@@ -381,21 +452,24 @@ void init_fft(nb::module_& parent_module) {
|
||||
[](const mx::array& a,
|
||||
const std::optional<mx::Shape>& n,
|
||||
const std::optional<std::vector<int>>& axes,
|
||||
const std::string& norm,
|
||||
mx::StreamOrDevice s) {
|
||||
auto fft_norm = parse_norm(norm, "rfftn");
|
||||
if (axes.has_value() && n.has_value()) {
|
||||
return mx::fft::rfftn(a, n.value(), axes.value(), s);
|
||||
return mx::fft::rfftn(a, n.value(), axes.value(), fft_norm, s);
|
||||
} else if (axes.has_value()) {
|
||||
return mx::fft::rfftn(a, axes.value(), s);
|
||||
return mx::fft::rfftn(a, axes.value(), fft_norm, s);
|
||||
} else if (n.has_value()) {
|
||||
throw std::invalid_argument(
|
||||
"[rfftn] `axes` should not be `None` if `s` is not `None`.");
|
||||
} else {
|
||||
return mx::fft::rfftn(a, s);
|
||||
return mx::fft::rfftn(a, fft_norm, s);
|
||||
}
|
||||
},
|
||||
"a"_a,
|
||||
"s"_a = nb::none(),
|
||||
"axes"_a = nb::none(),
|
||||
"norm"_a = "backward",
|
||||
"stream"_a = nb::none(),
|
||||
R"pbdoc(
|
||||
n-dimensional real discrete Fourier Transform.
|
||||
@@ -414,6 +488,8 @@ void init_fft(nb::module_& parent_module) {
|
||||
axes (list(int), optional): Axes along which to perform the FFT.
|
||||
The default is ``None`` in which case the FFT is over the last
|
||||
``len(s)`` axes or all axes if ``s`` is also ``None``.
|
||||
norm (str, optional): One of ``"backward"``, ``"ortho"``, or
|
||||
``"forward"``. Default is ``"backward"``.
|
||||
|
||||
Returns:
|
||||
array: The real DFT of the input along the given axes. The output
|
||||
@@ -423,21 +499,24 @@ void init_fft(nb::module_& parent_module) {
|
||||
[](const mx::array& a,
|
||||
const std::optional<mx::Shape>& n,
|
||||
const std::optional<std::vector<int>>& axes,
|
||||
const std::string& norm,
|
||||
mx::StreamOrDevice s) {
|
||||
auto fft_norm = parse_norm(norm, "irfftn");
|
||||
if (axes.has_value() && n.has_value()) {
|
||||
return mx::fft::irfftn(a, n.value(), axes.value(), s);
|
||||
return mx::fft::irfftn(a, n.value(), axes.value(), fft_norm, s);
|
||||
} else if (axes.has_value()) {
|
||||
return mx::fft::irfftn(a, axes.value(), s);
|
||||
return mx::fft::irfftn(a, axes.value(), fft_norm, s);
|
||||
} else if (n.has_value()) {
|
||||
throw std::invalid_argument(
|
||||
"[irfftn] `axes` should not be `None` if `s` is not `None`.");
|
||||
} else {
|
||||
return mx::fft::irfftn(a, s);
|
||||
return mx::fft::irfftn(a, fft_norm, s);
|
||||
}
|
||||
},
|
||||
"a"_a,
|
||||
"s"_a = nb::none(),
|
||||
"axes"_a = nb::none(),
|
||||
"norm"_a = "backward",
|
||||
"stream"_a = nb::none(),
|
||||
R"pbdoc(
|
||||
The inverse of :func:`rfftn`.
|
||||
@@ -456,6 +535,8 @@ void init_fft(nb::module_& parent_module) {
|
||||
axes (list(int), optional): Axes along which to perform the FFT.
|
||||
The default is ``None`` in which case the FFT is over the last
|
||||
``len(s)`` axes or all axes if ``s`` is also ``None``.
|
||||
norm (str, optional): One of ``"backward"``, ``"ortho"``, or
|
||||
``"forward"``. Default is ``"backward"``.
|
||||
|
||||
Returns:
|
||||
array: The real array containing the inverse of :func:`rfftn`.
|
||||
|
||||
Reference in New Issue
Block a user