Current section
Files
Jump to
Current section
Files
lib/integration/axon.ex
defmodule Scout.Integration.Axon do
@moduledoc """
Integration with Axon neural network library for Elixir.
Provides hyperparameter optimization for:
- Network architecture (layers, units, activation functions)
- Training parameters (learning rate, batch size, epochs)
- Regularization (dropout, weight decay)
- Optimizers (Adam, SGD, RMSprop parameters)
"""
alias Scout.{Study, Trial}
@doc """
Optimizes Axon model hyperparameters.
## Example
def create_model(params) do
Axon.input("input", shape: {nil, 784})
|> Axon.dense(params[:hidden1_units], activation: params[:activation1])
|> Axon.dropout(rate: params[:dropout1])
|> Axon.dense(params[:hidden2_units], activation: params[:activation2])
|> Axon.dropout(rate: params[:dropout2])
|> Axon.dense(10, activation: :softmax)
end
def train_model(model, params, data) do
model
|> Axon.Loop.trainer(:categorical_cross_entropy,
Axon.Optimizers.adam(params[:learning_rate]))
|> Axon.Loop.metric(:accuracy)
|> Axon.Loop.run(data, %{}, epochs: params[:epochs],
batch_size: params[:batch_size])
end
Scout.Integration.Axon.optimize(
&create_model/1,
&train_model/3,
train_data,
val_data,
search_space: %{
hidden1_units: {:int, 32, 512},
hidden2_units: {:int, 32, 512},
activation1: {:choice, [:relu, :tanh, :sigmoid]},
activation2: {:choice, [:relu, :tanh, :sigmoid]},
dropout1: {:uniform, 0.0, 0.5},
dropout2: {:uniform, 0.0, 0.5},
learning_rate: {:log_uniform, 1.0e-5, 1.0e-1},
batch_size: {:choice, [16, 32, 64, 128]},
epochs: {:int, 10, 100}
},
n_trials: 100,
pruner: Scout.Pruner.MedianPruner
)
"""
def optimize(model_fn, train_fn, train_data, val_data, opts \\ []) do
search_space = Keyword.fetch!(opts, :search_space)
n_trials = Keyword.get(opts, :n_trials, 100)
direction = Keyword.get(opts, :direction, :maximize)
pruner = Keyword.get(opts, :pruner)
sampler = Keyword.get(opts, :sampler, Scout.Sampler.TPE)
metric = Keyword.get(opts, :metric, :accuracy)
objective = fn trial ->
# Sample hyperparameters
params = sample_params(trial, search_space)
# Create model
model = model_fn.(params)
# Train with pruning support
result = train_with_pruning(
model,
params,
train_data,
val_data,
train_fn,
trial,
metric
)
result
end
study = %Study{
id: "axon_optimization_#{System.unique_integer([:positive])}",
goal: direction,
max_trials: n_trials,
parallelism: 1, # Add required field
search_space: fn _ix -> search_space end,
objective: objective,
sampler: sampler,
sampler_opts: %{},
pruner: pruner,
pruner_opts: %{},
metadata: %{framework: "axon"}
}
Scout.run(study)
end
defp sample_params(trial, search_space) do
for {param_name, spec} <- search_space, into: %{} do
value = case spec do
{:uniform, min, max} ->
Trial.suggest_float(trial, param_name, min, max)
{:log_uniform, min, max} ->
Trial.suggest_float(trial, param_name, min, max, log: true)
{:int, min, max} ->
Trial.suggest_int(trial, param_name, min, max)
{:choice, choices} ->
Trial.suggest_categorical(trial, param_name, choices)
end
{param_name, value}
end
end
defp train_with_pruning(model, params, train_data, val_data, train_fn, trial, metric) do
# Custom training loop with intermediate reporting
epochs = Map.get(params, :epochs, 10)
Enum.reduce_while(1..epochs, nil, fn epoch, _acc ->
# Train for one epoch
epoch_result = train_single_epoch(model, params, train_data, train_fn)
# Evaluate on validation set
val_metric = evaluate(model, val_data, metric)
# Report intermediate value
Trial.report(trial, val_metric, epoch)
# Check for pruning
if Trial.should_prune?(trial) do
{:halt, :pruned}
else
if epoch == epochs do
{:halt, val_metric}
else
{:cont, val_metric}
end
end
end)
end
defp train_single_epoch(_model, _params, _train_data, _train_fn) do
# Placeholder - would integrate with actual Axon training
%{loss: :rand.uniform()}
end
defp evaluate(_model, _val_data, _metric) do
# Placeholder - would integrate with actual Axon evaluation
:rand.uniform()
end
@doc """
Pruning callback for Axon training loops.
Use this in your Axon.Loop to enable pruning during training:
loop
|> Axon.Loop.handle_event(:epoch_completed,
&Scout.Integration.Axon.pruning_callback(&1, &2, trial))
"""
def pruning_callback(state, _metadata, trial) do
# Extract metric from state
metric_value = get_in(state, [:metrics, :validation, :accuracy])
# Report to trial
epoch = Map.get(state, :epoch, 0)
Trial.report(trial, metric_value, epoch)
# Check for pruning
if Trial.should_prune?(trial) do
{:halt_loop, state}
else
{:continue, state}
end
end
@doc """
Suggests optimal architecture based on dataset characteristics.
"""
def suggest_architecture(input_shape, output_shape, task_type \\ :classification) do
input_size = elem(input_shape, -1)
output_size = case output_shape do
n when is_integer(n) -> n
shape when is_tuple(shape) -> elem(shape, -1)
end
case task_type do
:classification ->
%{
hidden_layers: suggest_hidden_layers(input_size, output_size),
activation: :relu,
dropout: 0.2,
output_activation: if(output_size == 1, do: :sigmoid, else: :softmax)
}
:regression ->
%{
hidden_layers: suggest_hidden_layers(input_size, output_size),
activation: :relu,
dropout: 0.1,
output_activation: :linear
}
:autoencoder ->
%{
encoder_layers: suggest_encoder_layers(input_size),
decoder_layers: suggest_decoder_layers(input_size),
activation: :relu,
latent_dim: suggest_latent_dim(input_size)
}
end
end
defp suggest_hidden_layers(input_size, output_size) do
# Rule of thumb: geometric mean for first hidden layer
first_layer = round(:math.sqrt(input_size * output_size))
cond do
input_size < 100 ->
[first_layer]
input_size < 1000 ->
[first_layer, round(first_layer / 2)]
true ->
[first_layer, round(first_layer / 2), round(first_layer / 4)]
end
end
defp suggest_encoder_layers(input_size) do
[
round(input_size / 2),
round(input_size / 4),
round(input_size / 8)
]
|> Enum.filter(&(&1 > 10))
end
defp suggest_decoder_layers(input_size) do
suggest_encoder_layers(input_size)
|> Enum.reverse()
end
defp suggest_latent_dim(input_size) do
max(2, round(:math.log2(input_size)))
end
end