使用plt.imshow显示MNIST图像时遇Invalid shape (28,28,1)错误求助
Hey, I’ve dealt with this exact error before—let’s get it sorted out for you!
The problem here is that Matplotlib’s plt.imshow() doesn’t play nice with single-channel images that have a trailing channel dimension (like your (28,28,1) array). By default, it expects either a 2D array (height, width) for grayscale images, or a 3D array with 3/4 channels for RGB/RGBA images.
Here are two straightforward fixes you can use:
1. Strip the extra channel dimension with squeeze()
This method collapses that singleton channel axis, turning your (28,28,1) array into a clean (28,28) 2D array that imshow understands immediately:
# Basic fix plt.imshow(train_images[0].squeeze()) # Better: Add a grayscale colormap to make the digit look right plt.imshow(train_images[0].squeeze(), cmap='gray')
2. Tell imshow to treat it as grayscale directly
You can also skip modifying the array by specifying the colormap and letting imshow know to ignore the last channel axis:
plt.imshow(train_images[0], cmap='gray', interpolation='none')
Full Modified Code Example
Here’s how your code would look with the first fix applied (plus a plt.show() to actually display the image):
(train_images, train_labels), (test_images, test_labels) = datasets.mnist.load_data() train_images = train_images.reshape((60000, 28, 28, 1)) test_images = test_images.reshape((10000, 28, 28, 1)) train_images = train_images.astype('float32') / 255 test_images = test_images.astype('float32') / 255 # Fixed imshow call plt.imshow(train_images[0].squeeze(), cmap='gray') plt.show() # Don't forget this line—otherwise the image won't pop up!
Quick tip: Using cmap='gray' ensures your MNIST digit renders in proper grayscale instead of Matplotlib’s default blue-yellow viridis map, which makes the digit much easier to read.
内容的提问来源于stack exchange,提问作者juhyeok

