TF2.4中如何配置ModelCheckpoint以SavedModel格式保存完整模型
在TensorFlow 2.4中配置ModelCheckpoint保存SavedModel格式(save_weights_only=False时)
在TensorFlow 2.4版本中,当设置save_weights_only=False时,ModelCheckpoint默认会以HDF5(.h5)格式保存完整模型。要切换为SavedModel格式,只需在初始化回调时做好以下配置:
核心配置要点
- 指定
save_format="tf"参数:这是切换格式的关键,明确告知回调以TensorFlow原生的SavedModel格式保存模型。 - 文件路径不要加
.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
相关产品推荐
相关产品推荐

