如何在TensorFlow中加载最新检查点、续训并保存为完整模型?
TensorFlow加载最新检查点续训及保存完整模型
一、加载最新检查点从断点继续训练
你的代码中使用了save_weights_only=True的ModelCheckpoint,仅保存模型权重。续训时需先重建相同结构的模型,再加载最新权重后继续训练:
步骤及代码示例
- 获取最新检查点路径
用tf.train.latest_checkpoint自动定位指定目录下的最新检查点:
import os import tensorflow as tf # 你的checkpoint_path格式需包含占位符(如"chatbot/training/checkpoints/ckpt-{epoch}") latest_ckpt = tf.train.latest_checkpoint(os.path.dirname(checkpoint_path))
- 重建模型并加载权重
必须创建和训练时完全一致的模型结构,再加载最新权重:
# 重建与训练时相同结构的模型 model = create_model(training, output) # 加载最新检查点的权重 if latest_ckpt: model.load_weights(latest_ckpt) print(f"已加载最新检查点: {latest_ckpt}")
- 继续训练
直接调用model.fit,保持总epochs数不变(比如你原设置的500),模型会从断点处继续训练:
# 复用原回调函数 cp_callback = tf.keras.callbacks.ModelCheckpoint( filepath=checkpoint_path, verbose=1, save_weights_only=True, save_freq=1*batch_size) tb_callback = tf.keras.callbacks.TensorBoard(log_dir="chatbot/training/logs", histogram_freq=1, update_freq=1, profile_batch=1) # 启动续训 model.fit(training, output, epochs=500, batch_size=batch_size, validation_data=(training, output), callbacks=[cp_callback, tb_callback], verbose=1)
二、加载检查点并保存为完整模型
若要将权重转换为可直接加载的完整模型,步骤如下:
步骤及代码示例
- 重建模型并加载权重
和续训步骤一致,先创建相同结构的模型并加载最新权重:
model = create_model(training, output) latest_ckpt = tf.train.latest_checkpoint(os.path.dirname(checkpoint_path)) model.load_weights(latest_ckpt)
- 保存为完整模型
TensorFlow支持两种主流格式:
- SavedModel格式(推荐):TensorFlow原生格式,支持跨平台部署
model.save("chatbot/training/full_model") - HDF5格式:单文件存储,适合本地备份
model.save("chatbot/training/full_model.h5")
关键注意事项
- 确保
create_model返回的模型结构与训练时完全一致,包括层数量、参数、输入输出形状,自定义层需正确注册。 checkpoint_path必须包含占位符(如{epoch}),否则tf.train.latest_checkpoint无法识别检查点文件。
内容的提问来源于stack exchange,提问作者Dweller
相关产品推荐
相关产品推荐

