Current section
Files
Jump to
Current section
Files
src/mlx_param_server.erl
-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).