#include <erl_nif.h>
#include <mlx/mlx.h>
#include <mlx/ops.h>
#include <mlx/array.h>
#include <mlx/nn.h>
#include <memory>
#include <vector>
#include <string>
#include <map>

using namespace mlx::core;
namespace mx = mlx::core;
namespace nn = mlx::nn;

// Resource types for neural network components
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);
}

// 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;
}

// ==================== ACTIVATION FUNCTIONS ====================

static ERL_NIF_TERM mlx_relu(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array x;
    if (!get_array_resource(env, argv[0], x)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = relu(x);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "relu_error");
    }
}

static ERL_NIF_TERM mlx_gelu(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array x;
    if (!get_array_resource(env, argv[0], x)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = gelu(x);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "gelu_error");
    }
}

static ERL_NIF_TERM mlx_sigmoid(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array x;
    if (!get_array_resource(env, argv[0], x)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = sigmoid(x);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "sigmoid_error");
    }
}

static ERL_NIF_TERM mlx_softmax(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 2) return enif_make_badarg(env);
    
    array x;
    if (!get_array_resource(env, argv[0], x)) {
        return enif_make_badarg(env);
    }
    
    int axis = -1;
    if (argc == 2 && !enif_get_int(env, argv[1], &axis)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = softmax(x, axis);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "softmax_error");
    }
}

static ERL_NIF_TERM mlx_log_softmax(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 2) return enif_make_badarg(env);
    
    array x;
    if (!get_array_resource(env, argv[0], x)) {
        return enif_make_badarg(env);
    }
    
    int axis = -1;
    if (argc == 2 && !enif_get_int(env, argv[1], &axis)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = log_softmax(x, axis);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "log_softmax_error");
    }
}

static ERL_NIF_TERM mlx_swish(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array x;
    if (!get_array_resource(env, argv[0], x)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = swish(x);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "swish_error");
    }
}

static ERL_NIF_TERM mlx_leaky_relu(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);
    }
    
    double alpha;
    if (!enif_get_double(env, argv[1], &alpha)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = leaky_relu(x, alpha);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "leaky_relu_error");
    }
}

static ERL_NIF_TERM mlx_elu(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);
    }
    
    double alpha;
    if (!enif_get_double(env, argv[1], &alpha)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = elu(x, alpha);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "elu_error");
    }
}

static ERL_NIF_TERM mlx_selu(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array x;
    if (!get_array_resource(env, argv[0], x)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = selu(x);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "selu_error");
    }
}

static ERL_NIF_TERM mlx_mish(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array x;
    if (!get_array_resource(env, argv[0], x)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = mish(x);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "mish_error");
    }
}

// ==================== LOSS FUNCTIONS ====================

static ERL_NIF_TERM mlx_mse_loss(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array predictions, targets;
    if (!get_array_resource(env, argv[0], predictions) ||
        !get_array_resource(env, argv[1], targets)) {
        return enif_make_badarg(env);
    }
    
    try {
        array diff = subtract(predictions, targets);
        array squared_diff = square(diff);
        array result = mean(squared_diff);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "mse_loss_error");
    }
}

static ERL_NIF_TERM mlx_cross_entropy_loss(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array logits, targets;
    if (!get_array_resource(env, argv[0], logits) ||
        !get_array_resource(env, argv[1], targets)) {
        return enif_make_badarg(env);
    }
    
    try {
        array log_probs = log_softmax(logits, -1);
        array loss = negative(sum(multiply(targets, log_probs), -1));
        array result = mean(loss);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "cross_entropy_loss_error");
    }
}

