"""Neural Architecture Search optimizer."""
from __future__ import annotations
import logging
from typing import Any
import numpy as np
from .base import BaseOptimizer, OptimizationResult, register_optimizer
log = logging.getLogger(__name__)
[docs]
@register_optimizer("nas")
class NASOptimizer(BaseOptimizer):
"""Neural Architecture Search using Optuna."""
def __init__(
self,
n_trials: int = 50,
max_hidden_layers: int = 4,
min_hidden_dim: int = 16,
max_hidden_dim: int = 256,
) -> None:
self.n_trials = n_trials
self.max_hidden_layers = max_hidden_layers
self.min_hidden_dim = min_hidden_dim
self.max_hidden_dim = max_hidden_dim
[docs]
def optimize(
self,
model: Any,
train_data: tuple[Any, Any],
val_data: tuple[Any, Any] | None = None,
**kwargs: Any,
) -> OptimizationResult:
import optuna
train_features, train_labels = train_data
train_features = np.asarray(train_features)
train_labels = np.asarray(train_labels)
if val_data is not None:
val_features, val_labels = np.asarray(val_data[0]), np.asarray(val_data[1])
else:
val_features, val_labels = train_features, train_labels
input_dim = train_features.shape[1]
output_dim = int(np.max(train_labels)) + 1
activations = list(kwargs.get("activations", ["relu", "gelu", "tanh"]))
epochs = int(kwargs.get("epochs", 30))
batch_size = int(kwargs.get("batch_size", 128))
max_hidden_layers = self.max_hidden_layers
min_hidden_dim = self.min_hidden_dim
max_hidden_dim = self.max_hidden_dim
history: list[dict[str, Any]] = []
def objective(trial: optuna.Trial) -> float:
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader, TensorDataset
n_layers = trial.suggest_int("n_layers", 1, max_hidden_layers)
hidden_dims = []
for i in range(n_layers):
dim = trial.suggest_int(f"hidden_{i}", min_hidden_dim, max_hidden_dim, log=True)
hidden_dims.append(dim)
activation_name = trial.suggest_categorical("activation", activations)
dropout = trial.suggest_float("dropout", 0.0, 0.5)
lr = trial.suggest_float("lr", 1e-5, 1e-2, log=True)
act_fn = {"relu": nn.ReLU, "gelu": nn.GELU, "tanh": nn.Tanh}.get(activation_name, nn.ReLU)
layers: list[nn.Module] = []
prev = input_dim
for dim in hidden_dims:
layers.extend([nn.Linear(prev, dim), act_fn(), nn.Dropout(dropout)])
prev = dim
layers.append(nn.Linear(prev, output_dim))
student = nn.Sequential(*layers)
device = "cuda" if torch.cuda.is_available() else "cpu"
student = student.to(device)
opt = torch.optim.Adam(student.parameters(), lr=lr)
x_t = torch.tensor(train_features, dtype=torch.float32, device=device)
y_t = torch.tensor(train_labels, dtype=torch.long, device=device)
dl = DataLoader(TensorDataset(x_t, y_t), batch_size=batch_size, shuffle=True)
for _ in range(epochs):
student.train()
for xb, yb in dl:
opt.zero_grad()
F.cross_entropy(student(xb), yb).backward()
opt.step()
student.eval()
x_v = torch.tensor(val_features, dtype=torch.float32, device=device)
with torch.no_grad():
val_acc = float((student(x_v).argmax(1).cpu().numpy() == val_labels).mean())
n_params = sum(p.numel() for p in student.parameters())
history.append({
"hidden_dims": hidden_dims,
"activation": activation_name,
"dropout": dropout,
"lr": lr,
"val_accuracy": val_acc,
"n_parameters": n_params,
})
return val_acc
optuna.logging.set_verbosity(optuna.logging.WARNING)
study = optuna.create_study(
direction="maximize",
sampler=optuna.samplers.TPESampler(seed=kwargs.get("seed", 42)),
)
study.optimize(objective, n_trials=self.n_trials)
best = study.best_trial
best_arch = {
"hidden_dims": [best.params[f"hidden_{i}"] for i in range(best.params["n_layers"])],
"activation": best.params["activation"],
"dropout": best.params["dropout"],
"lr": best.params["lr"],
"val_accuracy": best.value,
}
pareto: list[dict[str, Any]] = []
sorted_h = sorted(history, key=lambda h: h.get("n_parameters", float("inf")))
best_so_far = 0.0
for arch in sorted_h:
if arch.get("val_accuracy", 0) > best_so_far:
best_so_far = arch["val_accuracy"]
pareto.append(arch)
best_n_params = next(
(h["n_parameters"] for h in history if h.get("val_accuracy") == best.value), 0,
)
return OptimizationResult(
model=model,
accuracy_after=best.value,
optimized_params=best_n_params,
pareto_points=pareto,
metadata={"best_architecture": best_arch, "search_history": history},
)