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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 10:20:35