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

分块训练大数据集时,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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 13:47:41