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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 09:25:44