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

同会话内ANN模型训练后与存读结果差异大(训练模型MAE过高)

训练后直接预测ANN模型结果极差,保存再加载后恢复正常的排查方案

核心问题

训练完成后直接用内存中的模型预测,MAE极高;但将模型保存后重新加载,预测结果恢复正常。需排查内存中模型与加载后模型的状态差异。

排查方向与验证方案

1. 模型训练/推理模式未切换

这是此类问题最常见的原因:Keras中Dropout、BatchNormalization等层在训练和推理时行为完全不同——

  • Dropout:训练时随机失活神经元,推理时需关闭该逻辑;
  • BatchNormalization:训练时实时更新移动均值/方差,推理时需使用训练阶段累积的统计值。

训练完成后,模型默认不会自动切换到推理模式,直接预测时仍沿用训练逻辑,导致结果偏差极大;而模型保存再加载后,默认会进入推理模式。

验证方法:在第一个代码块的predict前,手动强制模型进入推理模式:

# 遍历所有模型,确保predict时使用推理模式
for model in xyz_forecast.modelframe_dict.values():
    # 重写predict方法,显式指定training=False
    def predict_with_inference_mode(x):
        return model.model_object.predict(x, training=False)
    model.model_object.predict = predict_with_inference_mode

# 再执行预测
result_df = xyz_forecast.predict(
    [test_df], "xyz_DEMAND_POT", predict_config=pipeline_config
)

如果结果恢复正常,说明问题出在模式未切换上,需在xyzDemandForecast的predict方法中固定添加training=False参数。

2. 内存中模型权重与加载后模型权重不一致

理论上训练后的模型权重应与保存再加载后的完全一致,若不一致则说明训练或保存过程存在问题。

验证方法:对比两者的权重:

import numpy as np

# 提取训练后内存中模型的权重
original_weights = []
for model in xyz_forecast.modelframe_dict.values():
    original_weights.append(model.model_object.get_weights())

# 保存并加载模型
save_model(f"_forecaster_model_{EXPERIMENT_NAME}", xyz_forecast)
loaded_forecast = load_model(f"_forecaster_model_{EXPERIMENT_NAME}")

# 提取加载后模型的权重
loaded_weights = []
for model in loaded_forecast.modelframe_dict.values():
    loaded_weights.append(model.model_object.get_weights())

# 逐权重对比
for orig_w, load_w in zip(original_weights, loaded_weights):
    for o, l in zip(orig_w, load_w):
        print(f"权重一致:{np.array_equal(o, l)}")

若输出存在False,说明保存/加载过程中权重丢失或篡改:

  • 当前save_model函数中使用tf.keras.models.clone_model复制模型结构,但该方法不会复制权重,之后将原模型设为None再替换为克隆模型,会导致原forecaster_object的模型变为空权重模型(仅在调用save_model后生效)。不过你的第一个场景未调用save_model,所以此问题不影响,但需确认训练过程中权重是否正确更新。

3. 数据预处理逻辑不一致

训练和预测阶段的数据预处理(如标准化、归一化)必须使用相同的统计量(如训练集的均值/标准差),若预测时错误使用测试集的统计量,会导致结果偏差。

排查点:

  • 检查xyzDemandForecast类是否保存了训练阶段的预处理统计量(如均值、方差);
  • 确认predict方法是否使用训练阶段的统计量,而非实时计算测试集的统计量;
  • 对比训练和预测阶段的输入数据分布,确保预处理后的特征分布一致。

4. xyzDemandForecast类内部状态异常

训练过程中,类的某些内部状态变量可能被修改,导致直接预测时使用错误配置;而保存再加载时,这些状态被重置为正确值。

排查点:

  • 检查fit方法是否修改了类的全局配置(如predict_config的参数);
  • 对比fit后和加载后的xyz_forecast对象的关键属性(如预处理配置、模型参数)是否一致。

总结

优先排查模型训练/推理模式问题,这是此类现象最常见的原因;若排除该问题,再依次验证权重一致性、预处理逻辑和类内部状态。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 22:01:04