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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 03:51:24