Current section

Files

Jump to
emily c_src ops unary.cpp
Raw

c_src/ops/unary.cpp

// Unary elementwise ops. All take a Tensor and return a Tensor.
#include "../emily/async.hpp"
#include "../emily/tensor.hpp"
#include "../emily/worker.hpp"
#include <fine.hpp>
#include <mlx/mlx.h>
namespace mx = mlx::core;
using emily::async_encoded;
using emily::Tensor;
using emily::wrap;
using emily::WorkerThread;
// Anonymous namespace so our NIF names (log1p, sqrt, sin, etc.) don't
// clash with C math-library functions brought in by MLX headers.
namespace {
#define EMILY_UNARY(op_name, mlx_fn) \
fine::Term op_name##_nif( \
ErlNifEnv *env, \
fine::ResourcePtr<WorkerThread> w, \
fine::ResourcePtr<Tensor> a) { \
return async_encoded(env, w, [a](mx::Stream &s) { \
return wrap(mlx_fn(a->array, s)); \
}); \
} \
FINE_NIF(op_name##_nif, 0);
EMILY_UNARY(negative, mx::negative)
EMILY_UNARY(abs, mx::abs)
EMILY_UNARY(sign, mx::sign)
EMILY_UNARY(floor, mx::floor)
EMILY_UNARY(ceil, mx::ceil)
EMILY_UNARY(sqrt, mx::sqrt)
EMILY_UNARY(rsqrt, mx::rsqrt)
EMILY_UNARY(exp, mx::exp)
EMILY_UNARY(expm1, mx::expm1)
EMILY_UNARY(log, mx::log)
EMILY_UNARY(log1p, mx::log1p)
EMILY_UNARY(log2, mx::log2)
EMILY_UNARY(log10, mx::log10)
EMILY_UNARY(sin, mx::sin)
EMILY_UNARY(cos, mx::cos)
EMILY_UNARY(tan, mx::tan)
EMILY_UNARY(arcsin, mx::arcsin)
EMILY_UNARY(arccos, mx::arccos)
EMILY_UNARY(arctan, mx::arctan)
EMILY_UNARY(sinh, mx::sinh)
EMILY_UNARY(cosh, mx::cosh)
EMILY_UNARY(tanh, mx::tanh)
EMILY_UNARY(arcsinh, mx::arcsinh)
EMILY_UNARY(arccosh, mx::arccosh)
EMILY_UNARY(arctanh, mx::arctanh)
EMILY_UNARY(sigmoid, mx::sigmoid)
EMILY_UNARY(erf, mx::erf)
EMILY_UNARY(erfinv, mx::erfinv)
EMILY_UNARY(square, mx::square)
EMILY_UNARY(reciprocal, mx::reciprocal)
EMILY_UNARY(logical_not, mx::logical_not)
EMILY_UNARY(bitwise_invert, mx::bitwise_invert)
EMILY_UNARY(isnan, mx::isnan)
EMILY_UNARY(isinf, mx::isinf)
EMILY_UNARY(isfinite, mx::isfinite)
EMILY_UNARY(conjugate, mx::conjugate)
EMILY_UNARY(real, mx::real)
EMILY_UNARY(imag, mx::imag)
EMILY_UNARY(stop_gradient, mx::stop_gradient)
#undef EMILY_UNARY
fine::Term round_nif(
ErlNifEnv *env,
fine::ResourcePtr<WorkerThread> w,
fine::ResourcePtr<Tensor> a,
int64_t decimals) {
return async_encoded(env, w, [a, decimals](mx::Stream &s) {
return wrap(mx::round(a->array, static_cast<int>(decimals), s));
});
}
FINE_NIF(round_nif, 0);
} // namespace