自定义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等默认参数和Kerasfit中的一致(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)的设置正确(你已经做了这步),但要注意:Kerasfit会自动处理这些层的移动均值/方差更新,而自定义循环中,只要在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
相关产品推荐
相关产品推荐

