TensorFlow MNIST代码报错AttributeError: 'dict'无splits属性如何解决
报错原因
- 核心触发原因:你代码中的
mnist_info实际是字典类型对象,本身不存在splits属性。该问题通常出现在使用tensorflow_datasets加载MNIST数据集的阶段:调用tfds.load时没有设置with_info=True参数,或者没有正确拆包接口返回的(数据集集合,元信息)二元组,导致mnist_info被赋值为存储训练集、测试集的数据集字典,而非数据集元信息对象。 - 除此之外你的代码还存在几处拼写/调用错误,修复完核心问题后也会触发报错,需要同步修正:
- 变量名不一致:先定义了
num_validation_sample,后续tf.cast中误写为num_validation_samples多了后缀s - 数据类型拼写错误:
tf.init64应为tf.int64 - 数据集操作调用错误:打乱数据集需要显式调用
shuffle方法,原代码scaled_train_and_validation_data(BUFFER_SIZE)写法错误 - 属性名拼写错误:获取样本数的属性为
num_examples,原代码少写了后缀s,写为num_example
- 变量名不一致:先定义了
修复方案
首先修正数据集加载代码,确保正确获取元信息对象:
import tensorflow as tf import tensorflow_datasets as tfds # with_info=True会返回(数据集集合,元信息)二元组,as_supervised=True返回(图像,标签)的标准监督学习数据结构 mnist_dataset, mnist_info = tfds.load(name='mnist', with_info=True, as_supervised=True)
再修正后续业务代码的所有问题,完整可运行代码如下:
mnist_train = mnist_dataset['train'] mnist_test = mnist_dataset['test'] num_validation_sample = 0.1 * mnist_info.splits['train'].num_examples num_validation_sample = tf.cast(num_validation_sample, tf.int64) num_test_samples = mnist_info.splits['test'].num_examples num_test_samples = tf.cast(num_test_samples, tf.int64) def scale(image, label): image = tf.cast(image, tf.float32) image /= 255. return image, label scaled_train_and_validation_data = mnist_train.map(scale) test_data = mnist_test.map(scale) BUFFER_SIZE = 10000 shuffled_train_and_validation_data = scaled_train_and_validation_data.shuffle(BUFFER_SIZE) validation_data = shuffled_train_and_validation_data.take(num_validation_sample) train_data = shuffled_train_and_validation_data.skip(num_validation_sample) BATCH_SIZE = 100 train_data = train_data.batch(BATCH_SIZE) validation_data = validation_data.batch(num_validation_sample) validation_inputs, validation_targets = next(iter(validation_data))
内容的提问来源于stack exchange,提问作者Ashan
相关产品推荐
相关产品推荐

