Source code for cann.surrogate.mlp
from __future__ import annotations
import torch
from cann.config.model_config import ModelConfig
from cann.engine.training_config import TrainingConfig
[docs]
def build_mlp_network(
model_config: ModelConfig,
training_config: TrainingConfig,
) -> torch.nn.Module:
"""
Build a feedforward neural network (MLP) for the surrogate.
Args:
model_config (ModelConfig): Configuration for the underlying pricing model.
training_config (TrainingConfig): Hyperparameters and settings for training.
Returns:
torch.nn.Module: The constructed neural network model.
Raises:
ValueError: If the activation function specified in ``training_config`` is unsupported.
"""
layers = []
prev_dim = model_config.num_parameters + model_config.num_market_inputs
output_dim = model_config.num_outputs
activation_fn = {
"relu": torch.nn.ReLU(),
"tanh": torch.nn.Tanh(),
"gelu": torch.nn.GELU(),
}.get(training_config.activation.lower())
if activation_fn is None:
raise ValueError(
f"Unsupported activation function: {training_config.activation}"
)
for hidden_dim in training_config.hidden_layers:
layers.append(torch.nn.Linear(prev_dim, hidden_dim))
layers.append(activation_fn)
if training_config.batch_norm:
layers.append(torch.nn.BatchNorm1d(hidden_dim))
if training_config.dropout_rate > 0:
layers.append(torch.nn.Dropout(training_config.dropout_rate))
prev_dim = hidden_dim
layers.append(torch.nn.Linear(prev_dim, output_dim))
network = torch.nn.Sequential(*layers)
if training_config.weight_init == "uniform" or training_config.bias_init is not None:
for module in network.modules():
if not isinstance(module, torch.nn.Linear):
continue
if training_config.weight_init == "uniform":
torch.nn.init.uniform_(
module.weight,
a=training_config.weight_init_min,
b=training_config.weight_init_max,
)
if module.bias is not None and training_config.bias_init is not None:
torch.nn.init.constant_(module.bias, training_config.bias_init)
return network