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