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

PyTorch Lightning模型方法调用顺序验证及疑问

PyTorch Lightning训练/验证循环调用顺序疑问

我正尝试基于PyTorch Lightning重新实现训练流水线,原本理解的训练/验证循环方法调用顺序为:

  1. train_step()
  2. train_epoch_end()
  3. val_step()
  4. val_epoch_end()

为验证该顺序,我编写了如下示例代码:

import pytorch_lightning as pl
from torchmetrics import MeanMetric, SumMetric
from torch.utils.data import Dataset,DataLoader
import torch
import warnings
warnings.filterwarnings("ignore")

class DummyDataset(Dataset):
    def __init__(self):
        pass
    def __getitem__(self,idx):
        return torch.zeros([3,12,12]),torch.ones([3,12,12]) # Dummy image Like...
    def __len__(self):
        return 50

class DummyModel(pl.LightningModule):
    def __init__(self):
        super().__init__()
        self.conv = torch.nn.Conv2d(3,3,1,1) # Useless convolution
        self.mean = MeanMetric()
    def configure_optimizers(self):
        return torch.optim.Adam(self.parameters(),lr=1e-3)
    def training_step(self, batch,batch_idx):
        x,y=batch
        y_hat = self(x)
        loss = torch.sum((y-y_hat)**2)
        self.mean.update(2)
        return loss

    def training_epoch_end(self, outputs):
        mean_train = self.mean.compute()
        print(f"\nmean_train is : {mean_train}\n")
        self.mean.reset()

    def validation_step(self, batch,batch_idx):
        x,y=batch
        y_hat = self(x)
        loss = torch.sum((y-y_hat)**2)
        self.mean.update(4)
        return loss

    def validation_epoch_end(self, outputs):
        mean_val = self.mean.compute()
        print(f"\nmean_val is : {mean_val}\n")
        self.mean.reset()

    def forward(self,x):
        return self.conv(x)

if __name__=='__main__':
    dataset = DummyDataset()
    train_loader=DataLoader(dataset,batch_size=4,num_workers=0)
    val_loader=DataLoader(dataset,batch_size=4,num_workers=0)
    model = DummyModel()
    # We create trainer
    trainer = pl.Trainer(val_check_interval=None)
    # We fit model
    trainer.fit(model,train_dataloaders=train_loader,val_dataloaders=val_loader)

运行后输出为:

  • mean_val is : 3
  • mean_train is : nan

结合调试结果,实际调用顺序为:

  1. train_step()
  2. val_step()
    ...
  3. val_epoch_end()
  4. train_epoch_end()

请问:

  1. 实际情况是否如此?
  2. 我是否存在代码错误?
  3. 该机制的运行原理是什么?

解答

1. 实际调用顺序确实如此

PyTorch Lightning的默认执行逻辑是:完成整个训练epoch的所有training_step后,会先执行完整的验证流程(所有validation_step + validation_epoch_end),最后才执行training_epoch_end。这和你调试出来的顺序完全一致,你的初始理解是错误的。

2. 代码存在两个关键错误

(1)复用同一个Metric实例

你在__init__中只初始化了一个MeanMetric实例self.mean,同时用于训练和验证阶段的统计。由于验证阶段在训练epoch的training_step全部执行完后才开始,此时self.mean已经累积了训练阶段的所有更新值,验证阶段继续更新会导致数据混杂:

  • 训练阶段:50个样本,batch_size=4,共13个batch(12个满batch+1个2样本的batch),每个batch调用self.mean.update(2),累积13次2
  • 验证阶段:同样13个batch,每个调用self.mean.update(4),累积13次4
  • 验证阶段计算时,总共有26个值:(13*2 +13*4)/26 = 3,对应你看到的mean_val is :3
  • 而training_epoch_end执行时,self.mean已经被validation_epoch_end调用reset()清空,所以计算结果为nan

(2)val_check_interval=None的设置

这个参数设置为None时,PyTorch Lightning会在每个训练epoch结束后执行一次验证流程,正好触发了上述的顺序问题。如果设置为0-1之间的浮点数(比如0.5),则会在训练epoch进行到一半时就执行验证,但最终epoch结束时还是会先跑验证再执行training_epoch_end。

3. 运行原理

PyTorch Lightning的训练循环核心逻辑是将训练和验证流程解耦,但默认的epoch级执行顺序为:

  1. 遍历训练dataloader的所有batch,执行training_step,收集所有step的输出
  2. 遍历验证dataloader的所有batch,执行validation_step,收集所有step的输出
  3. 执行validation_epoch_end,处理验证step的输出
  4. 执行training_epoch_end,处理训练step的输出

这样设计的原因是:

  • 确保在训练epoch结束后先完成验证统计,方便用户在training_epoch_end中同时使用训练和验证的结果(比如记录对比日志、调整学习率等)
  • 保持验证流程的完整性,避免在训练epoch未完成时中断执行验证(除非手动设置val_check_interval)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 16:30:13