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

如何在fastai中提高验证环节的执行频率

fastai 步级验证与最佳模型保存回调实现

实现思路

fastai 原生SaveModelCallback默认仅在每个epoch结束后触发验证与模型保存逻辑,要实现每n步触发一次,我们可以通过继承SaveModelCallback自定义回调,在训练步结束的钩子中计数,达到指定步长阈值时主动触发验证流程,复用父类已有的最佳指标判定、模型持久化能力,避免重复开发。

完整代码实现

from fastai.callback.all import *
import numpy as np

class StepwiseSaveModelCallback(SaveModelCallback):
    def __init__(self, step_interval:int, monitor='valid_loss', comp=None, min_delta=0., 
                 with_opt=False, reset_on_fit=True, show_step_log:bool=True, **kwargs):
        super().__init__(monitor=monitor, comp=comp, min_delta=min_delta, 
                         with_opt=with_opt, reset_on_fit=reset_on_fit, **kwargs)
        self.step_interval = step_interval
        self.train_step_count = 0
        self.show_step_log = show_step_log
        # 未指定比较逻辑时根据监控指标自动判定:损失越小越好,其他指标默认越高越好
        if self.comp is None:
            self.comp = np.less if 'loss' in monitor else np.greater

    def after_batch(self):
        # 仅训练阶段统计步数
        if not self.training:
            return
        self.train_step_count += 1
        # 到达指定步长触发验证
        if self.train_step_count % self.step_interval == 0:
            # 主动运行验证集推理计算指标
            val_metrics = self.learn.validate()
            # 打印步级验证日志
            if self.show_step_log:
                metric_str = ' | '.join([f"{n}: {v:.4f}" for n,v in zip(self.learn.metrics_names, val_metrics)])
                print(f"Step {self.train_step_count} validation results: {metric_str}")
            # 复用父类逻辑判断是否为最优模型、执行保存
            self.after_epoch()

    def after_train(self):
        # 训练结束后自动加载最佳模型,和原生回调逻辑对齐
        if self.load_best:
            self.learn.load(self.fname, with_opt=self.with_opt, device=self.learn.device)

使用示例

创建回调实例后传入Learner的回调列表即可,示例为每100个训练步触发一次验证,跟踪验证集准确率,保存准确率最高的模型:

# 实例化步级保存回调
step_save_cb = StepwiseSaveModelCallback(
    step_interval=100, 
    monitor='accuracy',  # 填写要跟踪的指标名称,和传入Learner的metrics保持一致
    comp=np.greater,  # 指标越高越好用np.greater,越低越好用np.less
    load_best=True,  # 训练结束后自动加载最优权重
    with_opt=False  # 是否同步保存优化器状态,需要断点续训可设为True
)

# 回调传入Learner
learn = Learner(dls, model, metrics=[accuracy], cbs=[step_save_cb])

# 正常启动训练即可
learn.fit_one_cycle(10, 1e-3)

注意事项

  • 步长step_interval需要根据你的batch size、数据集规模调整,设得太小会导致验证过于频繁,大幅拖慢训练速度
  • 如果不需要保留默认的epoch级验证,可以在创建Learner时添加validate_every_epoch=False参数,避免重复运行验证浪费算力
  • 跟踪自定义指标时,monitor参数要和你传入Learner的metrics中的指标名称完全匹配

内容的提问来源于stack exchange,提问作者Aditya Bansal

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 02:06:08