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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 09:15:46