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

贝叶斯神经网络中model.evaluate()与手动RMSE结果差异的原因咨询

问题核心原因与解决方案

一、差异根源:Keras默认RMSE无法处理TFP分布输出

你的BNN模型输出是tfp.layers.IndependentNormal返回的概率分布对象,而非普通的预测值张量,但Keras的RootMeanSquaredError()指标是为普通张量设计的,它不知道如何正确处理分布对象,导致计算逻辑和你手动计算完全不同:

  • 你手动计算的是真实标签与分布均值的RMSE,这是衡量BNN预测准确性的合理方式(均值是最小化期望损失的最优预测值)。
  • model.evaluate()计算时,会将分布对象自动转换为张量——对于TFP的IndependentNormal,默认触发的是分布的单次采样值,而非均值。最终得到的是真实标签与随机采样值的RMSE,当分布方差较大时,采样值波动极强,结果自然和手动计算差异极大。

二、model.evaluate()在概率模型中的RMSE计算逻辑

Keras的指标函数是通用实现,不具备概率模型的感知能力:

  1. 当模型输出为TFP分布对象时,评估阶段会调用分布的__tensor__()方法将其转为张量。
  2. 对于IndependentNormal,该方法默认返回从当前分布中随机采样的一个样本,而非分布的均值或参数。
  3. 最终RMSE的计算逻辑变为:sqrt(mean((y_true - sample_from_distribution)^2)),这完全不是你需要的预测准确性指标。

三、解决方法:自定义适配概率分布的RMSE指标

你需要编写一个自定义指标,明确取分布的均值来计算RMSE,替换默认的RootMeanSquaredError():

class DistributionRMSE(keras.metrics.Metric):
    def __init__(self, name='distribution_rmse', **kwargs):
        super().__init__(name=name, **kwargs)
        self.mse = keras.metrics.MeanSquaredError()

    def update_state(self, y_true, y_pred, sample_weight=None):
        # 从分布对象中提取均值作为预测值
        y_pred_mean = y_pred.mean()
        self.mse.update_state(y_true, y_pred_mean, sample_weight)

    def result(self):
        return tf.sqrt(self.mse.result())

    def reset_state(self):
        self.mse.reset_state()

然后在编译模型时使用这个自定义指标:

bnn_model_probabilistic.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.01),
    loss=negative_loglikelihood,
    metrics=[DistributionRMSE()],
)

四、额外提示:获取模型的均值预测结果

如果你需要用model.predict()得到均值而非采样值,不要直接用默认返回值,而是调用分布的mean()方法:

# 获取测试集的分布预测结果
test_distributions = bnn_model_probabilistic.predict(feature_test)
# 提取均值作为预测值
test_pred_mean = test_distributions.mean()
# 手动计算RMSE的逻辑可以保持不变
BNN_residual_test = label_test - test_pred_mean
sq_BNN_residual_test = BNN_residual_test**2
BNN_MSE_test = np.mean(sq_BNN_residual_test)
rmse_test_bnn = np.sqrt(BNN_MSE_test)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 20:15:58