使用VGG16微调图像时训练/验证损失为NaN且精度不提升
VGG16微调时损失NaN、权重保存报错的解决思路
一、数据层面排查
- 检查图像完整性:遍历所有图像路径,用PIL/OpenCV尝试读取,剔除损坏、无法加载的图像。
- 修正标签数据类型:将标签从
float16改为float32,避免与Keras默认计算精度不匹配引发数值异常。 - 确认文件路径有效性:
directory=None要求file_path必须是绝对路径,若为相对路径,生成器无法读取图像会导致输入异常,进而出现NaN。 - 添加图像归一化:VGG16预训练时基于归一化后的图像训练,在
ImageDataGenerator中加入rescale=1./255,将像素值缩放到0-1区间,避免原始大数值引发梯度爆炸。
二、模型配置调整
- 冻结预训练层:直接训练整个VGG16+新增层容易引发梯度波动,先冻结预训练权重:
待新增层训练稳定后,再按需解冻上层微调。vgg16.trainable = False - 显式初始化输出层权重:避免默认初始化导致的数值异常,指定更稳定的初始化器:
Dense(1, activation='sigmoid', kernel_initializer='he_normal') - 验证标签合法性:确认所有标签仅为0/1,无缺失值或异常值。
三、训练流程优化
- 先移除早停调试:
restore_best_weights=True在所有轮次损失为NaN时会触发报错,先去掉早停回调,跑2-3轮看损失是否正常,再逐步恢复。 - 进一步降低学习率:尝试将学习率降至
1e-5,配合冻结预训练层,减少梯度波动。 - 调整批次大小:若当前批次无两类样本(全0或全1),会导致二分类损失计算异常,可将批次大小降至64,确保每个批次包含两类样本。
四、快速调试方法
- 单批次测试:用
train_image_generator.next()获取一个批次数据,手动调用model.train_on_batch(),观察是否出现NaN,快速定位是数据还是模型问题。 - 自定义损失函数打印中间值:在损失计算中加入打印逻辑,排查哪一步出现数值异常。
内容的提问来源于stack exchange,提问作者Shamayl
相关产品推荐
相关产品推荐

