Packages
Tensor library for Gleam/BEAM with a pure Gleam API, zero-copy views, and optional native acceleration
Current section
Files
Jump to
Current section
Files
src/viva_tensor@nn@optim.erl
-module(viva_tensor@nn@optim).
-compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]).
-define(FILEPATH, "src/viva_tensor/nn/optim.gleam").
-export([sgd/1, sgd_momentum/2, rmsprop/3, adam/1, adamw/2, step/3, zero_grad/1]).
-export_type([optimizer_kind/0, param/0, grad_pair/0, param_state/0, optimizer/0]).
-if(?OTP_RELEASE >= 27).
-define(MODULEDOC(Str), -moduledoc(Str)).
-define(DOC(Str), -doc(Str)).
-else.
-define(MODULEDOC(Str), -compile([])).
-define(DOC(Str), -compile([])).
-endif.
?MODULEDOC(false).
-type optimizer_kind() :: sgd | sgd_momentum | rmsprop | adam | adamw.
-type param() :: {param, binary(), viva_tensor@tensor:tensor()}.
-type grad_pair() :: {grad_pair, binary(), viva_tensor@tensor:tensor()}.
-type param_state() :: empty_state |
{momentum_state, viva_tensor@tensor:tensor()} |
{rmsprop_state, viva_tensor@tensor:tensor()} |
{adam_state,
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
integer()}.
-type optimizer() :: {optimizer,
optimizer_kind(),
float(),
float(),
float(),
float(),
float(),
float(),
gleam@dict:dict(binary(), param_state())}.
-file("src/viva_tensor/nn/optim.gleam", 89).
?DOC(false).
-spec sgd(float()) -> optimizer().
sgd(Lr) ->
{optimizer, sgd, Lr, +0.0, +0.0, +0.0, +0.0, +0.0, maps:new()}.
-file("src/viva_tensor/nn/optim.gleam", 110).
?DOC(false).
-spec sgd_momentum(float(), float()) -> optimizer().
sgd_momentum(Lr, Momentum) ->
{optimizer, sgd_momentum, Lr, Momentum, +0.0, +0.0, +0.0, +0.0, maps:new()}.
-file("src/viva_tensor/nn/optim.gleam", 134).
?DOC(false).
-spec rmsprop(float(), float(), float()) -> optimizer().
rmsprop(Lr, Alpha, Eps) ->
{optimizer, rmsprop, Lr, +0.0, +0.0, Alpha, Eps, +0.0, maps:new()}.
-file("src/viva_tensor/nn/optim.gleam", 161).
?DOC(false).
-spec adam(float()) -> optimizer().
adam(Lr) ->
{optimizer, adam, Lr, +0.0, 0.9, 0.999, 1.0e-8, +0.0, maps:new()}.
-file("src/viva_tensor/nn/optim.gleam", 184).
?DOC(false).
-spec adamw(float(), float()) -> optimizer().
adamw(Lr, Weight_decay) ->
{optimizer, adamw, Lr, +0.0, 0.9, 0.999, 1.0e-8, Weight_decay, maps:new()}.
-file("src/viva_tensor/nn/optim.gleam", 438).
?DOC(false).
-spec float_sqrt(float()) -> float().
float_sqrt(X) ->
case gleam@float:square_root(X) of
{ok, V} ->
V;
{error, _} ->
+0.0
end.
-file("src/viva_tensor/nn/optim.gleam", 445).
?DOC(false).
-spec pow_int(float(), integer()) -> float().
pow_int(Base, Exp) ->
case gleam@float:power(Base, erlang:float(Exp)) of
{ok, V} ->
V;
{error, _} ->
+0.0
end.
-file("src/viva_tensor/nn/optim.gleam", 394).
?DOC(false).
-spec adam_update(optimizer(), param(), viva_tensor@tensor:tensor(), boolean()) -> {ok,
{optimizer(), viva_tensor@tensor:tensor()}} |
{error, viva_tensor@core@error:tensor_error()}.
adam_update(Opt, Param, Grad, Decoupled_decay) ->
{Prev_m, Prev_v, Prev_t} = case gleam_stdlib:map_get(
erlang:element(9, Opt),
erlang:element(2, Param)
) of
{ok, {adam_state, M, V, T}} ->
{M, V, T};
_ ->
{viva_tensor@tensor:zeros_like(erlang:element(3, Param)),
viva_tensor@tensor:zeros_like(erlang:element(3, Param)),
0}
end,
T@1 = Prev_t + 1,
M_term = viva_tensor@tensor:scale(Prev_m, erlang:element(5, Opt)),
M_grad = viva_tensor@tensor:scale(Grad, 1.0 - erlang:element(5, Opt)),
gleam@result:'try'(
viva_tensor@tensor:add(M_term, M_grad),
fun(New_m) ->
gleam@result:'try'(
viva_tensor@tensor:mul(Grad, Grad),
fun(G_sq) ->
V_term = viva_tensor@tensor:scale(
Prev_v,
erlang:element(6, Opt)
),
V_grad = viva_tensor@tensor:scale(
G_sq,
1.0 - erlang:element(6, Opt)
),
gleam@result:'try'(
viva_tensor@tensor:add(V_term, V_grad),
fun(New_v) ->
Bc1 = 1.0 - pow_int(erlang:element(5, Opt), T@1),
Bc2 = 1.0 - pow_int(erlang:element(6, Opt), T@1),
M_hat = viva_tensor@tensor:scale(New_m, case Bc1 of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> 1.0 / Gleam@denominator
end),
V_hat = viva_tensor@tensor:scale(New_v, case Bc2 of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator@1 -> 1.0 / Gleam@denominator@1
end),
Denom = viva_tensor@tensor:map(
V_hat,
fun(X) ->
float_sqrt(X) + erlang:element(7, Opt)
end
),
gleam@result:'try'(
viva_tensor@tensor:'div'(M_hat, Denom),
fun(Ratio) ->
Step_t = viva_tensor@tensor:scale(
Ratio,
erlang:element(3, Opt)
),
gleam@result:'try'(case Decoupled_decay of
true ->
Decay = viva_tensor@tensor:scale(
erlang:element(3, Param),
erlang:element(3, Opt) * erlang:element(
8,
Opt
)
),
viva_tensor@tensor:sub(
erlang:element(3, Param),
Decay
);
false ->
{ok, erlang:element(3, Param)}
end, fun(Base) ->
gleam@result:'try'(
viva_tensor@tensor:sub(
Base,
Step_t
),
fun(New_value) ->
New_state = gleam@dict:insert(
erlang:element(9, Opt),
erlang:element(2, Param),
{adam_state,
New_m,
New_v,
T@1}
),
{ok,
{{optimizer,
erlang:element(
2,
Opt
),
erlang:element(
3,
Opt
),
erlang:element(
4,
Opt
),
erlang:element(
5,
Opt
),
erlang:element(
6,
Opt
),
erlang:element(
7,
Opt
),
erlang:element(
8,
Opt
),
New_state},
New_value}}
end
)
end)
end
)
end
)
end
)
end
).
-file("src/viva_tensor/nn/optim.gleam", 368).
?DOC(false).
-spec rmsprop_update(optimizer(), param(), viva_tensor@tensor:tensor()) -> {ok,
{optimizer(), viva_tensor@tensor:tensor()}} |
{error, viva_tensor@core@error:tensor_error()}.
rmsprop_update(Opt, Param, Grad) ->
Prev_s = case gleam_stdlib:map_get(
erlang:element(9, Opt),
erlang:element(2, Param)
) of
{ok, {rmsprop_state, S}} ->
S;
_ ->
viva_tensor@tensor:zeros_like(erlang:element(3, Param))
end,
Alpha = erlang:element(6, Opt),
gleam@result:'try'(
viva_tensor@tensor:mul(Grad, Grad),
fun(G_sq) ->
S_term = viva_tensor@tensor:scale(Prev_s, Alpha),
G_term = viva_tensor@tensor:scale(G_sq, 1.0 - Alpha),
gleam@result:'try'(
viva_tensor@tensor:add(S_term, G_term),
fun(New_s) ->
Denom = viva_tensor@tensor:map(
New_s,
fun(X) -> float_sqrt(X) + erlang:element(7, Opt) end
),
gleam@result:'try'(
viva_tensor@tensor:'div'(Grad, Denom),
fun(Ratio) ->
Update = viva_tensor@tensor:scale(
Ratio,
erlang:element(3, Opt)
),
gleam@result:'try'(
viva_tensor@tensor:sub(
erlang:element(3, Param),
Update
),
fun(New_value) ->
New_state = gleam@dict:insert(
erlang:element(9, Opt),
erlang:element(2, Param),
{rmsprop_state, New_s}
),
{ok,
{{optimizer,
erlang:element(2, Opt),
erlang:element(3, Opt),
erlang:element(4, Opt),
erlang:element(5, Opt),
erlang:element(6, Opt),
erlang:element(7, Opt),
erlang:element(8, Opt),
New_state},
New_value}}
end
)
end
)
end
)
end
).
-file("src/viva_tensor/nn/optim.gleam", 347).
?DOC(false).
-spec momentum_update(optimizer(), param(), viva_tensor@tensor:tensor()) -> {ok,
{optimizer(), viva_tensor@tensor:tensor()}} |
{error, viva_tensor@core@error:tensor_error()}.
momentum_update(Opt, Param, Grad) ->
Prev_v = case gleam_stdlib:map_get(
erlang:element(9, Opt),
erlang:element(2, Param)
) of
{ok, {momentum_state, V}} ->
V;
_ ->
viva_tensor@tensor:zeros_like(erlang:element(3, Param))
end,
Momentum_term = viva_tensor@tensor:scale(Prev_v, erlang:element(4, Opt)),
gleam@result:'try'(
viva_tensor@tensor:add(Momentum_term, Grad),
fun(New_v) ->
Step_term = viva_tensor@tensor:scale(New_v, erlang:element(3, Opt)),
gleam@result:'try'(
viva_tensor@tensor:sub(erlang:element(3, Param), Step_term),
fun(New_value) ->
New_state = gleam@dict:insert(
erlang:element(9, Opt),
erlang:element(2, Param),
{momentum_state, New_v}
),
{ok,
{{optimizer,
erlang:element(2, Opt),
erlang:element(3, Opt),
erlang:element(4, Opt),
erlang:element(5, Opt),
erlang:element(6, Opt),
erlang:element(7, Opt),
erlang:element(8, Opt),
New_state},
New_value}}
end
)
end
).
-file("src/viva_tensor/nn/optim.gleam", 335).
?DOC(false).
-spec sgd_update(optimizer(), param(), viva_tensor@tensor:tensor()) -> {ok,
{optimizer(), viva_tensor@tensor:tensor()}} |
{error, viva_tensor@core@error:tensor_error()}.
sgd_update(Opt, Param, Grad) ->
Scaled = viva_tensor@tensor:scale(Grad, erlang:element(3, Opt)),
gleam@result:'try'(
viva_tensor@tensor:sub(erlang:element(3, Param), Scaled),
fun(New_value) -> {ok, {Opt, New_value}} end
).
-file("src/viva_tensor/nn/optim.gleam", 319).
?DOC(false).
-spec update_param(optimizer(), param(), viva_tensor@tensor:tensor()) -> {ok,
{optimizer(), viva_tensor@tensor:tensor()}} |
{error, viva_tensor@core@error:tensor_error()}.
update_param(Opt, Param, Grad) ->
case erlang:element(2, Opt) of
sgd ->
sgd_update(Opt, Param, Grad);
sgd_momentum ->
momentum_update(Opt, Param, Grad);
rmsprop ->
rmsprop_update(Opt, Param, Grad);
adam ->
adam_update(Opt, Param, Grad, false);
adamw ->
adam_update(Opt, Param, Grad, true)
end.
-file("src/viva_tensor/nn/optim.gleam", 296).
?DOC(false).
-spec do_apply(
optimizer(),
list(param()),
gleam@dict:dict(binary(), viva_tensor@tensor:tensor()),
list(param())
) -> {ok, {optimizer(), list(param())}} |
{error, viva_tensor@core@error:tensor_error()}.
do_apply(Opt, Params, Grad_dict, Acc) ->
case Params of
[] ->
{ok, {Opt, lists:reverse(Acc)}};
[P | Rest] ->
case gleam_stdlib:map_get(Grad_dict, erlang:element(2, P)) of
{error, _} ->
do_apply(Opt, Rest, Grad_dict, [P | Acc]);
{ok, G} ->
gleam@result:'try'(
update_param(Opt, P, G),
fun(_use0) ->
{Opt2, New_value} = _use0,
do_apply(
Opt2,
Rest,
Grad_dict,
[{param, erlang:element(2, P), New_value} | Acc]
)
end
)
end
end.
-file("src/viva_tensor/nn/optim.gleam", 288).
?DOC(false).
-spec apply_updates(
optimizer(),
list(param()),
gleam@dict:dict(binary(), viva_tensor@tensor:tensor())
) -> {ok, {optimizer(), list(param())}} |
{error, viva_tensor@core@error:tensor_error()}.
apply_updates(Opt, Params, Grad_dict) ->
do_apply(Opt, Params, Grad_dict, []).
-file("src/viva_tensor/nn/optim.gleam", 266).
?DOC(false).
-spec validate_shapes(
list(param()),
gleam@dict:dict(binary(), viva_tensor@tensor:tensor())
) -> {ok, nil} | {error, viva_tensor@core@error:tensor_error()}.
validate_shapes(Params, Grad_dict) ->
case Params of
[] ->
{ok, nil};
[P | Rest] ->
case gleam_stdlib:map_get(Grad_dict, erlang:element(2, P)) of
{error, _} ->
validate_shapes(Rest, Grad_dict);
{ok, G} ->
case viva_tensor@tensor:shape(erlang:element(3, P)) =:= viva_tensor@tensor:shape(
G
) of
false ->
{error,
{shape_mismatch,
viva_tensor@tensor:shape(
erlang:element(3, P)
),
viva_tensor@tensor:shape(G)}};
true ->
validate_shapes(Rest, Grad_dict)
end
end
end.
-file("src/viva_tensor/nn/optim.gleam", 248).
?DOC(false).
-spec validate_pairing(
list(param()),
gleam@dict:dict(binary(), viva_tensor@tensor:tensor())
) -> {ok, nil} | {error, viva_tensor@core@error:tensor_error()}.
validate_pairing(Params, Grad_dict) ->
Names = gleam@list:map(Params, fun(P) -> erlang:element(2, P) end),
Unknown = begin
_pipe = maps:keys(Grad_dict),
gleam@list:filter(
_pipe,
fun(Name) -> not gleam@list:contains(Names, Name) end
)
end,
case Unknown of
[Missing | _] ->
{error,
{dimension_error,
<<<<"optim.step: gradient for unknown parameter '"/utf8,
Missing/binary>>/binary,
"'"/utf8>>}};
[] ->
validate_shapes(Params, Grad_dict)
end.
-file("src/viva_tensor/nn/optim.gleam", 242).
?DOC(false).
-spec grads_to_dict(list(grad_pair())) -> gleam@dict:dict(binary(), viva_tensor@tensor:tensor()).
grads_to_dict(Grads) ->
gleam@list:fold(
Grads,
maps:new(),
fun(Acc, Gp) ->
gleam@dict:insert(Acc, erlang:element(2, Gp), erlang:element(3, Gp))
end
).
-file("src/viva_tensor/nn/optim.gleam", 214).
?DOC(false).
-spec step(optimizer(), list(param()), list(grad_pair())) -> {ok,
{optimizer(), list(param())}} |
{error, viva_tensor@core@error:tensor_error()}.
step(Opt, Params, Grads) ->
Grad_dict = grads_to_dict(Grads),
gleam@result:'try'(
validate_pairing(Params, Grad_dict),
fun(_) -> apply_updates(Opt, Params, Grad_dict) end
).
-file("src/viva_tensor/nn/optim.gleam", 234).
?DOC(false).
-spec zero_grad(list(grad_pair())) -> list(grad_pair()).
zero_grad(Grads) ->
gleam@list:map(
Grads,
fun(Gp) ->
{grad_pair,
erlang:element(2, Gp),
viva_tensor@tensor:zeros_like(erlang:element(3, Gp))}
end
).