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

TensorFlow2中基于tf.GradientTape的模型保存与加载咨询

问题描述

我正在使用tf.GradientTape进行自编码器模型训练,目前已成功实现每个epoch保存检查点,训练代码如下:

with train_summary_writer.as_default():
    with tf.summary.record_if(True):
        for epoch in range(epochs):
            for train_id in range(train_start_id, train_end_id):
                batch_data_path= train_data_path + 'train_data_' + str(train_id).zfill(6) + ".npy"
                batch_data = np.load(data_path)
                batch_data = np.transpose(batch_data, (0, 2, 3, 1))
                x_inp = np.reshape(np.asarray(batch_data), [-1, 5, 5, 5, 3])
                train(loss, model, opt, x_inp)
                loss_values = loss(model, x_inp)
                reconstructed = np.reshape(model(x_inp), [1, sensor_n, sensor_n, scale_n])
                # if int(train_id) % 2000:
                tf.summary.scalar('loss',loss_values, step = train_id)
                tf.summary.image('original', tf.reshape(x_inp, (step_max, sensor_n, sensor_n, scale_n)), max_outputs=10, step=train_id)
                tf.summary.image('reconstructed', reconstructed, max_outputs=10, step=train_id)
                print("Epoch: {} ///// Step: {}/{} ===========================> Loss: {} ".format(epoch, train_id, train_end_id, loss_values))
                save_path = manager.save()
                print("Saved checkpoint for epoch {}: {}".format(epoch, save_path))
                print("loss : {}".format(loss_values.numpy()))

我的模型属于自编码器类型,需构建重构模型以对比查看误差,现咨询两个技术问题:

  1. 如何保存该模型?
  2. 后续如何加载该模型?

解决方案

针对你的自编码器模型保存与加载需求,结合TensorFlow的常用实践,给你两种主流方案,同时适配你已在使用的Checkpoint场景:

一、保存完整模型(推荐用于直接推理/部署)

TensorFlow支持两种常用的完整模型保存格式:SavedModel(官方推荐)和HDF5,两种方式都能保存模型的结构、权重和计算图,加载后可直接用于推理或继续训练。

1. SavedModel格式保存

这种格式是TensorFlow的原生格式,兼容性最好,支持跨平台部署(比如TensorFlow Serving、TensorFlow Lite)。你可以在训练结束后,或者每个epoch结束时调用model.save():

# 训练结束后保存完整模型到指定目录
model.save("autoencoder_saved_model")

# 或者在每个epoch循环内保存(保留每个epoch的模型版本)
model.save(f"./autoencoder_epoch_{epoch}")

2. HDF5格式保存

如果你习惯单文件存储,可以用HDF5格式,保存为.h5后缀的文件:

# 训练结束后保存为HDF5文件
model.save("autoencoder.h5")

# 每个epoch保存版本
model.save(f"./autoencoder_epoch_{epoch}.h5")

二、加载保存的模型

对应上面两种保存格式,加载方式也有所不同:

1. 加载SavedModel格式模型

from tensorflow.keras.models import load_model

# 加载SavedModel目录下的模型
loaded_model = load_model("autoencoder_saved_model")

# 验证加载成功:可以做一次推理测试
sample_input = np.random.rand(1, 5, 5, 5, 3)  # 和你的输入维度一致
reconstructed_sample = loaded_model(sample_input)
print("模型加载成功,已完成样本重构")

2. 加载HDF5格式模型

from tensorflow.keras.models import load_model

# 加载HDF5模型文件
loaded_model = load_model("autoencoder.h5")

注意:自定义组件的加载

如果你的模型包含自定义层、自定义损失函数或自定义指标,加载时需要通过custom_objects参数指定这些组件:

# 假设你有自定义层CustomLayer和自定义损失custom_loss
loaded_model = load_model("autoencoder.h5", 
                          custom_objects={"CustomLayer": CustomLayer, "custom_loss": custom_loss})

三、基于已有Checkpoint的加载(仅加载权重)

你目前已经在使用CheckpointManager保存检查点,这种方式默认保存的是模型的权重而非完整模型结构。如果要从检查点恢复,需要先重建模型结构,再加载权重:

# 1. 先重建和训练时完全一致的模型结构
model = build_your_autoencoder_model()  # 替换成你的模型构建函数

# 2. 初始化Checkpoint并关联模型
checkpoint = tf.train.Checkpoint(model=model)
manager = tf.train.CheckpointManager(checkpoint, directory="./checkpoints", max_to_keep=5)

# 3. 加载最新的检查点
if manager.latest_checkpoint:
    checkpoint.restore(manager.latest_checkpoint).expect_partial()  # expect_partial()避免未加载组件的警告
    print(f"已从检查点 {manager.latest_checkpoint} 恢复权重")
else:
    print("未找到检查点,将初始化新模型")

# 验证:查看模型权重或做推理测试
print(model.layers[0].weights)

内容的提问来源于stack exchange,提问作者Brian Lee

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 16:02:35