Current section
Files
Jump to
Current section
Files
c_src/mlx_nif_old.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 managing MLX objects
static ErlNifResourceType* ARRAY_RESOURCE_TYPE;
static ErlNifResourceType* STREAM_RESOURCE_TYPE;
// Resource wrappers
struct ArrayResource {
array arr;
ArrayResource(const array& a) : arr(a) {}
};
struct StreamResource {
Stream stream;
StreamResource(const Stream& s) : stream(s) {}
};
// 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);
}
// Array creation functions
static ERL_NIF_TERM mlx_zeros(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) {
return enif_make_badarg(env);
}
// Parse shape
unsigned int shape_len;
if (!enif_get_list_length(env, argv[0], &shape_len)) {
return enif_make_badarg(env);
}
std::vector<int> shape_vec(shape_len);
ERL_NIF_TERM head, tail = argv[0];
for (unsigned int i = 0; i < shape_len; i++) {
if (!enif_get_list_cell(env, tail, &head, &tail)) {
return enif_make_badarg(env);
}
if (!enif_get_int(env, head, &shape_vec[i])) {
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);
}
Dtype dtype = float32; // Default initialization
if (strcmp(dtype_str, "float32") == 0) dtype = float32;
else if (strcmp(dtype_str, "float16") == 0) dtype = float16;
else if (strcmp(dtype_str, "bfloat16") == 0) dtype = bfloat16;
else if (strcmp(dtype_str, "int32") == 0) dtype = int32;
else if (strcmp(dtype_str, "int16") == 0) dtype = int16;
else if (strcmp(dtype_str, "int8") == 0) dtype = int8;
else if (strcmp(dtype_str, "uint32") == 0) dtype = uint32;
else if (strcmp(dtype_str, "uint16") == 0) dtype = uint16;
else if (strcmp(dtype_str, "uint8") == 0) dtype = uint8;
else if (strcmp(dtype_str, "bool") == 0) dtype = bool_;
else return make_error(env, "invalid_dtype");
try {
array result = zeros(shape_vec, dtype);
ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
new(res) ArrayResource(result);
ERL_NIF_TERM term = enif_make_resource(env, res);
enif_release_resource(res);
return make_ok(env, term);
} catch (const std::exception& e) {
return make_error(env, "mlx_error");
}
}
static ERL_NIF_TERM mlx_ones(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) {
return enif_make_badarg(env);
}
// Parse shape
unsigned int shape_len;
if (!enif_get_list_length(env, argv[0], &shape_len)) {
return enif_make_badarg(env);
}
std::vector<int> shape_vec(shape_len);
ERL_NIF_TERM head, tail = argv[0];
for (unsigned int i = 0; i < shape_len; i++) {
if (!enif_get_list_cell(env, tail, &head, &tail)) {
return enif_make_badarg(env);
}
if (!enif_get_int(env, head, &shape_vec[i])) {
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);
}
Dtype dtype = float32; // Default initialization
if (strcmp(dtype_str, "float32") == 0) dtype = float32;
else if (strcmp(dtype_str, "float16") == 0) dtype = float16;
else if (strcmp(dtype_str, "bfloat16") == 0) dtype = bfloat16;
else if (strcmp(dtype_str, "int32") == 0) dtype = int32;
else if (strcmp(dtype_str, "int16") == 0) dtype = int16;
else if (strcmp(dtype_str, "int8") == 0) dtype = int8;
else if (strcmp(dtype_str, "uint32") == 0) dtype = uint32;
else if (strcmp(dtype_str, "uint16") == 0) dtype = uint16;
else if (strcmp(dtype_str, "uint8") == 0) dtype = uint8;
else if (strcmp(dtype_str, "bool") == 0) dtype = bool_;
else return make_error(env, "invalid_dtype");
try {
array result = ones(shape_vec, dtype);
ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
new(res) ArrayResource(result);
ERL_NIF_TERM term = enif_make_resource(env, res);
enif_release_resource(res);
return make_ok(env, term);
} catch (const std::exception& e) {
return make_error(env, "mlx_error");
}
}
// Array operations
static ERL_NIF_TERM mlx_add(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) {
return enif_make_badarg(env);
}
ArrayResource* a_res;
ArrayResource* b_res;
if (!enif_get_resource(env, argv[0], ARRAY_RESOURCE_TYPE, (void**)&a_res) ||
!enif_get_resource(env, argv[1], ARRAY_RESOURCE_TYPE, (void**)&b_res)) {
return enif_make_badarg(env);
}
try {
array result = add(a_res->arr, b_res->arr);
ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
new(res) ArrayResource(result);
ERL_NIF_TERM term = enif_make_resource(env, res);
enif_release_resource(res);
return make_ok(env, term);
} catch (const std::exception& e) {
return make_error(env, "mlx_error");
}
}
static ERL_NIF_TERM mlx_multiply(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) {
return enif_make_badarg(env);
}
ArrayResource* a_res;
ArrayResource* b_res;
if (!enif_get_resource(env, argv[0], ARRAY_RESOURCE_TYPE, (void**)&a_res) ||
!enif_get_resource(env, argv[1], ARRAY_RESOURCE_TYPE, (void**)&b_res)) {
return enif_make_badarg(env);
}
try {
array result = multiply(a_res->arr, b_res->arr);
ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
new(res) ArrayResource(result);
ERL_NIF_TERM term = enif_make_resource(env, res);
enif_release_resource(res);
return make_ok(env, term);
} catch (const std::exception& e) {
return make_error(env, "mlx_error");
}
}
static ERL_NIF_TERM mlx_matmul(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) {
return enif_make_badarg(env);
}
ArrayResource* a_res;
ArrayResource* b_res;
if (!enif_get_resource(env, argv[0], ARRAY_RESOURCE_TYPE, (void**)&a_res) ||
!enif_get_resource(env, argv[1], ARRAY_RESOURCE_TYPE, (void**)&b_res)) {
return enif_make_badarg(env);
}
try {
array result = matmul(a_res->arr, b_res->arr);
ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
new(res) ArrayResource(result);
ERL_NIF_TERM term = enif_make_resource(env, res);
enif_release_resource(res);
return make_ok(env, term);
} catch (const std::exception& e) {
return make_error(env, "mlx_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);
}
ArrayResource* res;
if (!enif_get_resource(env, argv[0], ARRAY_RESOURCE_TYPE, (void**)&res)) {
return enif_make_badarg(env);
}
try {
const std::vector<int>& shape = res->arr.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, "mlx_error");
}
}
static ERL_NIF_TERM mlx_eval(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) {
return enif_make_badarg(env);
}
ArrayResource* res;
if (!enif_get_resource(env, argv[0], ARRAY_RESOURCE_TYPE, (void**)&res)) {
return enif_make_badarg(env);
}
try {
eval(res->arr);
return make_atom(env, "ok");
} catch (const std::exception& e) {
return make_error(env, "mlx_error");
}
}
// Device management
static ERL_NIF_TERM mlx_set_default_device(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) {
return enif_make_badarg(env);
}
char device_str[32];
if (!enif_get_atom(env, argv[0], device_str, sizeof(device_str), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
if (strcmp(device_str, "cpu") == 0) {
set_default_device(Device::cpu);
} else if (strcmp(device_str, "gpu") == 0) {
set_default_device(Device::gpu);
} else {
return make_error(env, "invalid_device");
}
return make_atom(env, "ok");
} catch (const std::exception& e) {
std::cerr << "MLX device error: " << e.what() << std::endl;
return make_error(env, "mlx_error");
}
}
// Test function to verify MLX is working
static ERL_NIF_TERM mlx_test_basic(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 0) {
return enif_make_badarg(env);
}
try {
// Test basic MLX functionality
array a = ones({2, 3}, float32);
array b = zeros({2, 3}, float32);
array c = add(a, b);
eval(c); // Force evaluation
auto shape = c.shape();
auto dtype = c.dtype();
// Return shape and dtype info
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);
}
ERL_NIF_TERM dtype_atom;
if (dtype == float32) dtype_atom = make_atom(env, "float32");
else if (dtype == int32) dtype_atom = make_atom(env, "int32");
else dtype_atom = make_atom(env, "unknown");
ERL_NIF_TERM result = enif_make_tuple3(env,
make_atom(env, "test_passed"),
shape_list,
dtype_atom);
return make_ok(env, result);
} catch (const std::exception& e) {
std::cerr << "MLX test error: " << e.what() << std::endl;
return make_error(env, "mlx_test_failed");
}
}
// Get MLX version info
static ERL_NIF_TERM mlx_version(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 0) {
return enif_make_badarg(env);
}
try {
// Create a simple test to verify MLX is loaded
array test = ones({1}, float32);
eval(test);
return make_ok(env, make_atom(env, "mlx_loaded"));
} catch (const std::exception& e) {
std::cerr << "MLX version check error: " << e.what() << std::endl;
return make_error(env, "mlx_not_loaded");
}
}
// Resource destructors
static void array_resource_destructor(ErlNifEnv* env, void* obj) {
ArrayResource* res = (ArrayResource*)obj;
res->~ArrayResource();
}
static void stream_resource_destructor(ErlNifEnv* env, void* obj) {
StreamResource* res = (StreamResource*)obj;
res->~StreamResource();
}
// NIF function table using dirty schedulers for CPU-intensive operations
static ErlNifFunc nif_funcs[] = {
{"zeros", 2, mlx_zeros, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"ones", 2, mlx_ones, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"add", 2, mlx_add, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"multiply", 2, mlx_multiply, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"matmul", 2, mlx_matmul, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"shape", 1, mlx_shape, 0},
{"eval", 1, mlx_eval, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"set_default_device", 1, mlx_set_default_device, 0},
{"test_basic", 0, mlx_test_basic, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"version", 0, mlx_version, 0}
};
static int load(ErlNifEnv* env, void** priv_data, ERL_NIF_TERM load_info) {
// Initialize resource types
ErlNifResourceFlags flags = (ErlNifResourceFlags)(ERL_NIF_RT_CREATE | ERL_NIF_RT_TAKEOVER);
ErlNifResourceFlags* tried = NULL;
ARRAY_RESOURCE_TYPE = enif_open_resource_type(
env, NULL, "mlx_array", array_resource_destructor,
flags, tried);
STREAM_RESOURCE_TYPE = enif_open_resource_type(
env, NULL, "mlx_stream", stream_resource_destructor,
flags, tried);
if (ARRAY_RESOURCE_TYPE == NULL || STREAM_RESOURCE_TYPE == NULL) {
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)