分块训练大数据集时,ModelCheckpoint能否保存全局最优模型?
分块训练下ModelCheckpoint的全局最优保存问题
你当前采用分块加载数据、逐块训练的方式处理大数据集,示例代码如下:
checkpoint = tf.keras.callbacks.ModelCheckpoint(filepath='blabla.h5', monitor='val_loss', mode='min', save_best_only=True) for file_nb in range(100000): data = pd.read_csv('a_path/to/my/datas/files' + str(file_nb)) history = model.fit(x=data[:,:3], y = data[:, -1] , calbacks=checkpoint)
关于ModelCheckpoint的默认行为
默认的ModelCheckpoint只会跟踪当前单次fit()调用(即当前数据块训练周期)内的指标变化,不会保留之前数据块的训练指标历史。也就是说,它仅会保存当前训练块中的最优轮次模型,完全不会识别此前训练块里的更优模型——如果当前块的最优指标不如之前块,它不会回滚保存之前的最优模型;若当前块的指标更差,甚至可能不生成新的模型文件,但原文件也不会自动保留历史最优版本。
实现全局最优模型保存的方法
要保存全局真正的最优训练模型,可通过以下两种方式实现:
1. 手动跟踪全局最优指标
手动维护全局的最优指标值,每次训练完一个数据块后,对比当前块的最优指标和全局最优,若更优则保存模型:
import shutil import pandas as pd import tensorflow as tf # 初始化全局最优val_loss(监控min模式,初始设为无穷大) global_best_val_loss = float('inf') global_best_model_path = 'global_best_model.h5' # 定义回调,保存当前块的最优模型到临时路径 temp_checkpoint = tf.keras.callbacks.ModelCheckpoint( filepath='temp_block_best.h5', monitor='val_loss', mode='min', save_best_only=True ) for file_nb in range(100000): data = pd.read_csv(f'a_path/to/my/datas/files{file_nb}') # 必须划分验证集,否则val_loss等于训练loss,失去监控意义 history = model.fit( x=data[:, :3], y=data[:, -1], callbacks=[temp_checkpoint], validation_split=0.2 # 或传入单独的验证集数据 ) # 获取当前块训练的最优val_loss current_block_best_val = min(history.history['val_loss']) # 对比全局最优,更新并保存 if current_block_best_val < global_best_val_loss: global_best_val_loss = current_block_best_val shutil.copy('temp_block_best.h5', global_best_model_path)
2. 自定义全局最优回调类
继承ModelCheckpoint,在回调内部维护全局的最优指标状态,让每次训练块时都能对比全局最优:
import tensorflow as tf from tensorflow.keras.callbacks import ModelCheckpoint class GlobalBestCheckpoint(ModelCheckpoint): def __init__(self, filepath, monitor='val_loss', mode='min', **kwargs): super().__init__(filepath, monitor=monitor, mode=mode, save_best_only=True, **kwargs) # 根据监控模式初始化全局最优值 self.global_best = float('inf') if mode == 'min' else -float('inf') def on_epoch_end(self, epoch, logs=None): logs = logs or {} current_metric = logs.get(self.monitor) if current_metric is None: print(f"警告:未监控到{self.monitor}指标,请确认训练时传入了验证集") return # 判断当前指标是否优于全局最优 if (self.mode == 'min' and current_metric < self.global_best) or \ (self.mode == 'max' and current_metric > self.global_best): self.global_best = current_metric # 保存当前最优模型 self.model.save(self.filepath) # 使用自定义回调 global_checkpoint = GlobalBestCheckpoint( filepath='global_best_model.h5', monitor='val_loss', mode='min' ) for file_nb in range(100000): data = pd.read_csv(f'a_path/to/my/datas/files{file_nb}') history = model.fit( x=data[:, :3], y=data[:, -1], callbacks=[global_checkpoint], validation_split=0.2 )
关键注意事项
- 必须在
fit()中传入验证集(通过validation_split或validation_data),否则val_loss指标不存在,监控毫无意义。 - 如果监控的是准确率这类需要最大化的指标,要将
mode设为max,同时初始化全局最优值为-float('inf')。
内容的提问来源于stack exchange,提问作者Jonathan Roy
相关产品推荐
相关产品推荐

