Current section

Files

Jump to
viva_tensor src viva_tensor@nn@init.erl
Raw

src/viva_tensor@nn@init.erl

-module(viva_tensor@nn@init).
-compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]).
-define(FILEPATH, "src/viva_tensor/nn/init.gleam").
-export([zeros/1, ones/1, constant/2, identity/1, uniform/3, normal/3, truncated_normal/5, xavier_uniform/2, xavier_normal/2, kaiming_uniform/3, kaiming_normal/3, orthogonal/3, relu_gain/0, leaky_relu_gain/1, tanh_gain/0, linear_gain/0, sigmoid_gain/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).
-file("src/viva_tensor/nn/init.gleam", 51).
?DOC(false).
-spec sample_unit() -> float().
sample_unit() ->
case 2147483648.0 of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> erlang:float(gleam@int:random(2147483648)) / Gleam@denominator
end.
-file("src/viva_tensor/nn/init.gleam", 56).
?DOC(false).
-spec sample_uniform(float(), float()) -> float().
sample_uniform(Low, High) ->
Low + (sample_unit() * (High - Low)).
-file("src/viva_tensor/nn/init.gleam", 71).
?DOC(false).
-spec log_unsafe(float()) -> float().
log_unsafe(X) ->
V@1 = case gleam@float:logarithm(X) of
{ok, V} -> V;
_assert_fail ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/nn/init"/utf8>>,
function => <<"log_unsafe"/utf8>>,
line => 72,
value => _assert_fail,
start => 3015,
'end' => 3052,
pattern_start => 3026,
pattern_end => 3031})
end,
V@1.
-file("src/viva_tensor/nn/init.gleam", 62).
?DOC(false).
-spec sample_standard_normal() -> float().
sample_standard_normal() ->
U1 = gleam@float:max(sample_unit(), 1.0e-12),
U2 = sample_unit(),
R@1 = case gleam@float:square_root(-2.0 * log_unsafe(U1)) of
{ok, R} -> R;
_assert_fail ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/nn/init"/utf8>>,
function => <<"sample_standard_normal"/utf8>>,
line => 65,
value => _assert_fail,
start => 2752,
'end' => 2812,
pattern_start => 2763,
pattern_end => 2768})
end,
R@1 * gleam_community@maths:cos((2.0 * gleam_community@maths:pi()) * U2).
-file("src/viva_tensor/nn/init.gleam", 77).
?DOC(false).
-spec size_of(list(integer())) -> integer().
size_of(Shape) ->
gleam@list:fold(Shape, 1, fun(Acc, D) -> Acc * D end).
-file("src/viva_tensor/nn/init.gleam", 332).
?DOC(false).
-spec range_loop(integer(), integer(), list(integer())) -> list(integer()).
range_loop(From, To, Acc) ->
case From > To of
true ->
lists:reverse(Acc);
false ->
range_loop(From + 1, To, [From | Acc])
end.
-file("src/viva_tensor/nn/init.gleam", 328).
?DOC(false).
-spec range_int(integer(), integer()) -> list(integer()).
range_int(From, To) ->
range_loop(From, To, []).
-file("src/viva_tensor/nn/init.gleam", 83).
?DOC(false).
-spec build(list(integer()), fun(() -> float())) -> viva_tensor@tensor:tensor().
build(Shape, Gen) ->
Size = size_of(Shape),
Data = begin
_pipe = range_int(1, Size),
gleam@list:map(_pipe, fun(_) -> Gen() end)
end,
{tensor, Data, Shape}.
-file("src/viva_tensor/nn/init.gleam", 94).
?DOC(false).
-spec sample_truncated(float(), float(), float(), float(), integer()) -> float().
sample_truncated(Mean, Std, A, B, Iters_left) ->
case Iters_left =< 0 of
true ->
sample_uniform(A, B);
false ->
X = Mean + (Std * sample_standard_normal()),
case (X >= A) andalso (X =< B) of
true ->
X;
false ->
sample_truncated(Mean, Std, A, B, Iters_left - 1)
end
end.
-file("src/viva_tensor/nn/init.gleam", 119).
?DOC(false).
-spec zeros(list(integer())) -> viva_tensor@tensor:tensor().
zeros(Shape) ->
viva_tensor@tensor:zeros(Shape).
-file("src/viva_tensor/nn/init.gleam", 127).
?DOC(false).
-spec ones(list(integer())) -> viva_tensor@tensor:tensor().
ones(Shape) ->
viva_tensor@tensor:ones(Shape).
-file("src/viva_tensor/nn/init.gleam", 136).
?DOC(false).
-spec constant(list(integer()), float()) -> viva_tensor@tensor:tensor().
constant(Shape, Value) ->
viva_tensor@tensor:fill(Shape, Value).
-file("src/viva_tensor/nn/init.gleam", 145).
?DOC(false).
-spec identity(integer()) -> viva_tensor@tensor:tensor().
identity(N) ->
viva_tensor@tensor:eye(N).
-file("src/viva_tensor/nn/init.gleam", 159).
?DOC(false).
-spec uniform(list(integer()), float(), float()) -> viva_tensor@tensor:tensor().
uniform(Shape, Low, High) ->
build(Shape, fun() -> sample_uniform(Low, High) end).
-file("src/viva_tensor/nn/init.gleam", 169).
?DOC(false).
-spec normal(list(integer()), float(), float()) -> viva_tensor@tensor:tensor().
normal(Shape, Mean, Std) ->
build(Shape, fun() -> Mean + (Std * sample_standard_normal()) end).
-file("src/viva_tensor/nn/init.gleam", 185).
?DOC(false).
-spec truncated_normal(list(integer()), float(), float(), float(), float()) -> viva_tensor@tensor:tensor().
truncated_normal(Shape, Mean, Std, A, B) ->
build(Shape, fun() -> sample_truncated(Mean, Std, A, B, 100) end).
-file("src/viva_tensor/nn/init.gleam", 206).
?DOC(false).
-spec xavier_uniform(integer(), integer()) -> viva_tensor@tensor:tensor().
xavier_uniform(Fan_in, Fan_out) ->
Denom = erlang:float(Fan_in + Fan_out),
A@1 = case gleam@float:square_root(case Denom of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> 6.0 / Gleam@denominator
end) of
{ok, A} -> A;
_assert_fail ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/nn/init"/utf8>>,
function => <<"xavier_uniform"/utf8>>,
line => 208,
value => _assert_fail,
start => 7454,
'end' => 7504,
pattern_start => 7465,
pattern_end => 7470})
end,
uniform([Fan_in, Fan_out], +0.0 - A@1, A@1).
-file("src/viva_tensor/nn/init.gleam", 218).
?DOC(false).
-spec xavier_normal(integer(), integer()) -> viva_tensor@tensor:tensor().
xavier_normal(Fan_in, Fan_out) ->
Denom = erlang:float(Fan_in + Fan_out),
Std@1 = case gleam@float:square_root(case Denom of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> 2.0 / Gleam@denominator
end) of
{ok, Std} -> Std;
_assert_fail ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/nn/init"/utf8>>,
function => <<"xavier_normal"/utf8>>,
line => 220,
value => _assert_fail,
start => 7894,
'end' => 7946,
pattern_start => 7905,
pattern_end => 7912})
end,
normal([Fan_in, Fan_out], +0.0, Std@1).
-file("src/viva_tensor/nn/init.gleam", 232).
?DOC(false).
-spec kaiming_uniform(integer(), integer(), float()) -> viva_tensor@tensor:tensor().
kaiming_uniform(Fan_in, Fan_out, Gain) ->
S@1 = case gleam@float:square_root(case erlang:float(Fan_in) of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> 3.0 / Gleam@denominator
end) of
{ok, S} -> S;
_assert_fail ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/nn/init"/utf8>>,
function => <<"kaiming_uniform"/utf8>>,
line => 233,
value => _assert_fail,
start => 8429,
'end' => 8494,
pattern_start => 8440,
pattern_end => 8445})
end,
Bound = Gain * S@1,
uniform([Fan_in, Fan_out], +0.0 - Bound, Bound).
-file("src/viva_tensor/nn/init.gleam", 246).
?DOC(false).
-spec kaiming_normal(integer(), integer(), float()) -> viva_tensor@tensor:tensor().
kaiming_normal(Fan_in, Fan_out, Gain) ->
S@1 = case gleam@float:square_root(case erlang:float(Fan_in) of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> 1.0 / Gleam@denominator
end) of
{ok, S} -> S;
_assert_fail ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/nn/init"/utf8>>,
function => <<"kaiming_normal"/utf8>>,
line => 247,
value => _assert_fail,
start => 9004,
'end' => 9069,
pattern_start => 9015,
pattern_end => 9020})
end,
Std = Gain * S@1,
normal([Fan_in, Fan_out], +0.0, Std).
-file("src/viva_tensor/nn/init.gleam", 272).
?DOC(false).
-spec orthogonal(integer(), integer(), float()) -> {ok,
viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
orthogonal(Rows, Cols, Gain) ->
{M, N, Transpose_result} = case Rows < Cols of
true ->
{Cols, Rows, true};
false ->
{Rows, Cols, false}
end,
G = normal([M, N], +0.0, 1.0),
gleam@result:'try'(
viva_tensor@core@linalg:qr(G),
fun(_use0) ->
{Q, _} = _use0,
gleam@result:'try'(case Transpose_result of
true ->
viva_tensor@tensor:transpose(Q);
false ->
{ok, Q}
end, fun(Oriented) ->
Scaled = begin
_pipe = viva_tensor@tensor:to_list(Oriented),
gleam@list:map(_pipe, fun(X) -> X * Gain end)
end,
{ok, {tensor, Scaled, [Rows, Cols]}}
end)
end
).
-file("src/viva_tensor/nn/init.gleam", 297).
?DOC(false).
-spec relu_gain() -> float().
relu_gain() ->
G@1 = case gleam@float:square_root(2.0) of
{ok, G} -> G;
_assert_fail ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/nn/init"/utf8>>,
function => <<"relu_gain"/utf8>>,
line => 298,
value => _assert_fail,
start => 10837,
'end' => 10878,
pattern_start => 10848,
pattern_end => 10853})
end,
G@1.
-file("src/viva_tensor/nn/init.gleam", 304).
?DOC(false).
-spec leaky_relu_gain(float()) -> float().
leaky_relu_gain(Negative_slope) ->
Denom = 1.0 + (Negative_slope * Negative_slope),
G@1 = case gleam@float:square_root(case Denom of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> 2.0 / Gleam@denominator
end) of
{ok, G} -> G;
_assert_fail ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/nn/init"/utf8>>,
function => <<"leaky_relu_gain"/utf8>>,
line => 306,
value => _assert_fail,
start => 11127,
'end' => 11177,
pattern_start => 11138,
pattern_end => 11143})
end,
G@1.
-file("src/viva_tensor/nn/init.gleam", 312).
?DOC(false).
-spec tanh_gain() -> float().
tanh_gain() ->
5.0 / 3.0.
-file("src/viva_tensor/nn/init.gleam", 317).
?DOC(false).
-spec linear_gain() -> float().
linear_gain() ->
1.0.
-file("src/viva_tensor/nn/init.gleam", 324).
?DOC(false).
-spec sigmoid_gain() -> float().
sigmoid_gain() ->
1.0.