quant_test / models /mlp_model.py
lucky-loster's picture
Upload folder using huggingface_hub
590a501 verified
Raw
History Blame Contribute Delete
1.95 kB
"""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