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@distributed@trainer.erl
-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", 185).
?DOC(false).
-spec spawn_workers(integer(), fun((integer()) -> any())) -> list(worker()).
spawn_workers(Num, Worker_loop) ->
_pipe = gleam@list:range(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 = gleam@list:range(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(TEN), list(TEN), integer(), list(TEN)) -> list(TEN).
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(TEK), integer()) -> list(TEK).
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.