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

修改法英机器翻译Notebook后无法从Checkpoint恢复训练的解决方案咨询

Hey there! Let's work through this checkpoint resume issue together—since you're adapting the French-English machine translation notebook with your own dataset and extended training to 100 epochs, here are the exact code changes and configurations you need to make training pick up right where it left off:

1. Update Your Checkpoint Callback to Save Full State

The default checkpoint setup in many notebooks only saves model weights, but you need to capture the optimizer's state (like momentum and learning rate progress) and the current epoch number too. Replace your existing checkpoint callback with this:

import tensorflow as tf

# Define a checkpoint path that includes the epoch number
checkpoint_dir = "./mt_checkpoints"
checkpoint_path = f"{checkpoint_dir}/ckpt-{{epoch}}"

checkpoint_callback = tf.keras.callbacks.ModelCheckpoint(
    filepath=checkpoint_path,
    save_weights_only=False,  # Critical: saves model + optimizer state
    save_best_only=False,     # Save every epoch (adjust to save best if needed)
    save_freq="epoch",        # Trigger save at the end of each epoch
    verbose=1
)

Pro tip: Setting save_weights_only=False ensures all training state is preserved, not just the model's weights. This is the most common fix for resume issues.

2. Load the Latest Checkpoint Before Training

Add this block right before your model.fit() call to check for existing checkpoints and load the latest one:

import os

# Check if any checkpoints exist
latest_ckpt = tf.train.latest_checkpoint(checkpoint_dir)
if latest_ckpt:
    print(f"Resuming training from checkpoint: {latest_ckpt}")
    # Load the full model (architecture + weights + optimizer state)
    model = tf.keras.models.load_model(latest_ckpt)
    # Extract the last completed epoch from the checkpoint filename
    last_epoch = int(latest_ckpt.split("-")[-1])
else:
    print("No checkpoint found. Starting training from scratch.")
    last_epoch = 0

Note: Make sure you define your model architecture before loading the checkpoint, or use load_model() which will reconstruct the architecture automatically from the checkpoint file.

3. Set initial_epoch in model.fit()

Tell Keras exactly where to resume training by passing the last_epoch value to the initial_epoch parameter:

model.fit(
    train_dataset,
    epochs=100,
    initial_epoch=last_epoch,  # Starts training from the next epoch after last saved
    validation_data=val_dataset,
    callbacks=[checkpoint_callback]
)

Without this parameter, Keras will always start at epoch 1, even if you loaded a checkpoint.

4. Verify Learning Rate Scheduling (If Used)

If your notebook uses a learning rate scheduler (like ReduceLROnPlateau or a custom schedule), confirm it's restored correctly. After loading the checkpoint, print the current learning rate to check:

if latest_ckpt:
    print(f"Current learning rate: {model.optimizer.lr.numpy()}")

If it's reset to the initial value, you may need to save/load the scheduler state separately, but in most cases, save_weights_only=False handles this automatically.

5. Fix Common Pitfalls
  • Checkpoint path mismatches: Ensure the checkpoint_dir in your load code matches the path used in the ModelCheckpoint callback—typos here are a frequent culprit.
  • Overwriting the model: Don't redefine your model architecture after loading the checkpoint, as this will overwrite the loaded weights.
  • Old checkpoints with only weights: If you previously used save_weights_only=True, those checkpoints won't have optimizer state. Delete them and start fresh with the updated callback to generate full checkpoints.

Once you implement these steps, test by stopping training mid-run, restarting your script, and verifying it picks up at the correct epoch number. This should solve your issue of restarting from epoch 1 every time.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.28 23:59:05