def val_test(dataloader, model): # Get dataset size dataset_size = len(dataloader.dataset) # Turn off gradient calculation for validation with torch.no_grad(): # Loop over dataset correct = 0 wrong_preds = [] for (images, labels) in dataloader: images, labels = images.to(device), labels.to(device) # Get raw values from model output = model(images) # Derive prediction y_pred = output.argmax(1) # Count correct classifications over all batches correct += (y_pred == labels).type(torch.float32).sum().item() # Save wrong predictions (image, pred_lbl, true_lbl) for i, _ in enumerate(labels): if y_pred[i] != labels[i]: wrong_preds.append((images[i], y_pred[i], labels[i])) # Calculate accuracy acc = correct / dataset_size return acc, wrong_preds