You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

我的实现步骤

  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}

  1. 过滤指定标签样本
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}

  1. 数据预处理
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)
  1. 模型创建与训练
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],具体代码修改如下:

  1. 创建标签映射表
    在过滤数据集之后,添加标签映射逻辑:
# 创建原始标签到新标签的映射
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)
  1. 应用标签转换
    在数据集过滤之后、归一化之前,添加标签转换的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)
  1. 保持其余代码不变
    归一化、模型构建和训练的代码无需修改,此时标签范围已经是0-5,和模型输出的6个神经元完全匹配,不会再触发报错。

内容的提问来源于stack exchange,提问作者Mahendra Singh

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.19 16:45:25