TensorFlow加载MNIST数据遇警告,寻求官方替代方案
解决MNIST数据加载的弃用警告问题
那个旧的input_data模块确实已经被TensorFlow官方弃用了,所以才会弹出这个警告。现在新版TensorFlow推荐用Keras内置的数据集接口来加载MNIST,不仅更简洁,还能更好地兼容最新版本的框架。
给你替换后的完整代码,和原来的功能完全一致,还去掉了冗余的导入:
import numpy as np import tensorflow as tf import os # 屏蔽不必要的日志信息,和你原来的作用一样 os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3' # 用Keras接口加载MNIST数据 (x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data() # 把数据转换成和原来一样的浮点型数组(原来的images是0-1的浮点值) train_data = x_train.reshape(-1, 784).astype(np.float32) / 255.0 eval_data = x_test.reshape(-1, 784).astype(np.float32) / 255.0 # 处理one-hot标签,和你原来的one_hot=True对应 train_labels = tf.keras.utils.to_categorical(y_train, num_classes=10).astype(np.int32) eval_labels = tf.keras.utils.to_categorical(y_test, num_classes=10).astype(np.int32)
几点说明:
tf.keras.datasets.mnist.load_data()会自动下载并缓存MNIST数据,不需要手动指定路径(默认存在用户目录下的.keras/datasets里)- 原来的
images是扁平化的784维数组,所以这里用reshape(-1,784)把28x28的图片转成一维,再除以255归一化到0-1区间,和旧接口的输出完全一致 to_categorical函数直接实现了one-hot编码,完美替代原来的one_hot=True参数
如果之后还要用TensorFlow的高级API做训练,还可以把这些数组转换成数据集对象,方便后续批量、打乱等操作:
train_dataset = tf.data.Dataset.from_tensor_slices((train_data, train_labels)) train_dataset = train_dataset.shuffle(60000).batch(32)
这样就完全符合新版TensorFlow的最佳实践啦,不会再收到弃用警告~
内容的提问来源于stack exchange,提问作者Jorge Rodriguez Molinuevo
相关产品推荐
相关产品推荐

