Fix to cuda error
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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__(
|
||||
|
||||
Reference in New Issue
Block a user