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

PyTorch Lightning中UNet模型Loss不下降问题求助

PyTorch Lightning训练UNet时Loss不更新的排查方案

我搭了个基础UNet模型,用自定义训练函数训练时优化效果正常,但改用PyTorch Lightning的training_step训练时,Loss一直停在初始值,模型预测完全没提升。已经按照教程去掉了zero_grad/backward/step相关代码,问题出在哪?

# 自定义训练函数,优化正常
def train(dataloader, model, loss_fn, optimizer):
    size = len(dataloader.dataset)
    model.train()
    for batch, (X, y) in enumerate(dataloader):
        X, y = X.to('cuda',dtype=torch.float), y.to('cuda',dtype=torch.float)

        # 计算预测误差
        pred = model(X)
        loss = loss_fn(pred, y)

        # 反向传播
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

# 集成到UNet类中的training_step,喂给pytorch_lightning.Trainer后Loss不更新
def training_step(self, batch, batch_idx):       
    X,y = batch
    X, y = X.to(self.device,dtype=torch.float), y.to(self.device,dtype=torch.float)

    # 计算预测误差
    pred = self.forward(X)
    loss = self.loss_fn(pred, y)
    self.log("train_loss", loss)
    return loss

核心排查与解决方法

  • 必须实现configure_optimizers方法
    PyTorch Lightning不会自动创建优化器,你必须在模型类里显式定义这个方法返回优化器——这是最常见的问题。你的自定义训练函数手动传入了optimizer,但PL框架需要你主动提供优化器实例,否则训练时根本没有参数更新逻辑。示例代码:

    def configure_optimizers(self):
        # 用和自定义训练时一致的优化器配置
        optimizer = torch.optim.Adam(self.parameters(), lr=1e-3)
        return optimizer
    
  • 检查模型参数是否被正确追踪
    确认UNet的所有子模块都是通过self.xxx = nn.Module(...)的方式定义的。如果是手动实现的层没继承nn.Module,或者没赋值给self属性,这些参数不会被加入优化器的更新列表,自然无法产生梯度变化。

  • 删除多余的设备迁移代码
    PL会自动将batch数据同步到模型所在的设备上,你手动写的X.to(self.device)完全没必要,甚至可能在多卡训练场景下导致数据与模型设备不匹配,进而让计算出的loss没有梯度。直接删掉这两行设备迁移代码即可。

  • 验证Loss函数的初始化逻辑
    检查self.loss_fn是不是在模型__init__方法中正确初始化的:如果是带可学习参数的自定义Loss,需要将其移到对应设备(比如self.loss_fn = loss_fn.to(self.device))。不过你自定义训练时没问题,这个排查优先级可以靠后。

  • 确认模型处于训练模式
    PL在training_step执行时会自动将模型设为train模式,但如果你的模型类其他方法(比如validation_step)手动调用了self.eval()且未切回,或者Trainer参数误设了model_mode='eval',也会导致梯度无法更新。可以在training_step开头加一行self.train()强制确认。

内容的提问来源于stack exchange,提问作者illan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 10:56:06