如何在TensorFlow中本地加载已下载的MNIST.npz数据集?
解决MNIST数据集本地加载的问题
错误原因
你用tf.keras.utils.get_file直接解包导致报错,是因为这个函数只返回本地文件的路径字符串,而不是load_data()返回的((训练集数据,训练集标签),(测试集数据,测试集标签))结构,自然无法按你写的方式解包。
正确的本地加载方法
方法1:手动用numpy加载npz文件
MNIST的npz压缩包内置了x_train、y_train、x_test、y_test四个数组,直接用numpy读取并提取即可:
import numpy as np import tensorflow as tf # 加载本地mnist.npz文件(替换为你的文件路径,相对/绝对路径都可以) with np.load('mnist.npz', allow_pickle=True) as mnist_data: x_train = mnist_data['x_train'] y_train = mnist_data['y_train'] x_test = mnist_data['x_test'] y_test = mnist_data['y_test'] # 可选:对数据做归一化处理(和load_data返回的数据使用方式一致) x_train = x_train / 255.0 x_test = x_test / 255.0
方法2:直接用load_data()指定本地路径
tf.keras.datasets.mnist.load_data()本身支持传入本地文件路径,省去手动加载的步骤:
import tensorflow as tf # 传入本地mnist.npz的路径,直接获取标准格式的数据集 (x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data(path='mnist.npz')
注意事项
- 确保下载的
mnist.npz文件完整无损坏,正常文件大小约为11MB - 如果使用绝对路径,Windows系统注意用双反斜杠
\\或者原始字符串r'路径',避免转义问题
内容的提问来源于stack exchange,提问作者Tim Arazmedov
相关产品推荐
相关产品推荐

