尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

模型训练过程中引入“智能早停”功能

模型训练过程中引入“智能早停”功能 EarlyStoppingCallback本质上是训练过程中的钩子函数用于在特定时机如每个 epoch 结束执行额外逻辑。训练代码虽然是手动循环但完全可以实现同样的功能无需依赖 Hugging FaceTrainer。方案仿照 Hugging Face 回调接口更灵活如果希望将来能轻松添加更多回调如学习率监控、内存优化等可以定义一个简单的回调基类然后在训练循环中调用相应钩子。这样就能直接使用你样例中的EarlyStoppingCallback需将其适配为你的接口。1. 定义回调基类和早停实现classTrainerCallback:defon_epoch_end(self,epoch,val_loss,best_loss,model,**kwargs):passclassEarlyStoppingCallback(TrainerCallback):def__init__(self,patience5,threshold0.002):self.patiencepatience self.thresholdthreshold self.counter0self.stop_trainingFalsedefon_epoch_end(self,epoch,val_loss,best_loss,model,**kwargs):ifval_lossbest_loss-self.threshold:self.counter0else:self.counter1ifself.counterself.patience:self.stop_trainingTrueprint(fEarly stopping at epoch{epoch1})2. 在train()中集成回调deftrain(...,callbacksNone):ifcallbacksisNone:callbacks[]min_val_lossfloat(inf)best_state_dictNoneforepochinrange(num_epochs):# ... validation ...ifepoch_val_lossmin_val_loss:min_val_lossepoch_val_loss best_state_dictdeepcopy(policy.state_dict())torch.save(best_state_dict,os.path.join(ckpt_dir,policy_best.ckpt))# 调用所有回调的 on_epoch_endstopFalseforcbincallbacks:cb.on_epoch_end(epoch,epoch_val_loss,min_val_loss,policy)ifgetattr(cb,stop_training,False):stopTrueifstop:break# ... training ...然后调用时创建回调对象并传入early_stopEarlyStoppingCallback(patience5,threshold0.002)train(...,callbacks[early_stop])优点扩展性强可轻松添加AggressiveMemoryOptimizationCallback等自定义回调无需改动核心循环。
返回列表