TensorFlow中tf.keras.metrics.MeanSquaredError每次调用结果不同原因?
问题描述
使用TensorFlow的tf.keras.metrics.MeanSquaredError()指标评估两个numpy数组的均方误差,但每次调用mse()都会得到不同结果。
代码示例:
import numpy as np import tensorflow as tf import matplotlib.pyplot as plt a = np.random.random(size=(100,2000)) b = np.random.random(size=(100,2000)) mse = tf.keras.metrics.MeanSquaredError() for i in range(100): v = mse(a, b).numpy() plt.scatter(i,v) print(v)
输出结果呈现持续递增的趋势。
问题原因与解决方法
- 核心原因:
tf.keras.metrics.MeanSquaredError是累积型指标,它会持续累积每次调用时传入数据的计算结果,而非每次独立计算当前输入的MSE。每次调用mse(a,b)时,程序会把当前批次的误差加入到之前的累积值中,再计算整体均值,因此结果会随调用次数递增。 - 解决方法:
- 每次计算前重置指标状态:在循环内调用
mse.reset_states(),确保每次都基于当前输入重新计算MSE:for i in range(100): mse.reset_states() v = mse(a, b).numpy() plt.scatter(i,v) print(v) - 使用无状态的损失函数计算:如果仅需单次独立计算MSE,直接使用
tf.keras.losses.mean_squared_error,它不会累积历史数据,每次调用都返回当前输入的计算结果:mse_loss = tf.keras.losses.mean_squared_error for i in range(100): # 由于loss函数返回每个样本的误差,需取均值得到整体MSE v = mse_loss(a, b).numpy().mean() plt.scatter(i,v) print(v)
- 每次计算前重置指标状态:在循环内调用
内容的提问来源于stack exchange,提问作者Aks
相关产品推荐
相关产品推荐

