-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).