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

