diff --git a/Code/python/notebooks/hyperparameter_scan.py b/Code/python/notebooks/hyperparameter_scan.py index 81c8d96..0c9ce34 100644 --- a/Code/python/notebooks/hyperparameter_scan.py +++ b/Code/python/notebooks/hyperparameter_scan.py @@ -2,6 +2,7 @@ import torch import pathlib import pandas as pd from aiRNN import dataloader, models, losses +import copy MODE = 3 DEVICE = "cuda:2" @@ -76,8 +77,20 @@ def train_one_model(model, dataset, val_dataset, warm_up_steps, pred_steps, epoc criterion = losses.HaversineAltitudeLoss(alt_const=altitude_weight) for epoch in range(epochs): + # snapshot model + optimizer at start of epoch + model_state = copy.deepcopy(model.state_dict()) + opt_state = copy.deepcopy(optimizer.state_dict()) + batch_losses = [] - loader = torch.utils.data.DataLoader(dataset, batch_size=batch_size, shuffle=True, collate_fn=dataloader.collate_to_cpu, num_workers=4) + loader = torch.utils.data.DataLoader( + dataset, + batch_size=batch_size, + shuffle=True, + collate_fn=dataloader.collate_to_cpu, + num_workers=4 + ) + + nan_triggered = False for X_f, X_t, y, X_c in loader: X_f = X_f.to(DEVICE) @@ -85,18 +98,41 @@ def train_one_model(model, dataset, val_dataset, warm_up_steps, pred_steps, epoc y = y.to(DEVICE) if X_c is not None: X_c = X_c.to(DEVICE) + optimizer.zero_grad() y_pred, _ = model(X_t, X_f, X_c, warm_up_steps, pred_steps) loss = criterion(y_pred, y) + + # bail early if experiencing NaNs + if not torch.isfinite(loss): + print(f"Epoch {epoch}: NaN detected — rolling back.") + nan_triggered = True + break + loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) + + # check for NaNs in gradients + if any(torch.isnan(p.grad).any() for p in model.parameters() if p.grad is not None): + print(f"Epoch {epoch}: NaN in gradients detected — rolling back.") + nan_triggered = True + break + optimizer.step() batch_losses.append(loss.item()) + if nan_triggered: + # revert + model.load_state_dict(model_state) + optimizer.load_state_dict(opt_state) + # optionally lower LR or break training altogether + continue + mean_loss = sum(batch_losses) / len(batch_losses) scheduler.step(mean_loss) print(f"Epoch {epoch}: loss={mean_loss}, lr={optimizer.param_groups[0]['lr']}") + # Validation with torch.no_grad(): val_losses = []