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
相关产品推荐
相关产品推荐

