Current section
Files
Jump to
Current section
Files
src/mlx_linalg.erl
-module(mlx_linalg).
%% Linear algebra functions for MLX
-export([
%% Basic linear algebra
dot/2, matmul/2, inner/2, outer/2,
tensordot/2, tensordot/3,
%% Matrix decompositions
svd/1, svd/2,
eig/1, eigh/1,
qr/1, qr/2,
cholesky/1,
lu/1, lu/2,
%% Matrix properties
det/1, slogdet/1,
trace/1, diagonal/1, diagonal/2,
matrix_rank/1, matrix_rank/2,
norm/1, norm/2, norm/3,
%% Matrix operations
inv/1, pinv/1, pinv/2,
solve/2, lstsq/2,
matrix_power/2,
%% Vector operations
cross/2, cross/3,
%% Eigenvalues and eigenvectors
eigvals/1, eigvalsh/1,
%% Condition numbers
condition_number/1, condition_number/2
]).
%% Type definitions
-type array() :: reference().
-type shape() :: [integer()].
-type dtype() :: atom().
%% Basic linear algebra
-spec dot(array(), array()) -> {ok, array()} | {error, term()}.
dot(A, B) ->
mlx_linalg_nif:dot(A, B).
-spec matmul(array(), array()) -> {ok, array()} | {error, term()}.
matmul(A, B) ->
mlx_linalg_nif:matmul(A, B).
-spec inner(array(), array()) -> {ok, array()} | {error, term()}.
inner(A, B) ->
mlx_linalg_nif:inner(A, B).
-spec outer(array(), array()) -> {ok, array()} | {error, term()}.
outer(A, B) ->
mlx_linalg_nif:outer(A, B).
-spec tensordot(array(), array()) -> {ok, array()} | {error, term()}.
tensordot(A, B) ->
tensordot(A, B, 2).
-spec tensordot(array(), array(), integer()) -> {ok, array()} | {error, term()}.
tensordot(A, B, Axes) ->
mlx_linalg_nif:tensordot(A, B, Axes).
%% Matrix decompositions
-spec svd(array()) -> {ok, {array(), array(), array()}} | {error, term()}.
svd(A) ->
svd(A, true).
-spec svd(array(), boolean()) -> {ok, {array(), array(), array()}} | {error, term()}.
svd(A, FullMatrices) ->
mlx_linalg_nif:svd(A, FullMatrices).
-spec eig(array()) -> {ok, {array(), array()}} | {error, term()}.
eig(A) ->
mlx_linalg_nif:eig(A).
-spec eigh(array()) -> {ok, {array(), array()}} | {error, term()}.
eigh(A) ->
mlx_linalg_nif:eigh(A).
-spec qr(array()) -> {ok, {array(), array()}} | {error, term()}.
qr(A) ->
qr(A, reduced).
-spec qr(array(), atom()) -> {ok, {array(), array()}} | {error, term()}.
qr(A, Mode) ->
mlx_linalg_nif:qr(A, Mode).
-spec cholesky(array()) -> {ok, array()} | {error, term()}.
cholesky(A) ->
mlx_linalg_nif:cholesky(A).
-spec lu(array()) -> {ok, {array(), array(), array()}} | {error, term()}.
lu(A) ->
lu(A, true).
-spec lu(array(), boolean()) -> {ok, {array(), array(), array()}} | {error, term()}.
lu(A, Permute) ->
mlx_linalg_nif:lu(A, Permute).
%% Matrix properties
-spec det(array()) -> {ok, array()} | {error, term()}.
det(A) ->
mlx_linalg_nif:det(A).
-spec slogdet(array()) -> {ok, {array(), array()}} | {error, term()}.
slogdet(A) ->
mlx_linalg_nif:slogdet(A).
-spec trace(array()) -> {ok, array()} | {error, term()}.
trace(A) ->
mlx_linalg_nif:trace(A).
-spec diagonal(array()) -> {ok, array()} | {error, term()}.
diagonal(A) ->
diagonal(A, 0).
-spec diagonal(array(), integer()) -> {ok, array()} | {error, term()}.
diagonal(A, Offset) ->
mlx_linalg_nif:diagonal(A, Offset).
-spec matrix_rank(array()) -> {ok, array()} | {error, term()}.
matrix_rank(A) ->
matrix_rank(A, 1.0e-8).
-spec matrix_rank(array(), number()) -> {ok, array()} | {error, term()}.
matrix_rank(A, Tol) ->
mlx_linalg_nif:matrix_rank(A, Tol).
-spec norm(array()) -> {ok, array()} | {error, term()}.
norm(A) ->
norm(A, fro, []).
-spec norm(array(), atom() | number()) -> {ok, array()} | {error, term()}.
norm(A, Ord) ->
norm(A, Ord, []).
-spec norm(array(), atom() | number(), [integer()]) -> {ok, array()} | {error, term()}.
norm(A, Ord, Axis) ->
mlx_linalg_nif:norm(A, Ord, Axis).
%% Matrix operations
-spec inv(array()) -> {ok, array()} | {error, term()}.
inv(A) ->
mlx_linalg_nif:inv(A).
-spec pinv(array()) -> {ok, array()} | {error, term()}.
pinv(A) ->
pinv(A, 1.0e-15).
-spec pinv(array(), number()) -> {ok, array()} | {error, term()}.
pinv(A, Rcond) ->
mlx_linalg_nif:pinv(A, Rcond).
-spec solve(array(), array()) -> {ok, array()} | {error, term()}.
solve(A, B) ->
mlx_linalg_nif:solve(A, B).
-spec lstsq(array(), array()) -> {ok, {array(), array(), integer(), array()}} | {error, term()}.
lstsq(A, B) ->
mlx_linalg_nif:lstsq(A, B).
-spec matrix_power(array(), integer()) -> {ok, array()} | {error, term()}.
matrix_power(A, N) ->
mlx_linalg_nif:matrix_power(A, N).
%% Vector operations
-spec cross(array(), array()) -> {ok, array()} | {error, term()}.
cross(A, B) ->
cross(A, B, -1).
-spec cross(array(), array(), integer()) -> {ok, array()} | {error, term()}.
cross(A, B, Axis) ->
mlx_linalg_nif:cross(A, B, Axis).
%% Eigenvalues and eigenvectors
-spec eigvals(array()) -> {ok, array()} | {error, term()}.
eigvals(A) ->
mlx_linalg_nif:eigvals(A).
-spec eigvalsh(array()) -> {ok, array()} | {error, term()}.
eigvalsh(A) ->
mlx_linalg_nif:eigvalsh(A).
%% Condition numbers
-spec condition_number(array()) -> {ok, array()} | {error, term()}.
condition_number(A) ->
condition_number(A, 2).
-spec condition_number(array(), atom() | number()) -> {ok, array()} | {error, term()}.
condition_number(A, P) ->
mlx_linalg_nif:condition_number(A, P).