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

如何按自定义间隔保存TensorFlow神经网络模型?

实现指定Epoch列表的模型保存调度

完全可行,你可以通过自定义Keras回调函数来实现这种非均匀间隔的模型保存逻辑,无需依赖ModelCheckpoint的固定间隔设置。

核心实现思路

继承tf.keras.callbacks.Callback类,重写on_epoch_end方法——在每个epoch结束后,检查当前epoch是否在你指定的目标列表中,若是则触发模型保存。

代码示例

import tensorflow as tf

class CustomModelCheckpoint(tf.keras.callbacks.Callback):
    def __init__(self, save_epochs, save_path="model_epoch_{epoch}.h5"):
        super().__init__()
        self.save_epochs = set(save_epochs)  # 转为集合提升查找效率
        self.save_path = save_path

    def on_epoch_end(self, epoch, logs=None):
        # Keras的epoch从0开始计数,需+1匹配实际轮次
        current_epoch = epoch + 1
        if current_epoch in self.save_epochs:
            formatted_path = self.save_path.format(epoch=current_epoch)
            self.model.save(formatted_path)
            print(f"模型已保存至: {formatted_path}")

使用方式

  1. 定义你需要保存的epoch列表:
target_epochs = [1, 2, 4, 8, 16, 32, 64, 128]
  1. 实例化自定义回调:
custom_checkpoint = CustomModelCheckpoint(target_epochs)
  1. 训练模型时传入回调:
model.fit(
    x_train, y_train,
    epochs=200,  # 根据你的训练需求设置总轮次
    callbacks=[custom_checkpoint]
)

额外优化建议

  • 如果需要保存为TensorFlow原生的SavedModel格式,只需修改save_path为类似"model_epoch_{epoch}",并在save方法中指定格式:self.model.save(formatted_path, save_format="tf")。
  • 若想同时保留ModelCheckpoint的其他功能(如只保存最优权重),可以将自定义回调和ModelCheckpoint一起传入callbacks列表,两者互不冲突。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 04:45:41