Implement nan failsafe for training
This commit is contained in:
@@ -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 = []
|
||||
|
||||
Reference in New Issue
Block a user