TensorFlow从指定Checkpoint恢复部分变量至新模型的技术问询
刚好做过类似的需求,给你整理一个完整的实现方案,分TensorFlow 1.x和2.x两种情况,毕竟现在两个版本都有人用😎
完整实现方案
核心思路
先从Model1的Checkpoint中提取指定变量的值,再找到Model2中名称匹配的变量,将值赋值过去完成部分初始化。
TensorFlow 1.x 版本实现
步骤1:加载Model1的Checkpoint
首先用tf.train.load_checkpoint读取Checkpoint中的变量值:
import tensorflow as tf # 替换成你的Model1 Checkpoint路径(不需要后缀,比如.ckpt) ckpt_path = "path/to/model1_checkpoint" ckpt_reader = tf.train.load_checkpoint(ckpt_path)
步骤2:完善get_vars_by_name函数
这个函数用来从当前图中筛选出名称在目标列表里的变量,注意变量名称默认会带:0后缀,需要处理一下:
def get_vars_by_name(target_names): # 获取所有可训练变量 all_trainable_vars = tf.trainable_variables() # 筛选名称匹配的变量(去掉:0后缀) matched_vars = [] for var in all_trainable_vars: var_base_name = var.name.split(":")[0] if var_base_name in target_names: matched_vars.append(var) return matched_vars
步骤3:初始化Model2的指定参数
先定义好Model2的网络结构,然后执行赋值操作:
# 这里先定义你的Model2结构,比如: # ex1_model/fc2/b 对应Model2中的同名变量,确保名称一致 def build_model2(): inputs = tf.placeholder(tf.float32, shape=(None, 784)) fc1 = tf.layers.dense(inputs, 128, activation=tf.nn.relu, name="ex1_model/fc1") fc2 = tf.layers.dense(fc1, 10, name="ex1_model/fc2") # 和Model1的fc2对应 return inputs, fc2 # 构建Model2 model2_inputs, model2_outputs = build_model2() # 获取需要初始化的变量列表 var_list = ["ex1_model/fc2/b", "ex1_model/fc2/b/Adam", "ex1_model/fc2/b/Adam_1", "ex1_model/fc2/w", "ex1_model/fc2/w/Adam"] target_vars = get_vars_by_name(var_list) # 创建赋值操作 assign_ops = [] for var in target_vars: # 从Checkpoint中读取变量值 var_value = ckpt_reader.get_tensor(var.name.split(":")[0]) # 添加赋值操作 assign_ops.append(var.assign(var_value)) # 执行初始化 with tf.Session() as sess: # 先初始化Model2中其他未指定的变量 sess.run(tf.global_variables_initializer()) # 运行赋值操作,用Model1的参数覆盖指定变量 sess.run(assign_ops) # 保存初始化后的Model2 Checkpoint saver = tf.train.Saver() saver.save(sess, "path/to/initialized_model2")
TensorFlow 2.x 版本实现
TF2.x用动态图(Eager Execution),实现更简洁:
步骤1:定义并实例化Model2
import tensorflow as tf class Model2(tf.keras.Model): def __init__(self): super().__init__() # 注意这里的名称要和Model1对应,比如Model1是ex1_model/fc2,这里也要保持一致 self.fc1 = tf.keras.layers.Dense(128, activation="relu", name="ex1_model/fc1") self.fc2 = tf.keras.layers.Dense(10, name="ex1_model/fc2") # 对应Model1的fc2 def call(self, inputs): x = self.fc1(inputs) return self.fc2(x) # 实例化Model2,并调用一次以创建变量 model2 = Model2() _ = model2(tf.random.normal((1, 784)))
步骤2:加载Model1变量并赋值
# Model1的Checkpoint路径 ckpt_path = "path/to/model1_checkpoint" var_list = ["ex1_model/fc2/b", "ex1_model/fc2/b/Adam", "ex1_model/fc2/b/Adam_1", "ex1_model/fc2/w", "ex1_model/fc2/w/Adam"] for var_name in var_list: # 从Checkpoint中读取变量值 var_value = tf.train.load_variable(ckpt_path, var_name) # 找到Model2中对应的变量(去掉:0后缀) model2_var = [var for var in model2.trainable_variables if var.name.split(":")[0] == var_name][0] # 直接赋值 model2_var.assign(var_value) # 保存初始化后的权重 model2.save_weights("path/to/initialized_model2_weights")
注意事项
- 名称匹配:必须保证Model2中目标变量的名称(去掉
:0后缀)和Model1的完全一致,如果Model2的前缀不同(比如Model2是ex2_model),可以做名称映射,比如model2_var_name = var_name.replace("ex1_model", "ex2_model")。 - Adam优化器变量:你列表里包含了Adam的变量(
/Adam、/Adam_1),这些是优化器的动量和平方项,如果需要继续用Model1的优化器状态训练,就需要加载;如果只是初始化模型参数,可以只加载w和b。 - TF2.x的Checkpoint兼容:如果Model1是用TF1.x保存的,TF2.x依然可以用
tf.train.load_variable读取,不需要额外转换。
内容的提问来源于stack exchange,提问作者Gerry
相关产品推荐
相关产品推荐

