Packages

MLX machine learning framework bindings for Erlang

Current section

Files

Jump to
mlx src mlx_advanced.erl
Raw

src/mlx_advanced.erl

-module(mlx_advanced).
%% Advanced mathematical operations for SOTA MLX implementation
-export([
%% Einstein summation and tensor operations
einsum/2, einsum/3,
%% Statistical functions
median/1, median/2,
percentile/2, percentile/3,
histogram/2, histogram/3,
covariance/2, correlation/2,
%% Advanced linear algebra
svd/1, eig/1, qr/1, cholesky/1,
pinv/1, matrix_rank/1,
trace/1, det/1,
%% Signal processing
fft/1, ifft/1, fft2/1, ifft2/1,
convolve/2, correlate/2,
%% Complex number operations
real/1, imag/1, conj/1, angle/1,
complex/2,
%% Advanced activation functions
gelu/1, swish/1, mish/1, selu/1,
leaky_relu/1, leaky_relu/2,
elu/1, elu/2,
%% Normalization operations
layer_norm/1, layer_norm/2, layer_norm/3,
batch_norm/1, batch_norm/2, batch_norm/3,
group_norm/2, group_norm/3,
rms_norm/1, rms_norm/2,
%% Loss functions
cross_entropy/2, cross_entropy/3,
focal_loss/2, focal_loss/3, focal_loss/4,
kl_divergence/2, js_divergence/2,
huber_loss/2, huber_loss/3,
%% Optimization functions
gradient_norm/1, gradient_clip/2, gradient_clip/3,
spectral_norm/1, weight_norm/1
]).
%% Einstein summation - fundamental for tensor operations
einsum(Equation, Arrays) ->
einsum(Equation, Arrays, []).
einsum(Equation, Arrays, Options) ->
mlx_nif:einsum(Equation, Arrays, Options).
%% Statistical functions
median(Array) ->
median(Array, -1).
median(Array, Axis) ->
mlx_nif:median(Array, Axis).
percentile(Array, Q) ->
percentile(Array, Q, -1).
percentile(Array, Q, Axis) ->
mlx_nif:percentile(Array, Q, Axis).
histogram(Array, Bins) ->
histogram(Array, Bins, []).
histogram(Array, Bins, Options) ->
mlx_nif:histogram(Array, Bins, Options).
covariance(X, Y) ->
% Compute covariance matrix
XMean = mlx:mean(X, 0),
YMean = mlx:mean(Y, 0),
XCentered = mlx:subtract(X, XMean),
YCentered = mlx:subtract(Y, YMean),
N = element(1, mlx:shape(X)),
Cov = mlx:matmul(mlx:transpose(XCentered), YCentered),
mlx:divide(Cov, mlx:array(N - 1)).
correlation(X, Y) ->
% Pearson correlation coefficient
case covariance(X, Y) of
{ok, Cov} ->
{ok, StdX} = mlx:std(X, 0),
{ok, StdY} = mlx:std(Y, 0),
{ok, StdProd} = mlx:outer(StdX, StdY),
mlx:divide(Cov, StdProd);
Error -> Error
end.
%% Advanced linear algebra
svd(Array) ->
mlx_nif:svd(Array).
eig(Array) ->
mlx_nif:eig(Array).
qr(Array) ->
mlx_nif:qr(Array).
cholesky(Array) ->
mlx_nif:cholesky(Array).
pinv(Array) ->
% Pseudo-inverse using SVD
case svd(Array) of
{ok, {U, S, Vt}} ->
% Create S_inv with threshold for numerical stability
{ok, Threshold} = mlx:multiply(mlx:array(1.0e-15), mlx:max(S)),
{ok, SMask} = mlx:greater(S, Threshold),
{ok, SInv} = mlx:where(SMask, mlx:reciprocal(S), mlx:array(0.0)),
{ok, SInvDiag} = mlx:diag(SInv),
{ok, Temp} = mlx:matmul(Vt, SInvDiag),
mlx:matmul(mlx:transpose(Temp), mlx:transpose(U));
Error -> Error
end.
matrix_rank(Array) ->
case svd(Array) of
{ok, {_U, S, _Vt}} ->
{ok, Threshold} = mlx:multiply(mlx:array(1.0e-12), mlx:max(S)),
{ok, Mask} = mlx:greater(S, Threshold),
mlx:sum(mlx:cast(Mask, int32));
Error -> Error
end.
trace(Array) ->
mlx_nif:trace(Array).
det(Array) ->
mlx_nif:det(Array).
%% Signal processing
fft(Array) ->
mlx_nif:fft(Array).
ifft(Array) ->
mlx_nif:ifft(Array).
fft2(Array) ->
mlx_nif:fft2(Array).
ifft2(Array) ->
mlx_nif:ifft2(Array).
convolve(A, B) ->
mlx_nif:convolve(A, B).
correlate(A, B) ->
mlx_nif:correlate(A, B).
%% Complex number operations
real(Array) ->
mlx_nif:real(Array).
imag(Array) ->
mlx_nif:imag(Array).
conj(Array) ->
mlx_nif:conj(Array).
angle(Array) ->
mlx_nif:angle(Array).
complex(Real, Imag) ->
mlx_nif:complex(Real, Imag).
%% Advanced activation functions
gelu(X) ->
% Gaussian Error Linear Unit
{ok, Half} = mlx:array(0.5),
{ok, One} = mlx:array(1.0),
{ok, Sqrt2} = mlx:array(1.4142135623730951),
{ok, XDivSqrt2} = mlx:divide(X, Sqrt2),
{ok, Erf} = mlx:erf(XDivSqrt2),
{ok, OnePlusErf} = mlx:add(One, Erf),
{ok, HalfX} = mlx:multiply(Half, X),
mlx:multiply(HalfX, OnePlusErf).
swish(X) ->
% Swish activation: x * sigmoid(x)
case mlx:sigmoid(X) of
{ok, Sig} -> mlx:multiply(X, Sig);
Error -> Error
end.
mish(X) ->
% Mish activation: x * tanh(softplus(x))
{ok, Softplus} = mlx:log(mlx:add(mlx:array(1.0), mlx:exp(X))),
{ok, TanhSoftplus} = mlx:tanh(Softplus),
mlx:multiply(X, TanhSoftplus).
selu(X) ->
% Scaled Exponential Linear Unit
Alpha = 1.6732632423543772848170429916717,
Scale = 1.0507009873554804934193349852946,
{ok, Zero} = mlx:array(0.0),
{ok, AlphaArray} = mlx:array(Alpha),
{ok, ScaleArray} = mlx:array(Scale),
{ok, Mask} = mlx:greater(X, Zero),
{ok, ExpPart} = mlx:subtract(mlx:exp(X), mlx:array(1.0)),
{ok, NegPart} = mlx:multiply(AlphaArray, ExpPart),
{ok, Result} = mlx:where(Mask, X, NegPart),
mlx:multiply(ScaleArray, Result).
leaky_relu(X) ->
leaky_relu(X, 0.01).
leaky_relu(X, Alpha) ->
{ok, Zero} = mlx:array(0.0),
{ok, AlphaArray} = mlx:array(Alpha),
{ok, Mask} = mlx:greater(X, Zero),
{ok, LeakyPart} = mlx:multiply(AlphaArray, X),
mlx:where(Mask, X, LeakyPart).
elu(X) ->
elu(X, 1.0).
elu(X, Alpha) ->
{ok, Zero} = mlx:array(0.0),
{ok, AlphaArray} = mlx:array(Alpha),
{ok, One} = mlx:array(1.0),
{ok, Mask} = mlx:greater(X, Zero),
{ok, ExpPart} = mlx:subtract(mlx:exp(X), One),
{ok, EluPart} = mlx:multiply(AlphaArray, ExpPart),
mlx:where(Mask, X, EluPart).
%% Normalization operations
layer_norm(X) ->
layer_norm(X, [], 1.0e-5).
layer_norm(X, Axes) ->
layer_norm(X, Axes, 1.0e-5).
layer_norm(X, Axes, Eps) ->
% Layer normalization
{ok, Mean} = mlx:mean(X, Axes, true),
{ok, Var} = mlx:var(X, Axes, true),
{ok, EpsArray} = mlx:array(Eps),
{ok, VarEps} = mlx:add(Var, EpsArray),
{ok, Std} = mlx:sqrt(VarEps),
{ok, Centered} = mlx:subtract(X, Mean),
mlx:divide(Centered, Std).
batch_norm(X) ->
batch_norm(X, 0, 1.0e-5).
batch_norm(X, Axis) ->
batch_norm(X, Axis, 1.0e-5).
batch_norm(X, Axis, Eps) ->
% Batch normalization
{ok, Mean} = mlx:mean(X, Axis, true),
{ok, Var} = mlx:var(X, Axis, true),
{ok, EpsArray} = mlx:array(Eps),
{ok, VarEps} = mlx:add(Var, EpsArray),
{ok, Std} = mlx:sqrt(VarEps),
{ok, Centered} = mlx:subtract(X, Mean),
mlx:divide(Centered, Std).
group_norm(X, NumGroups) ->
group_norm(X, NumGroups, 1.0e-5).
group_norm(X, NumGroups, Eps) ->
% Group normalization - simplified implementation
{ok, Shape} = mlx:shape(X),
[N, C | Rest] = Shape,
GroupSize = C div NumGroups,
NewShape = [N, NumGroups, GroupSize | Rest],
{ok, Reshaped} = mlx:reshape(X, NewShape),
{ok, Normalized} = layer_norm(Reshaped, [2], Eps),
mlx:reshape(Normalized, Shape).
rms_norm(X) ->
rms_norm(X, 1.0e-8).
rms_norm(X, Eps) ->
% Root Mean Square normalization
{ok, Square} = mlx:square(X),
{ok, MeanSquare} = mlx:mean(Square, -1, true),
{ok, EpsArray} = mlx:array(Eps),
{ok, MeanSquareEps} = mlx:add(MeanSquare, EpsArray),
{ok, Rms} = mlx:sqrt(MeanSquareEps),
mlx:divide(X, Rms).
%% Loss functions
cross_entropy(Predictions, Targets) ->
cross_entropy(Predictions, Targets, -1).
cross_entropy(Predictions, Targets, Axis) ->
{ok, LogSoftmax} = log_softmax(Predictions, Axis),
{ok, NegLogProb} = mlx:negative(LogSoftmax),
{ok, Loss} = mlx:multiply(Targets, NegLogProb),
mlx:sum(Loss, Axis).
focal_loss(Predictions, Targets) ->
focal_loss(Predictions, Targets, 2.0, 0.25).
focal_loss(Predictions, Targets, Gamma) ->
focal_loss(Predictions, Targets, Gamma, 0.25).
focal_loss(Predictions, Targets, Gamma, Alpha) ->
% Focal loss for addressing class imbalance
{ok, CE} = cross_entropy(Predictions, Targets),
{ok, P} = mlx:exp(mlx:negative(CE)),
{ok, One} = mlx:array(1.0),
{ok, OneMiusP} = mlx:subtract(One, P),
{ok, GammaArray} = mlx:array(Gamma),
{ok, AlphaArray} = mlx:array(Alpha),
{ok, FocalWeight} = mlx:power(OneMiusP, GammaArray),
{ok, WeightedCE} = mlx:multiply(AlphaArray, CE),
mlx:multiply(FocalWeight, WeightedCE).
kl_divergence(P, Q) ->
% Kullback-Leibler divergence
{ok, LogP} = mlx:log(P),
{ok, LogQ} = mlx:log(Q),
{ok, LogRatio} = mlx:subtract(LogP, LogQ),
{ok, KL} = mlx:multiply(P, LogRatio),
mlx:sum(KL).
js_divergence(P, Q) ->
% Jensen-Shannon divergence
{ok, Half} = mlx:array(0.5),
{ok, M} = mlx:multiply(Half, mlx:add(P, Q)),
{ok, KL1} = kl_divergence(P, M),
{ok, KL2} = kl_divergence(Q, M),
{ok, Sum} = mlx:add(KL1, KL2),
mlx:multiply(Half, Sum).
huber_loss(Predictions, Targets) ->
huber_loss(Predictions, Targets, 1.0).
huber_loss(Predictions, Targets, Delta) ->
% Huber loss (smooth L1 loss)
{ok, Diff} = mlx:subtract(Predictions, Targets),
{ok, AbsDiff} = mlx:abs(Diff),
{ok, DeltaArray} = mlx:array(Delta),
{ok, Mask} = mlx:less_equal(AbsDiff, DeltaArray),
{ok, Half} = mlx:array(0.5),
{ok, QuadraticPart} = mlx:multiply(Half, mlx:square(Diff)),
{ok, LinearPart} = mlx:subtract(mlx:multiply(DeltaArray, AbsDiff),
mlx:multiply(Half, mlx:square(DeltaArray))),
mlx:where(Mask, QuadraticPart, LinearPart).
%% Optimization functions
gradient_norm(Gradients) ->
% Compute L2 norm of gradients
{ok, Squares} = lists:foldl(fun(Grad, {ok, Acc}) ->
{ok, Square} = mlx:square(Grad),
case Acc of
undefined -> {ok, Square};
_ -> {ok, Sum} = mlx:add(Acc, Square), {ok, Sum}
end
end, {ok, undefined}, Gradients),
mlx:sqrt(mlx:sum(Squares)).
gradient_clip(Gradients, MaxNorm) ->
gradient_clip(Gradients, MaxNorm, l2).
gradient_clip(Gradients, MaxNorm, NormType) ->
case NormType of
l2 ->
{ok, TotalNorm} = gradient_norm(Gradients),
{ok, MaxNormArray} = mlx:array(MaxNorm),
{ok, ClipCoeff} = mlx:minimum(mlx:divide(MaxNormArray, TotalNorm), mlx:array(1.0)),
lists:map(fun(Grad) ->
{ok, Clipped} = mlx:multiply(Grad, ClipCoeff),
Clipped
end, Gradients);
value ->
MaxNormArray = mlx:array(MaxNorm),
MinNormArray = mlx:array(-MaxNorm),
lists:map(fun(Grad) ->
{ok, Clipped} = mlx:clip(Grad, MinNormArray, MaxNormArray),
Clipped
end, Gradients)
end.
spectral_norm(Weight) ->
% Spectral normalization using power iteration
case svd(Weight) of
{ok, {_U, S, _Vt}} ->
{ok, MaxS} = mlx:max(S),
mlx:divide(Weight, MaxS);
Error -> Error
end.
weight_norm(Weight) ->
% Weight normalization
{ok, Norm} = mlx:sqrt(mlx:sum(mlx:square(Weight))),
mlx:divide(Weight, Norm).
%% Helper functions
log_softmax(X, Axis) ->
{ok, MaxX} = mlx:max(X, Axis, true),
{ok, Shifted} = mlx:subtract(X, MaxX),
{ok, LogSumExp} = mlx:log(mlx:sum(mlx:exp(Shifted), Axis, true)),
mlx:subtract(Shifted, LogSumExp).