为何显示多张MNIST图像无需执行reshape操作?
为什么显示多张MNIST图像时无需reshape?
先看两种场景的代码差异:
单张图像显示(需要reshape)
import tensorflow as tf import matplotlib.pyplot as plt mnist = tf.keras.datasets.mnist (x_train, y_train), (x_test, y_test) = mnist.load_data() x_train, x_test = x_train / 255.0, x_test / 255.0 sample = x_train[:1].reshape((28,28)) plt.imshow(sample, cmap="gray") plt.show()
多张图像显示(无需reshape)
import tensorflow as tf import matplotlib.pyplot as plt mnist = tf.keras.datasets.mnist (x_train, y_train), (x_test, y_test) = mnist.load_data() x_train, x_test = x_train / 255.0, x_test / 255.0 plt.figure(figsize=(10,10)) for i in range(25): plt.subplot(5,5,i+1) plt.imshow(x_train[i]) plt.show()
核心原因是两种取数方式得到的数组形状不一样:
- MNIST的
x_train本身是(60000, 28, 28)的三维数组,其中每个单张图像样本是(28,28)的二维数组。 - 第一种场景用
x_train[:1]是切片操作,切片会保留原数组的维度结构,得到的是(1,28,28)的三维数组——相当于把单张图像套在了一个额外的维度里。而plt.imshow处理灰度图像时,只接受(高,宽)的二维数组,或者(高,宽,3)/(高,宽,4)的多通道数组,这种(1,28,28)的结构不符合要求,必须用reshape去掉多余维度。 - 第二种场景用
x_train[i]是直接索引单个样本,得到的就是(28,28)的二维数组,完全符合plt.imshow的输入要求,自然不需要reshape。
内容的提问来源于stack exchange,提问作者xiaochuan fang
相关产品推荐
相关产品推荐

