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

训练神经网络时损失下降但自定义MSE指标随epoch上升是什么原因

训练损失下降但自定义MSE指标上升问题排查

训练神经网络时出现异常现象:训练过程中损失值持续下降,但MSE指标却随epoch增加不断上升,对应的变化趋势如下图:
损失与MSE变化趋势图

问题根因

通过分析你提供的自定义MSE指标代码,问题主要出在以下几个方面:

  • 指标累积逻辑错误
    Keras的Metric类会在整个epoch的所有batch之间累积计算结果,你当前的实现仅将每个batch的平均MSE直接累加到权重变量中,最终返回的是整个epoch所有batch的MSE之和,而非整个epoch的平均MSE。如果不同epoch的batch数量不一致,或是单batch MSE下降速度慢于batch数量带来的累加效应,就会出现指标随epoch持续上升的异常。
  • 样本权重分支逻辑错误
    代码中处理sample_weight的分支引用了未定义的values变量,你实际计算得到的损失变量为loss,如果传入样本权重会直接触发报错;即使未传入样本权重不执行该分支,也说明逻辑存在疏漏,可能影响计算准确性。
  • 指标与训练损失计算逻辑不一致
    你在自定义MSE中仅计算输出前2个通道的误差,且计算时先对axis=1做nanmean再求全局平均,需要确认训练时使用的损失函数是否和该逻辑完全匹配:如果训练损失计算了全部4个通道的误差,模型会优先降低4个通道的整体损失,很可能出现前2个通道误差上升、后2个通道误差下降更快的情况,最终表现为整体损失下降但前2通道的MSE指标上升。
  • NaN处理逻辑不匹配
    你使用tf.experimental.numpy.nanmean处理y_true中的NaN值,需要确认训练损失是否也使用了完全相同的NaN处理逻辑,如果训练损失对NaN的处理方式不同,也会导致两者的优化目标和计算结果不匹配。

修复后的参考实现

import tensorflow as tf
from tensorflow.keras import backend as K

class custom_MSE(tf.keras.metrics.Metric):

  def __init__(self, name='custom_mse', **kwargs):
    super(custom_MSE, self).__init__(name=name, **kwargs)
    # 累积所有batch的总MSE损失
    self.total_mse = self.add_weight(name='total_mse', initializer='zeros')
    # 累积有效batch数量
    self.batch_count = self.add_weight(name='batch_count', initializer='zeros')

  def update_state(self, y_true, y_pred, sample_weight=None):
    y_true = tf.convert_to_tensor(y_true)
    y_pred = tf.convert_to_tensor(y_pred)
    
    batch_size = tf.shape(y_true)[0]
    y_h = int(y_true.shape[1]//4)
    
    y_true_reshape = tf.reshape(y_true,shape=(batch_size,y_h,4))
    y_pred_reshape = tf.reshape(y_pred,shape=(batch_size,y_h,4))
    # 仅取前2个通道计算误差
    y_true_ = tf.cast(y_true_reshape[:,:,:2], tf.float32)
    y_pred_ = tf.cast(y_pred_reshape[:,:,:2], tf.float32)

    loss = K.square(y_true_ - y_pred_)  
    loss = tf.experimental.numpy.nanmean(loss,axis=1)
    batch_mse = tf.reduce_mean(loss)

    # 修复样本权重处理逻辑
    if sample_weight is not None:
        sample_weight = tf.cast(sample_weight, self.dtype)
        batch_mse = tf.multiply(batch_mse, sample_weight)
    
    # 累加总损失和batch计数
    self.total_mse.assign_add(batch_mse)
    self.batch_count.assign_add(1)

  def result(self):
    # 返回整个epoch的平均MSE,除零保护
    return tf.math.divide_no_nan(self.total_mse, self.batch_count)

  def reset_state(self):
    # 每个epoch重置计数
    self.total_mse.assign(0)
    self.batch_count.assign(0)

内容的提问来源于stack exchange,提问作者hoai lan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 18:00:03