"""Simple MLP model for factor combination.""" from __future__ import annotations import torch import torch.nn as nn from torch.utils.data import DataLoader, TensorDataset class FactorMLP(nn.Module): def __init__(self, n_features: int, hidden: int = 64): super().__init__() self.net = nn.Sequential( nn.Linear(n_features, hidden), nn.ReLU(), nn.Dropout(0.2), nn.Linear(hidden, hidden // 2), nn.ReLU(), nn.Linear(hidden // 2, 1), ) def forward(self, x): return self.net(x).squeeze(-1) def train_mlp(X_train, y_train, X_valid, y_valid, epochs: int = 20, lr: float = 1e-3): device = "cuda" if torch.cuda.is_available() else "cpu" model = FactorMLP(X_train.shape[1]).to(device) optimizer = torch.optim.Adam(model.parameters(), lr=lr) loss_fn = nn.MSELoss() train_loader = DataLoader( TensorDataset(torch.tensor(X_train, dtype=torch.float32), torch.tensor(y_train, dtype=torch.float32)), batch_size=4096, shuffle=True, ) X_valid_t = torch.tensor(X_valid, dtype=torch.float32, device=device) y_valid_t = torch.tensor(y_valid, dtype=torch.float32, device=device) best_loss, best_state = float("inf"), None for _ in range(epochs): model.train() for xb, yb in train_loader: xb, yb = xb.to(device), yb.to(device) optimizer.zero_grad() loss = loss_fn(model(xb), yb) loss.backward() optimizer.step() model.eval() with torch.no_grad(): valid_loss = loss_fn(model(X_valid_t), y_valid_t).item() if valid_loss < best_loss: best_loss = valid_loss best_state = {k: v.cpu().clone() for k, v in model.state_dict().items()} if best_state: model.load_state_dict(best_state) print(f"MLP valid MSE: {best_loss:.6f}") return model