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

TensorFlow Eager模式下如何存储与恢复自定义模型的可训练变量?

如何在TensorFlow Eager模式下仅存储和恢复可训练变量

我已经用TensorFlow Eager模式编写了一个自定义模型,希望仅存储和恢复模型中的可训练变量,之前在非Eager模式下我是这样实现的:

def store(self, sess_var, model_path):
    if model_path is not None:
        saver = tf.train.Saver(var_list=tf.trainable_variables())
        save_path = saver.save(sess_var, model_path)
        print("Model saved in path: %s" % save_path)
    else:
        print("Model path is None - Nothing to store")

# 对应的restore方法逻辑类似

请问在Eager模式下该怎么实现相同的功能?


在Eager模式下,TensorFlow推荐使用tf.train.Checkpoint(配合tf.train.CheckpointManager管理检查点文件)来实现变量的存储与恢复,完全可以对应你之前非Eager模式中仅保存可训练变量的需求。以下是具体的实现代码:

1. 存储可训练变量

你可以直接将所有可训练变量封装到tf.train.Checkpoint中,然后调用保存方法:

def store(self, model_path):
    if model_path is not None:
        # 创建Checkpoint对象,指定要保存的可训练变量集合
        checkpoint = tf.train.Checkpoint(trainable_vars=tf.trainable_variables())
        # 用CheckpointManager管理多个检查点,自动保留最新的3个(可按需调整)
        checkpoint_manager = tf.train.CheckpointManager(checkpoint, model_path, max_to_keep=3)
        save_path = checkpoint_manager.save()
        # 如果不需要管理多个检查点,直接调用checkpoint.save(model_path)也可以
        print("Model saved in path: %s" % save_path)
    else:
        print("Model path is None - Nothing to store")

2. 恢复可训练变量

恢复时需要创建结构一致的Checkpoint对象,再调用restore方法加载检查点:

def restore(self, model_path):
    if model_path is not None:
        # 同样映射到当前的可训练变量
        checkpoint = tf.train.Checkpoint(trainable_vars=tf.trainable_variables())
        # 获取路径下最新的检查点文件
        latest_checkpoint = tf.train.latest_checkpoint(model_path)
        if latest_checkpoint:
            # 加载检查点,expect_partial()用于忽略非目标变量的未恢复警告
            checkpoint.restore(latest_checkpoint).expect_partial()
            print("Model restored from path: %s" % latest_checkpoint)
        else:
            print("No checkpoint found in path: %s" % model_path)
    else:
        print("Model path is None - Nothing to restore")

补充说明

  • 如果你是在自定义模型类中实现这些方法,也可以直接将模型实例传入Checkpoint,比如checkpoint = tf.train.Checkpoint(model=self),这样会自动保存模型的所有可训练变量,效果和指定tf.trainable_variables()一致。
  • expect_partial()方法可以避免因模型存在非可训练变量而产生的警告,如果你确定检查点包含所有需要恢复的变量,也可以省略这个方法。
  • 相比旧的tf.train.Saver,tf.train.Checkpoint在Eager模式下更易用,支持自动追踪变量依赖,不需要手动维护变量列表。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 09:16:37