Packages

MLX machine learning framework bindings for Erlang

Current section

Files

Jump to
mlx c_src mlx_nif.cpp
Raw

c_src/mlx_nif.cpp

#include <erl_nif.h>
#include <mlx/mlx.h>
#include <mlx/ops.h>
#include <mlx/array.h>
#include <mlx/device.h>
#include <mlx/dtype.h>
#include <mlx/transforms.h>
#include <mlx/random.h>
#include <mlx/linalg.h>
#include <mlx/fft.h>
#include <mlx/io.h>
#include <memory>
#include <vector>
#include <map>
#include <string>
#include <iostream>
#include <functional>
#include <cmath>
using namespace mlx::core;
namespace mx = mlx::core;
// Resource types for comprehensive MLX functionality
static ErlNifResourceType* ARRAY_RESOURCE_TYPE;
// Array resource wrapper
struct ArrayResource {
array arr;
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;
}
// Parse dtype from string
static Dtype parse_dtype(const char* dtype_str) {
if (strcmp(dtype_str, "float32") == 0) return float32;
else if (strcmp(dtype_str, "float16") == 0) return float16;
else if (strcmp(dtype_str, "bfloat16") == 0) return bfloat16;
else if (strcmp(dtype_str, "float64") == 0) return float64;
else if (strcmp(dtype_str, "int32") == 0) return int32;
else if (strcmp(dtype_str, "int16") == 0) return int16;
else if (strcmp(dtype_str, "int8") == 0) return int8;
else if (strcmp(dtype_str, "int64") == 0) return int64;
else if (strcmp(dtype_str, "uint32") == 0) return uint32;
else if (strcmp(dtype_str, "uint16") == 0) return uint16;
else if (strcmp(dtype_str, "uint8") == 0) return uint8;
else if (strcmp(dtype_str, "uint64") == 0) return uint64;
else if (strcmp(dtype_str, "bool") == 0) return bool_;
else if (strcmp(dtype_str, "complex64") == 0) return complex64;
// complex128 is not supported in MLX, map to complex64
else if (strcmp(dtype_str, "complex128") == 0) return complex64;
else return float32;
}
// 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 term;
}
// 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;
}
// Get array from resource (by value)
static array get_array_resource_val(ErlNifEnv* env, ERL_NIF_TERM term, bool& success) {
ArrayResource* res;
if (!enif_get_resource(env, term, ARRAY_RESOURCE_TYPE, (void**)&res)) {
success = false;
return array({0.0f}); // Return a dummy array
}
success = true;
return res->arr;
}
// ==================== ARRAY CREATION OPERATIONS ====================
static ERL_NIF_TERM mlx_arange(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
double start, stop, step;
if (!enif_get_double(env, argv[0], &start) ||
!enif_get_double(env, argv[1], &stop) ||
!enif_get_double(env, argv[2], &step)) {
return enif_make_badarg(env);
}
try {
array result = arange(start, stop, step);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "arange_error");
}
}
static ERL_NIF_TERM mlx_linspace(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
double start, stop;
int num;
if (!enif_get_double(env, argv[0], &start) ||
!enif_get_double(env, argv[1], &stop) ||
!enif_get_int(env, argv[2], &num)) {
return enif_make_badarg(env);
}
try {
array result = linspace(start, stop, num);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "linspace_error");
}
}
static ERL_NIF_TERM mlx_zeros(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
auto shape = parse_shape(env, argv[0]);
if (shape.empty()) return enif_make_badarg(env);
char dtype_str[32];
if (!enif_get_atom(env, argv[1], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
Dtype dtype = parse_dtype(dtype_str);
array result = zeros(shape, dtype);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "zeros_error");
}
}
static ERL_NIF_TERM mlx_ones(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
auto shape = parse_shape(env, argv[0]);
if (shape.empty()) return enif_make_badarg(env);
char dtype_str[32];
if (!enif_get_atom(env, argv[1], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
Dtype dtype = parse_dtype(dtype_str);
array result = ones(shape, dtype);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "ones_error");
}
}
static ERL_NIF_TERM mlx_full(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
auto shape = parse_shape(env, argv[0]);
if (shape.empty()) return enif_make_badarg(env);
double value;
if (!enif_get_double(env, argv[1], &value)) {
return enif_make_badarg(env);
}
char dtype_str[32];
if (!enif_get_atom(env, argv[2], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
Dtype dtype = parse_dtype(dtype_str);
array result = full(shape, array(value), dtype);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "full_error");
}
}
static ERL_NIF_TERM mlx_eye(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
int n;
if (!enif_get_int(env, argv[0], &n)) {
return enif_make_badarg(env);
}
char dtype_str[32];
if (!enif_get_atom(env, argv[1], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
Dtype dtype = parse_dtype(dtype_str);
array result = eye(n, dtype);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "eye_error");
}
}
// Helper function to parse nested lists into a flat vector and determine shape
static bool parse_nested_list(ErlNifEnv* env, ERL_NIF_TERM term, std::vector<float>& data, std::vector<int>& shape, int depth = 0) {
unsigned int list_len;
if (!enif_get_list_length(env, term, &list_len)) {
// Not a list, must be a number
double val;
int int_val;
if (enif_get_double(env, term, &val)) {
data.push_back(static_cast<float>(val));
return true;
} else if (enif_get_int(env, term, &int_val)) {
data.push_back(static_cast<float>(int_val));
return true;
}
return false;
}
if (depth >= shape.size()) {
shape.push_back(list_len);
} else if (shape[depth] != static_cast<int>(list_len)) {
// Shape mismatch
return false;
}
ERL_NIF_TERM head, tail = term;
for (unsigned int i = 0; i < list_len; i++) {
if (!enif_get_list_cell(env, tail, &head, &tail)) {
return false;
}
if (!parse_nested_list(env, head, data, shape, depth + 1)) {
return false;
}
}
return true;
}
static ERL_NIF_TERM mlx_array(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
// Parse the data
std::vector<float> data;
std::vector<int> shape;
// Check if it's a scalar
double scalar_val;
int int_val;
if (enif_get_double(env, argv[0], &scalar_val)) {
data.push_back(static_cast<float>(scalar_val));
// Scalar has empty shape
} else if (enif_get_int(env, argv[0], &int_val)) {
data.push_back(static_cast<float>(int_val));
// Scalar has empty shape
} else {
// Parse nested list
if (!parse_nested_list(env, argv[0], data, shape)) {
return enif_make_badarg(env);
}
}
// Parse dtype
char dtype_str[32];
if (!enif_get_atom(env, argv[1], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
Dtype dtype = parse_dtype(dtype_str);
// Create array from data
if (shape.empty()) {
// Scalar
array result = array(data[0], dtype);
return make_array_resource(env, result);
} else {
// Create array with specific shape
array result = array(data.data(), shape, dtype);
return make_array_resource(env, result);
}
} catch (const std::exception& e) {
return make_error(env, "array_error");
}
}
// ==================== ARITHMETIC OPERATIONS ====================
static ERL_NIF_TERM mlx_add(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = add(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "add_error");
}
}
static ERL_NIF_TERM mlx_subtract(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = subtract(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "subtract_error");
}
}
static ERL_NIF_TERM mlx_multiply(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = multiply(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "multiply_error");
}
}
static ERL_NIF_TERM mlx_divide(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = divide(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "divide_error");
}
}
static ERL_NIF_TERM mlx_power(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = power(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "power_error");
}
}
static ERL_NIF_TERM mlx_negative(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = negative(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "negative_error");
}
}
// ==================== TRIGONOMETRIC FUNCTIONS ====================
static ERL_NIF_TERM mlx_sin(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = sin(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "sin_error");
}
}
static ERL_NIF_TERM mlx_cos(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = cos(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "cos_error");
}
}
static ERL_NIF_TERM mlx_tan(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = tan(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "tan_error");
}
}
static ERL_NIF_TERM mlx_arcsin(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = arcsin(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "arcsin_error");
}
}
static ERL_NIF_TERM mlx_arccos(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = arccos(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "arccos_error");
}
}
static ERL_NIF_TERM mlx_arctan(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = arctan(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "arctan_error");
}
}
// ==================== HYPERBOLIC FUNCTIONS ====================
static ERL_NIF_TERM mlx_sinh(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = sinh(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "sinh_error");
}
}
static ERL_NIF_TERM mlx_cosh(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = cosh(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "cosh_error");
}
}
static ERL_NIF_TERM mlx_tanh(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = tanh(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "tanh_error");
}
}
// ==================== EXPONENTIAL AND LOGARITHMIC ====================
static ERL_NIF_TERM mlx_exp(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = exp(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "exp_error");
}
}
static ERL_NIF_TERM mlx_log(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = log(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "log_error");
}
}
static ERL_NIF_TERM mlx_log2(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = log2(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "log2_error");
}
}
static ERL_NIF_TERM mlx_log10(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = log10(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "log10_error");
}
}
static ERL_NIF_TERM mlx_sqrt(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = sqrt(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "sqrt_error");
}
}
static ERL_NIF_TERM mlx_square(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = square(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "square_error");
}
}
// ==================== COMPARISON OPERATIONS ====================
static ERL_NIF_TERM mlx_equal(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = equal(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "equal_error");
}
}
static ERL_NIF_TERM mlx_not_equal(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = not_equal(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "not_equal_error");
}
}
static ERL_NIF_TERM mlx_greater(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = greater(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "greater_error");
}
}
static ERL_NIF_TERM mlx_greater_equal(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = greater_equal(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "greater_equal_error");
}
}
static ERL_NIF_TERM mlx_less(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = less(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "less_error");
}
}
static ERL_NIF_TERM mlx_less_equal(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = less_equal(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "less_equal_error");
}
}
// ==================== LOGICAL OPERATIONS ====================
static ERL_NIF_TERM mlx_logical_and(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = logical_and(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "logical_and_error");
}
}
static ERL_NIF_TERM mlx_logical_or(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = logical_or(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "logical_or_error");
}
}
static ERL_NIF_TERM mlx_logical_not(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = logical_not(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "logical_not_error");
}
}
// ==================== REDUCTION OPERATIONS ====================
static ERL_NIF_TERM mlx_sum(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = sum(a); // Default case
if (argc == 1) {
// already initialized with sum(a)
} else if (argc == 2) {
// Sum along specific axes
auto axes = parse_shape(env, argv[1]);
result = sum(a, axes);
} else {
// Sum with keepdims
auto axes = parse_shape(env, argv[1]);
int keepdims;
if (!enif_get_int(env, argv[2], &keepdims)) {
return enif_make_badarg(env);
}
result = sum(a, axes, keepdims != 0);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "sum_error");
}
}
static ERL_NIF_TERM mlx_mean(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = mean(a); // Default case
if (argc == 1) {
// already initialized with mean(a)
} else if (argc == 2) {
auto axes = parse_shape(env, argv[1]);
result = mean(a, axes);
} else {
auto axes = parse_shape(env, argv[1]);
int keepdims;
if (!enif_get_int(env, argv[2], &keepdims)) {
return enif_make_badarg(env);
}
result = mean(a, axes, keepdims != 0);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "mean_error");
}
}
static ERL_NIF_TERM mlx_max(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = max(a); // Default case
if (argc == 1) {
// already initialized with max(a)
} else if (argc == 2) {
auto axes = parse_shape(env, argv[1]);
result = max(a, axes);
} else {
auto axes = parse_shape(env, argv[1]);
int keepdims;
if (!enif_get_int(env, argv[2], &keepdims)) {
return enif_make_badarg(env);
}
result = max(a, axes, keepdims != 0);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "max_error");
}
}
static ERL_NIF_TERM mlx_min(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = min(a); // Initialize with default case
if (argc == 1) {
// Already initialized
} else if (argc == 2) {
auto axes = parse_shape(env, argv[1]);
result = min(a, axes);
} else {
auto axes = parse_shape(env, argv[1]);
int keepdims;
if (!enif_get_int(env, argv[2], &keepdims)) {
return enif_make_badarg(env);
}
result = min(a, axes, keepdims != 0);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "min_error");
}
}
static ERL_NIF_TERM mlx_var(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 4) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = var(a); // Initialize with default case
if (argc == 1) {
// Already initialized
} else {
auto axes = parse_shape(env, argv[1]);
int keepdims = 0, ddof = 0;
if (argc > 2 && !enif_get_int(env, argv[2], &keepdims)) {
return enif_make_badarg(env);
}
if (argc > 3 && !enif_get_int(env, argv[3], &ddof)) {
return enif_make_badarg(env);
}
result = var(a, axes, keepdims != 0, ddof);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "var_error");
}
}
static ERL_NIF_TERM mlx_std(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 4) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = mx::std(a); // Initialize with default case
if (argc == 1) {
// Already initialized
} else {
auto axes = parse_shape(env, argv[1]);
int keepdims = 0, ddof = 0;
if (argc > 2 && !enif_get_int(env, argv[2], &keepdims)) {
return enif_make_badarg(env);
}
if (argc > 3 && !enif_get_int(env, argv[3], &ddof)) {
return enif_make_badarg(env);
}
result = mx::std(a, axes, keepdims != 0, ddof);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "std_error");
}
}
// ==================== SHAPE MANIPULATION ====================
static ERL_NIF_TERM mlx_reshape(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
auto shape = parse_shape(env, argv[1]);
if (shape.empty()) return enif_make_badarg(env);
try {
array result = reshape(a, shape);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "reshape_error");
}
}
static ERL_NIF_TERM mlx_transpose(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 2) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = transpose(a); // Initialize with default case
if (argc == 1) {
// Already initialized
} else {
auto axes = parse_shape(env, argv[1]);
result = transpose(a, axes);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "transpose_error");
}
}
static ERL_NIF_TERM mlx_squeeze(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 2) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = squeeze(a); // Initialize with default case
if (argc == 1) {
// Already initialized
} else {
auto axes = parse_shape(env, argv[1]);
result = squeeze(a, axes);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "squeeze_error");
}
}
static ERL_NIF_TERM mlx_expand_dims(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
auto axes = parse_shape(env, argv[1]);
if (axes.empty()) return enif_make_badarg(env);
try {
array result = expand_dims(a, axes);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "expand_dims_error");
}
}
// ==================== ARRAY CONCATENATION AND STACKING ====================
static ERL_NIF_TERM mlx_concatenate(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
// Parse array list
unsigned int array_count;
if (!enif_get_list_length(env, argv[0], &array_count)) {
return enif_make_badarg(env);
}
std::vector<array> arrays;
arrays.reserve(array_count);
ERL_NIF_TERM head, tail = argv[0];
for (unsigned int i = 0; i < array_count; i++) {
if (!enif_get_list_cell(env, tail, &head, &tail)) {
return enif_make_badarg(env);
}
bool success;
array arr = get_array_resource_val(env, head, success);
if (!success) {
return enif_make_badarg(env);
}
arrays.push_back(arr);
}
int axis;
if (!enif_get_int(env, argv[1], &axis)) {
return enif_make_badarg(env);
}
try {
array result = concatenate(arrays, axis);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "concatenate_error");
}
}
static ERL_NIF_TERM mlx_stack(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
// Parse array list
unsigned int array_count;
if (!enif_get_list_length(env, argv[0], &array_count)) {
return enif_make_badarg(env);
}
std::vector<array> arrays;
arrays.reserve(array_count);
ERL_NIF_TERM head, tail = argv[0];
for (unsigned int i = 0; i < array_count; i++) {
if (!enif_get_list_cell(env, tail, &head, &tail)) {
return enif_make_badarg(env);
}
bool success;
array arr = get_array_resource_val(env, head, success);
if (!success) {
return enif_make_badarg(env);
}
arrays.push_back(arr);
}
int axis;
if (!enif_get_int(env, argv[1], &axis)) {
return enif_make_badarg(env);
}
try {
array result = stack(arrays, axis);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "stack_error");
}
}
// ==================== LINEAR ALGEBRA ====================
static ERL_NIF_TERM mlx_matmul(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
bool success_a, success_b;
array a = get_array_resource_val(env, argv[0], success_a);
array b = get_array_resource_val(env, argv[1], success_b);
if (!success_a || !success_b) {
return enif_make_badarg(env);
}
try {
array result = matmul(a, b);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "matmul_error");
}
}
// ==================== SORTING AND SEARCHING ====================
static ERL_NIF_TERM mlx_sort(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 2) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = (argc == 1) ? sort(a) : array({0.0f}); // Initialize with dummy
if (argc == 1) {
// Already initialized above
} else {
int axis;
if (!enif_get_int(env, argv[1], &axis)) {
return enif_make_badarg(env);
}
result = sort(a, axis);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "sort_error");
}
}
static ERL_NIF_TERM mlx_argsort(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 2) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = (argc == 1) ? argsort(a) : array({0.0f}); // Initialize with dummy
if (argc == 1) {
// Already initialized above
} else {
int axis;
if (!enif_get_int(env, argv[1], &axis)) {
return enif_make_badarg(env);
}
result = argsort(a, axis);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "argsort_error");
}
}
// ==================== UTILITY FUNCTIONS ====================
static ERL_NIF_TERM mlx_where(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
bool success_cond, success_x, success_y;
array condition = get_array_resource_val(env, argv[0], success_cond);
array x = get_array_resource_val(env, argv[1], success_x);
array y = get_array_resource_val(env, argv[2], success_y);
if (!success_cond || !success_x || !success_y) {
return enif_make_badarg(env);
}
try {
array result = where(condition, x, y);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "where_error");
}
}
static ERL_NIF_TERM mlx_abs(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = abs(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "abs_error");
}
}
static ERL_NIF_TERM mlx_sign(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
array result = sign(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "sign_error");
}
}
// ==================== ARRAY INFO FUNCTIONS ====================
static ERL_NIF_TERM mlx_shape(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
const std::vector<int>& shape = a.shape();
ERL_NIF_TERM shape_list = enif_make_list(env, 0);
for (int i = shape.size() - 1; i >= 0; i--) {
shape_list = enif_make_list_cell(env, enif_make_int(env, shape[i]), shape_list);
}
return make_ok(env, shape_list);
} catch (const std::exception& e) {
return make_error(env, "shape_error");
}
}
static ERL_NIF_TERM mlx_size(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
size_t size = a.size();
return make_ok(env, enif_make_ulong(env, size));
} catch (const std::exception& e) {
return make_error(env, "size_error");
}
}
static ERL_NIF_TERM mlx_ndim(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
int ndim = a.ndim();
return make_ok(env, enif_make_int(env, ndim));
} catch (const std::exception& e) {
return make_error(env, "ndim_error");
}
}
static ERL_NIF_TERM mlx_dtype_str(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
Dtype dtype = a.dtype();
const char* dtype_name;
if (dtype == float32) dtype_name = "float32";
else if (dtype == float16) dtype_name = "float16";
else if (dtype == bfloat16) dtype_name = "bfloat16";
else if (dtype == float64) dtype_name = "float64";
else if (dtype == int32) dtype_name = "int32";
else if (dtype == int16) dtype_name = "int16";
else if (dtype == int8) dtype_name = "int8";
else if (dtype == int64) dtype_name = "int64";
else if (dtype == uint32) dtype_name = "uint32";
else if (dtype == uint16) dtype_name = "uint16";
else if (dtype == uint8) dtype_name = "uint8";
else if (dtype == uint64) dtype_name = "uint64";
else if (dtype == bool_) dtype_name = "bool";
else if (dtype == complex64) dtype_name = "complex64";
else if (dtype == complex64) dtype_name = "complex64";
else dtype_name = "unknown";
return make_ok(env, make_atom(env, dtype_name));
} catch (const std::exception& e) {
return make_error(env, "dtype_error");
}
}
// ==================== EVALUATION ====================
static ERL_NIF_TERM mlx_eval(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
eval(a);
return make_atom(env, "ok");
} catch (const std::exception& e) {
return make_error(env, "eval_error");
}
}
// Helper function to convert array to nested list
static ERL_NIF_TERM array_to_list_recursive(ErlNifEnv* env, const float* data, const std::vector<int>& shape, int dim, size_t& offset) {
if (dim == shape.size() - 1) {
// Base case: create list of values
ERL_NIF_TERM* terms = new ERL_NIF_TERM[shape[dim]];
for (int i = 0; i < shape[dim]; i++) {
terms[i] = enif_make_double(env, data[offset++]);
}
ERL_NIF_TERM list = enif_make_list_from_array(env, terms, shape[dim]);
delete[] terms;
return list;
} else {
// Recursive case: create list of lists
ERL_NIF_TERM* terms = new ERL_NIF_TERM[shape[dim]];
for (int i = 0; i < shape[dim]; i++) {
terms[i] = array_to_list_recursive(env, data, shape, dim + 1, offset);
}
ERL_NIF_TERM list = enif_make_list_from_array(env, terms, shape[dim]);
delete[] terms;
return list;
}
}
static ERL_NIF_TERM mlx_to_list(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
bool success_a;
array a = get_array_resource_val(env, argv[0], success_a);
if (!success_a) {
return enif_make_badarg(env);
}
try {
// Evaluate the array to ensure data is computed
eval(a);
// Get shape
std::vector<int> shape = a.shape();
// Handle scalar case
if (shape.empty()) {
return enif_make_double(env, a.item<float>());
}
// Convert to float32 if needed
if (a.dtype() != float32) {
a = astype(a, float32);
eval(a);
}
// Get data pointer
const float* data = a.data<float>();
// Convert to nested list
size_t offset = 0;
return array_to_list_recursive(env, data, shape, 0, offset);
} catch (const std::exception& e) {
return make_error(env, "to_list_error");
}
}
// Resource destructor
static void array_resource_destructor(ErlNifEnv* env, void* obj) {
ArrayResource* res = (ArrayResource*)obj;
res->~ArrayResource();
}
// Complete MLX NIF function table
static ErlNifFunc nif_funcs[] = {
// Array creation
{"array", 2, mlx_array, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"arange", 3, mlx_arange, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"linspace", 3, mlx_linspace, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"zeros", 2, mlx_zeros, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"ones", 2, mlx_ones, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"full", 3, mlx_full, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"eye", 2, mlx_eye, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Arithmetic operations
{"add", 2, mlx_add, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"subtract", 2, mlx_subtract, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"multiply", 2, mlx_multiply, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"divide", 2, mlx_divide, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"power", 2, mlx_power, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"negative", 1, mlx_negative, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Trigonometric functions
{"sin", 1, mlx_sin, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"cos", 1, mlx_cos, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"tan", 1, mlx_tan, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"arcsin", 1, mlx_arcsin, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"arccos", 1, mlx_arccos, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"arctan", 1, mlx_arctan, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Hyperbolic functions
{"sinh", 1, mlx_sinh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"cosh", 1, mlx_cosh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"tanh", 1, mlx_tanh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Exponential and logarithmic
{"exp", 1, mlx_exp, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"log", 1, mlx_log, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"log2", 1, mlx_log2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"log10", 1, mlx_log10, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"sqrt", 1, mlx_sqrt, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"square", 1, mlx_square, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Comparison operations
{"equal", 2, mlx_equal, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"not_equal", 2, mlx_not_equal, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"greater", 2, mlx_greater, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"greater_equal", 2, mlx_greater_equal, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"less", 2, mlx_less, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"less_equal", 2, mlx_less_equal, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Logical operations
{"logical_and", 2, mlx_logical_and, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"logical_or", 2, mlx_logical_or, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"logical_not", 1, mlx_logical_not, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Reduction operations
{"sum", 1, mlx_sum, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"sum", 2, mlx_sum, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"sum", 3, mlx_sum, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"mean", 1, mlx_mean, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"mean", 2, mlx_mean, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"mean", 3, mlx_mean, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"max", 1, mlx_max, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"max", 2, mlx_max, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"max", 3, mlx_max, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"min", 1, mlx_min, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"min", 2, mlx_min, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"min", 3, mlx_min, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"var", 1, mlx_var, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"var", 2, mlx_var, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"var", 3, mlx_var, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"var", 4, mlx_var, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"std", 1, mlx_std, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"std", 2, mlx_std, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"std", 3, mlx_std, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"std", 4, mlx_std, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Shape manipulation
{"reshape", 2, mlx_reshape, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"transpose", 1, mlx_transpose, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"transpose", 2, mlx_transpose, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"squeeze", 1, mlx_squeeze, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"squeeze", 2, mlx_squeeze, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"expand_dims", 2, mlx_expand_dims, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Array concatenation and stacking
{"concatenate", 2, mlx_concatenate, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"stack", 2, mlx_stack, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Linear algebra
{"matmul", 2, mlx_matmul, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Sorting and searching
{"sort", 1, mlx_sort, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"sort", 2, mlx_sort, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"argsort", 1, mlx_argsort, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"argsort", 2, mlx_argsort, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Utility functions
{"where", 3, mlx_where, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"abs", 1, mlx_abs, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"sign", 1, mlx_sign, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Array info
{"shape", 1, mlx_shape, 0},
{"size", 1, mlx_size, 0},
{"ndim", 1, mlx_ndim, 0},
{"dtype_str", 1, mlx_dtype_str, 0},
// Evaluation
{"eval", 1, mlx_eval, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"to_list", 1, mlx_to_list, 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_complete_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_nif, nif_funcs, load, NULL, upgrade, unload)