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