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

如何在TensorFlow中加载最新检查点、续训并保存为完整模型?

TensorFlow加载最新检查点续训及保存完整模型

一、加载最新检查点从断点继续训练

你的代码中使用了save_weights_only=True的ModelCheckpoint,仅保存模型权重。续训时需先重建相同结构的模型,再加载最新权重后继续训练:

步骤及代码示例

  1. 获取最新检查点路径
    用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))
  1. 重建模型并加载权重
    必须创建和训练时完全一致的模型结构,再加载最新权重:
# 重建与训练时相同结构的模型
model = create_model(training, output)

# 加载最新检查点的权重
if latest_ckpt:
    model.load_weights(latest_ckpt)
    print(f"已加载最新检查点: {latest_ckpt}")
  1. 继续训练
    直接调用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)

二、加载检查点并保存为完整模型

若要将权重转换为可直接加载的完整模型,步骤如下:

步骤及代码示例

  1. 重建模型并加载权重
    和续训步骤一致,先创建相同结构的模型并加载最新权重:
model = create_model(training, output)
latest_ckpt = tf.train.latest_checkpoint(os.path.dirname(checkpoint_path))
model.load_weights(latest_ckpt)
  1. 保存为完整模型
    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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 13:41:39