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

TensorFlow 2.0 如何加载模型并从最新检查点恢复训练

TensorFlow 2.0 从权重检查点恢复训练方案

核心实现逻辑

你当前使用的ModelCheckpoint配置为仅保存权重,恢复训练只需按以下步骤操作即可:

  • 保持模型结构、编译配置和原有代码完全一致,否则会出现权重维度不匹配报错
  • 训练前检查检查点目录下的已有权重文件,存在则加载到模型中
  • 调整model.fit的initial_epoch参数,传入已经完成训练的epoch数,避免重复训练

优化后可断点续训的完整代码

import tensorflow as tf 
from tensorflow.keras import models, layers
import matplotlib.pyplot as plt
from tensorflow.python.keras.metrics import acc
import datetime
from tensorflow.keras.callbacks import TensorBoard
import os

IMAGE_SIZE = 224
CHANNELS = 3

from tensorflow.keras.preprocessing.image import ImageDataGenerator

train_datagen = ImageDataGenerator(
    rescale=1./255,
    rotation_range=10,
    horizontal_flip=True
 )
train_generator = train_datagen.flow_from_directory(
    'data/train/',
    color_mode="rgb",
    target_size=(IMAGE_SIZE,IMAGE_SIZE),
    batch_size=32,
    class_mode="sparse",

)
print(train_generator.class_indices)

class_names = list(train_generator.class_indices.keys())
print(class_names)

validation_datagen = ImageDataGenerator(
    rescale=1./255,
    rotation_range=10,
    horizontal_flip=True)
validation_generator = validation_datagen.flow_from_directory(
    'data/validation/',
    target_size=(IMAGE_SIZE,IMAGE_SIZE),
    batch_size=32,
    class_mode="sparse"
)

test_datagen = ImageDataGenerator(
    rescale=1./255,
    rotation_range=10,
    horizontal_flip=True)

test_generator = test_datagen.flow_from_directory(
    'data/test/',
    target_size=(IMAGE_SIZE,IMAGE_SIZE),
    batch_size=32,
    class_mode="sparse"
 )

input_shape = (IMAGE_SIZE, IMAGE_SIZE, CHANNELS)
n_classes = 2

model = models.Sequential([
    layers.InputLayer(input_shape=input_shape),
    layers.Conv2D(32, kernel_size = (3,3), activation='relu'),
    layers.MaxPooling2D((2, 2)),
    layers.Conv2D(64,  kernel_size = (3,3), activation='relu'),
    layers.MaxPooling2D((2, 2)),
    layers.Conv2D(64,  kernel_size = (3,3), activation='relu'),
    layers.MaxPooling2D((2, 2)),
    layers.Conv2D(64, (3, 3), activation='relu'),
    layers.MaxPooling2D((2, 2)),
    layers.Conv2D(64, (3, 3), activation='relu'),
    layers.MaxPooling2D((2, 2)),
    layers.Conv2D(64, (3, 3), activation='relu'),
    layers.MaxPooling2D((2, 2)),
    layers.Flatten(),
    layers.Dense(64, activation='relu'),
    layers.Dense(n_classes, activation='softmax'),
])
model.summary()
model.compile(
    optimizer='adam',
    loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=False),
    metrics=['accuracy']
)

checkpoint_path = "teta/cp-{epoch:02d}.ckpt" # 加入epoch编号占位符,避免覆盖旧检查点
checkpoint_dir = os.path.dirname(checkpoint_path)
os.makedirs(checkpoint_dir, exist_ok=True) # 自动创建检查点目录
cp_callback = tf.keras.callbacks.ModelCheckpoint(
    filepath=checkpoint_path,
    save_weights_only=True,
    verbose=1
)

# ========== 新增:加载已有检查点 ==========
initial_epoch = 0
latest_checkpoint = tf.train.latest_checkpoint(checkpoint_dir)
if latest_checkpoint:
    print(f"加载已有的检查点:{latest_checkpoint}")
    model.load_weights(latest_checkpoint)
    # 从检查点文件名中提取已完成的epoch数
    initial_epoch = int(latest_checkpoint.split('-')[-1].split('.')[0])

# 总共需要训练的epoch数,比如要跑30轮就填30,会自动从上次结束的位置开始
TOTAL_EPOCHS = 30
history = model.fit(
    train_generator,
    steps_per_epoch=30,
    batch_size=32,
    validation_data=validation_generator,
    validation_steps=22,
    verbose=1,
    callbacks=[cp_callback],
    epochs=TOTAL_EPOCHS,
    initial_epoch=initial_epoch # 传入已完成的epoch数
)

注意事项

  • 如果你继续使用原有覆盖式的检查点命名(固定为cp.ckpt),则无法自动读取已完成的epoch数,需要手动把initial_epoch修改为你已经跑完的epoch数量
  • 如果需要完整恢复优化器的状态(比如Adam的动量参数),则需要把ModelCheckpoint的save_weights_only参数改为False,直接保存整个模型,加载时使用tf.keras.models.load_model加载完整模型即可
  • 恢复训练前不要修改模型结构、损失函数、优化器类型,否则会出现加载失败或者训练异常的问题

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 17:15:03