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
相关产品推荐
相关产品推荐

