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

