# Test loader has a batch size of 1 img = next(iter(test_loader))[0].to(device) print(f"nImage has shape: {img.shape}n") # Plot image img_copy = img.to('cpu') plt.imshow(img_copy.reshape(400,700), cmap="gray") plt.axis("off") plt.show()