| """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 |
|
|