Pytorch Lightning训练NLP模型时ValueError问题求助
训练NLP模型续训时OneCycleLR触发ValueError问题排查与解决
问题描述
训练NLP模型时,续训到第50轮触发以下ValueError:
File "train/mrc_ner_trainer.py", line 431, in <module> if __name__ == '__main__': File "train/mrc_ner_trainer.py", line 418, in main File "/home/pkusam/anaconda3/envs/mrcNer3_7/lib/python3.7/site-packages/pytorch_lightning/trainer/states.py", line 48, in wrapped_fn result = fn(self, *args, **kwargs) File "/home/pkusam/anaconda3/envs/mrcNer3_7/lib/python3.7/site-packages/pytorch_lightning/trainer/trainer.py", line 1073, in fit results = self.accelerator_backend.train(model) File "/home/pkusam/anaconda3/envs/mrcNer3_7/lib/python3.7/site-packages/pytorch_lightning/accelerators/gpu_backend.py", line 51, in train results = self.trainer.run_pretrain_routine(model) File "/home/pkusam/anaconda3/envs/mrcNer3_7/lib/python3.7/site-packages/pytorch_lightning/trainer/trainer.py", line 1239, in run_pretrain_routine self.train() File "/home/pkusam/anaconda3/envs/mrcNer3_7/lib/python3.7/site-packages/pytorch_lightning/trainer/training_loop.py", line 394, in train self.run_training_epoch() File "/home/pkusam/anaconda3/envs/mrcNer3_7/lib/python3.7/site-packages/pytorch_lightning/trainer/training_loop.py", line 531, in run_training_epoch self.update_train_loop_lr_schedulers(monitor_metrics=monitor_metrics) File "/home/pkusam/anaconda3/envs/mrcNer3_7/lib/python3.7/site-packages/pytorch_lightning/trainer/training_loop.py", line 599, in update_train_loop_lr_schedulers self.update_learning_rates(interval='step', monitor_metrics=monitor_metrics) File "/home/pkusam/anaconda3/envs/mrcNer3_7/lib/python3.7/site-packages/pytorch_lightning/trainer/training_loop.py", line 1306, in update_learning_rates lr_scheduler['scheduler'].step() File "/home/pkusam/anaconda3/envs/mrcNer3_7/lib/python3.7/site-packages/torch/optim/lr_scheduler.py", line 154, in step values = self.get_lr() File "/home/pkusam/anaconda3/envs/mrcNer3_7/lib/python3.7/site-packages/torch/optim/lr_scheduler.py", line 1248, in get_lr .format(step_num + 1, self.total_steps)) ValueError: Tried to step 192652 times. The specified number of total steps is 192650
背景信息
- 此前已用相同代码训练至45轮并保存checkpoint,尝试从该checkpoint续训至55轮,第46-49轮无异常。
- 全新训练正常,仅续训时报错。
- OneCycleLR的
t_total计算方式:
t_total = (len(self.train_dataloader()) // (self.args.accumulate_grad_batches * num_gpus) + 1) * self.args.max_epochs
- 加载checkpoint的代码:
model = BertLabeling(args) if args.pretrained_checkpoint: model.load_state_dict(torch.load(args.pretrained_checkpoint, map_location=torch.device('cpu'))["state_dict"]) checkpoint_callback = ModelCheckpoint( filepath=args.default_root_dir, save_top_k=args.max_keep_ckpt, verbose=True, monitor="span_f1", period=-1, mode="max", ) trainer.fit(model)
BertLabeling的configure_optimizers函数:
def configure_optimizers(self): t_total = (len(self.train_dataloader()) // (self.args.accumulate_grad_batches * num_gpus) + 1) * self.args.max_epochs if self.args.lr_scheduler == "onecycle": scheduler = torch.optim.lr_scheduler.OneCycleLR( optimizer, max_lr=self.args.lr, pct_start=float(self.args.warmup_steps/t_total), final_div_factor=self.args.final_div_factor, total_steps=t_total, anneal_strategy='linear' ) # other scheduler.....
- 硬编码
t_total=192650*100仍报错,怀疑续训时加载了旧优化器配置。
解决方案
问题根源在于PyTorch Lightning默认会从checkpoint中恢复优化器和调度器的状态,包括旧的total_steps参数,导致续训时调度器的步数超过预设值。以下是重置优化器和调度器的具体方法:
1. 仅加载模型权重,跳过优化器/调度器状态
修改加载逻辑,只加载模型参数,不加载优化器和调度器的状态,让configure_optimizers重新初始化新的优化器和调度器:
model = BertLabeling(args) if args.pretrained_checkpoint: # 仅读取模型state_dict,忽略优化器相关状态 checkpoint = torch.load(args.pretrained_checkpoint, map_location=torch.device('cpu')) model.load_state_dict(checkpoint["state_dict"]) # 初始化Trainer时不要指定resume_from_checkpoint trainer = Trainer( gpus=num_gpus, accumulate_grad_batches=args.accumulate_grad_batches, max_epochs=55, # 设置总目标轮数 checkpoint_callback=checkpoint_callback, # 其他训练参数保持不变 ) trainer.fit(model)
2. 重新计算续训的t_total并跳过已训练步数
如果需要保留优化器的参数状态,需在configure_optimizers中重新计算总步数,并让调度器跳过已完成的训练步数:
def configure_optimizers(self): # 总目标轮数设为续训后的最终轮数(如55) total_epochs = self.args.max_epochs # 传入已训练的轮数(可通过命令行参数或checkpoint读取) trained_epochs = getattr(self.args, 'trained_epochs', 45) # 计算单轮训练步数 per_epoch_steps = len(self.train_dataloader()) // (self.args.accumulate_grad_batches * num_gpus) + 1 # 总步数基于最终轮数重新计算 t_total = per_epoch_steps * total_epochs # 已完成的训练步数 trained_steps = trained_epochs * per_epoch_steps optimizer = torch.optim.AdamW(model.parameters(), lr=self.args.lr) if self.args.lr_scheduler == "onecycle": scheduler = torch.optim.lr_scheduler.OneCycleLR( optimizer, max_lr=self.args.lr, pct_start=float(self.args.warmup_steps/t_total), final_div_factor=self.args.final_div_factor, total_steps=t_total, anneal_strategy='linear' ) # 跳过已训练的步数 if trained_steps > 0: scheduler.step(epoch=trained_steps) # 返回优化器和调度器 return [optimizer], [scheduler]
3. 使用Trainer的resume_from_checkpoint正确续训
若要完整恢复训练状态,需用Trainer的resume_from_checkpoint参数加载checkpoint,同时确保总步数匹配最终训练目标:
trainer = Trainer( resume_from_checkpoint=args.pretrained_checkpoint, gpus=num_gpus, accumulate_grad_batches=args.accumulate_grad_batches, max_epochs=55, # 设置最终训练轮数 checkpoint_callback=checkpoint_callback, ) trainer.fit(model)
此时需在configure_optimizers中动态调整t_total,基于最终轮数重新计算,覆盖checkpoint中的旧配置。
内容的提问来源于stack exchange,提问作者user19631495
相关产品推荐
相关产品推荐