static ERL_NIF_TERM mlx_binary_cross_entropy_loss(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array predictions, targets;
    if (!get_array_resource(env, argv[0], predictions) ||
        !get_array_resource(env, argv[1], targets)) {
        return enif_make_badarg(env);
    }
    
    try {
        array eps = array(1e-12f);
        array clipped_pred = clip(predictions, eps, subtract(ones_like(predictions), eps));
        array log_pred = log(clipped_pred);
        array log_one_minus_pred = log(subtract(ones_like(clipped_pred), clipped_pred));
        
        array loss = negative(add(
            multiply(targets, log_pred),
            multiply(subtract(ones_like(targets), targets), log_one_minus_pred)
        ));
        array result = mean(loss);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "binary_cross_entropy_loss_error");
    }
}

static ERL_NIF_TERM mlx_l1_loss(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array predictions, targets;
    if (!get_array_resource(env, argv[0], predictions) ||
        !get_array_resource(env, argv[1], targets)) {
        return enif_make_badarg(env);
    }
    
    try {
        array diff = subtract(predictions, targets);
        array abs_diff = abs(diff);
        array result = mean(abs_diff);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "l1_loss_error");
    }
}

static ERL_NIF_TERM mlx_huber_loss(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    array predictions, targets;
    if (!get_array_resource(env, argv[0], predictions) ||
        !get_array_resource(env, argv[1], targets)) {
        return enif_make_badarg(env);
    }
    
    double delta;
    if (!enif_get_double(env, argv[2], &delta)) {
        return enif_make_badarg(env);
    }
    
    try {
        array diff = subtract(predictions, targets);
        array abs_diff = abs(diff);
        array delta_arr = array(static_cast<float>(delta));
        
        array condition = less_equal(abs_diff, delta_arr);
        array quadratic = multiply(array(0.5f), square(diff));
        array linear = subtract(multiply(delta_arr, abs_diff), multiply(array(0.5f), square(delta_arr)));
        
        array loss = where(condition, quadratic, linear);
        array result = mean(loss);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "huber_loss_error");
    }
}

// ==================== NORMALIZATION LAYERS ====================

static ERL_NIF_TERM mlx_layer_norm(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    array x, weight, bias;
    if (!get_array_resource(env, argv[0], x) ||
        !get_array_resource(env, argv[1], weight) ||
        !get_array_resource(env, argv[2], bias)) {
        return enif_make_badarg(env);
    }
    
    double eps;
    if (!enif_get_double(env, argv[3], &eps)) {
        return enif_make_badarg(env);
    }
    
    try {
        // Layer normalization implementation
        std::vector<int> axes = {-1}; // Normalize over last dimension
        array mean_x = mean(x, axes, true);
        array var_x = var(x, axes, true);
        array normalized = divide(subtract(x, mean_x), sqrt(add(var_x, array(static_cast<float>(eps)))));
        array result = add(multiply(normalized, weight), bias);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "layer_norm_error");
    }
}

static ERL_NIF_TERM mlx_batch_norm(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 6) return enif_make_badarg(env);
    
    array x, weight, bias, running_mean, running_var;
    if (!get_array_resource(env, argv[0], x) ||
        !get_array_resource(env, argv[1], weight) ||
        !get_array_resource(env, argv[2], bias) ||
        !get_array_resource(env, argv[3], running_mean) ||
        !get_array_resource(env, argv[4], running_var)) {
        return enif_make_badarg(env);
    }
    
    double eps;
    if (!enif_get_double(env, argv[5], &eps)) {
        return enif_make_badarg(env);
    }
    
    try {
        // Batch normalization implementation
        array normalized = divide(subtract(x, running_mean), sqrt(add(running_var, array(static_cast<float>(eps)))));
        array result = add(multiply(normalized, weight), bias);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "batch_norm_error");
    }
}

