Packages

Tensor library for Gleam/BEAM with a pure Gleam API, zero-copy views, and optional native acceleration

Retired package: Release invalid

Current section

Files

Jump to
viva_tensor src viva_tensor@examples@training.erl
Raw

src/viva_tensor@examples@training.erl

-module(viva_tensor@examples@training).
-compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]).
-define(FILEPATH, "src/viva_tensor/examples/training.gleam").
-export([main/0]).
-export_type([training_state/0]).
-type training_state() :: {training_state,
viva_tensor@nn@autograd:tape(),
viva_tensor@nn@layers:linear(),
viva_tensor@nn@layers:linear()}.
-file("src/viva_tensor/examples/training.gleam", 59).
-spec train_step(
training_state(),
integer(),
viva_tensor@core@tensor:tensor(),
viva_tensor@core@tensor:tensor()
) -> training_state().
train_step(State, Epoch, X_data, Y_data) ->
{traced, X, Tape1} = viva_tensor@nn@autograd:new_variable(
erlang:element(2, State),
X_data
),
{traced, Target, Tape2} = viva_tensor@nn@autograd:new_variable(
Tape1,
Y_data
),
{L1_out@1, Tape3@1} = case viva_tensor@nn@layers:linear_forward(
Tape2,
erlang:element(3, State),
X
) of
{ok, {traced, L1_out, Tape3}} -> {L1_out, Tape3};
_assert_fail ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"train_step"/utf8>>,
line => 70,
value => _assert_fail,
start => 1878,
'end' => 1962,
pattern_start => 1889,
pattern_end => 1914})
end,
{traced, Hidden_act, Tape4} = viva_tensor@nn@layers:relu(Tape3@1, L1_out@1),
{Output@1, Tape5@1} = case viva_tensor@nn@layers:linear_forward(
Tape4,
erlang:element(4, State),
Hidden_act
) of
{ok, {traced, Output, Tape5}} -> {Output, Tape5};
_assert_fail@1 ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"train_step"/utf8>>,
line => 73,
value => _assert_fail@1,
start => 2022,
'end' => 2115,
pattern_start => 2033,
pattern_end => 2058})
end,
{Loss_var@1, Tape6@1} = case viva_tensor@nn@layers:mse_loss(
Tape5@1,
Output@1,
Target
) of
{ok, {traced, Loss_var, Tape6}} -> {Loss_var, Tape6};
_assert_fail@2 ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"train_step"/utf8>>,
line => 77,
value => _assert_fail@2,
start => 2139,
'end' => 2214,
pattern_start => 2150,
pattern_end => 2177})
end,
Grads@1 = case viva_tensor@nn@autograd:backward(Tape6@1, Loss_var@1) of
{ok, Grads} -> Grads;
_assert_fail@3 ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"train_step"/utf8>>,
line => 80,
value => _assert_fail@3,
start => 2237,
'end' => 2294,
pattern_start => 2248,
pattern_end => 2257})
end,
case (Epoch rem 100) =:= 0 of
true ->
Loss_val = begin
_pipe = viva_tensor@core@tensor:to_list(
erlang:element(3, Loss_var@1)
),
_pipe@1 = gleam@list:first(_pipe),
gleam@result:unwrap(_pipe@1, +0.0)
end,
gleam_stdlib:println(
<<<<<<"Epoch "/utf8, (erlang:integer_to_binary(Epoch))/binary>>/binary,
" | Loss: "/utf8>>/binary,
(gleam_stdlib:float_to_string(Loss_val))/binary>>
);
false ->
nil
end,
Gw1@1 = case gleam_stdlib:map_get(
Grads@1,
erlang:element(2, erlang:element(2, erlang:element(3, State)))
) of
{ok, Gw1} -> Gw1;
_assert_fail@4 ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"train_step"/utf8>>,
line => 98,
value => _assert_fail@4,
start => 2664,
'end' => 2719,
pattern_start => 2675,
pattern_end => 2682})
end,
Gb1@1 = case gleam_stdlib:map_get(
Grads@1,
erlang:element(2, erlang:element(3, erlang:element(3, State)))
) of
{ok, Gb1} -> Gb1;
_assert_fail@5 ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"train_step"/utf8>>,
line => 99,
value => _assert_fail@5,
start => 2722,
'end' => 2777,
pattern_start => 2733,
pattern_end => 2740})
end,
New_w1_data@1 = case viva_tensor@core@ops:sub(
erlang:element(3, erlang:element(2, erlang:element(3, State))),
viva_tensor@core@ops:scale(Gw1@1, 0.01)
) of
{ok, New_w1_data} -> New_w1_data;
_assert_fail@6 ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"train_step"/utf8>>,
line => 100,
value => _assert_fail@6,
start => 2780,
'end' => 2872,
pattern_start => 2791,
pattern_end => 2806})
end,
New_b1_data@1 = case viva_tensor@core@ops:sub(
erlang:element(3, erlang:element(3, erlang:element(3, State))),
viva_tensor@core@ops:scale(Gb1@1, 0.01)
) of
{ok, New_b1_data} -> New_b1_data;
_assert_fail@7 ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"train_step"/utf8>>,
line => 102,
value => _assert_fail@7,
start => 2875,
'end' => 2967,
pattern_start => 2886,
pattern_end => 2901})
end,
Gw2@1 = case gleam_stdlib:map_get(
Grads@1,
erlang:element(2, erlang:element(2, erlang:element(4, State)))
) of
{ok, Gw2} -> Gw2;
_assert_fail@8 ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"train_step"/utf8>>,
line => 105,
value => _assert_fail@8,
start => 2971,
'end' => 3026,
pattern_start => 2982,
pattern_end => 2989})
end,
Gb2@1 = case gleam_stdlib:map_get(
Grads@1,
erlang:element(2, erlang:element(3, erlang:element(4, State)))
) of
{ok, Gb2} -> Gb2;
_assert_fail@9 ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"train_step"/utf8>>,
line => 106,
value => _assert_fail@9,
start => 3029,
'end' => 3084,
pattern_start => 3040,
pattern_end => 3047})
end,
New_w2_data@1 = case viva_tensor@core@ops:sub(
erlang:element(3, erlang:element(2, erlang:element(4, State))),
viva_tensor@core@ops:scale(Gw2@1, 0.01)
) of
{ok, New_w2_data} -> New_w2_data;
_assert_fail@10 ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"train_step"/utf8>>,
line => 107,
value => _assert_fail@10,
start => 3087,
'end' => 3179,
pattern_start => 3098,
pattern_end => 3113})
end,
New_b2_data@1 = case viva_tensor@core@ops:sub(
erlang:element(3, erlang:element(3, erlang:element(4, State))),
viva_tensor@core@ops:scale(Gb2@1, 0.01)
) of
{ok, New_b2_data} -> New_b2_data;
_assert_fail@11 ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"train_step"/utf8>>,
line => 109,
value => _assert_fail@11,
start => 3182,
'end' => 3274,
pattern_start => 3193,
pattern_end => 3208})
end,
Next_tape = viva_tensor@nn@autograd:new_tape(),
{traced, Nw1, Nt1} = viva_tensor@nn@autograd:new_variable(
Next_tape,
New_w1_data@1
),
{traced, Nb1, Nt2} = viva_tensor@nn@autograd:new_variable(
Nt1,
New_b1_data@1
),
{traced, Nw2, Nt3} = viva_tensor@nn@autograd:new_variable(
Nt2,
New_w2_data@1
),
{traced, Nb2, Nt4} = viva_tensor@nn@autograd:new_variable(
Nt3,
New_b2_data@1
),
{training_state, Nt4, {linear, Nw1, Nb1}, {linear, Nw2, Nb2}}.
-file("src/viva_tensor/examples/training.gleam", 28).
-spec main() -> nil.
main() ->
gleam_stdlib:println(<<"🚀 Starting Mycelial Training Demo..."/utf8>>),
Tape = viva_tensor@nn@autograd:new_tape(),
X_data = viva_tensor@core@tensor:from_list([1.0, 2.0, 3.0, 4.0, 5.0]),
X_data@2 = case viva_tensor@core@shape:reshape(X_data, [5, 1]) of
{ok, X_data@1} -> X_data@1;
_assert_fail ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"main"/utf8>>,
line => 36,
value => _assert_fail,
start => 775,
'end' => 828,
pattern_start => 786,
pattern_end => 796})
end,
Y_data = viva_tensor@core@tensor:from_list([2.1, 3.9, 6.2, 8.1, 10.3]),
Y_data@2 = case viva_tensor@core@shape:reshape(Y_data, [5, 1]) of
{ok, Y_data@1} -> Y_data@1;
_assert_fail@1 ->
erlang:error(#{gleam_error => let_assert,
message => <<"Pattern match failed, no pattern matched the value."/utf8>>,
file => <<?FILEPATH/utf8>>,
module => <<"viva_tensor/examples/training"/utf8>>,
function => <<"main"/utf8>>,
line => 39,
value => _assert_fail@1,
start => 892,
'end' => 945,
pattern_start => 903,
pattern_end => 913})
end,
{traced, _, Tape1} = viva_tensor@nn@autograd:new_variable(Tape, X_data@2),
{traced, _, Tape2} = viva_tensor@nn@autograd:new_variable(Tape1, Y_data@2),
{traced, Layer1, Tape3} = viva_tensor@nn@layers:linear(Tape2, 1, 4),
{traced, Layer2, Tape4} = viva_tensor@nn@layers:linear(Tape3, 4, 1),
State = {training_state, Tape4, Layer1, Layer2},
_ = gleam@list:fold(
gleam@list:range(0, 500 - 1),
State,
fun(Acc_state, Epoch) ->
train_step(Acc_state, Epoch, X_data@2, Y_data@2)
end
),
gleam_stdlib:println(<<"✅ Training finished!"/utf8>>).