Packages

MLX machine learning framework bindings for Erlang

Current section

Files

Jump to
mlx c_src mlx_optimizers_nif.cpp
Raw

c_src/mlx_optimizers_nif.cpp

#include <erl_nif.h>
#include <mlx/mlx.h>
#include <mlx/ops.h>
#include <mlx/array.h>
#include <mlx/optimizers.h>
#include <memory>
#include <vector>
#include <map>
#include <string>
using namespace mlx::core;
namespace mx = mlx::core;
namespace opt = mlx::optimizers;
// Resource types for optimizers
static ErlNifResourceType* ARRAY_RESOURCE_TYPE;
static ErlNifResourceType* OPTIMIZER_RESOURCE_TYPE;
struct ArrayResource {
array arr;
std::string name;
ArrayResource(const array& a, const std::string& n = "") : arr(a), name(n) {}
};
struct OptimizerResource {
std::unique_ptr<opt::Optimizer> optimizer;
std::string type;
OptimizerResource(std::unique_ptr<opt::Optimizer> opt, const std::string& t)
: optimizer(std::move(opt)), type(t) {}
};
// 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);
}
// ==================== SGD OPTIMIZER ====================
static ERL_NIF_TERM mlx_sgd_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
double learning_rate, momentum, weight_decay;
if (!enif_get_double(env, argv[0], &learning_rate) ||
!enif_get_double(env, argv[1], &momentum) ||
!enif_get_double(env, argv[2], &weight_decay)) {
return enif_make_badarg(env);
}
try {
auto optimizer = std::make_unique<opt::SGD>(learning_rate, momentum, weight_decay);
OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
new(res) OptimizerResource(std::move(optimizer), "SGD");
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, "sgd_create_error");
}
}
// ==================== ADAM OPTIMIZER ====================
static ERL_NIF_TERM mlx_adam_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 5) return enif_make_badarg(env);
double learning_rate, beta1, beta2, eps, weight_decay;
if (!enif_get_double(env, argv[0], &learning_rate) ||
!enif_get_double(env, argv[1], &beta1) ||
!enif_get_double(env, argv[2], &beta2) ||
!enif_get_double(env, argv[3], &eps) ||
!enif_get_double(env, argv[4], &weight_decay)) {
return enif_make_badarg(env);
}
try {
auto optimizer = std::make_unique<opt::Adam>(learning_rate, beta1, beta2, eps, weight_decay);
OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
new(res) OptimizerResource(std::move(optimizer), "Adam");
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, "adam_create_error");
}
}
// ==================== ADAMW OPTIMIZER ====================
static ERL_NIF_TERM mlx_adamw_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 5) return enif_make_badarg(env);
double learning_rate, beta1, beta2, eps, weight_decay;
if (!enif_get_double(env, argv[0], &learning_rate) ||
!enif_get_double(env, argv[1], &beta1) ||
!enif_get_double(env, argv[2], &beta2) ||
!enif_get_double(env, argv[3], &eps) ||
!enif_get_double(env, argv[4], &weight_decay)) {
return enif_make_badarg(env);
}
try {
auto optimizer = std::make_unique<opt::AdamW>(learning_rate, beta1, beta2, eps, weight_decay);
OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
new(res) OptimizerResource(std::move(optimizer), "AdamW");
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, "adamw_create_error");
}
}
// ==================== RMSPROP OPTIMIZER ====================
static ERL_NIF_TERM mlx_rmsprop_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 4) return enif_make_badarg(env);
double learning_rate, alpha, eps, weight_decay;
if (!enif_get_double(env, argv[0], &learning_rate) ||
!enif_get_double(env, argv[1], &alpha) ||
!enif_get_double(env, argv[2], &eps) ||
!enif_get_double(env, argv[3], &weight_decay)) {
return enif_make_badarg(env);
}
try {
auto optimizer = std::make_unique<opt::RMSprop>(learning_rate, alpha, eps, weight_decay);
OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
new(res) OptimizerResource(std::move(optimizer), "RMSprop");
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, "rmsprop_create_error");
}
}
// ==================== ADAGRAD OPTIMIZER ====================
static ERL_NIF_TERM mlx_adagrad_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
double learning_rate, eps, weight_decay;
if (!enif_get_double(env, argv[0], &learning_rate) ||
!enif_get_double(env, argv[1], &eps) ||
!enif_get_double(env, argv[2], &weight_decay)) {
return enif_make_badarg(env);
}
try {
auto optimizer = std::make_unique<opt::Adagrad>(learning_rate, eps, weight_decay);
OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
new(res) OptimizerResource(std::move(optimizer), "Adagrad");
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, "adagrad_create_error");
}
}
// ==================== ADADELTA OPTIMIZER ====================
static ERL_NIF_TERM mlx_adadelta_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
double learning_rate, rho, eps;
if (!enif_get_double(env, argv[0], &learning_rate) ||
!enif_get_double(env, argv[1], &rho) ||
!enif_get_double(env, argv[2], &eps)) {
return enif_make_badarg(env);
}
try {
auto optimizer = std::make_unique<opt::Adadelta>(learning_rate, rho, eps);
OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
new(res) OptimizerResource(std::move(optimizer), "Adadelta");
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, "adadelta_create_error");
}
}
// ==================== LION OPTIMIZER ====================
static ERL_NIF_TERM mlx_lion_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 4) return enif_make_badarg(env);
double learning_rate, beta1, beta2, weight_decay;
if (!enif_get_double(env, argv[0], &learning_rate) ||
!enif_get_double(env, argv[1], &beta1) ||
!enif_get_double(env, argv[2], &beta2) ||
!enif_get_double(env, argv[3], &weight_decay)) {
return enif_make_badarg(env);
}
try {
auto optimizer = std::make_unique<opt::Lion>(learning_rate, beta1, beta2, weight_decay);
OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
new(res) OptimizerResource(std::move(optimizer), "Lion");
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, "lion_create_error");
}
}
// ==================== OPTIMIZER OPERATIONS ====================
static ERL_NIF_TERM mlx_optimizer_update(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
OptimizerResource* opt_res;
if (!enif_get_resource(env, argv[0], OPTIMIZER_RESOURCE_TYPE, (void**)&opt_res)) {
return enif_make_badarg(env);
}
array parameters, gradients;
if (!get_array_resource(env, argv[1], parameters) ||
!get_array_resource(env, argv[2], gradients)) {
return enif_make_badarg(env);
}
try {
// Apply optimizer update
auto updates = opt_res->optimizer->apply_gradients({gradients}, {parameters});
if (updates.empty()) {
return make_error(env, "no_updates");
}
// Return updated parameters
return make_array_resource(env, updates[0]);
} catch (const std::exception& e) {
return make_error(env, "optimizer_update_error");
}
}
static ERL_NIF_TERM mlx_optimizer_step(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
OptimizerResource* opt_res;
if (!enif_get_resource(env, argv[0], OPTIMIZER_RESOURCE_TYPE, (void**)&opt_res)) {
return enif_make_badarg(env);
}
try {
// Update optimizer internal state
opt_res->optimizer->update();
return make_atom(env, "ok");
} catch (const std::exception& e) {
return make_error(env, "optimizer_step_error");
}
}
// ==================== LEARNING RATE SCHEDULERS ====================
static ERL_NIF_TERM mlx_scheduler_step_decay(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 4) return enif_make_badarg(env);
double initial_lr, decay_rate;
int step, decay_steps;
if (!enif_get_double(env, argv[0], &initial_lr) ||
!enif_get_double(env, argv[1], &decay_rate) ||
!enif_get_int(env, argv[2], &step) ||
!enif_get_int(env, argv[3], &decay_steps)) {
return enif_make_badarg(env);
}
try {
double current_lr = initial_lr * pow(decay_rate, step / decay_steps);
return make_ok(env, enif_make_double(env, current_lr));
} catch (const std::exception& e) {
return make_error(env, "step_decay_error");
}
}
static ERL_NIF_TERM mlx_scheduler_exponential_decay(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
double initial_lr, decay_rate;
int step;
if (!enif_get_double(env, argv[0], &initial_lr) ||
!enif_get_double(env, argv[1], &decay_rate) ||
!enif_get_int(env, argv[2], &step)) {
return enif_make_badarg(env);
}
try {
double current_lr = initial_lr * pow(decay_rate, step);
return make_ok(env, enif_make_double(env, current_lr));
} catch (const std::exception& e) {
return make_error(env, "exponential_decay_error");
}
}
static ERL_NIF_TERM mlx_scheduler_cosine_decay(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
double initial_lr;
int step, total_steps;
if (!enif_get_double(env, argv[0], &initial_lr) ||
!enif_get_int(env, argv[1], &step) ||
!enif_get_int(env, argv[2], &total_steps)) {
return enif_make_badarg(env);
}
try {
double progress = static_cast<double>(step) / total_steps;
double current_lr = initial_lr * 0.5 * (1.0 + cos(M_PI * progress));
return make_ok(env, enif_make_double(env, current_lr));
} catch (const std::exception& e) {
return make_error(env, "cosine_decay_error");
}
}
static ERL_NIF_TERM mlx_scheduler_linear_decay(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 4) return enif_make_badarg(env);
double initial_lr, final_lr;
int step, total_steps;
if (!enif_get_double(env, argv[0], &initial_lr) ||
!enif_get_double(env, argv[1], &final_lr) ||
!enif_get_int(env, argv[2], &step) ||
!enif_get_int(env, argv[3], &total_steps)) {
return enif_make_badarg(env);
}
try {
double progress = static_cast<double>(step) / total_steps;
double current_lr = initial_lr + (final_lr - initial_lr) * progress;
return make_ok(env, enif_make_double(env, current_lr));
} catch (const std::exception& e) {
return make_error(env, "linear_decay_error");
}
}
// ==================== GRADIENT CLIPPING ====================
static ERL_NIF_TERM mlx_clip_grad_norm(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
array gradients;
if (!get_array_resource(env, argv[0], gradients)) {
return enif_make_badarg(env);
}
double max_norm;
if (!enif_get_double(env, argv[1], &max_norm)) {
return enif_make_badarg(env);
}
try {
// Calculate gradient norm
array grad_norm = sqrt(sum(square(gradients)));
// Clip gradients if norm exceeds max_norm
array max_norm_arr = array(static_cast<float>(max_norm));
array clip_coeff = minimum(divide(max_norm_arr, grad_norm), ones_like(grad_norm));
array clipped_gradients = multiply(gradients, clip_coeff);
return make_array_resource(env, clipped_gradients);
} catch (const std::exception& e) {
return make_error(env, "clip_grad_norm_error");
}
}
static ERL_NIF_TERM mlx_clip_grad_value(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
array gradients;
if (!get_array_resource(env, argv[0], gradients)) {
return enif_make_badarg(env);
}
double clip_value;
if (!enif_get_double(env, argv[1], &clip_value)) {
return enif_make_badarg(env);
}
try {
array clip_val = array(static_cast<float>(clip_value));
array neg_clip_val = negative(clip_val);
array clipped_gradients = clip(gradients, neg_clip_val, clip_val);
return make_array_resource(env, clipped_gradients);
} catch (const std::exception& e) {
return make_error(env, "clip_grad_value_error");
}
}
// ==================== OPTIMIZER UTILITIES ====================
static ERL_NIF_TERM mlx_get_optimizer_state(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
OptimizerResource* opt_res;
if (!enif_get_resource(env, argv[0], OPTIMIZER_RESOURCE_TYPE, (void**)&opt_res)) {
return enif_make_badarg(env);
}
try {
// Return optimizer type and learning rate info
ERL_NIF_TERM type_atom = make_atom(env, opt_res->type.c_str());
ERL_NIF_TERM state_info = enif_make_tuple2(env,
make_atom(env, "type"),
type_atom);
return make_ok(env, state_info);
} catch (const std::exception& e) {
return make_error(env, "get_optimizer_state_error");
}
}
static ERL_NIF_TERM mlx_set_learning_rate(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
OptimizerResource* opt_res;
if (!enif_get_resource(env, argv[0], OPTIMIZER_RESOURCE_TYPE, (void**)&opt_res)) {
return enif_make_badarg(env);
}
double new_lr;
if (!enif_get_double(env, argv[1], &new_lr)) {
return enif_make_badarg(env);
}
try {
// Set new learning rate (simplified - actual implementation would depend on optimizer type)
opt_res->optimizer->set_learning_rate(new_lr);
return make_atom(env, "ok");
} catch (const std::exception& e) {
return make_error(env, "set_learning_rate_error");
}
}
// Resource destructors
static void array_resource_destructor(ErlNifEnv* env, void* obj) {
ArrayResource* res = (ArrayResource*)obj;
res->~ArrayResource();
}
static void optimizer_resource_destructor(ErlNifEnv* env, void* obj) {
OptimizerResource* res = (OptimizerResource*)obj;
res->~OptimizerResource();
}
// Optimizers NIF function table
static ErlNifFunc nif_funcs[] = {
// Optimizer creation
{"sgd_create", 3, mlx_sgd_create, 0},
{"adam_create", 5, mlx_adam_create, 0},
{"adamw_create", 5, mlx_adamw_create, 0},
{"rmsprop_create", 4, mlx_rmsprop_create, 0},
{"adagrad_create", 3, mlx_adagrad_create, 0},
{"adadelta_create", 3, mlx_adadelta_create, 0},
{"lion_create", 4, mlx_lion_create, 0},
// Optimizer operations
{"optimizer_update", 3, mlx_optimizer_update, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"optimizer_step", 1, mlx_optimizer_step, 0},
// Learning rate schedulers
{"scheduler_step_decay", 4, mlx_scheduler_step_decay, 0},
{"scheduler_exponential_decay", 3, mlx_scheduler_exponential_decay, 0},
{"scheduler_cosine_decay", 3, mlx_scheduler_cosine_decay, 0},
{"scheduler_linear_decay", 4, mlx_scheduler_linear_decay, 0},
// Gradient clipping
{"clip_grad_norm", 2, mlx_clip_grad_norm, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"clip_grad_value", 2, mlx_clip_grad_value, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Optimizer utilities
{"get_optimizer_state", 1, mlx_get_optimizer_state, 0},
{"set_learning_rate", 2, mlx_set_learning_rate, 0}
};
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_optimizers_array", array_resource_destructor, flags, tried);
OPTIMIZER_RESOURCE_TYPE = enif_open_resource_type(
env, NULL, "mlx_optimizer", optimizer_resource_destructor, flags, tried);
if (!ARRAY_RESOURCE_TYPE || !OPTIMIZER_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_optimizers_nif, nif_funcs, load, NULL, upgrade, unload)