static ERL_NIF_TERM mlx_group_norm(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 5) return enif_make_badarg(env);
    
    array x, weight, bias;
    if (!get_array_resource(env, argv[0], x) ||
        !get_array_resource(env, argv[1], weight) ||
        !get_array_resource(env, argv[2], bias)) {
        return enif_make_badarg(env);
    }
    
    int num_groups;
    double eps;
    if (!enif_get_int(env, argv[3], &num_groups) ||
        !enif_get_double(env, argv[4], &eps)) {
        return enif_make_badarg(env);
    }
    
    try {
        // Group normalization implementation (simplified)
        auto shape = x.shape();
        int batch_size = shape[0];
        int channels = shape[1];
        int group_size = channels / num_groups;
        
        // Reshape for group-wise normalization
        std::vector<int> group_shape = {batch_size, num_groups, group_size};
        if (shape.size() > 2) {
            for (size_t i = 2; i < shape.size(); i++) {
                group_shape.push_back(shape[i]);
            }
        }
        
        array reshaped = reshape(x, group_shape);
        std::vector<int> norm_axes = {2}; // Normalize over group dimension
        array mean_x = mean(reshaped, norm_axes, true);
        array var_x = var(reshaped, norm_axes, true);
        array normalized = divide(subtract(reshaped, mean_x), sqrt(add(var_x, array(static_cast<float>(eps)))));
        
        array restored = reshape(normalized, shape);
        array result = add(multiply(restored, weight), bias);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "group_norm_error");
    }
}

// ==================== DROPOUT ====================

static ERL_NIF_TERM mlx_dropout(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    array x;
    if (!get_array_resource(env, argv[0], x)) {
        return enif_make_badarg(env);
    }
    
    double p;
    int training;
    if (!enif_get_double(env, argv[1], &p) ||
        !enif_get_int(env, argv[2], &training)) {
        return enif_make_badarg(env);
    }
    
    try {
        if (!training || p == 0.0) {
            return make_array_resource(env, x);
        }
        
        if (p == 1.0) {
            array result = zeros_like(x);
            return make_array_resource(env, result);
        }
        
        // Generate random mask
        array keep_prob = array(1.0f - static_cast<float>(p));
        array random_vals = random::uniform(0.0f, 1.0f, x.shape());
        array mask = greater(random_vals, array(static_cast<float>(p)));
        array result = divide(multiply(x, mask), keep_prob);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "dropout_error");
    }
}

// ==================== CONVOLUTION OPERATIONS ====================

static ERL_NIF_TERM mlx_conv1d(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 5) return enif_make_badarg(env);
    
    array input, weight;
    if (!get_array_resource(env, argv[0], input) ||
        !get_array_resource(env, argv[1], weight)) {
        return enif_make_badarg(env);
    }
    
    int stride, padding, dilation;
    if (!enif_get_int(env, argv[2], &stride) ||
        !enif_get_int(env, argv[3], &padding) ||
        !enif_get_int(env, argv[4], &dilation)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = conv1d(input, weight, stride, padding, dilation);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "conv1d_error");
    }
}

static ERL_NIF_TERM mlx_conv2d(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 5) return enif_make_badarg(env);
    
    array input, weight;
    if (!get_array_resource(env, argv[0], input) ||
        !get_array_resource(env, argv[1], weight)) {
        return enif_make_badarg(env);
    }
    
    auto stride = parse_shape(env, argv[2]);
    auto padding = parse_shape(env, argv[3]);
    auto dilation = parse_shape(env, argv[4]);
    
    if (stride.size() != 2 || padding.size() != 2 || dilation.size() != 2) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = conv2d(input, weight, {stride[0], stride[1]}, {padding[0], padding[1]}, {dilation[0], dilation[1]});
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "conv2d_error");
    }
}

// ==================== POOLING OPERATIONS ====================

static ERL_NIF_TERM mlx_max_pool1d(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    array x;
    if (!get_array_resource(env, argv[0], x)) {
        return enif_make_badarg(env);
    }
    
    int kernel_size, stride, padding;
    if (!enif_get_int(env, argv[1], &kernel_size) ||
        !enif_get_int(env, argv[2], &stride) ||
        !enif_get_int(env, argv[3], &padding)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = max_pool1d(x, kernel_size, stride, padding);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "max_pool1d_error");
    }
}

