def visualize_model(model, num_images=6): was_training = model.training model.eval() images_so_far = 0 fig = plt.figure() with [torch.no_grad](https://docs.pytorch.org/docs/stable/generated/torch.no_grad.html#torch.no_grad "torch.no_grad")(): for i, ([inputs](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), labels) in enumerate(dataloaders['val']): [inputs](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [inputs](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").to(device) labels = labels.to(device) outputs = model([inputs](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) _, preds = [torch.max](https://docs.pytorch.org/docs/stable/generated/torch.max.html#torch.max "torch.max")(outputs, 1) for j in range([inputs](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").size()[0]): images_so_far += 1 ax = plt.subplot(num_images//2, 2, images_so_far) ax.axis('off') ax.set_title(f'predicted: {class_names[preds[j]]}') imshow([inputs](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").cpu().data[j]) if images_so_far == num_images: model.train(mode=was_training) return model.train(mode=was_training)