TensorFlow报错‘Label IDs must < n_classes’,但标签ID看似合规
解决TensorFlow的InvalidArgumentError: Label IDs must < n_classes问题
嘿,这个错误指向的问题其实很直白——你的模型认为只有2个类别(错误里的y=2),但你的标签里出现了ID为4的数值,显然超出了0-1的合法范围,所以TensorFlow的断言检查直接失败了。结合你用Scikit-Learn的LabelEncoder()这件事,大概率是模型的类别数设置和实际标签的类别数不匹配导致的,下面是具体的解决步骤:
先确认你实际的类别总数
用LabelEncoder编码后,直接查看它识别到的类别数量:print(len(label_encoder.classes_))这个数字就是你需要分类的真实类别数,比如输出是5,就说明你有5个不同类别,标签ID应该是0-4。
同步模型的
n_classes参数
不管你用TensorFlow的预制估算器(比如LinearClassifier)还是自定义模型,一定要把类别数参数设置成上面得到的数值。举个例子,如果你之前的模型代码是这样的:classifier = tf.estimator.LinearClassifier( feature_columns=my_feature_columns, n_classes=2 # 这里是错误根源! )改成:
num_classes = len(label_encoder.classes_) classifier = tf.estimator.LinearClassifier( feature_columns=my_feature_columns, n_classes=num_classes )如果是自定义模型的最后一层(比如用
Dense做分类),也要确保单元数等于类别数:model.add(tf.keras.layers.Dense(num_classes, activation='softmax'))额外排查点
要是确认了类别数还是有问题,那得检查下:- 是不是划分训练/测试集时,测试集中出现了训练集没有的类别?LabelEncoder遇到未见过的类别会报错,但如果是手动编码时操作失误也有可能。
- 有没有在编码后不小心修改过标签数值?比如做了不必要的加减操作,导致标签ID超出合法范围。
按照这个思路调整后,应该就能解决这个断言错误了。
内容的提问来源于stack exchange,提问作者Rudy
相关产品推荐
相关产品推荐

