如何通过MonitoredTrainingSession开启tf.estimator.Estimator式的global_step/sec日志
这个问题我之前也碰到过!其实原因很简单:global_step/sec这条日志并不是由LoggingTensorHook输出的,而是另一个专门的钩子——StepCounterHook负责的。tf.estimator.Estimator在训练流程里默认帮你加了这个钩子,所以自动有那行日志;而你手动用MonitoredTrainingSession的时候,需要自己显式添加它才行。
具体解决步骤
导入StepCounterHook
根据你的TensorFlow版本选择对应的导入方式:- TF1.x 环境:
from tensorflow.train import StepCounterHook - TF2.x 兼容模式(使用
tf.compat.v1):from tensorflow.compat.v1.train import StepCounterHook
- TF1.x 环境:
创建StepCounterHook实例
你可以通过every_n_steps参数控制日志的输出频率(默认每100步输出一次):# 每100步输出一次global_step/sec统计 step_counter_hook = StepCounterHook(every_n_steps=100)要是需要把统计数据写入TensorBoard,还可以指定
output_dir参数;如果只需要控制台日志,这个参数可以忽略。在MonitoredTrainingSession中同时传入两个钩子
将StepCounterHook和你已有的LoggingTensorHook一起放入hooks参数列表中:完整代码示例:
import tensorflow as tf # 根据TF版本导入对应的钩子 from tensorflow.compat.v1.train import LoggingTensorHook, StepCounterHook, MonitoredTrainingSession # 假设你已经定义了全局步数、损失、准确率等张量 global_step = tf.compat.v1.train.get_or_create_global_step() loss = ... # 你的损失张量 accuracy = ... # 你的准确率张量 # 定义LoggingTensorHook,输出你关心的自定义张量 logging_hook = LoggingTensorHook( tensors={'step': global_step, 'loss': loss, 'acc': accuracy}, every_n_steps=100 # 和StepCounterHook保持相同频率,日志更整齐 ) # 定义StepCounterHook,负责输出global_step/sec step_counter_hook = StepCounterHook(every_n_steps=100) # 启动训练会话,同时传入两个钩子 with MonitoredTrainingSession( checkpoint_dir='./checkpoints', hooks=[logging_hook, step_counter_hook] ) as sess: while not sess.should_stop(): sess.run(...) # 你的训练操作
效果验证
运行上述代码后,每到指定步数,你就会看到和tf.estimator.Estimator一样的两行日志:
INFO:tensorflow:global_step/sec: 1110.33
INFO:tensorflow:step = 9376, loss = 0.00026583532, acc = 0.999 (0.090 sec)
小提示:建议把两个钩子的every_n_steps设为相同值,这样日志输出会更同步整齐。
内容的提问来源于stack exchange,提问作者imhuay

