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

Tensorflow中fit()返回训练Loss与控制台打印值不一致的问题

问题描述

训练TensorFlow Recommenders模型时,fit()方法返回的训练loss与控制台打印的loss数值存在明显差异:

  • 训练过程中控制台实时打印的loss和total loss初始值接近7000(见训练过程截图)
  • 训练完成后查看fit返回的history字典,total_loss的初始值约为400(见返回结果截图)
  • 验证loss的数值在控制台打印和history字典中完全一致

训练代码如下:

# Fitting
model = RecommendationModel(output_layer_size=output_layer_size, hidden_layer_sizes=hidden_layer_sizes)
model.compile(run_eagerly=False,
              optimizer=tf.keras.optimizers.legacy.Adagrad(learning_rate=learning_rate)) # 0.001

model_fitted = model.fit(
        x=cached_train,
        epochs=epochs,
        verbose=True,
        batch_size=batch_size,
        validation_data=cached_val,
        callbacks=[
            tf.keras.callbacks.EarlyStopping(
                monitor='val_loss',  
                patience=1,          
                min_delta=1.0,
                mode='min',
                ),
            ]
    )

训练过程打印的loss
fit返回的loss字典
补充截图

原因分析与解决方向

这种差异的核心原因是控制台打印的是单批次计算的loss值,而fit返回的history字典中存储的是整个epoch内所有批次loss的平均值:

  • 初始训练批次的数据分布可能存在极端情况,导致单批次loss值偏高(接近7000),控制台会实时输出该批次的计算结果
  • history.history['total_loss']中记录的是每个epoch结束后,对所有训练批次的loss取平均后的结果,因此初始值远低于单批次的峰值(约400)
  • 验证loss是在每个epoch结束时对整个验证集计算的整体结果,不存在批次拆分的情况,所以控制台打印和history中的数值完全一致

如果需要在训练后查看每一批次的训练loss,可以通过自定义回调函数实现:

class BatchLossLogger(tf.keras.callbacks.Callback):
    def on_train_begin(self, logs=None):
        self.batch_losses = []

    def on_train_batch_end(self, batch, logs=None):
        self.batch_losses.append(logs['total_loss'])

# 在fit中添加该回调
batch_logger = BatchLossLogger()
model_fitted = model.fit(
        x=cached_train,
        epochs=epochs,
        verbose=True,
        batch_size=batch_size,
        validation_data=cached_val,
        callbacks=[
            tf.keras.callbacks.EarlyStopping(
                monitor='val_loss',  
                patience=1,          
                min_delta=1.0,
                mode='min',
                ),
            batch_logger
            ]
    )

# 查看所有训练批次的loss
print(batch_logger.batch_losses)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 05:00:02