Current section
Files
Jump to
Current section
Files
c_src/mlx_linalg_nif.cpp
#include <erl_nif.h>
#include <mlx/mlx.h>
#include <mlx/ops.h>
#include <mlx/array.h>
#include <mlx/linalg.h>
#include <memory>
#include <vector>
using namespace mlx::core;
namespace mx = mlx::core;
namespace linalg = mlx::core::linalg;
// Resource types for linear algebra operations
static ErlNifResourceType* ARRAY_RESOURCE_TYPE;
struct ArrayResource {
array arr = array({1.0f}); // Initialize with dummy value
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);
}
// Create tuple of array resources
static ERL_NIF_TERM make_array_tuple(ErlNifEnv* env, const std::vector<array>& arrays) {
std::vector<ERL_NIF_TERM> terms;
for (const auto& arr : arrays) {
ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
new(res) ArrayResource(arr);
ERL_NIF_TERM term = enif_make_resource(env, res);
enif_release_resource(res);
terms.push_back(term);
}
ERL_NIF_TERM tuple;
if (arrays.size() == 2) {
tuple = enif_make_tuple2(env, terms[0], terms[1]);
} else if (arrays.size() == 3) {
tuple = enif_make_tuple3(env, terms[0], terms[1], terms[2]);
} else {
// For arbitrary number of arrays, create a list
ERL_NIF_TERM list = enif_make_list_from_array(env, terms.data(), terms.size());
tuple = list;
}
return make_ok(env, tuple);
}
// ==================== SINGULAR VALUE DECOMPOSITION ====================
static ERL_NIF_TERM mlx_svd(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 2) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
bool full_matrices = true;
if (argc > 1) {
char atom_name[32];
if (!enif_get_atom(env, argv[1], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
full_matrices = (strcmp(atom_name, "true") == 0);
}
try {
auto result = linalg::svd(a, full_matrices, {});
// result is already a std::vector<array>
return make_array_tuple(env, result);
} catch (const std::exception& e) {
return make_error(env, "svd_error");
}
}
// ==================== EIGENVALUE DECOMPOSITION ====================
static ERL_NIF_TERM mlx_eig(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
// TODO: MLX only has eigh (hermitian), not general eig
return make_error(env, "eig_not_available_use_eigh");
} catch (const std::exception& e) {
return make_error(env, "eig_error");
}
}
static ERL_NIF_TERM mlx_eigvals(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
// TODO: MLX only has eigvalsh (hermitian), not general eigvals
return make_error(env, "eigvals_not_available_use_eigvalsh");
} catch (const std::exception& e) {
return make_error(env, "eigvals_error");
}
}
static ERL_NIF_TERM mlx_eigh(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 2) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
std::string uplo = "L";
if (argc > 1) {
char atom_name[32];
if (!enif_get_atom(env, argv[1], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
uplo = std::string(atom_name);
}
try {
auto result = linalg::eigh(a, uplo);
std::vector<array> arrays = {std::get<0>(result), std::get<1>(result)};
return make_array_tuple(env, arrays);
} catch (const std::exception& e) {
return make_error(env, "eigh_error");
}
}
static ERL_NIF_TERM mlx_eigvalsh(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 2) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
std::string uplo = "L";
if (argc > 1) {
char atom_name[32];
if (!enif_get_atom(env, argv[1], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
uplo = std::string(atom_name);
}
try {
array result = linalg::eigvalsh(a, uplo);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "eigvalsh_error");
}
}
// ==================== QR DECOMPOSITION ====================
static ERL_NIF_TERM mlx_qr(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 2) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
std::string mode = "reduced";
if (argc > 1) {
char atom_name[32];
if (!enif_get_atom(env, argv[1], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
mode = std::string(atom_name);
}
try {
// MLX QR doesn't support mode parameter
auto result = linalg::qr(a);
std::vector<array> arrays = {std::get<0>(result), std::get<1>(result)};
return make_array_tuple(env, arrays);
} catch (const std::exception& e) {
return make_error(env, "qr_error");
}
}
// ==================== CHOLESKY DECOMPOSITION ====================
static ERL_NIF_TERM mlx_cholesky(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 2) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
bool upper = false;
if (argc > 1) {
char atom_name[32];
if (!enif_get_atom(env, argv[1], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
upper = (strcmp(atom_name, "true") == 0);
}
try {
array result = linalg::cholesky(a, upper);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "cholesky_error");
}
}
// ==================== MATRIX INVERSION ====================
static ERL_NIF_TERM mlx_inv(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
array result = linalg::inv(a);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "inv_error");
}
}
static ERL_NIF_TERM mlx_pinv(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 2) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
double rcond = 1e-15;
if (argc > 1) {
if (!enif_get_double(env, argv[1], &rcond)) {
return enif_make_badarg(env);
}
}
try {
array result = linalg::pinv(a); // MLX pinv doesn't take rcond parameter
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "pinv_error");
}
}
// ==================== MATRIX NORMS ====================
static ERL_NIF_TERM mlx_norm(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
// Default parameters
std::string ord = "fro";
std::vector<int> axis;
bool keepdims = false;
if (argc > 1) {
char atom_name[32];
if (enif_get_atom(env, argv[1], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
ord = std::string(atom_name);
} else {
int ord_int;
if (enif_get_int(env, argv[1], &ord_int)) {
ord = std::to_string(ord_int);
} else {
return enif_make_badarg(env);
}
}
}
if (argc > 2) {
// Parse axis - can be integer or list of integers
int single_axis;
if (enif_get_int(env, argv[2], &single_axis)) {
axis = {single_axis};
} else {
unsigned int axis_len;
if (enif_get_list_length(env, argv[2], &axis_len)) {
axis.resize(axis_len);
ERL_NIF_TERM head, tail = argv[2];
for (unsigned int i = 0; i < axis_len; i++) {
if (!enif_get_list_cell(env, tail, &head, &tail) ||
!enif_get_int(env, head, &axis[i])) {
return enif_make_badarg(env);
}
}
} else {
return enif_make_badarg(env);
}
}
}
try {
array result = array({1.0f}); // Initialize with dummy value
if (axis.empty()) {
result = linalg::norm(a, ord, std::nullopt, keepdims);
} else {
result = linalg::norm(a, ord, axis, keepdims);
}
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "norm_error");
}
}
// ==================== MATRIX DETERMINANT ====================
static ERL_NIF_TERM mlx_det(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
// TODO: MLX may add det in future versions
return make_error(env, "det_not_yet_available");
} catch (const std::exception& e) {
return make_error(env, "det_error");
}
}
static ERL_NIF_TERM mlx_slogdet(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
try {
// TODO: MLX may add slogdet in future versions
return make_error(env, "slogdet_not_yet_available");
} catch (const std::exception& e) {
return make_error(env, "slogdet_error");
}
}
// ==================== MATRIX TRACE ====================
static ERL_NIF_TERM mlx_trace(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 1 || argc > 4) return enif_make_badarg(env);
array a = array({1.0f}); // Initialize with dummy value
if (!get_array_resource(env, argv[0], a)) {
return enif_make_badarg(env);
}
int offset = 0;
int axis1 = -2;
int axis2 = -1;
if (argc > 1 && !enif_get_int(env, argv[1], &offset)) {
return enif_make_badarg(env);
}
if (argc > 2 && !enif_get_int(env, argv[2], &axis1)) {
return enif_make_badarg(env);
}
if (argc > 3 && !enif_get_int(env, argv[3], &axis2)) {
return enif_make_badarg(env);
}
try {
// TODO: MLX may add trace in future versions
return make_error(env, "trace_not_yet_available");
} catch (const std::exception& e) {
return make_error(env, "trace_error");
}
}
// ==================== CROSS PRODUCT ====================
static ERL_NIF_TERM mlx_cross(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 2 || argc > 3) return enif_make_badarg(env);
array a = array({1.0f}), b = array({1.0f}); // Initialize with dummy values
if (!get_array_resource(env, argv[0], a) ||
!get_array_resource(env, argv[1], b)) {
return enif_make_badarg(env);
}
int axis = -1;
if (argc > 2 && !enif_get_int(env, argv[2], &axis)) {
return enif_make_badarg(env);
}
try {
array result = linalg::cross(a, b, axis);
return make_array_resource(env, result);
} catch (const std::exception& e) {
return make_error(env, "cross_error");
}
}
// ==================== TRIANGULAR SOLVE ====================
static ERL_NIF_TERM mlx_tri_solve(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
if (argc < 2 || argc > 4) return enif_make_badarg(env);
array a = array({1.0f}), b = array({1.0f}); // Initialize with dummy values
if (!get_array_resource(env, argv[0], a) ||
!get_array_resource(env, argv[1], b)) {
return enif_make_badarg(env);
}
bool upper = false;
bool transpose = false;
if (argc > 2) {
char atom_name[32];
if (!enif_get_atom(env, argv[2], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
upper = (strcmp(atom_name, "true") == 0);
}
if (argc > 3) {
char atom_name[32];
if (!enif_get_atom(env, argv[3], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
return enif_make_badarg(env);
}
transpose = (strcmp(atom_name, "true") == 0);
}
try {
// TODO: MLX may add tri_solve in future versions
return make_error(env, "tri_solve_not_yet_available");
} catch (const std::exception& e) {
return make_error(env, "tri_solve_error");
}
}
// Resource destructor
static void array_resource_destructor(ErlNifEnv* env, void* obj) {
ArrayResource* res = (ArrayResource*)obj;
res->~ArrayResource();
}
// Linear algebra NIF function table
static ErlNifFunc nif_funcs[] = {
// SVD
{"svd", 1, mlx_svd, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"svd", 2, mlx_svd, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Eigenvalue decomposition
{"eig", 1, mlx_eig, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"eigvals", 1, mlx_eigvals, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"eigh", 1, mlx_eigh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"eigh", 2, mlx_eigh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"eigvalsh", 1, mlx_eigvalsh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"eigvalsh", 2, mlx_eigvalsh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// QR decomposition
{"qr", 1, mlx_qr, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"qr", 2, mlx_qr, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Cholesky decomposition
{"cholesky", 1, mlx_cholesky, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"cholesky", 2, mlx_cholesky, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Matrix inversion
{"inv", 1, mlx_inv, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"pinv", 1, mlx_pinv, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"pinv", 2, mlx_pinv, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Matrix norms
{"norm", 1, mlx_norm, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"norm", 2, mlx_norm, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"norm", 3, mlx_norm, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Matrix determinant
{"det", 1, mlx_det, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"slogdet", 1, mlx_slogdet, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Matrix trace
{"trace", 1, mlx_trace, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"trace", 2, mlx_trace, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"trace", 3, mlx_trace, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"trace", 4, mlx_trace, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Cross product
{"cross", 2, mlx_cross, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"cross", 3, mlx_cross, ERL_NIF_DIRTY_JOB_CPU_BOUND},
// Triangular solve
{"tri_solve", 2, mlx_tri_solve, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"tri_solve", 3, mlx_tri_solve, ERL_NIF_DIRTY_JOB_CPU_BOUND},
{"tri_solve", 4, mlx_tri_solve, 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_linalg_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_linalg_nif, nif_funcs, load, NULL, upgrade, unload)