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@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", 52).
?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", 57).
?DOC(false).
-spec sample_uniform(float(), float()) -> float().
sample_uniform(Low, High) ->
Low + (sample_unit() * (High - Low)).
-file("src/viva_tensor/nn/init.gleam", 72).
?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 => 73,
value => _assert_fail,
start => 3075,
'end' => 3112,
pattern_start => 3086,
pattern_end => 3091})
end,
V@1.
-file("src/viva_tensor/nn/init.gleam", 63).
?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 => 66,
value => _assert_fail,
start => 2803,
'end' => 2863,
pattern_start => 2814,
pattern_end => 2819})
end,
R@1 * math:cos((2.0 * 3.141592653589793) * U2).
-file("src/viva_tensor/nn/init.gleam", 78).
?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", 333).
?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", 329).
?DOC(false).
-spec range_int(integer(), integer()) -> list(integer()).
range_int(From, To) ->
range_loop(From, To, []).
-file("src/viva_tensor/nn/init.gleam", 84).
?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", 95).
?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", 120).
?DOC(false).
-spec zeros(list(integer())) -> viva_tensor@tensor:tensor().
zeros(Shape) ->
viva_tensor@tensor:zeros(Shape).
-file("src/viva_tensor/nn/init.gleam", 128).
?DOC(false).
-spec ones(list(integer())) -> viva_tensor@tensor:tensor().
ones(Shape) ->
viva_tensor@tensor:ones(Shape).
-file("src/viva_tensor/nn/init.gleam", 137).
?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", 146).
?DOC(false).
-spec identity(integer()) -> viva_tensor@tensor:tensor().
identity(N) ->
viva_tensor@tensor:eye(N).
-file("src/viva_tensor/nn/init.gleam", 160).
?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", 170).
?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", 186).
?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", 207).
?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 => 209,
value => _assert_fail,
start => 7514,
'end' => 7564,
pattern_start => 7525,
pattern_end => 7530})
end,
uniform([Fan_in, Fan_out], +0.0 - A@1, A@1).
-file("src/viva_tensor/nn/init.gleam", 219).
?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 => 221,
value => _assert_fail,
start => 7954,
'end' => 8006,
pattern_start => 7965,
pattern_end => 7972})
end,
normal([Fan_in, Fan_out], +0.0, Std@1).
-file("src/viva_tensor/nn/init.gleam", 233).
?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 => 234,
value => _assert_fail,
start => 8489,
'end' => 8554,
pattern_start => 8500,
pattern_end => 8505})
end,
Bound = Gain * S@1,
uniform([Fan_in, Fan_out], +0.0 - Bound, Bound).
-file("src/viva_tensor/nn/init.gleam", 247).
?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 => 248,
value => _assert_fail,
start => 9064,
'end' => 9129,
pattern_start => 9075,
pattern_end => 9080})
end,
Std = Gain * S@1,
normal([Fan_in, Fan_out], +0.0, Std).
-file("src/viva_tensor/nn/init.gleam", 273).
?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", 298).
?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 => 299,
value => _assert_fail,
start => 10897,
'end' => 10938,
pattern_start => 10908,
pattern_end => 10913})
end,
G@1.
-file("src/viva_tensor/nn/init.gleam", 305).
?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 => 307,
value => _assert_fail,
start => 11187,
'end' => 11237,
pattern_start => 11198,
pattern_end => 11203})
end,
G@1.
-file("src/viva_tensor/nn/init.gleam", 313).
?DOC(false).
-spec tanh_gain() -> float().
tanh_gain() ->
5.0 / 3.0.
-file("src/viva_tensor/nn/init.gleam", 318).
?DOC(false).
-spec linear_gain() -> float().
linear_gain() ->
1.0.
-file("src/viva_tensor/nn/init.gleam", 325).
?DOC(false).
-spec sigmoid_gain() -> float().
sigmoid_gain() ->
1.0.