Current section
Files
Jump to
Current section
Files
c_src/mlx_random_nif.cpp
#include <erl_nif.h>
#include <mlx/mlx.h>
#include <mlx/ops.h>
#include <mlx/array.h>
#include <mlx/random.h>
#include <memory>
#include <vector>
using namespace mlx::core;
namespace mx = mlx::core;
namespace rnd = mlx::core::random;
// Resource types for random 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);
}
// 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 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;
}
// ==================== RANDOM SEED MANAGEMENT ====================
static ERL_NIF_TERM mlx_random_seed(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
unsigned long seed;
if (!enif_get_ulong(env, argv[0], &seed)) {
return enif_make_badarg(env);
}
try {
rnd::seed(static_cast<uint64_t>(seed));
return make_atom(env, "ok");
} catch (const std::exception& e) {
return make_error(env, "random_seed_error");
}
}
// ==================== UNIFORM DISTRIBUTION ====================
static ERL_NIF_TERM mlx_random_uniform(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 4) return enif_make_badarg(env);
double low, high;
if (!enif_get_double(env, argv[0], &low) ||
!enif_get_double(env, argv[1], &high)) {
return enif_make_badarg(env);
}
auto shape = parse_shape(env, argv[2]);
if (shape.empty()) return enif_make_badarg(env);
char dtype_str[32];
if (!enif_get_atom(env, argv[3], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
Dtype dtype = parse_dtype(dtype_str);
array result = rnd::uniform(low, high, shape, dtype);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "random_uniform_error");
}
}
// ==================== NORMAL DISTRIBUTION ====================
static ERL_NIF_TERM mlx_random_normal(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 4) return enif_make_badarg(env);
double mean, std;
if (!enif_get_double(env, argv[0], &mean) ||
!enif_get_double(env, argv[1], &std)) {
return enif_make_badarg(env);
}
auto shape = parse_shape(env, argv[2]);
if (shape.empty()) return enif_make_badarg(env);
char dtype_str[32];
if (!enif_get_atom(env, argv[3], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
Dtype dtype = parse_dtype(dtype_str);
array result = rnd::normal(mean, std, shape, dtype);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "random_normal_error");
}
}
// ==================== BERNOULLI DISTRIBUTION ====================
static ERL_NIF_TERM mlx_random_bernoulli(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
double p;
if (!enif_get_double(env, argv[0], &p)) {
return enif_make_badarg(env);
}
auto shape = parse_shape(env, argv[1]);
if (shape.empty()) return enif_make_badarg(env);
try {
array result = rnd::bernoulli(p, shape);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "random_bernoulli_error");
}
}
// ==================== CATEGORICAL DISTRIBUTION ====================
static ERL_NIF_TERM mlx_random_categorical(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
array logits;
if (!get_array_resource(env, argv[0], logits)) {
return enif_make_badarg(env);
}
int axis, num_samples;
if (!enif_get_int(env, argv[1], &axis) ||
!enif_get_int(env, argv[2], &num_samples)) {
return enif_make_badarg(env);
}
try {
array result = rnd::categorical(logits, axis, num_samples);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "random_categorical_error");
}
}
// ==================== GUMBEL DISTRIBUTION ====================
static ERL_NIF_TERM mlx_random_gumbel(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 = rnd::gumbel(shape, dtype);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "random_gumbel_error");
}
}
// ==================== LAPLACE DISTRIBUTION ====================
static ERL_NIF_TERM mlx_random_laplace(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 4) return enif_make_badarg(env);
double loc, scale;
if (!enif_get_double(env, argv[0], &loc) ||
!enif_get_double(env, argv[1], &scale)) {
return enif_make_badarg(env);
}
auto shape = parse_shape(env, argv[2]);
if (shape.empty()) return enif_make_badarg(env);
char dtype_str[32];
if (!enif_get_atom(env, argv[3], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
Dtype dtype = parse_dtype(dtype_str);
array result = rnd::laplace(loc, scale, shape, dtype);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "random_laplace_error");
}
}
// ==================== RANDINT ====================
static ERL_NIF_TERM mlx_random_randint(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 4) return enif_make_badarg(env);
int low, high;
if (!enif_get_int(env, argv[0], &low) ||
!enif_get_int(env, argv[1], &high)) {
return enif_make_badarg(env);
}
auto shape = parse_shape(env, argv[2]);
if (shape.empty()) return enif_make_badarg(env);
char dtype_str[32];
if (!enif_get_atom(env, argv[3], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
Dtype dtype = parse_dtype(dtype_str);
array result = rnd::randint(low, high, shape, dtype);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "random_randint_error");
}
}
// ==================== SAMPLING OPERATIONS ====================
static ERL_NIF_TERM mlx_random_choice(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 4) return enif_make_badarg(env);
array a;
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
int size;
int replace, axis;
if (!enif_get_int(env, argv[1], &size) ||
!enif_get_int(env, argv[2], &replace) ||
!enif_get_int(env, argv[3], &axis)) {
return enif_make_badarg(env);
}
try {
array result = rnd::choice(a, size, replace != 0, axis);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "random_choice_error");
}
}
static ERL_NIF_TERM mlx_random_permutation(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
array x;
if (!get_array_resource(env, argv[0], x)) {
return enif_make_badarg(env);
}
int axis;
if (!enif_get_int(env, argv[1], &axis)) {
return enif_make_badarg(env);
}
try {
array result = rnd::permutation(x, axis);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "random_permutation_error");
}
}
static ERL_NIF_TERM mlx_random_shuffle(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
array x;
if (!get_array_resource(env, argv[0], x)) {
return enif_make_badarg(env);
}
int axis;
if (!enif_get_int(env, argv[1], &axis)) {
return enif_make_badarg(env);
}
try {
// In-place shuffle
array result = rnd::permutation(x, axis);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "random_shuffle_error");
}
}
// ==================== ADVANCED SAMPLING ====================
static ERL_NIF_TERM mlx_random_multivariate_normal(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
array mean, cov;
if (!get_array_resource(env, argv[0], mean) ||
!get_array_resource(env, argv[1], cov)) {
return enif_make_badarg(env);
}
auto shape = parse_shape(env, argv[2]);
if (shape.empty()) return enif_make_badarg(env);
try {
array result = rnd::multivariate_normal(mean, cov, shape);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "random_multivariate_normal_error");
}
}
static ERL_NIF_TERM mlx_random_truncated_normal(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 6) return enif_make_badarg(env);
double lower, upper, mean, std;
if (!enif_get_double(env, argv[0], &lower) ||
!enif_get_double(env, argv[1], &upper) ||
!enif_get_double(env, argv[2], &mean) ||
!enif_get_double(env, argv[3], &std)) {
return enif_make_badarg(env);
}
auto shape = parse_shape(env, argv[4]);
if (shape.empty()) return enif_make_badarg(env);
char dtype_str[32];
if (!enif_get_atom(env, argv[5], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
try {
Dtype dtype = parse_dtype(dtype_str);
array result = rnd::truncated_normal(lower, upper, mean, std, shape, dtype);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "random_truncated_normal_error");
}
}
// ==================== RANDOM KEY MANAGEMENT ====================
static ERL_NIF_TERM mlx_random_key(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
unsigned long seed;
if (!enif_get_ulong(env, argv[0], &seed)) {
return enif_make_badarg(env);
}
try {
array key = rnd::key(static_cast<uint64_t>(seed));
return make_array_resource(env, key);
} catch (const std::exception& e) {
return make_error(env, "random_key_error");
}
}
static ERL_NIF_TERM mlx_random_split(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
array key;
if (!get_array_resource(env, argv[0], key)) {
return enif_make_badarg(env);
}
int num_keys;
if (!enif_get_int(env, argv[1], &num_keys)) {
return enif_make_badarg(env);
}
try {
auto keys = rnd::split(key, num_keys);
// Return list of keys
ERL_NIF_TERM key_list = enif_make_list(env, 0);
for (int i = keys.size() - 1; i >= 0; i--) {
auto key_result = make_array_resource(env, keys[i]);
// Extract the array term from the {ok, Array} tuple
const ERL_NIF_TERM* tuple_elements;
int tuple_arity;
ERL_NIF_TERM key_term;
if (enif_get_tuple(env, key_result, &tuple_arity, &tuple_elements) && tuple_arity == 2) {
key_term = tuple_elements[1];
} else {
key_term = key_result;
}
key_list = enif_make_list_cell(env, key_term, key_list);
}
return make_ok(env, key_list);
} catch (const std::exception& e) {
return make_error(env, "random_split_error");
}
}
// ==================== RANDOM UTILITIES ====================
static ERL_NIF_TERM mlx_random_bits(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);
int width;
if (!enif_get_int(env, argv[1], &width)) {
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 = rnd::bits(shape, width, dtype);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "random_bits_error");
}
}
// Resource destructor
static void array_resource_destructor(ErlNifEnv* env, void* obj) {
ArrayResource* res = (ArrayResource*)obj;
res->~ArrayResource();
}
// Random NIF function table
static ErlNifFunc nif_funcs[] = {
// Seed management
{"seed", 1, mlx_random_seed, 0},
// Basic distributions
{"uniform", 4, mlx_random_uniform, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"normal", 4, mlx_random_normal, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"bernoulli", 2, mlx_random_bernoulli, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"categorical", 3, mlx_random_categorical, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"gumbel", 2, mlx_random_gumbel, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"laplace", 4, mlx_random_laplace, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"randint", 4, mlx_random_randint, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Sampling operations
{"choice", 4, mlx_random_choice, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"permutation", 2, mlx_random_permutation, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"shuffle", 2, mlx_random_shuffle, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Advanced sampling
{"multivariate_normal", 3, mlx_random_multivariate_normal, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"truncated_normal", 6, mlx_random_truncated_normal, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Key management
{"key", 1, mlx_random_key, 0},
{"split", 2, mlx_random_split, 0},
// Utilities
{"bits", 3, mlx_random_bits, 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_random_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_random_nif, nif_funcs, load, NULL, upgrade, unload)