K fold cross validation异常 RESNET-50模型准确率恒定不收敛
问题根因
你观察到的0.6931是二分类场景下模型对两类输出概率均为0.5时的二分类交叉熵损失值,准确率卡在0.5左右说明模型完全没有学习到有效特征,参数全程未有效更新,核心是代码存在多处逻辑错误,和你预期的ResNet-50迁移学习实现完全不符。
具体错误点排查
- 根本没有搭建ResNet-50模型结构
你在getModel()函数中完全没有引入ResNet-50预训练骨干,第一层直接用Flatten(),且未定义输入形状,本质是一个结构残缺的全连接网络,和你描述的ResNet-50迁移学习模型完全无关。直接将224×224×3的图像展平后接全连接层,参数量高达3800万以上,本身就极难训练。 - 损失函数与标签格式不匹配
你配置训练生成器时设置class_mode='categorical',该模式会返回one-hot格式的标签(二分类下格式为[1,0]/[0,1]),但模型编译时使用的损失函数是binary_crossentropy,该损失要求标签为0/1单值格式、输出层为1个神经元+sigmoid激活。二者不匹配会导致梯度计算完全错误,参数无法正常更新。 - K折交叉验证逻辑错误
你将model=getModel()写在了K折循环外部,第一折训练后的权重会直接带入后续折次训练,完全不符合交叉验证“每折单独初始化模型”的要求;同时你依赖手动移动文件切分数据集,但X/Y标签数组只在循环外初始化一次,后续文件移动后数组没有同步更新,会出现索引和文件不匹配的问题。 - 训练流程配置错误
调用fit_generator时没有传入validation_data参数,训练过程不会做验证集评估;验证集生成器错误设置class_mode=None,不会返回标签;训练生成器多余添加了subset='training'参数,但你没有为ImageDataGenerator配置validation_split,会导致生成器采样逻辑异常。 - 学习率配置不合理
你设置的Adam优化器学习率为2e-5,对于随机初始化的全连接头部来说数值过小,即使损失计算正确,前期参数更新速度也会极慢。
修复方案
- 正确构建ResNet-50迁移学习模型,替换原有残缺的
getModel()实现:
from tensorflow.keras.applications import ResNet50 def getModel(): # 加载imagenet预训练的ResNet50骨干,去掉顶部全连接层 base_model = ResNet50(weights='imagenet', include_top=False, input_shape=(224,224,3)) base_model.trainable = False # 初始阶段冻住预训练层,先训练分类头 model = Sequential([ base_model, layers.GlobalAveragePooling2D(), # 用全局平均池化替代Flatten,大幅减少参数量 Dense(256, activation='relu', name='fc1'), Dense(128, activation='relu', name='fc2'), layers.Dropout(0.5), Dense(2, activation='softmax') ]) # 2个神经元softmax输出对应categorical_crossentropy损失 model.compile(optimizer=optimizers.Adam(learning_rate=1e-4), loss='categorical_crossentropy', metrics=['accuracy']) return model
如果你坚持使用
binary_crossentropy,需要将输出层改为1个神经元+sigmoid激活,同时把所有数据生成器的class_mode改为binary,两种配置二选一即可,不能混合使用。
- 修正K折循环逻辑:将
model=getModel()移动到K折循环内部,保证每一折都重新初始化模型,避免上一折权重残留。 - 修正数据生成器配置:
- 去掉训练生成器中多余的
subset='training'参数 - 将验证集生成器的
class_mode从None改为'categorical',和训练集保持一致 - 调用
fit_generator(新版本TF直接用fit即可)时传入validation_data=validation_generator参数
- 去掉训练生成器中多余的
- 替换文件移动的数据集切分逻辑:不要在训练过程中来回移动文件,建议初始就将数据集按固定比例划分好train/test目录,K折交叉验证只在训练集内部通过索引切分训练/验证子集,避免文件操作带来的路径、标签不匹配问题。
- 数据校验:训练前先打印
train_generator.class_indices确认类别和标签映射正确,随机抽取一个batch的图像和标签检查,排除图片损坏、标签错位、类别极端不平衡的问题。
内容的提问来源于stack exchange,提问作者farabee
相关产品推荐
相关产品推荐

