PyTorch调整max_steps后训练异常的代码修改方案咨询
PyTorch Lightning训练异常:loss飙升至100+且准确率为0的修复方案
我使用PyTorch Lightning训练模型时,原本设置max_steps=50000,训练到第5轮左右会报错终止;将max_steps调整为100000后,出现loss超过100、准确率(acc)为0的异常情况。模型代码如下:
class mymodel(pl.LightningModule): def __init__(self, config , learning_rate = 1e-4, max_steps = 100000//2): super(mymodel, self).__init__() self.config = config self.save_hyperparameters() self.training_losses = [] self.validation_losses = [] self.max_steps = max_steps def configure_optimizers(self): return torch.optim.AdamW(self.parameters(), lr = self.hparams['learning_rate']) def forward(self, batch_dict): return answer_vector def calculate_metrics(self, prediction, labels): batch_size = len(prediction) ac_score = 0 for (pred, gt) in zip(prediction, labels): ac_score+= calculate_acc_score(pred.detach().cpu(), gt.detach().cpu()) ac_score = ac_score/batch_size return ac_score def training_step(self, batch, batch_idx): answer_vector = self.forward(batch) loss = nn.CrossEntropyLoss()(answer_vector.reshape(-1,self.config['classes']), batch['answer'].reshape(-1)) _, preds = torch.max(answer_vector, dim = -1) train_acc = self.calculate_metrics(preds, batch['answer']) train_acc = torch.tensor(train_acc) return loss def validation_step(self, batch, batch_idx): logits = self.forward(batch) loss = nn.CrossEntropyLoss()(logits.reshape(-1,self.config['classes']), batch['answer'].reshape(-1)) _, preds = torch.max(logits, dim = -1) ## Validation Accuracy val_acc = self.calculate_metrics(preds.cpu(), batch['answer'].cpu()) val_acc = torch.tensor(val_acc) ## Logging self.log('val_ce_loss', loss, prog_bar = True) self.log('val_acc', val_acc, prog_bar = True) return {'val_loss': loss, 'val_acc': val_acc} def optimizer_step(self, epoch_nb, batch_nb, optimizer, optimizer_i, opt_closure = None, on_tpu=False, using_native_amp=False, using_lbfgs=False): ## Warmup for 1000 steps if self.trainer.global_step < 1000: lr_scale = min(1., float(self.trainer.global_step + 1) / 1000.) for pg in optimizer.param_groups: pg['lr'] = lr_scale * self.hparams.learning_rate ## Linear Decay else: for pg in optimizer.param_groups: pg['lr'] = polynomial(self.hparams.learning_rate, self.trainer.global_step, max_iter = self.max_steps) optimizer.step(opt_closure) optimizer.zero_grad()
问题核心分析
异常的根源是学习率调整逻辑错误,具体包括:
- 未定义
polynomial衰减函数或函数实现导致学习率突变(过大或趋近于0) max_steps修改后,衰减计算的基准值未适配新步数,导致学习率异常- 训练过程未记录关键指标,无法及时发现学习率问题
具体修改方案
1. 实现正确的多项式学习率衰减函数
添加多项式衰减实现,确保学习率随步数平滑下降:
def polynomial(base_lr, current_step, max_iter, power=1.0): """多项式学习率衰减,power=1时为线性衰减""" return base_lr * ((1 - float(current_step) / max_iter) ** power)
将此函数放在模型类外部,或作为类方法实现。
2. 修正optimizer_step中的学习率调整逻辑
确保max_steps使用self.hparams.max_steps(避免参数不一致),并调整衰减起始步数:
def optimizer_step(self, epoch_nb, batch_nb, optimizer, optimizer_i, opt_closure = None, on_tpu=False, using_native_amp=False, using_lbfgs=False): ## Warmup for 1000 steps if self.trainer.global_step < 1000: lr_scale = min(1., float(self.trainer.global_step + 1) / 1000.) for pg in optimizer.param_groups: pg['lr'] = lr_scale * self.hparams.learning_rate ## 多项式衰减(含线性衰减) else: # 从warmup结束后开始计算衰减步数 current_step = self.trainer.global_step - 1000 max_decay_steps = self.hparams.max_steps - 1000 if max_decay_steps <= 0: max_decay_steps = 1 # 避免除以0 for pg in optimizer.param_groups: pg['lr'] = polynomial(self.hparams.learning_rate, current_step, max_decay_steps) optimizer.step(opt_closure) optimizer.zero_grad()
3. 添加训练过程日志记录
在training_step中记录训练loss和acc,方便监控状态:
def training_step(self, batch, batch_idx): answer_vector = self.forward(batch) loss = nn.CrossEntropyLoss()(answer_vector.reshape(-1,self.config['classes']), batch['answer'].reshape(-1)) _, preds = torch.max(answer_vector, dim = -1) train_acc = self.calculate_metrics(preds, batch['answer']) train_acc = torch.tensor(train_acc, device=self.device) # 确保设备一致 # 记录训练指标 self.log('train_ce_loss', loss, prog_bar=True, logger=True) self.log('train_acc', train_acc, prog_bar=True, logger=True) return loss
4. 统一max_steps参数的使用
删除__init__中手动赋值的self.max_steps,直接使用self.hparams.max_steps:
def __init__(self, config , learning_rate = 1e-4, max_steps = 100000//2): super(mymodel, self).__init__() self.config = config self.save_hyperparameters() # 自动保存learning_rate、max_steps等参数 self.training_losses = [] self.validation_losses = [] # 移除self.max_steps = max_steps,统一用self.hparams.max_steps
5. 验证基础逻辑正确性
- 检查
calculate_acc_score函数实现,确保单样本准确率计算正确 - 验证
answer_vector与batch['answer']的维度匹配,避免损失计算时的形状错误
内容的提问来源于stack exchange,提问作者diamond
相关产品推荐
相关产品推荐

