-module(mlx_param_server). -behaviour(gen_server). %% API -export([start_link/0, start_link/1, initialize_parameters/1, get_parameters/0, get_parameters/1, update_gradients/2, aggregate_gradients/1, apply_optimizer_step/1, checkpoint/1, restore/1, get_stats/0]). %% gen_server callbacks -export([init/1, handle_call/3, handle_cast/2, handle_info/2, terminate/2, code_change/3]). -record(state, { parameters = #{}, % Model parameters gradients = [], % Accumulated gradients from workers optimizer_state = #{}, % Optimizer state (momentum, etc.) config = #{}, % Training configuration iteration = 0, stats = #{}, checkpoints = [] }). %%==================================================================== %% API %%==================================================================== start_link() -> start_link(#{}). start_link(Config) -> gen_server:start_link({local, ?MODULE}, ?MODULE, Config, []). %% Initialize model parameters initialize_parameters(ModelSpec) -> gen_server:call(?MODULE, {initialize_parameters, ModelSpec}). %% Get all parameters get_parameters() -> gen_server:call(?MODULE, get_parameters). %% Get specific parameter by key get_parameters(Key) -> gen_server:call(?MODULE, {get_parameters, Key}). %% Update gradients from a worker update_gradients(WorkerId, Gradients) -> gen_server:call(?MODULE, {update_gradients, WorkerId, Gradients}). %% Aggregate gradients from multiple workers aggregate_gradients(GradientsList) -> gen_server:call(?MODULE, {aggregate_gradients, GradientsList}). %% Apply optimizer step with current gradients apply_optimizer_step(OptimizerConfig) -> gen_server:call(?MODULE, {apply_optimizer_step, OptimizerConfig}). %% Save checkpoint checkpoint(Path) -> gen_server:call(?MODULE, {checkpoint, Path}). %% Restore from checkpoint restore(Path) -> gen_server:call(?MODULE, {restore, Path}). %% Get training statistics get_stats() -> gen_server:call(?MODULE, get_stats). %%==================================================================== %% gen_server callbacks %%==================================================================== init(Config) -> io:format("MLX Parameter Server started~n"), application:start(mlx), {ok, #state{config = Config}}. handle_call({initialize_parameters, ModelSpec}, _From, State) -> try Parameters = initialize_model_parameters(ModelSpec), OptimizerState = initialize_optimizer_state(Parameters, State#state.config), NewState = State#state{ parameters = Parameters, optimizer_state = OptimizerState }, {reply, ok, NewState} catch Error:Reason -> {reply, {error, {Error, Reason}}, State} end; handle_call(get_parameters, _From, State) -> {reply, {ok, State#state.parameters}, State}; handle_call({get_parameters, Key}, _From, State) -> case maps:find(Key, State#state.parameters) of {ok, Value} -> {reply, {ok, Value}, State}; error -> {reply, {error, not_found}, State} end; handle_call({update_gradients, WorkerId, Gradients}, _From, State) -> %% Store gradients from worker NewGradients = [{WorkerId, Gradients} | State#state.gradients], {reply, ok, State#state{gradients = NewGradients}}; handle_call({aggregate_gradients, GradientsList}, _From, State) -> %% Aggregate gradients using different strategies Strategy = maps:get(aggregation_strategy, State#state.config, average), AggregatedGrads = case Strategy of average -> average_gradients(GradientsList); federated_avg -> federated_average(GradientsList); weighted -> weighted_average(GradientsList) end, {reply, {ok, AggregatedGrads}, State}; handle_call({apply_optimizer_step, OptimizerConfig}, _From, State) -> %% Apply optimizer update Optimizer = maps:get(optimizer, OptimizerConfig, sgd), {NewParams, NewOptState} = case Optimizer of sgd -> apply_sgd(State#state.parameters, State#state.gradients, OptimizerConfig, State#state.optimizer_state); adam -> apply_adam(State#state.parameters, State#state.gradients, OptimizerConfig, State#state.optimizer_state); momentum -> apply_momentum(State#state.parameters, State#state.gradients, OptimizerConfig, State#state.optimizer_state) end, NewState = State#state{ parameters = NewParams, optimizer_state = NewOptState, gradients = [], % Clear gradients after update iteration = State#state.iteration + 1 }, {reply, ok, update_stats(NewState)}; handle_call({checkpoint, Path}, _From, State) -> %% Save model checkpoint Checkpoint = #{ parameters => State#state.parameters, optimizer_state => State#state.optimizer_state, iteration => State#state.iteration, stats => State#state.stats }, case save_checkpoint(Path, Checkpoint) of ok -> NewCheckpoints = [{erlang:system_time(), Path} | State#state.checkpoints], {reply, ok, State#state{checkpoints = NewCheckpoints}}; Error -> {reply, Error, State} end; handle_call({restore, Path}, _From, State) -> case load_checkpoint(Path) of {ok, Checkpoint} -> NewState = State#state{ parameters = maps:get(parameters, Checkpoint), optimizer_state = maps:get(optimizer_state, Checkpoint), iteration = maps:get(iteration, Checkpoint), stats = maps:get(stats, Checkpoint, #{}) }, {reply, ok, NewState}; Error -> {reply, Error, State} end; handle_call(get_stats, _From, State) -> Stats = maps:merge(State#state.stats, #{ iteration => State#state.iteration, num_parameters => count_parameters(State#state.parameters) }), {reply, {ok, Stats}, State}; handle_call(_Request, _From, State) -> {reply, {error, unknown_request}, State}. handle_cast(_Msg, State) -> {noreply, State}. handle_info(_Info, State) -> {noreply, State}. terminate(_Reason, _State) -> ok. code_change(_OldVsn, State, _Extra) -> {ok, State}. %%==================================================================== %% Internal functions %%==================================================================== initialize_model_parameters(ModelSpec) -> %% Initialize parameters based on model specification maps:map(fun(_Key, Spec) -> case Spec of {shape, Shape, init_method} -> initialize_tensor(Shape, init_method); {shape, Shape} -> initialize_tensor(Shape, xavier); _ -> Spec end end, ModelSpec). initialize_tensor(Shape, Method) -> case Method of zeros -> mlx:zeros(Shape); ones -> mlx:ones(Shape); xavier -> xavier_initialization(Shape); he -> he_initialization(Shape); normal -> normal_initialization(Shape); _ -> mlx:zeros(Shape) end. xavier_initialization(Shape) -> %% Xavier/Glorot initialization FanIn = hd(Shape), FanOut = lists:last(Shape), Scale = math:sqrt(6.0 / (FanIn + FanOut)), %% Random uniform [-scale, scale] Random = create_random_array(Shape), Centered = mlx:subtract(mlx:multiply(Random, mlx:array(2)), mlx:array(1)), mlx:multiply(Centered, mlx:array(Scale)). he_initialization(Shape) -> %% He initialization for ReLU networks FanIn = hd(Shape), Scale = math:sqrt(2.0 / FanIn), %% Random normal with std = scale Random = create_random_array(Shape), mlx:multiply(Random, mlx:array(Scale)). normal_initialization(Shape) -> %% Standard normal initialization create_random_array(Shape). create_random_array([Rows, Cols]) -> Data = [[rand:normal() || _ <- lists:seq(1, Cols)] || _ <- lists:seq(1, Rows)], mlx:array(Data); create_random_array(Shape) -> %% For other shapes, flatten then reshape TotalSize = lists:foldl(fun(X, Acc) -> X * Acc end, 1, Shape), Data = [rand:normal() || _ <- lists:seq(1, TotalSize)], mlx:reshape(mlx:array(Data), Shape). initialize_optimizer_state(Parameters, Config) -> %% Initialize optimizer-specific state Optimizer = maps:get(optimizer, Config, sgd), case Optimizer of sgd -> #{}; momentum -> %% Initialize momentum buffers maps:map(fun(_K, V) -> #{velocity => mlx:zeros(mlx:shape(V))} end, Parameters); adam -> %% Initialize Adam buffers (first and second moments) maps:map(fun(_K, V) -> Shape = mlx:shape(V), #{ m => mlx:zeros(Shape), % First moment v => mlx:zeros(Shape), % Second moment t => 0 % Timestep } end, Parameters) end. average_gradients(GradientsList) when is_list(GradientsList) -> %% Simple averaging of gradients case GradientsList of [] -> #{}; [First | Rest] -> NumWorkers = length(GradientsList), %% Sum all gradients Summed = lists:foldl(fun(Grads, Acc) -> maps:merge_with(fun(_, G1, G2) -> mlx:add(G1, G2) end, Acc, Grads) end, First, Rest), %% Divide by number of workers maps:map(fun(_, GradSum) -> mlx:divide(GradSum, mlx:array(NumWorkers)) end, Summed) end. federated_average(GradientsList) -> %% Federated averaging - can weight by data size %% For now, same as simple average average_gradients(GradientsList). weighted_average(GradientsList) -> %% Weighted average - would need weights %% For now, same as simple average average_gradients(GradientsList). apply_sgd(Parameters, Gradients, Config, OptState) -> LR = maps:get(learning_rate, Config, 0.01), %% Average gradients first AvgGrads = average_gradients([G || {_, G} <- Gradients]), %% Update parameters: theta = theta - lr * grad NewParams = maps:merge_with(fun(_, Param, Grad) -> mlx:subtract(Param, mlx:multiply(Grad, mlx:array(LR))) end, Parameters, AvgGrads), {NewParams, OptState}. apply_momentum(Parameters, Gradients, Config, OptState) -> LR = maps:get(learning_rate, Config, 0.01), Momentum = maps:get(momentum, Config, 0.9), %% Average gradients AvgGrads = average_gradients([G || {_, G} <- Gradients]), %% Update with momentum {NewParams, NewOptState} = maps:fold(fun(Key, Param, {ParamsAcc, StateAcc}) -> Grad = maps:get(Key, AvgGrads, mlx:zeros(mlx:shape(Param))), VelState = maps:get(Key, OptState, #{velocity => mlx:zeros(mlx:shape(Param))}), %% v = momentum * v - lr * grad OldVel = maps:get(velocity, VelState), NewVel = mlx:subtract( mlx:multiply(OldVel, mlx:array(Momentum)), mlx:multiply(Grad, mlx:array(LR)) ), %% param = param + v NewParam = mlx:add(Param, NewVel), {maps:put(Key, NewParam, ParamsAcc), maps:put(Key, #{velocity => NewVel}, StateAcc)} end, {#{}, #{}}, Parameters), {NewParams, NewOptState}. apply_adam(Parameters, Gradients, Config, OptState) -> LR = maps:get(learning_rate, Config, 0.001), Beta1 = maps:get(beta1, Config, 0.9), Beta2 = maps:get(beta2, Config, 0.999), Epsilon = maps:get(epsilon, Config, 0.00000001), %% Average gradients AvgGrads = average_gradients([G || {_, G} <- Gradients]), %% Update with Adam {NewParams, NewOptState} = maps:fold(fun(Key, Param, {ParamsAcc, StateAcc}) -> Grad = maps:get(Key, AvgGrads, mlx:zeros(mlx:shape(Param))), State = maps:get(Key, OptState, #{m => mlx:zeros(mlx:shape(Param)), v => mlx:zeros(mlx:shape(Param)), t => 0}), T = maps:get(t, State) + 1, M = maps:get(m, State), V = maps:get(v, State), %% Update biased first moment NewM = mlx:add( mlx:multiply(M, mlx:array(Beta1)), mlx:multiply(Grad, mlx:array(1 - Beta1)) ), %% Update biased second moment NewV = mlx:add( mlx:multiply(V, mlx:array(Beta2)), mlx:multiply(mlx:square(Grad), mlx:array(1 - Beta2)) ), %% Bias correction MHat = mlx:divide(NewM, mlx:array(1 - math:pow(Beta1, T))), VHat = mlx:divide(NewV, mlx:array(1 - math:pow(Beta2, T))), %% Update parameters NewParam = mlx:subtract(Param, mlx:multiply( mlx:divide(MHat, mlx:add(mlx:sqrt(VHat), mlx:array(Epsilon))), mlx:array(LR) ) ), NewState = #{m => NewM, v => NewV, t => T}, {maps:put(Key, NewParam, ParamsAcc), maps:put(Key, NewState, StateAcc)} end, {#{}, #{}}, Parameters), {NewParams, NewOptState}. count_parameters(Parameters) -> maps:fold(fun(_, Param, Acc) -> Size = mlx:size(Param), {ok, SizeVal} = Size, Acc + SizeVal end, 0, Parameters). update_stats(State) -> %% Update training statistics State#state{ stats = maps:merge(State#state.stats, #{ last_update => erlang:system_time(millisecond), total_updates => maps:get(total_updates, State#state.stats, 0) + 1 }) }. save_checkpoint(Path, Checkpoint) -> %% Save checkpoint to disk %% Convert MLX arrays to lists for serialization try SerializedCheckpoint = serialize_checkpoint(Checkpoint), file:write_file(Path, term_to_binary(SerializedCheckpoint)), ok catch Error:Reason -> {error, {Error, Reason}} end. load_checkpoint(Path) -> %% Load checkpoint from disk case file:read_file(Path) of {ok, Binary} -> try SerializedCheckpoint = binary_to_term(Binary), Checkpoint = deserialize_checkpoint(SerializedCheckpoint), {ok, Checkpoint} catch Error:Reason -> {error, {Error, Reason}} end; Error -> Error end. serialize_checkpoint(Checkpoint) -> %% Convert MLX arrays to lists maps:map(fun (parameters, Params) -> maps:map(fun(_, Array) -> mlx:to_list(Array) end, Params); (optimizer_state, OptState) -> maps:map(fun(_, State) -> maps:map(fun (K, V) when K == m; K == v; K == velocity -> mlx:to_list(V); (_, V) -> V end, State) end, OptState); (_, V) -> V end, Checkpoint). deserialize_checkpoint(SerializedCheckpoint) -> %% Convert lists back to MLX arrays maps:map(fun (parameters, Params) -> maps:map(fun(_, List) -> mlx:array(List) end, Params); (optimizer_state, OptState) -> maps:map(fun(_, State) -> maps:map(fun (K, V) when K == m; K == v; K == velocity -> mlx:array(V); (_, V) -> V end, State) end, OptState); (_, V) -> V end, SerializedCheckpoint).