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

Keras训练中如何在单个epoch内定时保存模型权重?

解决Keras中每半小时保存CNN模型权重的问题

首先得指出你之前的思路误区:ModelCheckpoint里的period参数是用来指定每N个完整epoch保存一次的,而且必须是正整数,所以设置period=0.0125完全达不到你想要的效果——这个参数根本不支持按样本比例或时间间隔来触发保存。

针对你每半小时保存一次权重的需求,这里有两个实用的解决方案:

方案一:按训练步数保存(简单直接)

既然你每个epoch要跑8400步、耗时18小时,那半小时对应的步数大概是 8400 * (30/1080) ≈ 233 步。tf.keras的ModelCheckpoint提供了save_freq参数,可以直接指定每多少步保存一次权重,不用等整epoch。

修改后的代码如下:

cp_callback = tf.keras.callbacks.ModelCheckpoint(
    checkpoint_path,
    verbose=1,
    save_weights_only=True,
    save_freq=233  # 每233步保存一次权重
)
# 注意:如果你的TF版本较新,推荐用model.fit替代fit_generator
model.fit(
    training_set,
    steps_per_epoch=8400,
    epochs=25,
    callbacks=[cp_callback],
    validation_data=test_set,
    validation_steps=2165
)

这个方法的优点是实现简单,不需要额外代码;缺点是如果训练速度有波动(比如某些batch处理慢),实际的时间间隔可能会略有偏差。

方案二:自定义时间间隔回调(严格按时间触发)

如果想要严格按照每半小时保存一次,不管训练步数的快慢,你可以写一个自定义的Callback类,通过监测时间差来触发保存操作:

import time
from tensorflow.keras.callbacks import Callback

class TimeBasedCheckpoint(Callback):
    def __init__(self, filepath, interval_minutes=30, save_weights_only=True, verbose=1):
        super().__init__()
        self.filepath = filepath
        self.interval = interval_minutes * 60  # 转换为秒
        self.save_weights_only = save_weights_only
        self.verbose = verbose
        self.last_save_time = None

    def on_train_begin(self, logs=None):
        # 训练开始时先保存一次权重
        self.last_save_time = time.time()
        if self.verbose > 0:
            print(f"Initial weight save at {time.ctime()}")
        self.model.save_weights(self.filepath) if self.save_weights_only else self.model.save(self.filepath)

    def on_batch_end(self, batch, logs=None):
        current_time = time.time()
        # 检查是否达到时间间隔
        if current_time - self.last_save_time >= self.interval:
            if self.verbose > 0:
                print(f"\nSaving weights at {time.ctime()} (batch {batch})")
            self.model.save_weights(self.filepath) if self.save_weights_only else self.model.save(self.filepath)
            self.last_save_time = current_time

# 使用自定义回调
time_cp_callback = TimeBasedCheckpoint(
    checkpoint_path,
    interval_minutes=30,
    save_weights_only=True,
    verbose=1
)

model.fit(
    training_set,
    steps_per_epoch=8400,
    epochs=25,
    callbacks=[time_cp_callback],
    validation_data=test_set,
    validation_steps=2165
)

这个回调会在训练开始时先保存一次,之后每个batch结束时检查时间,只要距离上次保存超过30分钟,就自动保存权重,完全符合你的时间间隔需求。

额外小提示

  • 如果你使用的是较新版本的TensorFlow,建议用model.fit()替代fit_generator()——现在fit()已经完美支持生成器、Dataset等输入类型,用法和fit_generator()几乎一致。
  • 频繁保存权重可能会有轻微的IO开销,但半小时一次的频率完全在可接受范围内,不用担心影响训练效率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 08:08:51