Current section
Files
Jump to
Current section
Files
c_src/mlx_fft_nif.cpp
#include <erl_nif.h>
#include <mlx/mlx.h>
#include <mlx/ops.h>
#include <mlx/array.h>
#include <mlx/fft.h>
#include <memory>
#include <vector>
using namespace mlx::core;
namespace mx = mlx::core;
namespace fft = mlx::core::fft;
// Resource types for FFT operations
static ErlNifResourceType* ARRAY_RESOURCE_TYPE;
struct ArrayResource {
array arr = array({1.0f}); // Initialize with dummy value
std::string name;
ArrayResource(const array& a, const std::string& n = "") : arr(a), name(n) {}
};
// Helper functions
static ERL_NIF_TERM make_atom(ErlNifEnv* env, const char* name) {
ERL_NIF_TERM ret;
if (enif_make_existing_atom(env, name, &ret, ERL_NIF_LATIN1)) {
return ret;
}
return enif_make_atom(env, name);
}
static ERL_NIF_TERM make_error(ErlNifEnv* env, const char* reason) {
return enif_make_tuple2(env, make_atom(env, "error"), make_atom(env, reason));
}
static ERL_NIF_TERM make_ok(ErlNifEnv* env, ERL_NIF_TERM term) {
return enif_make_tuple2(env, make_atom(env, "ok"), term);
}
// Parse shape from Erlang list
static std::vector<int> parse_shape(ErlNifEnv* env, ERL_NIF_TERM shape_term) {
unsigned int shape_len;
if (!enif_get_list_length(env, shape_term, &shape_len)) {
return {};
}
std::vector<int> shape_vec(shape_len);
ERL_NIF_TERM head, tail = shape_term;
for (unsigned int i = 0; i < shape_len; i++) {
if (!enif_get_list_cell(env, tail, &head, &tail)) {
return {};
}
if (!enif_get_int(env, head, &shape_vec[i])) {
return {};
}
}
return shape_vec;
}
// Get array from resource
static bool get_array_resource(ErlNifEnv* env, ERL_NIF_TERM term, array& arr) {
ArrayResource* res;
if (!enif_get_resource(env, term, ARRAY_RESOURCE_TYPE, (void**)&res)) {
return false;
}
arr = res->arr;
return true;
}
// Create array resource
static ERL_NIF_TERM make_array_resource(ErlNifEnv* env, const array& arr, const std::string& name = "") {
ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
new(res) ArrayResource(arr, name);
ERL_NIF_TERM term = enif_make_resource(env, res);
enif_release_resource(res);
return make_ok(env, term);
}
// ==================== 1D FFT OPERATIONS ====================
static ERL_NIF_TERM mlx_fft(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
int n = -1; // Default: use input size
int axis = -1; // Default: last axis
if (argc > 1 && !enif_get_int(env, argv[1], &n)) {
return enif_make_badarg(env);
}
if (argc > 2 && !enif_get_int(env, argv[2], &axis)) {
return enif_make_badarg(env);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (n == -1 && axis == -1) {
result = fft::fft(a);
} else if (n == -1) {
result = fft::fft(a, axis);
} else if (axis == -1) {
result = fft::fft(a, n);
} else {
result = fft::fft(a, n, axis);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "fft_error");
}
}
static ERL_NIF_TERM mlx_ifft(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
int n = -1;
int axis = -1;
if (argc > 1 && !enif_get_int(env, argv[1], &n)) {
return enif_make_badarg(env);
}
if (argc > 2 && !enif_get_int(env, argv[2], &axis)) {
return enif_make_badarg(env);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (n == -1 && axis == -1) {
result = fft::ifft(a);
} else if (n == -1) {
result = fft::ifft(a, axis);
} else if (axis == -1) {
result = fft::ifft(a, n);
} else {
result = fft::ifft(a, n, axis);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "ifft_error");
}
}
// ==================== REAL FFT OPERATIONS ====================
static ERL_NIF_TERM mlx_rfft(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
int n = -1;
int axis = -1;
if (argc > 1 && !enif_get_int(env, argv[1], &n)) {
return enif_make_badarg(env);
}
if (argc > 2 && !enif_get_int(env, argv[2], &axis)) {
return enif_make_badarg(env);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (n == -1 && axis == -1) {
result = fft::rfft(a);
} else if (n == -1) {
result = fft::rfft(a, axis);
} else if (axis == -1) {
result = fft::rfft(a, n);
} else {
result = fft::rfft(a, n, axis);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "rfft_error");
}
}
static ERL_NIF_TERM mlx_irfft(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
int n = -1;
int axis = -1;
if (argc > 1 && !enif_get_int(env, argv[1], &n)) {
return enif_make_badarg(env);
}
if (argc > 2 && !enif_get_int(env, argv[2], &axis)) {
return enif_make_badarg(env);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (n == -1 && axis == -1) {
result = fft::irfft(a);
} else if (n == -1) {
result = fft::irfft(a, axis);
} else if (axis == -1) {
result = fft::irfft(a, n);
} else {
result = fft::irfft(a, n, axis);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "irfft_error");
}
}
// ==================== 2D FFT OPERATIONS ====================
static ERL_NIF_TERM mlx_fft2(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
std::vector<int> s;
std::vector<int> axes;
if (argc > 1) {
s = parse_shape(env, argv[1]);
}
if (argc > 2) {
axes = parse_shape(env, argv[2]);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (s.empty() && axes.empty()) {
result = fft::fft2(a);
} else if (axes.empty()) {
result = fft::fft2(a, s);
} else if (s.empty()) {
result = fft::fft2(a, {}, axes);
} else {
result = fft::fft2(a, s, axes);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "fft2_error");
}
}
static ERL_NIF_TERM mlx_ifft2(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
std::vector<int> s;
std::vector<int> axes;
if (argc > 1) {
s = parse_shape(env, argv[1]);
}
if (argc > 2) {
axes = parse_shape(env, argv[2]);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (s.empty() && axes.empty()) {
result = fft::ifft2(a);
} else if (axes.empty()) {
result = fft::ifft2(a, s);
} else if (s.empty()) {
result = fft::ifft2(a, {}, axes);
} else {
result = fft::ifft2(a, s, axes);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "ifft2_error");
}
}
static ERL_NIF_TERM mlx_rfft2(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
std::vector<int> s;
std::vector<int> axes;
if (argc > 1) {
s = parse_shape(env, argv[1]);
}
if (argc > 2) {
axes = parse_shape(env, argv[2]);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (s.empty() && axes.empty()) {
result = fft::rfft2(a);
} else if (axes.empty()) {
result = fft::rfft2(a, s);
} else if (s.empty()) {
result = fft::rfft2(a, {}, axes);
} else {
result = fft::rfft2(a, s, axes);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "rfft2_error");
}
}
static ERL_NIF_TERM mlx_irfft2(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
std::vector<int> s;
std::vector<int> axes;
if (argc > 1) {
s = parse_shape(env, argv[1]);
}
if (argc > 2) {
axes = parse_shape(env, argv[2]);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (s.empty() && axes.empty()) {
result = fft::irfft2(a);
} else if (axes.empty()) {
result = fft::irfft2(a, s);
} else if (s.empty()) {
result = fft::irfft2(a, {}, axes);
} else {
result = fft::irfft2(a, s, axes);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "irfft2_error");
}
}
// ==================== N-D FFT OPERATIONS ====================
static ERL_NIF_TERM mlx_fftn(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
std::vector<int> s;
std::vector<int> axes;
if (argc > 1) {
s = parse_shape(env, argv[1]);
}
if (argc > 2) {
axes = parse_shape(env, argv[2]);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (s.empty() && axes.empty()) {
result = fft::fftn(a);
} else if (axes.empty()) {
result = fft::fftn(a, s);
} else if (s.empty()) {
result = fft::fftn(a, {}, axes);
} else {
result = fft::fftn(a, s, axes);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "fftn_error");
}
}
static ERL_NIF_TERM mlx_ifftn(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
std::vector<int> s;
std::vector<int> axes;
if (argc > 1) {
s = parse_shape(env, argv[1]);
}
if (argc > 2) {
axes = parse_shape(env, argv[2]);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (s.empty() && axes.empty()) {
result = fft::ifftn(a);
} else if (axes.empty()) {
result = fft::ifftn(a, s);
} else if (s.empty()) {
result = fft::ifftn(a, {}, axes);
} else {
result = fft::ifftn(a, s, axes);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "ifftn_error");
}
}
static ERL_NIF_TERM mlx_rfftn(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
std::vector<int> s;
std::vector<int> axes;
if (argc > 1) {
s = parse_shape(env, argv[1]);
}
if (argc > 2) {
axes = parse_shape(env, argv[2]);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (s.empty() && axes.empty()) {
result = fft::rfftn(a);
} else if (axes.empty()) {
result = fft::rfftn(a, s);
} else if (s.empty()) {
result = fft::rfftn(a, {}, axes);
} else {
result = fft::rfftn(a, s, axes);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "rfftn_error");
}
}
static ERL_NIF_TERM mlx_irfftn(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
std::vector<int> s;
std::vector<int> axes;
if (argc > 1) {
s = parse_shape(env, argv[1]);
}
if (argc > 2) {
axes = parse_shape(env, argv[2]);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (s.empty() && axes.empty()) {
result = fft::irfftn(a);
} else if (axes.empty()) {
result = fft::irfftn(a, s);
} else if (s.empty()) {
result = fft::irfftn(a, {}, axes);
} else {
result = fft::irfftn(a, s, axes);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "irfftn_error");
}
}
// ==================== FFT FREQUENCY UTILITIES ====================
static ERL_NIF_TERM mlx_fftfreq(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
int n;
double d;
if (!enif_get_int(env, argv[0], &n) ||
!enif_get_double(env, argv[1], &d)) {
return enif_make_badarg(env);
}
try {
// TODO: MLX may add fftfreq in future versions
return make_error(env, "fftfreq_not_yet_available");
} catch (const std::exception& e) {
return make_error(env, "fftfreq_error");
}
}
static ERL_NIF_TERM mlx_rfftfreq(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
int n;
double d;
if (!enif_get_int(env, argv[0], &n) ||
!enif_get_double(env, argv[1], &d)) {
return enif_make_badarg(env);
}
try {
// TODO: MLX may add rfftfreq in future versions
return make_error(env, "rfftfreq_not_yet_available");
} catch (const std::exception& e) {
return make_error(env, "rfftfreq_error");
}
}
// ==================== FFT SHIFT OPERATIONS ====================
static ERL_NIF_TERM mlx_fftshift(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 2) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
std::vector<int> axes;
if (argc > 1) {
axes = parse_shape(env, argv[1]);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (axes.empty()) {
result = fft::fftshift(a);
} else {
result = fft::fftshift(a, axes);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "fftshift_error");
}
}
static ERL_NIF_TERM mlx_ifftshift(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 2) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
std::vector<int> axes;
if (argc > 1) {
axes = parse_shape(env, argv[1]);
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (axes.empty()) {
result = fft::ifftshift(a);
} else {
result = fft::ifftshift(a, axes);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "ifftshift_error");
}
}
// Resource destructor
static void array_resource_destructor(ErlNifEnv* env, void* obj) {
ArrayResource* res = (ArrayResource*)obj;
res->~ArrayResource();
}
// FFT NIF function table
static ErlNifFunc nif_funcs[] = {
// 1D FFT
{"fft", 1, mlx_fft, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"fft", 2, mlx_fft, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"fft", 3, mlx_fft, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"ifft", 1, mlx_ifft, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"ifft", 2, mlx_ifft, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"ifft", 3, mlx_ifft, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Real FFT
{"rfft", 1, mlx_rfft, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"rfft", 2, mlx_rfft, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"rfft", 3, mlx_rfft, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"irfft", 1, mlx_irfft, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"irfft", 2, mlx_irfft, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"irfft", 3, mlx_irfft, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// 2D FFT
{"fft2", 1, mlx_fft2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"fft2", 2, mlx_fft2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"fft2", 3, mlx_fft2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"ifft2", 1, mlx_ifft2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"ifft2", 2, mlx_ifft2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"ifft2", 3, mlx_ifft2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"rfft2", 1, mlx_rfft2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"rfft2", 2, mlx_rfft2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"rfft2", 3, mlx_rfft2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"irfft2", 1, mlx_irfft2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"irfft2", 2, mlx_irfft2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"irfft2", 3, mlx_irfft2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// N-D FFT
{"fftn", 1, mlx_fftn, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"fftn", 2, mlx_fftn, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"fftn", 3, mlx_fftn, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"ifftn", 1, mlx_ifftn, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"ifftn", 2, mlx_ifftn, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"ifftn", 3, mlx_ifftn, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"rfftn", 1, mlx_rfftn, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"rfftn", 2, mlx_rfftn, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"rfftn", 3, mlx_rfftn, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"irfftn", 1, mlx_irfftn, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"irfftn", 2, mlx_irfftn, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"irfftn", 3, mlx_irfftn, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Frequency utilities
{"fftfreq", 2, mlx_fftfreq, 0},
{"rfftfreq", 2, mlx_rfftfreq, 0},
// Shift operations
{"fftshift", 1, mlx_fftshift, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"fftshift", 2, mlx_fftshift, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"ifftshift", 1, mlx_ifftshift, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"ifftshift", 2, mlx_ifftshift, ERL_NIF_DIRTY_JOB_CPU_BOUND}
};
static int load(ErlNifEnv* env, void** priv_data, ERL_NIF_TERM load_info) {
ErlNifResourceFlags flags = (ErlNifResourceFlags)(ERL_NIF_RT_CREATE | ERL_NIF_RT_TAKEOVER);
ErlNifResourceFlags* tried = NULL;
ARRAY_RESOURCE_TYPE = enif_open_resource_type(
env, NULL, "mlx_fft_array", array_resource_destructor, flags, tried);
if (!ARRAY_RESOURCE_TYPE) {
return 1;
}
return 0;
}
static int upgrade(ErlNifEnv* env, void** priv_data, void** old_priv_data, ERL_NIF_TERM load_info) {
return load(env, priv_data, load_info);
}
static void unload(ErlNifEnv* env, void* priv_data) {
// Cleanup if needed
}
ERL_NIF_INIT(mlx_fft_nif, nif_funcs, load, NULL, upgrade, unload)