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

如何为TensorFlow Keras的CSVLogger生成的CSV文件添加参数与学习率?

解答:在TensorFlow Keras中扩展CSVLogger的输出内容

问题1:如何为CSVLogger生成的CSV文件添加自定义参数?

CSVLogger的核心逻辑是在每个epoch结束时,将logs字典中的键值对写入CSV文件。要添加自定义参数,只需要在CSVLogger执行前,把你的参数注入到logs字典中,具体可以通过自定义Callback实现:

  • 编写一个自定义Callback,在on_epoch_end(和CSVLogger的执行时机对齐)将参数添加到logs里
  • 训练时同时加载这个自定义Callback和CSVLogger,CSVLogger就会自动把新参数写入文件

问题2:如何将学习率写入CSVLogger的输出文件?

你已经实现了把学习率记录到训练历史中,但问题在于CSVLogger不直接读取model.history.history,而是依赖on_epoch_end时传入的logs字典。下面是修改后的完整方案:

步骤1:修改自定义Callback,将学习率注入到logs中

import tensorflow as tf

class LRCsvLogger(tf.keras.callbacks.Callback):
    def on_epoch_end(self, epoch, logs=None):
        logs = logs or {}
        # 获取当前优化器的学习率(如果用了学习率调度器,这里会自动获取更新后的值)
        current_lr = self.model.optimizer.lr.numpy()
        # 将学习率添加到logs字典,CSVLogger会自动读取这个键值对
        logs['lr'] = current_lr
        # 同时继续把学习率存入history,方便后续直接查看
        if 'lr' not in self.model.history.history:
            self.model.history.history['lr'] = []
        self.model.history.history['lr'].append(current_lr)

步骤2:训练时同时使用LRCsvLogger和CSVLogger

# 初始化CSVLogger,指定输出文件路径
csv_logger = tf.keras.callbacks.CSVLogger('training_logs.csv')
# 初始化自定义的学习率日志回调
lr_logger = LRCsvLogger()

# 训练模型时,将两个回调都加入callbacks列表
model.fit(
    x_train, y_train,
    epochs=10,
    validation_data=(x_val, y_val),
    callbacks=[csv_logger, lr_logger]
)

为什么原来的代码无法写入CSV?

你之前的回调是在on_epoch_begin中把学习率存入model.history.history,但:

  • CSVLogger的写入逻辑触发于on_epoch_end,它读取的是该方法传入的logs参数,而非model.history.history
  • model.history.history是Keras用来存储训练历史的独立容器,和CSVLogger依赖的logs字典是分离的

这样修改后,你的CSV文件里就会多出一列lr,记录每个epoch的学习率值了。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.09 06:37:47