训练神经网络时损失下降但自定义MSE指标随epoch上升是什么原因
训练损失下降但自定义MSE指标上升问题排查
训练神经网络时出现异常现象:训练过程中损失值持续下降,但MSE指标却随epoch增加不断上升,对应的变化趋势如下图:
问题根因
通过分析你提供的自定义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
相关产品推荐
相关产品推荐

