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

自定义Keras训练循环与Keras Fit结果不一致的原因排查

排查自定义Keras训练循环与Fit结果差异的建议

看起来你已经做了很多基础的确定性保障工作,这种小batch下的差异确实很让人困惑。结合你的代码和Keras的内部机制,给你几个具体的排查方向:

1. 确认损失函数的计算逻辑完全一致

你在自定义训练步骤中每次创建MeanSquaredError实例,虽然看起来没问题,但要明确Keras fit中使用的"mse"对应的是默认的损失归约方式:Reduction.SUM_OVER_BATCH_SIZE(即每个batch内样本的损失平均值)。建议你:

  • 将损失函数的定义移到training_step外部,避免重复创建可能带来的隐性差异:
    mseLoss = keras.losses.MeanSquaredError(reduction=keras.losses.Reduction.SUM_OVER_BATCH_SIZE)
    
    @tf.function
    def training_step(x, y, model, opt):
        with tf.GradientTape() as tape:
            predictions = model(x, training=True)
            loss = mseLoss(y, predictions)
        # ... 其余代码
    
  • 对比两种模式下第一个batch的损失值:在fit中用LambdaCallback打印每个batch的损失,和自定义循环的第一个step loss对比。如果初始损失就不一样,那大概率是损失计算或数据加载的问题。

2. 检查优化器的状态与参数一致性

Adam优化器依赖动量(m)和二阶动量(v)的累积状态,哪怕微小的初始化差异都会导致后续更新偏离:

  • 确认两种模式下优化器的所有参数完全一致:你现在只设置了learning_rate、beta_1、beta_2,要确保epsilon等默认参数和Keras fit中的一致(Keras Adam默认epsilon=1e-7)。
  • 打印模型和优化器的初始状态:在两种RUN_TYPE下,创建模型和优化器后,分别打印model.get_weights()和optimizer.get_weights(),确认初始权重和优化器状态完全相同。

3. 验证数据加载的顺序与批次划分

虽然你设置了shuffle=False,但要确保自定义循环的tf.data.Dataset和fit的批次划分完全一致:

  • 检查样本总数是否能被batch_size整除:如果不能,确认两种模式下对最后一个不完整batch的处理逻辑一致(tf.data.Dataset.batch()默认drop_remainder=False,fit也默认保留最后一个batch)。
  • 打印第一个batch的输入数据:在自定义循环中打印x_batch_train的前几个值,在fit中用LambdaCallback打印第一个batch的输入,确保两者完全相同。

4. 排查模型层的训练状态更新

如果你的DNN包含BatchNormalization、Dropout等依赖训练状态的层,自定义循环需要确保这些层的状态被正确更新:

  • 确认model(x, training=True)的设置正确(你已经做了这步),但要注意:Keras fit会自动处理这些层的移动均值/方差更新,而自定义循环中,只要在training=True下调用模型,这些状态也会被更新,但可以打印这些层的状态(比如model.layers[X].moving_mean)来对比两种模式下的更新是否一致。

5. 检查梯度计算与应用的正确性

自定义循环中梯度的计算和应用可能存在隐性问题:

  • 确保tape.gradient(loss, model.trainable_variables)捕获了所有可训练变量的梯度:可以打印grads的数量和对应变量的名称,和model.trainable_variables的数量对比,确认没有遗漏。
  • 尝试在自定义循环中手动复制fit的梯度更新逻辑:比如,打印优化器应用梯度前后的权重变化,和fit中对应batch后的权重变化对比。

6. 排除tf.function的编译差异

@tf.function的编译可能会带来一些隐性的行为变化:

  • 尝试去掉@tf.function装饰器,运行自定义循环看结果是否和fit一致(虽然速度会慢,但可以排除图编译带来的问题)。
  • 检查training_step中的print语句:tf.function中的print会被转换为图操作,可能不会实时输出,建议改用tf.print来确保打印的是图内的张量值。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 07:57:51