如何用TensorFlow SavedModel API实现迭代式模型保存及版本管理?
用SavedModel API实现批次触发保存、版本管理与手动触发
当然可以!SavedModel API完全能替代tf.train.Saver()实现你要的所有功能——包括按批次间隔保存、限制保留的模型副本数,以及自定义版本命名格式。虽然它不像tf.train.Saver()那样直接提供max_to_keep参数,但我们可以通过手动维护版本号+清理旧模型的逻辑轻松实现,同时还能保持SavedModel的标准化优势。
下面是具体的实现思路和完整示例:
核心实现要点
- 版本号管理:维护一个全局步数/版本变量,每次保存时格式化三位数字(如
001、002)作为模型文件夹后缀 - 自动批次触发:在训练循环中检查当前批次是否达到间隔值,触发保存
- 手动触发保存:提供独立的保存函数,可在任意时机调用
- 旧模型清理:每次保存后,遍历模型目录,只保留最新的20个版本
完整示例代码
import tensorflow as tf import os import glob # 配置参数 SAVE_DIR = "./saved_model" MAX_TO_KEEP = 20 SAVE_EVERY_N_BATCHES = 100 # 每100个批次保存一次 # 初始化全局步数(可从检查点恢复,这里示例从零开始) global_step = tf.Variable(0, trainable=False, dtype=tf.int64) # 定义你的模型(示例用简单的全连接模型) class MyModel(tf.keras.Model): def __init__(self): super().__init__() self.dense1 = tf.keras.layers.Dense(64, activation='relu') self.dense2 = tf.keras.layers.Dense(10) def call(self, inputs): x = self.dense1(inputs) return self.dense2(x) model = MyModel() # 定义SignatureDef(所有模型共享同一签名) @tf.function(input_signature=[tf.TensorSpec(shape=(None, 32), dtype=tf.float32)]) def serving_fn(inputs): return model(inputs) signature_def_map = { tf.saved_model.DEFAULT_SERVING_SIGNATURE_DEF_KEY: tf.saved_model.signature_def_utils.predict_signature_def( inputs={'inputs': serving_fn.inputs[0]}, outputs={'outputs': serving_fn.outputs[0]} ) } def save_model(step): """保存模型并清理旧版本""" # 格式化版本号为三位数字 version_str = f"model-{step:03d}" save_path = os.path.join(SAVE_DIR, version_str) # 保存模型(复用统一的signature def) tf.saved_model.save( model, save_path, signatures=signature_def_map ) print(f"模型已保存至: {save_path}") # 清理旧模型:只保留最新的MAX_TO_KEEP个版本 # 遍历所有模型文件夹,提取版本号 model_dirs = glob.glob(os.path.join(SAVE_DIR, "model-*")) if not model_dirs: return # 按版本号排序(从旧到新) model_dirs.sort(key=lambda x: int(x.split("-")[-1])) # 删除超出数量的旧模型 while len(model_dirs) > MAX_TO_KEEP: old_dir = model_dirs.pop(0) # 删除文件夹及内容 tf.io.gfile.rmtree(old_dir) print(f"已删除旧模型: {old_dir}") # 模拟训练循环 for batch in range(1, 501): # 模拟训练步骤(这里省略实际训练逻辑) x = tf.random.normal((32, 32)) y = model(x) # 更新全局步数 global_step.assign_add(1) # 每N个批次自动保存 if batch % SAVE_EVERY_N_BATCHES == 0: save_model(global_step.numpy()) # 手动触发保存示例(比如训练结束后或用户触发时) save_model(global_step.numpy() + 1)
关键细节说明
- 版本命名:通过
f"{step:03d}"实现三位数字的版本号格式,确保文件夹排序正确 - 签名复用:提前定义好
signature_def_map,每次保存时传入,保证所有模型的签名一致 - 旧模型清理:使用
glob匹配模型文件夹,按版本号排序后删除超出数量的旧模型,这里用tf.io.gfile.rmtree是为了兼容TensorFlow的文件系统(包括本地和云存储) - 全局步数维护:使用
tf.Variable存储全局步数,方便后续从检查点恢复(如果需要断点续训,可以结合tf.train.Checkpoint来保存和恢复global_step和模型权重)
这样实现后,你就能完全替代tf.train.Saver()的功能,同时享受SavedModel API带来的标准化、跨平台部署优势。
内容的提问来源于stack exchange,提问作者bluesummers
相关产品推荐
相关产品推荐

