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

如何在TensorFlow微调模型时延续loss/step信息?

如何在TensorFlow使用tf.train.Supervisor微调时延续基础模型的loss/step信息?

我在TensorFlow中使用tf.train.Supervisor对基础模型进行微调,相关代码如下:

sv = tf.train.Supervisor( 
    logdir=args.checkpoint_dir, 
    save_summaries_secs=args.summary_interval, 
    save_model_secs=args.checkpoint_interval, 
    init_fn=load_initial_weights_insess
) 
#using moving average to update loss and psnr 
update_ma = ema.apply([loss, psnr]) 

def load_initial_weights_insess(sess): 
    log.info("------------------------------------------") 
    if len(initial_ckpt) <= 0: 
        return 
    log.info('Prepare to load initial weights from {}'.format(initial_ckpt)) 
    log.info("Variables to load initial weights are:") 
    for v in initial_variables: 
        log.info("----{}".format(v.name)) 
    initial_saver.restore(sess, initial_ckpt) 
    log.info("Finished load initial weights") 

代码可以正常运行,但无法延续基础模型的loss/step信息——微调过程中这些信息总会被初始化,而不是从基础模型的最后状态延续。我担心loss初始值为0会导致不良的优化方向。

基础模型训练日志(末尾步骤)及微调日志如下:

#Training log for the basic model, the following log is from the end step:
Step 88689 | loss = 0.0551 | psnr = 30.8 dB
#Finetune log:
Step 0 | loss = 0.0 | psnr = 0.0 dB
Step 52 | loss = 0.0012 | psnr = 4.1 dB
Step 103 | loss = 0.0029 | psnr = 7.2 dB

请问如何在微调模型时延续loss/step信息?


问题根源分析

你遇到的问题核心在于:tf.train.Supervisor默认会初始化全局步骤(global step)、损失(loss)和PSNR的滑动平均变量,而你的init_fn只加载了模型权重,没有从基础模型的检查点中恢复这些非权重类的状态变量。

解决方案步骤

1. 明确需要恢复的额外变量

除了模型权重,你还需要恢复三类关键变量:

  • 全局步骤变量(控制日志中step计数的核心变量,通常是tf.train.get_global_step()对应的变量)
  • 损失和PSNR的滑动平均影子变量(由ema.apply()生成,对应日志里显示的loss/psnr值)
  • 若你的优化器有自身状态(如动量项),也需要一起恢复才能保证优化延续性

2. 修改初始化函数,加载完整训练状态

调整load_initial_weights_insess函数,创建包含所有需要恢复变量的saver,而不是只加载模型权重:

def load_initial_weights_insess(sess): 
    log.info("------------------------------------------") 
    if len(initial_ckpt) <= 0: 
        return 
    log.info('Prepare to load initial weights and training state from {}'.format(initial_ckpt)) 
    
    # 获取需要恢复的变量集合
    global_step = tf.train.get_global_step()
    # 模型权重变量(你原有的initial_variables)
    # EMA影子变量
    ema_restore_vars = ema.variables_to_restore()
    
    # 合并所有需要恢复的变量
    all_restore_vars = initial_variables + [global_step] + list(ema_restore_vars.values())
    
    # 创建覆盖所有变量的saver
    full_saver = tf.train.Saver(var_list=all_restore_vars)
    
    log.info("Variables to load:") 
    for v in all_restore_vars: 
        log.info("----{}".format(v.name)) 
    full_saver.restore(sess, initial_ckpt) 
    log.info("Finished loading initial weights and training state") 

3. 显式指定Supervisor的全局步骤

避免tf.train.Supervisor自动创建新的全局步骤变量,需要显式传入已定义的全局步骤:

global_step = tf.train.get_global_step()
sv = tf.train.Supervisor( 
    logdir=args.checkpoint_dir, 
    save_summaries_secs=args.summary_interval, 
    save_model_secs=args.checkpoint_interval, 
    init_fn=load_initial_weights_insess,
    global_step=global_step  # 显式绑定全局步骤变量
) 

4. 验证效果

完成修改后,微调时的step会从基础模型的最后一步(88689)开始递增,loss和PSNR也会从基础模型的最终值(0.0551、30.8dB)开始更新,不会再被初始化为0。

补充说明

  • 如果你的基础模型检查点已经包含了所有状态变量,直接使用tf.train.Saver()(不指定var_list)加载所有变量也可以,但显式指定变量列表能更精准地控制加载范围。
  • 若你不需要恢复EMA变量,只想延续step计数,仅恢复全局步骤变量即可,但loss的初始值需要你在微调前先计算一次基础模型在微调数据集上的loss,或者从检查点中恢复loss的滑动平均变量。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.27 09:32:42