-module(viva_tensor@distributed@trainer). -compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]). -define(FILEPATH, "src/viva_tensor/distributed/trainer.gleam"). -export([distribute_grads/2, synchronous_train_step/4, spawn_workers/2, send_batch_to_worker/4, receive_grads_from_worker/2, all_reduce_grads/3, train_synchronous/6]). -export_type([grad_aggregation/0, train_config/0, train_result/0, worker/0, worker_message/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). -type grad_aggregation() :: average_grads | sum_grads. -type train_config() :: {train_config, integer(), integer(), grad_aggregation()}. -type train_result() :: {train_result, list(viva_tensor@nn@optim:param()), viva_tensor@nn@optim:optimizer(), float(), integer()}. -type worker() :: {worker, integer(), gleam@erlang@process:pid_(), gleam@erlang@process:subject(worker_message()), gleam@erlang@process:subject({integer(), integer(), list(viva_tensor@nn@optim:grad_pair())})}. -type worker_message() :: {run_batch, integer(), viva_tensor@data@dataloader:batch(), list(viva_tensor@nn@optim:param())} | stop. -file("src/viva_tensor/distributed/trainer.gleam", 441). ?DOC(false). -spec add_grad_lists( list(viva_tensor@nn@optim:grad_pair()), list(viva_tensor@nn@optim:grad_pair()) ) -> {ok, list(viva_tensor@nn@optim:grad_pair())} | {error, viva_tensor@core@error:tensor_error()}. add_grad_lists(A, B) -> B_dict = gleam@list:fold( B, maps:new(), fun(Acc, Gp) -> gleam@dict:insert(Acc, erlang:element(2, Gp), erlang:element(3, Gp)) end ), gleam@list:try_map( A, fun(Gp@1) -> case gleam_stdlib:map_get(B_dict, erlang:element(2, Gp@1)) of {error, _} -> {error, {dimension_error, <<<<"distributed.add_grad_lists: missing gradient for '"/utf8, (erlang:element(2, Gp@1))/binary>>/binary, "'"/utf8>>}}; {ok, Other} -> case viva_tensor@tensor:shape(erlang:element(3, Gp@1)) =:= viva_tensor@tensor:shape( Other ) of false -> {error, {shape_mismatch, viva_tensor@tensor:shape( erlang:element(3, Gp@1) ), viva_tensor@tensor:shape(Other)}}; true -> gleam@result:'try'( viva_tensor@tensor:add( erlang:element(3, Gp@1), Other ), fun(Summed) -> {ok, {grad_pair, erlang:element(2, Gp@1), Summed}} end ) end end end ). -file("src/viva_tensor/distributed/trainer.gleam", 472). ?DOC(false). -spec sum_grad_lists( list(viva_tensor@nn@optim:grad_pair()), list(list(viva_tensor@nn@optim:grad_pair())) ) -> {ok, list(viva_tensor@nn@optim:grad_pair())} | {error, viva_tensor@core@error:tensor_error()}. sum_grad_lists(First, Rest) -> gleam@list:try_fold( Rest, First, fun(Acc, Other) -> add_grad_lists(Acc, Other) end ). -file("src/viva_tensor/distributed/trainer.gleam", 479). ?DOC(false). -spec validate_same_shape_lists( list(viva_tensor@nn@optim:grad_pair()), list(list(viva_tensor@nn@optim:grad_pair())) ) -> {ok, nil} | {error, viva_tensor@core@error:tensor_error()}. validate_same_shape_lists(First, Rest) -> First_names = gleam@list:map(First, fun(Gp) -> erlang:element(2, Gp) end), First_shapes = gleam@list:fold( First, maps:new(), fun(Acc, Gp@1) -> gleam@dict:insert( Acc, erlang:element(2, Gp@1), viva_tensor@tensor:shape(erlang:element(3, Gp@1)) ) end ), _pipe = gleam@list:try_fold( Rest, nil, fun(_, Other) -> case erlang:length(Other) =:= erlang:length(First) of false -> {error, {dimension_error, <<"distribute_grads: worker grad lists have different lengths"/utf8>>}}; true -> gleam@list:try_fold( Other, nil, fun(_, Gp@2) -> case gleam@list:contains( First_names, erlang:element(2, Gp@2) ) of false -> {error, {dimension_error, <<<<"distribute_grads: parameter name '"/utf8, (erlang:element(2, Gp@2))/binary>>/binary, "' missing from first worker's grads"/utf8>>}}; true -> case gleam_stdlib:map_get( First_shapes, erlang:element(2, Gp@2) ) of {error, _} -> {ok, nil}; {ok, Expected} -> case viva_tensor@tensor:shape( erlang:element(3, Gp@2) ) =:= Expected of false -> {error, {shape_mismatch, Expected, viva_tensor@tensor:shape( erlang:element( 3, Gp@2 ) )}}; true -> {ok, nil} end end end end ) end end ), gleam@result:map(_pipe, fun(_) -> nil end). -file("src/viva_tensor/distributed/trainer.gleam", 128). ?DOC(false). -spec distribute_grads( list(list(viva_tensor@nn@optim:grad_pair())), grad_aggregation() ) -> {ok, list(viva_tensor@nn@optim:grad_pair())} | {error, viva_tensor@core@error:tensor_error()}. distribute_grads(Per_worker_grads, Aggregation) -> case Per_worker_grads of [] -> {error, {dimension_error, <<"distribute_grads: need at least one worker's grads"/utf8>>}}; [First | Rest] -> gleam@result:'try'( validate_same_shape_lists(First, Rest), fun(_) -> Num_workers = erlang:length(Per_worker_grads), gleam@result:'try'( sum_grad_lists(First, Rest), fun(Summed) -> case Aggregation of sum_grads -> {ok, Summed}; average_grads -> Denom = erlang:float(Num_workers), {ok, gleam@list:map( Summed, fun(Gp) -> {grad_pair, erlang:element(2, Gp), viva_tensor@tensor:scale( erlang:element(3, Gp), case Denom of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator -> 1.0 / Gleam@denominator end )} end )} end end ) end ) end. -file("src/viva_tensor/distributed/trainer.gleam", 162). ?DOC(false). -spec synchronous_train_step( viva_tensor@nn@optim:optimizer(), list(viva_tensor@nn@optim:param()), list(list(viva_tensor@nn@optim:grad_pair())), grad_aggregation() ) -> {ok, {viva_tensor@nn@optim:optimizer(), list(viva_tensor@nn@optim:param())}} | {error, viva_tensor@core@error:tensor_error()}. synchronous_train_step(Opt, Params, Per_worker_grads, Aggregation) -> gleam@result:'try'( distribute_grads(Per_worker_grads, Aggregation), fun(Aggregated) -> viva_tensor@nn@optim:step(Opt, Params, Aggregated) end ). -file("src/viva_tensor/distributed/trainer.gleam", 561). ?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/distributed/trainer.gleam", 557). ?DOC(false). -spec range_int(integer(), integer()) -> list(integer()). range_int(From, To) -> range_loop(From, To, []). -file("src/viva_tensor/distributed/trainer.gleam", 185). ?DOC(false). -spec spawn_workers(integer(), fun((integer()) -> any())) -> list(worker()). spawn_workers(Num, Worker_loop) -> _pipe = range_int(0, Num - 1), gleam@list:map( _pipe, fun(Id) -> Inbox = gleam@erlang@process:new_subject(), Outbox = gleam@erlang@process:new_subject(), Pid = proc_lib:spawn_link( fun() -> _ = Worker_loop(Id), nil end ), {worker, Id, Pid, Inbox, Outbox} end ). -file("src/viva_tensor/distributed/trainer.gleam", 206). ?DOC(false). -spec send_batch_to_worker( worker(), integer(), viva_tensor@data@dataloader:batch(), list(viva_tensor@nn@optim:param()) ) -> nil. send_batch_to_worker(Worker, Batch_id, Batch, Params) -> gleam@erlang@process:send( erlang:element(4, Worker), {run_batch, Batch_id, Batch, Params} ). -file("src/viva_tensor/distributed/trainer.gleam", 221). ?DOC(false). -spec receive_grads_from_worker(worker(), integer()) -> {ok, list(viva_tensor@nn@optim:grad_pair())} | {error, viva_tensor@core@error:tensor_error()}. receive_grads_from_worker(Worker, Timeout_ms) -> case gleam@erlang@process:'receive'(erlang:element(5, Worker), Timeout_ms) of {ok, {_, _, Grads}} -> {ok, Grads}; {error, _} -> {error, {dimension_error, <<<<"receive_grads_from_worker: timeout after "/utf8, (erlang:integer_to_binary(Timeout_ms))/binary>>/binary, "ms"/utf8>>}} end. -file("src/viva_tensor/distributed/trainer.gleam", 246). ?DOC(false). -spec all_reduce_grads( list(worker()), list(list(viva_tensor@nn@optim:grad_pair())), grad_aggregation() ) -> {ok, list(viva_tensor@nn@optim:grad_pair())} | {error, viva_tensor@core@error:tensor_error()}. all_reduce_grads(Workers, Local_grads, Aggregation) -> case erlang:length(Workers) =:= erlang:length(Local_grads) of false -> {error, {dimension_error, <<"all_reduce_grads: number of workers must equal number of local grad lists"/utf8>>}}; true -> distribute_grads(Local_grads, Aggregation) end. -file("src/viva_tensor/distributed/trainer.gleam", 421). ?DOC(false). -spec accumulate_grads_for_worker( list(viva_tensor@data@dataloader:batch()), list(viva_tensor@nn@optim:param()), fun((viva_tensor@data@dataloader:batch(), list(viva_tensor@nn@optim:param())) -> {ok, list(viva_tensor@nn@optim:grad_pair())} | {error, viva_tensor@core@error:tensor_error()}) ) -> {ok, list(viva_tensor@nn@optim:grad_pair())} | {error, viva_tensor@core@error:tensor_error()}. accumulate_grads_for_worker(Batches, Params, Compute_grads) -> case Batches of [] -> {ok, gleam@list:map( Params, fun(P) -> {grad_pair, erlang:element(2, P), viva_tensor@tensor:zeros_like(erlang:element(3, P))} end )}; [First | Rest] -> gleam@result:'try'( Compute_grads(First, Params), fun(First_grads) -> gleam@list:try_fold( Rest, First_grads, fun(Acc, B) -> gleam@result:'try'( Compute_grads(B, Params), fun(G) -> add_grad_lists(Acc, G) end ) end ) end ) end. -file("src/viva_tensor/distributed/trainer.gleam", 402). ?DOC(false). -spec do_collect_workers( list(list(viva_tensor@data@dataloader:batch())), list(viva_tensor@nn@optim:param()), fun((viva_tensor@data@dataloader:batch(), list(viva_tensor@nn@optim:param())) -> {ok, list(viva_tensor@nn@optim:grad_pair())} | {error, viva_tensor@core@error:tensor_error()}), list(list(viva_tensor@nn@optim:grad_pair())) ) -> {ok, list(list(viva_tensor@nn@optim:grad_pair()))} | {error, viva_tensor@core@error:tensor_error()}. do_collect_workers(Per_worker_batches, Params, Compute_grads, Acc) -> case Per_worker_batches of [] -> {ok, lists:reverse(Acc)}; [Worker_batches | Rest] -> gleam@result:'try'( accumulate_grads_for_worker( Worker_batches, Params, Compute_grads ), fun(Grads) -> do_collect_workers( Rest, Params, Compute_grads, [Grads | Acc] ) end ) end. -file("src/viva_tensor/distributed/trainer.gleam", 523). ?DOC(false). -spec assign_batches(list(viva_tensor@data@dataloader:batch()), integer()) -> list(list(viva_tensor@data@dataloader:batch())). assign_batches(Batches, Num_workers) -> Indexed = gleam@list:index_map(Batches, fun(B, I) -> {I, B} end), _pipe = range_int(0, Num_workers - 1), gleam@list:map(_pipe, fun(Worker_id) -> _pipe@1 = Indexed, _pipe@2 = gleam@list:filter( _pipe@1, fun(Pair) -> (case Num_workers of 0 -> 0; Gleam@denominator -> erlang:element(1, Pair) rem Gleam@denominator end) =:= Worker_id end ), gleam@list:map( _pipe@2, fun(Pair@1) -> erlang:element(2, Pair@1) end ) end). -file("src/viva_tensor/distributed/trainer.gleam", 390). ?DOC(false). -spec collect_per_worker_grads( list(viva_tensor@data@dataloader:batch()), integer(), list(viva_tensor@nn@optim:param()), fun((viva_tensor@data@dataloader:batch(), list(viva_tensor@nn@optim:param())) -> {ok, list(viva_tensor@nn@optim:grad_pair())} | {error, viva_tensor@core@error:tensor_error()}) ) -> {ok, list(list(viva_tensor@nn@optim:grad_pair()))} | {error, viva_tensor@core@error:tensor_error()}. collect_per_worker_grads(Batches, Num_workers, Params, Compute_grads) -> Assigned = assign_batches(Batches, Num_workers), do_collect_workers(Assigned, Params, Compute_grads, []). -file("src/viva_tensor/distributed/trainer.gleam", 541). ?DOC(false). -spec do_take_cycle(list(WOJ), list(WOJ), integer(), list(WOJ)) -> list(WOJ). do_take_cycle(Remaining, Full, N, Acc) -> case N =< 0 of true -> lists:reverse(Acc); false -> case Remaining of [] -> do_take_cycle(Full, Full, N, Acc); [Head | Rest] -> do_take_cycle(Rest, Full, N - 1, [Head | Acc]) end end. -file("src/viva_tensor/distributed/trainer.gleam", 534). ?DOC(false). -spec take_cycle(list(WOG), integer()) -> list(WOG). take_cycle(Xs, N) -> case Xs of [] -> []; _ -> do_take_cycle(Xs, Xs, N, []) end. -file("src/viva_tensor/distributed/trainer.gleam", 340). ?DOC(false). -spec run_steps( train_config(), list(viva_tensor@nn@optim:param()), viva_tensor@nn@optim:optimizer(), list(viva_tensor@data@dataloader:batch()), fun((viva_tensor@data@dataloader:batch(), list(viva_tensor@nn@optim:param())) -> {ok, list(viva_tensor@nn@optim:grad_pair())} | {error, viva_tensor@core@error:tensor_error()}), integer(), integer() ) -> {ok, train_result()} | {error, viva_tensor@core@error:tensor_error()}. run_steps(Config, Params, Opt, Batches, Compute_grads, Remaining, Done) -> case Remaining =< 0 of true -> {ok, {train_result, Params, Opt, +0.0, Done}}; false -> Step_batches = take_cycle(Batches, erlang:element(3, Config)), gleam@result:'try'( collect_per_worker_grads( Step_batches, erlang:element(2, Config), Params, Compute_grads ), fun(Per_worker) -> gleam@result:'try'( synchronous_train_step( Opt, Params, Per_worker, erlang:element(4, Config) ), fun(_use0) -> {Opt2, Params2} = _use0, run_steps( Config, Params2, Opt2, Batches, Compute_grads, Remaining - 1, Done + 1 ) end ) end ) end. -file("src/viva_tensor/distributed/trainer.gleam", 289). ?DOC(false). -spec train_synchronous( train_config(), list(viva_tensor@nn@optim:param()), viva_tensor@nn@optim:optimizer(), viva_tensor@data@dataloader:data_loader(), fun((viva_tensor@data@dataloader:batch(), list(viva_tensor@nn@optim:param())) -> {ok, list(viva_tensor@nn@optim:grad_pair())} | {error, viva_tensor@core@error:tensor_error()}), integer() ) -> {ok, train_result()} | {error, viva_tensor@core@error:tensor_error()}. train_synchronous( Config, Initial_params, Initial_optimizer, Data_loader, Compute_grads, Num_steps ) -> case erlang:element(2, Config) =< 0 of true -> {error, {dimension_error, <<"train_synchronous: num_workers must be > 0"/utf8>>}}; false -> case erlang:element(3, Config) =< 0 of true -> {error, {dimension_error, <<"train_synchronous: batches_per_step must be > 0"/utf8>>}}; false -> case Num_steps < 0 of true -> {error, {dimension_error, <<"train_synchronous: num_steps must be >= 0"/utf8>>}}; false -> gleam@result:'try'( viva_tensor@data@dataloader:data_loader_batches( Data_loader ), fun(Batches) -> case Batches of [] -> {ok, {train_result, Initial_params, Initial_optimizer, +0.0, 0}}; _ -> run_steps( Config, Initial_params, Initial_optimizer, Batches, Compute_grads, Num_steps, 0 ) end end ) end end end.