在TF回调中使用model.predict后,YOLOv3训练出现NaN损失问题
YOLOv3 TPU训练中callback调用predict后损失NaN的解决方案
针对你在Kaggle TPU VM(TensorFlow 2.12)上遇到的问题——自定义callback的on_epoch_end中调用model.predict计算mAP后,第二个epoch损失变为NaN,以下是具体原因和解决办法:
1. 分布式策略下的状态同步问题
TPU的分布式策略会对模型权重和优化器状态进行分片管理,model.predict内部的跨设备数据流转可能破坏状态同步。
- 修复方式:在预测时强制使用模型的分布式策略作用域:
def on_epoch_end(self, epoch, logs=None): with self.model.strategy.scope(): pred = self.model.predict(...) metric = self.compute(pred, ...) tf.print(metric) - 替代方案:直接使用前向传播代替
predict,避免内部的分布式处理逻辑:def on_epoch_end(self, epoch, logs=None): val_preds = [] for x_batch in your_validation_dataset: # 无需梯度,直接执行前向传播 pred = self.model(x_batch, training=False) val_preds.append(pred) metric = self.compute(val_preds, ...) tf.print(metric)
2. BatchNormalization层状态冲突
YOLOv3中的BN层在predict时默认使用training=False(滑动均值/方差),与训练阶段的training=True(批次均值/方差)状态冲突,可能导致模型输出异常。
- 临时修复(不影响mAP计算):预测时强制BN层使用训练模式:
pred = self.model(x_batch, training=True) - 长期解决方案:若数据集足够大,训练前固定BN层的滑动统计:
with strategy.scope(): model = model_def() # 冻结BN层 for layer in model.layers: if isinstance(layer, tf.keras.layers.BatchNormalization): layer.trainable = False model.compile(loss=loss, optimizer=optimizer)
3. 输入数据预处理不一致
若验证集的预处理逻辑与训练集不同,可能产生NaN/Inf值,导致模型输出异常,进而引发后续训练损失NaN。
- 修复方式:统一训练和验证数据的预处理管道,确保归一化、数据裁剪等逻辑完全一致;同时在
compute函数中添加异常值检查:def compute(self, preds, ground_truth): # 检查预测结果是否存在异常值 for pred in preds: if tf.reduce_any(tf.math.is_nan(pred)) or tf.reduce_any(tf.math.is_inf(pred)): tf.print("预测结果包含NaN/Inf,终止mAP计算") return 0.0 # 正常计算mAP ...
4. 优化器状态被破坏
TPU的分布式优化器状态可能因predict调用被意外修改,导致后续梯度更新出现NaN。
- 修复方式:预测前后保存并恢复优化器状态:
def on_epoch_end(self, epoch, logs=None): # 保存优化器权重 opt_weights = self.model.optimizer.get_weights() # 执行预测 pred = self.model.predict(...) metric = self.compute(pred, ...) tf.print(metric) # 恢复优化器状态 self.model.optimizer.set_weights(opt_weights)
5. 损失函数数值稳定性不足
YOLOv3的损失函数包含log、除法等易产生NaN的操作,若模型输出极端值(如置信度为0),会直接导致损失NaN。
- 修复方式:在损失函数中添加数值稳定处理:
def yolo_loss(y_true, y_pred): # 置信度损失:避免log(0) conf_pos_loss = -y_true[..., 0] * tf.math.log(tf.maximum(y_pred[..., 0], 1e-8)) conf_neg_loss = -(1 - y_true[..., 0]) * tf.math.log(tf.maximum(1 - y_pred[..., 0], 1e-8)) conf_loss = tf.reduce_mean(conf_pos_loss + conf_neg_loss) # 边界框回归损失:避免除以0 box_loss = tf.reduce_mean( y_true[..., 1:5] * tf.square(y_pred[..., 1:5] - y_true[..., 5:9]) / tf.maximum(y_true[..., 9:10], 1e-4) ) # 分类损失同理处理 cls_loss = ... return conf_loss + box_loss + cls_loss
内容的提问来源于stack exchange,提问作者wd violet
相关产品推荐
相关产品推荐

