如何解决ImageDataGenerator中训练与验证集类别不匹配问题
解决验证集与训练集类别映射不匹配的维度错误
错误根源
你遇到的报错核心原因是:
- 训练集生成器
train_gen识别出206个类别,模型最后一层输出维度为206; - 验证集生成器
val_gen默认遍历自身文件夹生成了189个类别的独立映射,输出的one-hot标签维度为189;
两者维度不匹配,导致交叉熵计算时无法完成广播操作,触发报错。
能否沿用训练集的类别映射?
完全可以,这是此类场景下的标准处理方式——验证集必须和训练集使用完全一致的类别ID映射,这样模型输出的logits才能和验证集标签正确对应(缺失类别的标签位始终为0,不影响损失计算)。
具体解决步骤
修改验证集生成器的创建代码,强制它使用训练集的类别列表:
- 提取训练集的所有类别名称列表(顺序与
train_gen的类别ID映射完全一致)
train_classes = list(train_gen.class_indices.keys())
- 创建验证集生成器时,通过
classes参数指定训练集的类别列表
val_gen = idg.flow_from_directory( TEST_DATA_PATH, target_size=(ROWS, COLS), batch_size=32, classes=train_classes, # 强制沿用训练集的所有类别体系 class_mode='categorical' # 保持分类模式与模型损失函数匹配 )
原理说明
指定classes=train_classes后,val_gen会严格按照训练集的类别顺序分配ID:
- 即使验证集没有某个类别,也会保留该类别对应的标签维度;
- 验证集样本的标签会生成206维的one-hot向量,仅对应自身类别的位置为1,其余(包括缺失类别)为0,与模型输出的206维logits维度完全匹配,交叉熵计算即可正常进行。
内容的提问来源于stack exchange,提问作者Mateusz Dorobek
相关产品推荐
相关产品推荐

