diff --git a/Code/python/notebooks/hyperparameter_scan.py b/Code/python/notebooks/hyperparameter_scan.py index 875a4f2..f10578c 100644 --- a/Code/python/notebooks/hyperparameter_scan.py +++ b/Code/python/notebooks/hyperparameter_scan.py @@ -77,9 +77,14 @@ def train_one_model(model, dataset, val_dataset, warm_up_steps, pred_steps, epoc for epoch in range(epochs): batch_losses = [] - loader = torch.utils.data.DataLoader(dataset, batch_size=256, shuffle=True, collate_fn=dataloader.collate_to_cuda(DEVICE), num_workers=4) + loader = torch.utils.data.DataLoader(dataset, batch_size=256, shuffle=True, collate_fn=dataloader.collate_to_cpu, num_workers=4) for X_f, X_t, y, X_c in loader: + X_f = X_f.to(DEVICE) + X_t = X_t.to(DEVICE) + 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) @@ -95,8 +100,13 @@ def train_one_model(model, dataset, val_dataset, warm_up_steps, pred_steps, epoc # Validation with torch.no_grad(): val_losses = [] - val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=256, collate_fn=dataloader.collate_to_cuda(DEVICE), num_workers=4) + val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=256, collate_fn=dataloader.collate_to_cpu, num_workers=4) for X_f, X_t, y, X_c in val_loader: + X_f = X_f.to(DEVICE) + X_t = X_t.to(DEVICE) + y = y.to(DEVICE) + if X_c is not None: + X_c = X_c.to(DEVICE) y_pred, _ = model(X_t, X_f, X_c, warm_up_steps, pred_steps) val_losses.append(criterion(y_pred, y).item()) val_loss = sum(val_losses) / len(val_losses) diff --git a/Code/python/src/aiRNN/dataloader.py b/Code/python/src/aiRNN/dataloader.py index 02da58c..66a9083 100644 --- a/Code/python/src/aiRNN/dataloader.py +++ b/Code/python/src/aiRNN/dataloader.py @@ -172,6 +172,20 @@ def collate_to_cuda(device): return Xf, Xt, Y, Xc return _collate +def collate_to_cpu(batch): + Xf, Xt, Y, Xc = zip(*batch) + + Xf = torch.stack(Xf).cpu() + Xt = torch.stack(Xt).cpu() + Y = torch.stack(Y).cpu() + + if Xc[0] is not None: + Xc = torch.stack(Xc).cpu() + else: + Xc = None + + return Xf, Xt, Y, Xc + class EvenlySpacedDataset(BaseDataset): def __init__(