Packages

MLX machine learning framework bindings for Erlang

Current section

Files

Jump to
mlx c_src mlx_io_nif.cpp
Raw

c_src/mlx_io_nif.cpp

#include <erl_nif.h>
#include <mlx/mlx.h>
#include <mlx/ops.h>
#include <mlx/array.h>
#include <mlx/io.h>
#include <memory>
#include <vector>
#include <string>
#include <fstream>
#include <map>
using namespace mlx::core;
namespace mx = mlx::core;
namespace io = mlx::core::io;
// Resource types for I/O operations
static ErlNifResourceType* ARRAY_RESOURCE_TYPE;
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);
}
// Get array from resource
static bool get_array_resource(ErlNifEnv* env, ERL_NIF_TERM term, array& arr) {
ArrayResource* res;
if (!enif_get_resource(env, term, ARRAY_RESOURCE_TYPE, (void**)&res)) {
return false;
}
arr = res->arr;
return true;
}
// Create array resource
static ERL_NIF_TERM make_array_resource(ErlNifEnv* env, const array& arr, const std::string& name = "") {
ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
new(res) ArrayResource(arr, name);
ERL_NIF_TERM term = enif_make_resource(env, res);
enif_release_resource(res);
return make_ok(env, term);
}
// ==================== ARRAY SAVING ====================
static ERL_NIF_TERM mlx_save_array(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
char filename[1024];
if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
array arr;
if (!get_array_resource(env, argv[1], arr)) {
return enif_make_badarg(env);
}
try {
io::save(filename, arr);
return make_atom(env, "ok");
} catch (const std::exception& e) {
return make_error(env, "save_array_error");
}
}
static ERL_NIF_TERM mlx_load_array(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
char filename[1024];
if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
array result = io::load(filename);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "load_array_error");
}
}
// ==================== SAFETENSORS FORMAT ====================
static ERL_NIF_TERM mlx_save_safetensors(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
char filename[1024];
if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
// Parse array dictionary from Erlang
unsigned int dict_size;
if (!enif_get_map_size(env, argv[1], &dict_size)) {
return enif_make_badarg(env);
}
try {
std::map<std::string, array> array_dict;
ErlNifMapIterator iter;
enif_map_iterator_create(env, argv[1], &iter, ERL_NIF_MAP_ITERATOR_FIRST);
ERL_NIF_TERM key, value;
while (enif_map_iterator_get_pair(env, &iter, &key, &value)) {
char key_str[256];
if (!enif_get_string(env, key, key_str, sizeof(key_str), ERL_NIF_LATIN1)) {
enif_map_iterator_destroy(env, &iter);
return enif_make_badarg(env);
}
array arr;
if (!get_array_resource(env, value, arr)) {
enif_map_iterator_destroy(env, &iter);
return enif_make_badarg(env);
}
array_dict[std::string(key_str)] = arr;
enif_map_iterator_next(env, &iter);
}
enif_map_iterator_destroy(env, &iter);
io::save_safetensors(filename, array_dict);
return make_atom(env, "ok");
} catch (const std::exception& e) {
return make_error(env, "save_safetensors_error");
}
}
static ERL_NIF_TERM mlx_load_safetensors(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
char filename[1024];
if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
auto array_dict = io::load_safetensors(filename);
// Convert to Erlang map
ERL_NIF_TERM result_map = enif_make_new_map(env);
for (const auto& pair : array_dict) {
ERL_NIF_TERM key = enif_make_string(env, pair.first.c_str(), ERL_NIF_LATIN1);
auto value_result = make_array_resource(env, pair.second);
// Extract array term from {ok, Array} tuple
const ERL_NIF_TERM* tuple_elements;
int tuple_arity;
ERL_NIF_TERM value_term;
if (enif_get_tuple(env, value_result, &tuple_arity, &tuple_elements) && tuple_arity == 2) {
value_term = tuple_elements[1];
} else {
value_term = value_result;
}
enif_make_map_put(env, result_map, key, value_term, &result_map);
}
return make_ok(env, result_map);
} catch (const std::exception& e) {
return make_error(env, "load_safetensors_error");
}
}
// ==================== NPY FORMAT ====================
static ERL_NIF_TERM mlx_save_npy(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
char filename[1024];
if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
array arr;
if (!get_array_resource(env, argv[1], arr)) {
return enif_make_badarg(env);
}
try {
io::save_npy(filename, arr);
return make_atom(env, "ok");
} catch (const std::exception& e) {
return make_error(env, "save_npy_error");
}
}
static ERL_NIF_TERM mlx_load_npy(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
char filename[1024];
if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
array result = io::load_npy(filename);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "load_npy_error");
}
}
// ==================== NPZ FORMAT ====================
static ERL_NIF_TERM mlx_save_npz(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
char filename[1024];
if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
// Parse array dictionary
unsigned int dict_size;
if (!enif_get_map_size(env, argv[1], &dict_size)) {
return enif_make_badarg(env);
}
try {
std::map<std::string, array> array_dict;
ErlNifMapIterator iter;
enif_map_iterator_create(env, argv[1], &iter, ERL_NIF_MAP_ITERATOR_FIRST);
ERL_NIF_TERM key, value;
while (enif_map_iterator_get_pair(env, &iter, &key, &value)) {
char key_str[256];
if (!enif_get_string(env, key, key_str, sizeof(key_str), ERL_NIF_LATIN1)) {
enif_map_iterator_destroy(env, &iter);
return enif_make_badarg(env);
}
array arr;
if (!get_array_resource(env, value, arr)) {
enif_map_iterator_destroy(env, &iter);
return enif_make_badarg(env);
}
array_dict[std::string(key_str)] = arr;
enif_map_iterator_next(env, &iter);
}
enif_map_iterator_destroy(env, &iter);
io::save_npz(filename, array_dict);
return make_atom(env, "ok");
} catch (const std::exception& e) {
return make_error(env, "save_npz_error");
}
}
static ERL_NIF_TERM mlx_load_npz(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
char filename[1024];
if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
auto array_dict = io::load_npz(filename);
// Convert to Erlang map
ERL_NIF_TERM result_map = enif_make_new_map(env);
for (const auto& pair : array_dict) {
ERL_NIF_TERM key = enif_make_string(env, pair.first.c_str(), ERL_NIF_LATIN1);
auto value_result = make_array_resource(env, pair.second);
// Extract array term from {ok, Array} tuple
const ERL_NIF_TERM* tuple_elements;
int tuple_arity;
ERL_NIF_TERM value_term;
if (enif_get_tuple(env, value_result, &tuple_arity, &tuple_elements) && tuple_arity == 2) {
value_term = tuple_elements[1];
} else {
value_term = value_result;
}
enif_make_map_put(env, result_map, key, value_term, &result_map);
}
return make_ok(env, result_map);
} catch (const std::exception& e) {
return make_error(env, "load_npz_error");
}
}
// ==================== GGUF FORMAT ====================
static ERL_NIF_TERM mlx_save_gguf(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
char filename[1024];
if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
// Parse array dictionary
unsigned int dict_size;
if (!enif_get_map_size(env, argv[1], &dict_size)) {
return enif_make_badarg(env);
}
try {
std::map<std::string, array> array_dict;
ErlNifMapIterator iter;
enif_map_iterator_create(env, argv[1], &iter, ERL_NIF_MAP_ITERATOR_FIRST);
ERL_NIF_TERM key, value;
while (enif_map_iterator_get_pair(env, &iter, &key, &value)) {
char key_str[256];
if (!enif_get_string(env, key, key_str, sizeof(key_str), ERL_NIF_LATIN1)) {
enif_map_iterator_destroy(env, &iter);
return enif_make_badarg(env);
}
array arr;
if (!get_array_resource(env, value, arr)) {
enif_map_iterator_destroy(env, &iter);
return enif_make_badarg(env);
}
array_dict[std::string(key_str)] = arr;
enif_map_iterator_next(env, &iter);
}
enif_map_iterator_destroy(env, &iter);
io::save_gguf(filename, array_dict);
return make_atom(env, "ok");
} catch (const std::exception& e) {
return make_error(env, "save_gguf_error");
}
}
static ERL_NIF_TERM mlx_load_gguf(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
char filename[1024];
if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
auto array_dict = io::load_gguf(filename);
// Convert to Erlang map
ERL_NIF_TERM result_map = enif_make_new_map(env);
for (const auto& pair : array_dict) {
ERL_NIF_TERM key = enif_make_string(env, pair.first.c_str(), ERL_NIF_LATIN1);
auto value_result = make_array_resource(env, pair.second);
// Extract array term from {ok, Array} tuple
const ERL_NIF_TERM* tuple_elements;
int tuple_arity;
ERL_NIF_TERM value_term;
if (enif_get_tuple(env, value_result, &tuple_arity, &tuple_elements) && tuple_arity == 2) {
value_term = tuple_elements[1];
} else {
value_term = value_result;
}
enif_make_map_put(env, result_map, key, value_term, &result_map);
}
return make_ok(env, result_map);
} catch (const std::exception& e) {
return make_error(env, "load_gguf_error");
}
}
// ==================== ARRAY SERIALIZATION ====================
static ERL_NIF_TERM mlx_array_to_bytes(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
array arr;
if (!get_array_resource(env, argv[0], arr)) {
return enif_make_badarg(env);
}
try {
// Convert array to bytes (simplified - would use actual MLX serialization)
eval(arr); // Ensure array is materialized
auto shape = arr.shape();
auto dtype = arr.dtype();
size_t byte_size = arr.nbytes();
// Create byte representation (simplified)
ErlNifBinary binary;
if (!enif_alloc_binary(byte_size, &binary)) {
return make_error(env, "binary_alloc_error");
}
// Copy array data (simplified - would use proper MLX data access)
memset(binary.data, 0, byte_size); // Placeholder
ERL_NIF_TERM binary_term = enif_make_binary(env, &binary);
return make_ok(env, binary_term);
} catch (const std::exception& e) {
return make_error(env, "array_to_bytes_error");
}
}
static ERL_NIF_TERM mlx_bytes_to_array(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
ErlNifBinary binary;
if (!enif_inspect_binary(env, argv[0], &binary)) {
return enif_make_badarg(env);
}
// Parse shape and dtype
auto shape = parse_shape(env, argv[1]);
if (shape.empty()) 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 {
// Create array from bytes (simplified - would use actual MLX deserialization)
Dtype dtype = parse_dtype(dtype_str);
array result = zeros(shape, dtype); // Placeholder
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "bytes_to_array_error");
}
}
// ==================== FILE UTILITIES ====================
static ERL_NIF_TERM mlx_file_exists(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
char filename[1024];
if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
std::ifstream file(filename);
bool exists = file.good();
file.close();
return make_ok(env, exists ? make_atom(env, "true") : make_atom(env, "false"));
}
// 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 return float32;
}
// 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;
}
// Resource destructor
static void array_resource_destructor(ErlNifEnv* env, void* obj) {
ArrayResource* res = (ArrayResource*)obj;
res->~ArrayResource();
}
// I/O NIF function table
static ErlNifFunc nif_funcs[] = {
// Basic array I/O
{"save_array", 2, mlx_save_array, ERL_NIF_DIRTY_JOB_IO_BOUND},
{"load_array", 1, mlx_load_array, ERL_NIF_DIRTY_JOB_IO_BOUND},
// SafeTensors format
{"save_safetensors", 2, mlx_save_safetensors, ERL_NIF_DIRTY_JOB_IO_BOUND},
{"load_safetensors", 1, mlx_load_safetensors, ERL_NIF_DIRTY_JOB_IO_BOUND},
// NPY format
{"save_npy", 2, mlx_save_npy, ERL_NIF_DIRTY_JOB_IO_BOUND},
{"load_npy", 1, mlx_load_npy, ERL_NIF_DIRTY_JOB_IO_BOUND},
// NPZ format
{"save_npz", 2, mlx_save_npz, ERL_NIF_DIRTY_JOB_IO_BOUND},
{"load_npz", 1, mlx_load_npz, ERL_NIF_DIRTY_JOB_IO_BOUND},
// GGUF format
{"save_gguf", 2, mlx_save_gguf, ERL_NIF_DIRTY_JOB_IO_BOUND},
{"load_gguf", 1, mlx_load_gguf, ERL_NIF_DIRTY_JOB_IO_BOUND},
// Array serialization
{"array_to_bytes", 1, mlx_array_to_bytes, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"bytes_to_array", 3, mlx_bytes_to_array, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// File utilities
{"file_exists", 1, mlx_file_exists, ERL_NIF_DIRTY_JOB_IO_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_io_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_io_nif, nif_funcs, load, NULL, upgrade, unload)