如何按自定义间隔保存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}")
使用方式
- 定义你需要保存的epoch列表:
target_epochs = [1, 2, 4, 8, 16, 32, 64, 128]
- 实例化自定义回调:
custom_checkpoint = CustomModelCheckpoint(target_epochs)
- 训练模型时传入回调:
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
相关产品推荐
相关产品推荐

