如何在TensorFlow微调模型时延续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

