TF2自定义训练循环下如何保存所有变量以恢复训练?
在TF2自定义训练循环中保存和恢复训练状态(含模型、优化器及自定义变量)
在TF2里写自定义训练循环时,确实会遇到这种困惑——TF1里的tf.train.Saver一键存所有变量,但TF2官方主推的Kerascompile/fit流程又用不上,别急,其实TF2专门给这种场景准备了tf.train.Checkpoint工具,完全能搞定保存模型参数、优化器状态和自定义变量的需求,而且根本不需要调用compile。
核心思路:用tf.train.Checkpoint追踪所有需要保存的对象
tf.train.Checkpoint可以自动追踪TensorFlow对象(比如tf.keras.Model、tf.keras.optimizers.Optimizer,以及自定义的tf.Variable)的状态,你只需要把所有需要保存的对象都传给它即可,不用手动列举每个变量。
1. 初始化Checkpoint对象
假设你有这些训练核心组件:
- 自定义模型
model = MyCustomModel() - 优化器
optimizer = tf.keras.optimizers.Adam(learning_rate=1e-3) - 记录训练步数的自定义变量
step = tf.Variable(0, dtype=tf.int64) - 其他需要保存的变量(比如累计损失)
total_loss = tf.Variable(0.0, dtype=tf.float32)
只需要把这些对象打包传给Checkpoint:
checkpoint = tf.train.Checkpoint( model=model, optimizer=optimizer, step=step, total_loss=total_loss )
2. 保存训练状态
在训练循环中,你可以在合适的时机(比如每N步、每轮结束后)调用保存方法:
# 比如每100步保存一次状态 if step % 100 == 0: # 指定保存路径,比如存到./checkpoints目录下 save_path = checkpoint.save("./checkpoints/training_checkpoint") print(f"Checkpoint saved to {save_path}")
如果想避免保存太多旧Checkpoint占用空间,可以用tf.train.CheckpointManager管理,只保留最近N个:
checkpoint_manager = tf.train.CheckpointManager( checkpoint, directory="./checkpoints", max_to_keep=5 # 仅保留最近5个Checkpoint ) # 保存时调用manager的save方法 if step % 100 == 0: save_path = checkpoint_manager.save() print(f"Step {step.numpy()}: Checkpoint saved to {save_path}")
3. 恢复训练状态
当需要重启训练时,先创建和之前完全一致结构的训练组件,再用Checkpoint加载最近的状态:
# 重新初始化所有训练组件 model = MyCustomModel() optimizer = tf.keras.optimizers.Adam(learning_rate=1e-3) step = tf.Variable(0, dtype=tf.int64) total_loss = tf.Variable(0.0, dtype=tf.float32) # 必须创建和保存时结构完全一致的Checkpoint对象 checkpoint = tf.train.Checkpoint( model=model, optimizer=optimizer, step=step, total_loss=total_loss ) # 加载最新的Checkpoint latest_checkpoint = tf.train.latest_checkpoint("./checkpoints") if latest_checkpoint: # .expect_partial()用于避免部分未匹配对象的警告(可选,按需使用) checkpoint.restore(latest_checkpoint).expect_partial() print(f"Restored from checkpoint: {latest_checkpoint}") print(f"Resuming training from step {step.numpy()}") else: print("No checkpoint found, starting training from scratch")
关键注意事项
- 对象结构要完全一致:保存和恢复时,Checkpoint里的键名(比如
model、optimizer)必须完全相同,模型的层结构也不能修改,否则无法正确恢复。 - 自定义变量要显式传入:所有需要保存的自定义
tf.Variable都要加到Checkpoint里,不然不会被保存。 - 完全脱离compile流程:整个保存/恢复过程不需要调用
model.compile(),完美适配自定义训练循环场景。
完整示例代码片段
import tensorflow as tf # 自定义模型示例 class MyCustomModel(tf.keras.Model): def __init__(self): super().__init__() self.dense1 = tf.keras.layers.Dense(64, activation='relu') self.dense2 = tf.keras.layers.Dense(10) def call(self, x): x = self.dense1(x) return self.dense2(x) # 初始化训练组件 model = MyCustomModel() optimizer = tf.keras.optimizers.Adam(1e-3) step = tf.Variable(0, dtype=tf.int64) total_loss = tf.Variable(0.0, dtype=tf.float32) # 初始化Checkpoint和管理器 checkpoint = tf.train.Checkpoint(model=model, optimizer=optimizer, step=step, total_loss=total_loss) checkpoint_manager = tf.train.CheckpointManager(checkpoint, "./checkpoints", max_to_keep=5) # 尝试恢复之前的训练状态 latest_checkpoint = tf.train.latest_checkpoint("./checkpoints") if latest_checkpoint: checkpoint.restore(latest_checkpoint).expect_partial() print(f"Resumed from checkpoint: {latest_checkpoint}, starting at step {step.numpy()}") # 自定义训练循环 while step < 1000: # 模拟训练步骤 with tf.GradientTape() as tape: x = tf.random.normal((32, 10)) y_pred = model(x) y_true = tf.random.uniform((32, 10), maxval=10, dtype=tf.int32) loss = tf.keras.losses.sparse_categorical_crossentropy(y_true, y_pred, from_logits=True) loss = tf.reduce_mean(loss) grads = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables)) step.assign_add(1) total_loss.assign_add(loss) # 每100步保存一次 if step % 100 == 0: save_path = checkpoint_manager.save() print(f"Step {step.numpy()}: Checkpoint saved to {save_path}, total loss: {total_loss.numpy()}")
这样不管中途停止多少次,再次运行代码时都会自动恢复到最近的训练状态——包括模型参数、优化器的动量/学习率状态,甚至你自定义的训练步数、累计损失等变量。
内容的提问来源于stack exchange,提问作者user209974
相关产品推荐
相关产品推荐

