-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 => <>, 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 => <>, 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 => <>, 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 => <>, 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 => <>, 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 => <>, 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 => <>, 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 => <>, 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.