TensorFlow中model.fit返回的History损失曲线异常嘈杂问题
问题核心:对History对象中loss的记录维度误解
你遇到的问题根源在于:hist.history["loss"]存储的是每个训练batch的损失值,而非每个epoch的平均损失值,而你从日志中提取的loss是每个epoch结束时输出的该epoch所有batch损失的平均值,两者维度不同,自然曲线表现差异极大。
具体细节
- Keras的
model.fit()在默认配置下,History.history["loss"]会记录训练过程中每一个batch的损失,由于每个batch的样本分布存在差异,单个batch的损失会有明显波动,绘制出来的曲线就会显得嘈杂。 - 训练日志中默认输出的loss(比如
Epoch 1/300 - loss: 0.5234这类内容),是该epoch内所有batch损失的平均值,平滑了单个batch的波动,所以曲线更平稳。
验证方法
你可以打印len(hist.history["loss"])的值来验证:
print(len(hist.history["loss"]))
假设你的训练样本量len(xs)是10000,batch_size=100,那么每个epoch有100个batch;如果训练了100个epoch(EarlyStopping提前终止),这个长度会是100*100=10000,远大于epoch数,这就证明它是按batch记录的。
修正代码:绘制epoch级平均损失曲线
要得到和日志一致的平稳曲线,需要手动计算每个epoch的平均损失:
import numpy as np # 计算每个epoch的batch数量 batch_size = 100 steps_per_epoch = len(xs) // batch_size if len(xs) % batch_size != 0: steps_per_epoch += 1 # 将batch级loss按epoch分组,计算平均值 epoch_loss = [] total_steps = len(hist.history["loss"]) for i in range(total_steps // steps_per_epoch): start = i * steps_per_epoch end = start + steps_per_epoch epoch_loss.append(np.mean(hist.history["loss"][start:end])) # 同理处理验证集损失 epoch_val_loss = [] for i in range(len(hist.history["val_loss"]) // steps_per_epoch): start = i * steps_per_epoch end = start + steps_per_epoch epoch_val_loss.append(np.mean(hist.history["val_loss"][start:end])) # 绘制平稳的epoch级损失曲线 plt.figure(dpi=200) plt.plot(epoch_loss) plt.plot(epoch_val_loss) plt.legend(["loss", "val_loss"]) plt.xlabel("Epoch") plt.ylabel("Loss") plt.show()
更简洁的方案:自定义Callback记录epoch损失
如果不想手动计算,也可以通过自定义Callback直接记录每个epoch的平均损失:
from tensorflow.keras.callbacks import Callback class EpochLossLogger(Callback): def on_train_begin(self, logs=None): self.epoch_loss = [] self.epoch_val_loss = [] def on_epoch_end(self, epoch, logs=None): self.epoch_loss.append(logs["loss"]) self.epoch_val_loss.append(logs["val_loss"]) # 训练时添加该Callback loss_logger = EpochLossLogger() hist = model.fit( xs, ys, epochs=300, batch_size=100, validation_split=0.1, callbacks=[K.callbacks.EarlyStopping(patience=30), loss_logger] ) # 直接用logger中的数据绘图 plt.figure(dpi=200) plt.plot(loss_logger.epoch_loss) plt.plot(loss_logger.epoch_val_loss) plt.legend(["loss", "val_loss"]) plt.show()
内容的提问来源于stack exchange,提问作者Alberto
相关产品推荐
相关产品推荐

