You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.23 11:27:45