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