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)
|
||||
|
||||
Reference in New Issue
Block a user