如何为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
相关产品推荐
相关产品推荐

