使用TensorFlow训练MNIST指定标签分类器遇标签范围错误求助
MNIST部分标签训练报错:标签值超出有效范围
错误信息
res = tf.nn.sparse_softmax_cross_entropy_with_logits( Node: 'sparse_categorical_crossentropy/SparseSoftmaxCrossEntropyWithLogits/SparseSoftmaxCrossEntropyWithLogits' Received a label value of 9 which is outside the valid range of [0, 6). Label values: 9 5 9 9 6 1 4 4 6 6 9 4 9 1 8 5 9 5 4 8 9 9 1 8 6 4 4 9 9 4 4 8 8 6 6 5 9 4 1 5 5 6 4 1 1 8 9 6 8 5 6 1 6 6 4 6 1 4 4 4 1 1 1 6 9 8 8 8 5 1 8 8 6 6 5 1 1 5 1 6 9 8 1 8 4 6 4 9 8 1 6 5 5 9 1 6 8 1 5 5 6 9 1 9 9 6 4 6 6 4 8 6 6 4 5 4 4 5 8 1 8 6 1 5 4 5 8 1
我的实现步骤
- 加载并划分数据集
import tensorflow_datasets as tfds import tensorflow as tf val_split = 20 # percent of training data (ds_test, ds_valid, ds_train), ds_info = tfds.load( 'mnist', split=['test', f'train[0%:{val_split}%]', f'train[{val_split}%:]'], as_supervised=True, with_info=True )
过滤前ds_train标签分布:{0: 4705, 1: 5433, 2: 4772, 3: 4936, 4: 4681, 5: 4333, 6: 4728, 7: 4966, 8: 4703, 9: 4743}
- 过滤指定标签样本
known_classes = [1, 4, 5, 6, 8, 9] kc = tf.constant(known_classes, dtype=tf.int64) def predicate(image, label): isallowed = tf.equal(kc, label) reduced = tf.reduce_sum(tf.cast(isallowed, tf.int64)) return tf.greater(reduced, tf.constant(0, dtype=tf.int64)) ds_test = ds_test.filter(predicate) ds_valid = ds_valid.filter(predicate) ds_train = ds_train.filter(predicate)
过滤后ds_train标签分布:{0: 0, 1: 5433, 2: 0, 3: 0, 4: 4681, 5: 4333, 6: 4728, 7: 0, 8: 4703, 9: 4743}
- 数据预处理
def normalize_img(image, label): """Normalizes images: `uint8` -> `float32`.""" return tf.cast(image, tf.float32) / 255., label # 训练集预处理 ds_train = ds_train.map(normalize_img, num_parallel_calls=tf.data.AUTOTUNE) ds_train = ds_train.shuffle(ds_info.splits['train'].num_examples) ds_train = ds_train.batch(128, drop_remainder=True) ds_train = ds_train.prefetch(tf.data.AUTOTUNE) # 测试集与验证集预处理 ds_test = ds_test.map(normalize_img, num_parallel_calls=tf.data.AUTOTUNE) ds_test = ds_test.batch(128, drop_remainder=True) ds_test = ds_test.prefetch(tf.data.AUTOTUNE) ds_valid = ds_valid.map(normalize_img, num_parallel_calls=tf.data.AUTOTUNE) ds_valid = ds_valid.batch(128, drop_remainder=True) ds_valid = ds_valid.prefetch(tf.data.AUTOTUNE)
- 模型创建与训练
model = tf.keras.models.Sequential([ tf.keras.layers.Flatten(input_shape=(28, 28)), tf.keras.layers.Dense(128, activation='relu'), tf.keras.layers.Dense(len(known_classes)) # 对应6个类别 ]) model.compile( optimizer=tf.keras.optimizers.Adam(0.001), loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=[tf.keras.metrics.SparseCategoricalAccuracy()], ) history = model.fit( ds_train, epochs=5, validation_data=ds_valid, )
解决方案
问题根源
SparseCategoricalCrossentropy要求标签必须是0到n_classes-1的连续整数。你的模型最后一层输出6个神经元(对应6个类别),但标签还是原始的MNIST数值(比如9),9超出了0-5的有效范围,因此报错。
修复步骤
需要把原始标签[1,4,5,6,8,9]映射为连续的[0,1,2,3,4,5],具体代码修改如下:
- 创建标签映射表
在过滤数据集之后,添加标签映射逻辑:
# 创建原始标签到新标签的映射 original_labels = tf.constant(known_classes, dtype=tf.int64) new_labels = tf.constant(range(len(known_classes)), dtype=tf.int64) # 构建哈希表用于快速映射 table = tf.lookup.StaticHashTable( tf.lookup.KeyValueTensorInitializer(original_labels, new_labels), default_value=-1 # 不在已知标签里的样本会被映射为-1,后续可过滤(已提前过滤过) ) # 定义标签转换函数 def convert_label(image, label): return image, table.lookup(label)
- 应用标签转换
在数据集过滤之后、归一化之前,添加标签转换的map操作:
# 过滤后立即转换标签 ds_train = ds_train.filter(predicate).map(convert_label, num_parallel_calls=tf.data.AUTOTUNE) ds_valid = ds_valid.filter(predicate).map(convert_label, num_parallel_calls=tf.data.AUTOTUNE) ds_test = ds_test.filter(predicate).map(convert_label, num_parallel_calls=tf.data.AUTOTUNE)
- 保持其余代码不变
归一化、模型构建和训练的代码无需修改,此时标签范围已经是0-5,和模型输出的6个神经元完全匹配,不会再触发报错。
内容的提问来源于stack exchange,提问作者Mahendra Singh
相关产品推荐
相关产品推荐

