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

如何用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)

关键细节说明

  1. 版本命名:通过f"{step:03d}"实现三位数字的版本号格式,确保文件夹排序正确
  2. 签名复用:提前定义好signature_def_map,每次保存时传入,保证所有模型的签名一致
  3. 旧模型清理:使用glob匹配模型文件夹,按版本号排序后删除超出数量的旧模型,这里用tf.io.gfile.rmtree是为了兼容TensorFlow的文件系统(包括本地和云存储)
  4. 全局步数维护:使用tf.Variable存储全局步数,方便后续从检查点恢复(如果需要断点续训,可以结合tf.train.Checkpoint来保存和恢复global_step和模型权重)

这样实现后,你就能完全替代tf.train.Saver()的功能,同时享受SavedModel API带来的标准化、跨平台部署优势。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 06:40:39