static ERL_NIF_TERM mlx_max_pool2d(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    array x;
    if (!get_array_resource(env, argv[0], x)) {
        return enif_make_badarg(env);
    }
    
    auto kernel_size = parse_shape(env, argv[1]);
    auto stride = parse_shape(env, argv[2]);
    auto padding = parse_shape(env, argv[3]);
    
    if (kernel_size.size() != 2 || stride.size() != 2 || padding.size() != 2) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = max_pool2d(x, {kernel_size[0], kernel_size[1]}, {stride[0], stride[1]}, {padding[0], padding[1]});
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "max_pool2d_error");
    }
}

static ERL_NIF_TERM mlx_avg_pool1d(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    array x;
    if (!get_array_resource(env, argv[0], x)) {
        return enif_make_badarg(env);
    }
    
    int kernel_size, stride, padding;
    if (!enif_get_int(env, argv[1], &kernel_size) ||
        !enif_get_int(env, argv[2], &stride) ||
        !enif_get_int(env, argv[3], &padding)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = avg_pool1d(x, kernel_size, stride, padding);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "avg_pool1d_error");
    }
}

static ERL_NIF_TERM mlx_avg_pool2d(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    array x;
    if (!get_array_resource(env, argv[0], x)) {
        return enif_make_badarg(env);
    }
    
    auto kernel_size = parse_shape(env, argv[1]);
    auto stride = parse_shape(env, argv[2]);
    auto padding = parse_shape(env, argv[3]);
    
    if (kernel_size.size() != 2 || stride.size() != 2 || padding.size() != 2) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = avg_pool2d(x, {kernel_size[0], kernel_size[1]}, {stride[0], stride[1]}, {padding[0], padding[1]});
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "avg_pool2d_error");
    }
}

// ==================== EMBEDDING ====================

static ERL_NIF_TERM mlx_embedding(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array indices, weights;
    if (!get_array_resource(env, argv[0], indices) ||
        !get_array_resource(env, argv[1], weights)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = take(weights, indices, 0);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "embedding_error");
    }
}

// Resource destructor
static void array_resource_destructor(ErlNifEnv* env, void* obj) {
    ArrayResource* res = (ArrayResource*)obj;
    res->~ArrayResource();
}

// Neural network NIF function table
static ErlNifFunc nif_funcs[] = {
    // Activation functions
    {"relu", 1, mlx_relu, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"gelu", 1, mlx_gelu, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"sigmoid", 1, mlx_sigmoid, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"softmax", 1, mlx_softmax, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"softmax", 2, mlx_softmax, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"log_softmax", 1, mlx_log_softmax, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"log_softmax", 2, mlx_log_softmax, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"swish", 1, mlx_swish, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"leaky_relu", 2, mlx_leaky_relu, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"elu", 2, mlx_elu, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"selu", 1, mlx_selu, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"mish", 1, mlx_mish, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Loss functions
    {"mse_loss", 2, mlx_mse_loss, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"cross_entropy_loss", 2, mlx_cross_entropy_loss, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"binary_cross_entropy_loss", 2, mlx_binary_cross_entropy_loss, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"l1_loss", 2, mlx_l1_loss, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"huber_loss", 3, mlx_huber_loss, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Normalization
    {"layer_norm", 4, mlx_layer_norm, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"batch_norm", 6, mlx_batch_norm, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"group_norm", 5, mlx_group_norm, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Dropout
    {"dropout", 3, mlx_dropout, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Convolution
    {"conv1d", 5, mlx_conv1d, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"conv2d", 5, mlx_conv2d, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Pooling
    {"max_pool1d", 4, mlx_max_pool1d, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"max_pool2d", 4, mlx_max_pool2d, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"avg_pool1d", 4, mlx_avg_pool1d, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"avg_pool2d", 4, mlx_avg_pool2d, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Embedding
    {"embedding", 2, mlx_embedding, 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_nn_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_nn_nif, nif_funcs, load, NULL, upgrade, unload)