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

TF2.4中如何配置ModelCheckpoint以SavedModel格式保存完整模型

在TensorFlow 2.4中配置ModelCheckpoint保存SavedModel格式(save_weights_only=False时)

在TensorFlow 2.4版本中,当设置save_weights_only=False时,ModelCheckpoint默认会以HDF5(.h5)格式保存完整模型。要切换为SavedModel格式,只需在初始化回调时做好以下配置:

核心配置要点

  1. 指定save_format="tf"参数:这是切换格式的关键,明确告知回调以TensorFlow原生的SavedModel格式保存模型。
  2. 文件路径不要加.h5后缀:SavedModel是以文件夹形式存储的,若路径包含.h5或.hdf5后缀,即使设置了save_format="tf",仍可能被强制保存为HDF5格式(这是2.4版本的特定行为)。

代码示例

import tensorflow as tf
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Dense
from tensorflow.keras.callbacks import ModelCheckpoint

# 构建示例模型
model = Sequential([
    Dense(64, activation='relu', input_shape=(32,)),
    Dense(10, activation='softmax')
])
model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'])

# 配置ModelCheckpoint回调
checkpoint = ModelCheckpoint(
    filepath="./best_model",  # 路径为文件夹名,无.h5后缀
    save_weights_only=False,
    save_format="tf",  # 关键参数:指定SavedModel格式
    monitor="val_accuracy",
    save_best_only=True,
    verbose=1
)

# 模拟训练数据
x_train = tf.random.normal((1000, 32))
y_train = tf.random.uniform((1000,), maxval=10, dtype=tf.int32)
x_val = tf.random.normal((200, 32))
y_val = tf.random.uniform((200,), maxval=10, dtype=tf.int32)

# 训练并启用回调
model.fit(
    x_train, y_train,
    validation_data=(x_val, y_val),
    epochs=10,
    callbacks=[checkpoint]
)

加载SavedModel模型

训练完成后,可直接通过tf.keras.models.load_model()加载保存的模型:

loaded_model = tf.keras.models.load_model("./best_model")
loaded_model.evaluate(x_val, y_val)

额外注意

  • SavedModel保存的是一个文件夹,包含模型结构、权重、计算图签名等完整信息,比HDF5更适配TensorFlow生态的部署和迁移场景。
  • 若之前在同一路径下保存过HDF5格式的模型,建议先删除旧文件/文件夹,避免加载或保存时出现冲突。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 10:52:15