使用tf.train.MonitoredTrainingSession如何保存并获取global_step
解决MonitoredTrainingSession恢复global_step的问题
首先,咱们先明确tf.train.MonitoredTrainingSession的核心特性:它会自动处理检查点的保存与恢复流程,包括global_step变量——只要你是用tf.train.get_or_create_global_step()创建的这个变量,它默认就会被包含在检查点里。你现在遇到的恢复后拿不到global_step的问题,大概率是因为对MonitoredSession的恢复逻辑理解有误,下面给你针对性的解决步骤和修正代码:
关键注意点
- 不需要手动调用
saver.restore():MonitoredTrainingSession在启动时,会自动检测指定目录下的检查点,若存在则自动恢复所有已注册的变量(包括global_step)。 global_step必须是通过tf.train.get_or_create_global_step()创建的:这个方法会把变量注册到TensorFlow的全局变量集合,确保被CheckpointSaverHook捕获并保存。- 训练op要关联
global_step:确保优化器的minimize方法指定了global_step=global_step,这样训练时global_step才会自动递增,保存的检查点值才符合预期。
修正后的完整代码示例
import tensorflow as tf # 1. 在图构建阶段全局创建global_step global_step = tf.train.get_or_create_global_step() # 2. 定义模型、损失、优化器(示例) loss = tf.reduce_mean(tf.square(tf.random.normal([100]) - tf.random.normal([100]))) optimizer = tf.train.GradientDescentOptimizer(learning_rate=0.01) # 关键:将训练op与global_step关联,确保训练时自动递增 train_op = optimizer.minimize(loss, global_step=global_step) # 3. 创建CheckpointSaverHook,指定保存规则 save_checkpoint_hook = tf.train.CheckpointSaverHook( checkpoint_dir=checkpoints_abs_path, save_steps=5, checkpoint_basename=f"{checkpoints_prefix}.ckpt" ) # 4. 启动MonitoredTrainingSession with tf.train.MonitoredTrainingSession( master=server.target, is_chief=is_chief, hooks=[sync_replicas_hook, save_checkpoint_hook], config=config ) as session: # 会话启动后,若存在检查点,global_step已被自动恢复 restored_gstep = session.run(global_step) print(f"恢复后的初始global_step: {restored_gstep}") # 训练循环中正常执行训练,同步获取当前global_step while not session.should_stop(): _, current_gstep = session.run([train_op, global_step], feed_dict=feed_dict_train) print(f"当前训练global_step: {current_gstep}")
排查验证步骤
如果还是无法获取恢复后的global_step,可以做以下验证:
- 检查检查点是否真的保存了
global_step:用以下代码查看检查点内的变量
如果输出显示from tensorflow.python.tools.inspect_checkpoint import print_tensors_in_checkpoint_file print_tensors_in_checkpoint_file( tf.train.latest_checkpoint(checkpoints_abs_path), tensor_name='global_step', all_tensors=False )global_step存在,说明保存环节没问题,问题出在恢复后的获取逻辑;如果不存在,检查global_step的创建和注册是否正确。
内容的提问来源于stack exchange,提问作者chesschi
相关产品推荐
相关产品推荐

