Current section
Files
Jump to
Current section
Files
src/mlx_dist_coordinator.erl
-module(mlx_dist_coordinator).
-behaviour(gen_server).
%% API
-export([start_link/0, start_link/1,
add_worker/1, remove_worker/1, list_workers/0,
distribute_data/2, gather_gradients/0,
broadcast_parameters/1, get_parameters/0,
start_training/3, stop_training/0,
get_status/0]).
%% gen_server callbacks
-export([init/1, handle_call/3, handle_cast/2, handle_info/2,
terminate/2, code_change/3]).
-record(state, {
workers = [], % List of worker nodes
parameters = undefined, % Current model parameters
gradients = [], % Accumulated gradients from workers
training_config = undefined,
status = idle,
start_time = undefined,
iterations = 0
}).
%%====================================================================
%% API
%%====================================================================
start_link() ->
start_link([]).
start_link(Options) ->
gen_server:start_link({local, ?MODULE}, ?MODULE, Options, []).
%% Add a worker node
add_worker(Node) ->
gen_server:call(?MODULE, {add_worker, Node}).
%% Remove a worker node
remove_worker(Node) ->
gen_server:call(?MODULE, {remove_worker, Node}).
%% List all worker nodes
list_workers() ->
gen_server:call(?MODULE, list_workers).
%% Distribute data to workers
distribute_data(Data, Strategy) ->
gen_server:call(?MODULE, {distribute_data, Data, Strategy}, 60000).
%% Gather gradients from all workers
gather_gradients() ->
gen_server:call(?MODULE, gather_gradients, 60000).
%% Broadcast parameters to all workers
broadcast_parameters(Parameters) ->
gen_server:call(?MODULE, {broadcast_parameters, Parameters}).
%% Get current parameters
get_parameters() ->
gen_server:call(?MODULE, get_parameters).
%% Start distributed training
start_training(ModelConfig, DataConfig, TrainingConfig) ->
gen_server:call(?MODULE, {start_training, ModelConfig, DataConfig, TrainingConfig}).
%% Stop training
stop_training() ->
gen_server:call(?MODULE, stop_training).
%% Get training status
get_status() ->
gen_server:call(?MODULE, get_status).
%%====================================================================
%% gen_server callbacks
%%====================================================================
init(_Options) ->
io:format("MLX Distributed Coordinator started on ~p~n", [node()]),
%% Start MLX application
application:start(mlx),
%% Monitor nodes
net_kernel:monitor_nodes(true),
{ok, #state{}}.
handle_call({add_worker, Node}, _From, State) ->
case net_adm:ping(Node) of
pong ->
case rpc:call(Node, mlx_dist_worker, register_with_coordinator, [node()]) of
ok ->
NewWorkers = lists:usort([Node | State#state.workers]),
io:format("Added worker: ~p (total: ~p workers)~n",
[Node, length(NewWorkers)]),
{reply, ok, State#state{workers = NewWorkers}};
Error ->
{reply, {error, Error}, State}
end;
pang ->
{reply, {error, node_not_reachable}, State}
end;
handle_call({remove_worker, Node}, _From, State) ->
NewWorkers = lists:delete(Node, State#state.workers),
io:format("Removed worker: ~p (remaining: ~p workers)~n",
[Node, length(NewWorkers)]),
{reply, ok, State#state{workers = NewWorkers}};
handle_call(list_workers, _From, State) ->
{reply, State#state.workers, State};
handle_call({distribute_data, Data, Strategy}, _From, State) ->
case State#state.workers of
[] ->
{reply, {error, no_workers}, State};
Workers ->
Result = distribute_data_to_workers(Data, Workers, Strategy),
{reply, Result, State}
end;
handle_call(gather_gradients, _From, State) ->
case State#state.workers of
[] ->
{reply, {error, no_workers}, State};
Workers ->
Gradients = gather_from_workers(Workers),
%% Average gradients
AvgGradients = average_gradients(Gradients),
{reply, {ok, AvgGradients}, State#state{gradients = Gradients}}
end;
handle_call({broadcast_parameters, Parameters}, _From, State) ->
case State#state.workers of
[] ->
{reply, {error, no_workers}, State};
Workers ->
broadcast_to_workers(Workers, Parameters),
{reply, ok, State#state{parameters = Parameters}}
end;
handle_call(get_parameters, _From, State) ->
{reply, State#state.parameters, State};
handle_call({start_training, ModelConfig, DataConfig, TrainingConfig}, _From, State) ->
case State#state.workers of
[] ->
{reply, {error, no_workers}, State};
Workers ->
%% Initialize training on all workers
Results = [rpc:call(Worker, mlx_dist_worker, initialize_training,
[ModelConfig, TrainingConfig]) || Worker <- Workers],
case lists:all(fun(ok) -> true; (_) -> false end, Results) of
true ->
%% Distribute data
distribute_data_to_workers(DataConfig, Workers, TrainingConfig),
%% Start training loop
self() ! training_step,
NewState = State#state{
training_config = TrainingConfig,
status = training,
start_time = erlang:system_time(millisecond),
iterations = 0
},
{reply, ok, NewState};
false ->
{reply, {error, worker_initialization_failed}, State}
end
end;
handle_call(stop_training, _From, State) ->
%% Stop all workers
[rpc:call(Worker, mlx_dist_worker, stop_training, []) || Worker <- State#state.workers],
{reply, ok, State#state{status = idle}};
handle_call(get_status, _From, State) ->
Status = #{
status => State#state.status,
workers => length(State#state.workers),
iterations => State#state.iterations,
training_time => case State#state.start_time of
undefined -> 0;
Start -> erlang:system_time(millisecond) - Start
end
},
{reply, Status, State};
handle_call(_Request, _From, State) ->
{reply, {error, unknown_request}, State}.
handle_cast(_Msg, State) ->
{noreply, State}.
handle_info({nodedown, Node}, State) ->
case lists:member(Node, State#state.workers) of
true ->
io:format("Worker node ~p went down!~n", [Node]),
NewWorkers = lists:delete(Node, State#state.workers),
{noreply, State#state{workers = NewWorkers}};
false ->
{noreply, State}
end;
handle_info(training_step, State = #state{status = training}) ->
%% Execute one training step
case execute_training_step(State) of
{ok, NewState} ->
%% Schedule next step
erlang:send_after(10, self(), training_step),
{noreply, NewState};
{done, NewState} ->
io:format("Training completed after ~p iterations~n",
[NewState#state.iterations]),
{noreply, NewState#state{status = idle}};
{error, Reason} ->
io:format("Training error: ~p~n", [Reason]),
{noreply, State#state{status = error}}
end;
handle_info(_Info, State) ->
{noreply, State}.
terminate(_Reason, _State) ->
ok.
code_change(_OldVsn, State, _Extra) ->
{ok, State}.
%%====================================================================
%% Internal functions
%%====================================================================
distribute_data_to_workers(Data, Workers, Strategy) ->
NumWorkers = length(Workers),
case Strategy of
{split, DataList} ->
%% Split data among workers
ChunkSize = length(DataList) div NumWorkers,
distribute_chunks(DataList, Workers, ChunkSize);
{replicate, DataList} ->
%% Each worker gets all data but different batch
[rpc:call(Worker, mlx_dist_worker, set_data, [DataList, Index])
|| {Index, Worker} <- lists:zip(lists:seq(1, NumWorkers), Workers)];
_ ->
{error, unknown_strategy}
end.
distribute_chunks([], [], _) -> ok;
distribute_chunks(Data, [Worker|Workers], ChunkSize) ->
{Chunk, Rest} = case length(Data) > ChunkSize of
true -> lists:split(ChunkSize, Data);
false -> {Data, []}
end,
rpc:call(Worker, mlx_dist_worker, set_data, [Chunk, 0]),
distribute_chunks(Rest, Workers, ChunkSize).
gather_from_workers(Workers) ->
%% Gather gradients from all workers in parallel
Parent = self(),
Refs = [spawn_monitor(fun() ->
Result = rpc:call(Worker, mlx_dist_worker, get_gradients, []),
Parent ! {gradient, self(), Worker, Result}
end) || Worker <- Workers],
gather_results(Refs, []).
gather_results([], Results) -> Results;
gather_results(Refs, Results) ->
receive
{gradient, Pid, Worker, {ok, Gradient}} ->
{Ref, _} = lists:keyfind(Pid, 2, Refs),
NewRefs = lists:delete({Ref, Pid}, Refs),
gather_results(NewRefs, [{Worker, Gradient} | Results]);
{'DOWN', Ref, process, Pid, _Reason} ->
NewRefs = lists:delete({Ref, Pid}, Refs),
gather_results(NewRefs, Results)
after 30000 ->
{error, timeout}
end.
average_gradients(Gradients) when is_list(Gradients) ->
%% Convert to MLX arrays and average
case Gradients of
[] -> undefined;
[{_, FirstGrad} | _] ->
%% Assuming gradients are already MLX arrays
NumWorkers = length(Gradients),
GradArrays = [Grad || {_, Grad} <- Gradients],
%% Sum and average
Sum = lists:foldl(fun(Grad, Acc) ->
mlx:add(Acc, Grad)
end, FirstGrad, tl(GradArrays)),
mlx:divide(Sum, mlx:array(NumWorkers))
end.
broadcast_to_workers(Workers, Parameters) ->
[rpc:call(Worker, mlx_dist_worker, update_parameters, [Parameters])
|| Worker <- Workers].
execute_training_step(State) ->
%% Check if we should continue training
case should_continue_training(State) of
true ->
%% Tell workers to compute forward and backward pass
Workers = State#state.workers,
%% Execute forward-backward pass on all workers
[rpc:cast(Worker, mlx_dist_worker, compute_gradients, []) || Worker <- Workers],
%% Wait a bit for computation
timer:sleep(100),
%% Gather gradients
case gather_from_workers(Workers) of
{error, _} = Error ->
Error;
Gradients ->
%% Average gradients
AvgGradients = average_gradients(Gradients),
%% Update parameters (simplified gradient descent)
NewParams = update_parameters(State#state.parameters, AvgGradients,
State#state.training_config),
%% Broadcast new parameters
broadcast_to_workers(Workers, NewParams),
NewState = State#state{
parameters = NewParams,
iterations = State#state.iterations + 1
},
if
State#state.iterations rem 10 == 0 ->
io:format("Iteration ~p completed~n", [State#state.iterations]);
true -> ok
end,
{ok, NewState}
end;
false ->
{done, State}
end.
should_continue_training(#state{iterations = Iter, training_config = Config}) ->
MaxIter = maps:get(max_iterations, Config, 100),
Iter < MaxIter.
update_parameters(undefined, _Gradients, _Config) ->
undefined;
update_parameters(Parameters, Gradients, Config) ->
LearningRate = maps:get(learning_rate, Config, 0.01),
%% Simple gradient descent: params = params - lr * gradients
ScaledGrads = mlx:multiply(Gradients, mlx:array(LearningRate)),
mlx:subtract(Parameters, ScaledGrads).