Implement nan failsafe for training

This commit is contained in:
2025-11-27 10:24:02 +01:00
parent 3931c14791
commit 51094364b2
+37 -1
View File
@@ -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 = []