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