如何在SessionRunHook中使用tf.train.Saver?预训练子模型初始化失败求助
用预训练子模型初始化最终模型参数的解决方案
我之前也遇到过类似的问题——用预训练子模型拼接最终模型时,直接用SessionRunHook加载总是踩坑,后来换了更靠谱的方法,给你分享下:
核心问题分析
你遇到的报错大概率是变量名不匹配或者SessionRunHook的执行时机不对导致的:子模型训练时的变量名,和最终模型里对应组件的变量名(比如多了前缀)往往不一样,直接加载肯定会找不到变量;另外SessionRunHook如果时机没选对,也会在graph未完全构建好时就执行加载,导致失败。
方法一:用tf.train.init_from_checkpoint(官方推荐,最简洁)
这是TensorFlow专门为从预训练模型初始化参数设计的API,特别适合子模型组合的场景。
步骤:
- 创建变量名映射字典:把最终模型中对应子模型的变量前缀,映射到子模型ckpt里的变量前缀。比如子模型预训练时变量是
conv1/kernel,最终模型里这个子模型的组件变量是sub_model_1/conv1/kernel,那映射就是{"sub_model_1/": ""}(表示最终模型sub_model_1/下的变量,对应ckpt根路径的变量)。 - 在graph构建后调用初始化:在你搭建完最终模型的graph之后、启动训练之前,调用
tf.train.init_from_checkpoint传入每个子模型的ckpt路径和映射。
示例代码:
# 假设两个子模型的ckpt路径 sub_model_1_ckpt = "./sub_model_1.ckpt" sub_model_2_ckpt = "./sub_model_2.ckpt" # 构建最终模型(包含两个子模型组件) def build_final_model(): with tf.variable_scope("sub_model_1"): # 完全复用子模型1预训练时的网络结构 x = tf.layers.conv2d(inputs, 32, 3, name="conv1") ... with tf.variable_scope("sub_model_2"): # 完全复用子模型2预训练时的网络结构 y = tf.layers.dense(inputs, 64, name="fc1") ... # 最终模型的其他自定义层 ... build_final_model() # 为子模型1创建变量映射 init_map_1 = {"sub_model_1/": ""} # 最终模型前缀 -> ckpt前缀 tf.train.init_from_checkpoint(sub_model_1_ckpt, init_map_1) # 子模型2如果预训练时有自己的前缀(比如训练时用了"model/"),调整映射 init_map_2 = {"sub_model_2/": "model/"} tf.train.init_from_checkpoint(sub_model_2_ckpt, init_map_2) # 之后正常启动训练即可(比如用Estimator或Session)
方法二:手动加载变量(灵活应对复杂场景)
如果官方API满足不了你的特殊需求(比如需要对参数做额外处理),可以手动加载ckpt变量并赋值给最终模型。
步骤:
- 筛选并映射变量:获取最终模型中属于子模型的变量,创建变量名映射(把最终模型的变量名去掉前缀,对应到ckpt里的变量名)。
- 分批次初始化:先初始化不需要预训练的变量,再加载子模型的参数。
示例代码:
build_final_model() # 获取最终模型中sub_model_1的所有可训练变量 sub_model_1_vars = tf.get_collection(tf.GraphKeys.TRAINABLE_VARIABLES, scope="sub_model_1") # 构建映射:ckpt里的变量名 -> 最终模型的变量对象 var_map_1 = {var.name.replace("sub_model_1/", "").split(":")[0]: var for var in sub_model_1_vars} # 创建子模型1的Saver saver_1 = tf.train.Saver(var_map_1) # 启动Session加载参数 with tf.Session() as sess: # 先初始化非预训练的变量 non_pretrain_vars = [v for v in tf.global_variables() if v not in sub_model_1_vars + sub_model_2_vars] sess.run(tf.variables_initializer(non_pretrain_vars)) # 加载子模型1的参数 saver_1.restore(sess, sub_model_1_ckpt) # 同理加载子模型2的参数 ... # 开始训练 ...
如果你一定要用SessionRunHook
如果坚持要用SessionRunHook,得注意执行时机和变量映射:
- 必须在
after_create_session方法里执行restore(这个时机是Session创建后、训练前,graph已经完全构建)。 - 必须用正确的变量映射创建Saver。
示例Hook代码:
class SubModelLoadHook(tf.train.SessionRunHook): def __init__(self, ckpt_path, var_map): self.ckpt_path = ckpt_path self.var_map = var_map self.saver = tf.train.Saver(self.var_map) def after_create_session(self, session, coord): # 在这里执行加载操作 self.saver.restore(session, self.ckpt_path) # 使用时,把Hook传入Estimator的RunConfig var_map_1 = {var.name.replace("sub_model_1/", "").split(":")[0]: var for var in sub_model_1_vars} hook = SubModelLoadHook(sub_model_1_ckpt, var_map_1) config = tf.estimator.RunConfig(hooks=[hook]) estimator = tf.estimator.Estimator(model_fn=your_model_fn, config=config)
关键排查点
- 检查变量名匹配:用
tf.train.list_variables(ckpt_path)查看ckpt里的所有变量名,对比最终模型的变量名(可以用[var.name for var in tf.global_variables()]查看),确保映射正确。 - 避免重复初始化:不要用
tf.global_variables_initializer()初始化已经被预训练参数覆盖的变量,否则会把预训练值冲掉。 - 确认ckpt完整性:确保子模型的ckpt文件完整(包含
.index、.data-00000-of-00001等文件)。
内容的提问来源于stack exchange,提问作者Jason Zhou
相关产品推荐
相关产品推荐

