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()))
我的模型属于自编码器类型,需构建重构模型以对比查看误差,现咨询两个技术问题:
- 如何保存该模型?
- 后续如何加载该模型?
解决方案
针对你的自编码器模型保存与加载需求,结合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
相关产品推荐
相关产品推荐

