Current section
Files
Jump to
Current section
Files
c_src/mlx_nif_backup.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;
else if (strcmp(dtype_str, "complex128") == 0) return complex128;
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 make_ok(env, 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;
}
// ==================== 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");
}
}
// ==================== 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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
array result;
if (argc == 1) {
result = 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);
array a;
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
array result;
if (argc == 1) {
result = 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);
array a;
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
array result;
if (argc == 1) {
result = 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);
array a;
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
array result;
if (argc == 1) {
result = min(a);
} 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);
array a;
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
array result;
if (argc == 1) {
result = var(a);
} 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);
array a;
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
array result;
if (argc == 1) {
result = std(a);
} 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 = 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
array result;
if (argc == 1) {
result = transpose(a);
} 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);
array a;
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
array result;
if (argc == 1) {
result = squeeze(a);
} 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);
array a;
if (!get_array_resource(env, argv[0], 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);
}
array arr;
if (!get_array_resource(env, head, arr)) {
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);
}
array arr;
if (!get_array_resource(env, head, arr)) {
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);
array a, b;
if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], 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);
array a;
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
array result;
if (argc == 1) {
result = sort(a);
} 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);
array a;
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
array result;
if (argc == 1) {
result = argsort(a);
} 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);
array condition, x, y;
if (!get_array_resource(env, argv[0], condition) ||
!get_array_resource(env, argv[1], x) ||
!get_array_resource(env, argv[2], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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);
array a;
if (!get_array_resource(env, argv[0], 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 == complex128) dtype_name = "complex128";
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);
array a;
if (!get_array_resource(env, argv[0], 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");
}
}
// 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
{"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}
};
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_complete_nif, nif_funcs, load, NULL, upgrade, unload)