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

PyTorch Lightning加载模型后同验证集测试性能骤降排查求助

模型加载后PSNR骤降问题排查建议

问题现象

使用PyTorch Lightning训练模型后,加载最优checkpoint并使用训练阶段的同一验证集做sanity check,预期PSNR值应为37dB,但实际仅得到25dB。已确认模型结构、测试数据与训练时一致,且此前测试脚本运行正常,怀疑新增的自定义BatchNorm类是性能异常的根源。

相关代码

模型加载代码

checkpoint = "/.../my_location/model_name.ckpt"
model = LightningResunet2.load_from_checkpoint(checkpoint, strict=True)

Lightning主模型定义

class LightningResunet2(pl.LightningModule):
    def __init__(self):
        super().__init__()
        self.save_hyperparameters(logger=False)
        self.model = network(in_nc=6, out_nc=3, nc=96, nb=20, act_mode='BR')
        self.custom_loss = fixed_loss()

    def forward(self, x):
        prediction, noise_map = self.model(x)
        return prediction, noise_map

    def training_step(self, batch, batch_idx):
        noisy, clean_img, sigma_img, if_asym = batch
        output, predicted_noise = self.model(noisy)
        loss = self.custom_loss(output, clean_img, predicted_noise, sigma_img, if_asym)
        self.log("train_loss", loss, on_step=True, on_epoch=True, prog_bar=True, logger=True)
        return loss

    def configure_optimizers(self):
        optimizer = torch.optim.Adam(self.parameters(), lr=2e-4, eps=1e-8)
        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, 4000,
                                                                        eta_min=float(1e-6))
        scheduler.step()
        return [optimizer], [scheduler]

    def validation_step(self, batch, batch_idx):
        noisy, clean = batch
        output, _ = self.model(noisy)
        val_loss = F.l1_loss(output, clean)
        psnr = test_psnr(clean, output)
        
        self.log("validation_loss", val_loss, on_epoch=True, on_step=False, sync_dist=True, prog_bar=True, logger=True)
        self.log("validation_psnr", psnr, on_epoch=True, on_step=False, sync_dist=True, prog_bar=True, logger=True)
        return {'val_loss': val_loss, 'val_psnr': psnr}

    def test_step(self, batch, batch_idx):
        noisy, clean = batch
        output, _ = self.model(noisy)
        psnr = test_psnr(clean, output)
        self.log("validation_psnr", psnr, on_epoch=True, on_step=False, sync_dist=True, prog_bar=True, logger=True)
        return {'val_psnr': psnr}

自定义网络结构

class DnCNN_lightning(pl.LightningModule):
    def __init__(self, in_nc=6, out_nc=3, nc=64, nb=20, act_mode='BR'):
        super(DnCNN_lightning, self).__init__()
        self.save_hyperparameters(logger=False)
        assert 'R' in act_mode or 'L' in act_mode

        bias = False
        self.fcn = FCN()
        dncnn1 = dncnn_block(nc, nc, nc)
        head1 = conv(in_nc, nc, mode='C'+act_mode[-1], bias=bias, kernel_size=5,padding=2,)
        # head1 = conv(in_nc, nc, mode='C'+act_mode[-1], bias=bias)
        tail1 = conv(nc, out_nc, mode='C', bias=bias)
        self.model1 = sequential(head1, dncnn1, tail1)

    def forward(self, x, train_mode=True):
        noise_level = self.fcn(x)
        concat_img = torch.cat([x, noise_level], 1)
        level1_out = self.model1(concat_img) + x
        return level1_out, noise_level

自定义BatchNorm类

class BFBatchNorm2d(nn.BatchNorm2d):
    def __init__(self, num_features, eps=1e-5, momentum=0.1, use_bias = False, affine=True):
        super(BFBatchNorm2d, self).__init__(num_features, eps, momentum)
        self.use_bias = use_bias

    def forward(self, x):
        self._check_input_dim(x)
        y = x.transpose(0,1)
        return_shape = y.shape
        y = y.contiguous().view(x.size(1), -1)
        if self.use_bias:
            mu = y.mean(dim=1)
        sigma2 = y.var(dim=1)

        if self.training is not True:
            if self.use_bias:
                y = y - self.running_mean.view(-1, 1)
            y = y / ( self.running_var.view(-1, 1)**0.5 + self.eps)
        else:
            if self.track_running_stats is True:
                with torch.no_grad():
                    if self.use_bias:
                        self.running_mean = (1-self.momentum)*self.running_mean + self.momentum * mu
                    self.running_var = (1-self.momentum)*self.running_var + self.momentum * sigma2
            if self.use_bias:
                y = y - mu.view(-1,1)
            y = y / (sigma2.view(-1,1)**.5 + self.eps)

        if self.affine:
            y = self.weight.view(-1, 1) * y;
            if self.use_bias:
                y += self.bias.view(-1, 1)

        return y.view(return_shape).transpose(0,1)

排查建议

  • 强制切换模型到评估模式:加载checkpoint后必须调用model.eval(),并在推理时使用torch.no_grad()。自定义BatchNorm依赖self.training判断是否使用训练阶段统计的running_mean/running_var,若模型处于训练模式,会用当前batch的均值方差做归一化,导致结果异常。
  • 检查自定义BatchNorm的均值处理逻辑:
    • 当前代码中,当use_bias=False时,测试阶段不会减去running_mean,训练阶段也不会减去当前batch的均值mu,这和标准BatchNorm的行为完全不符。标准BatchNorm无论是否启用bias,都会执行均值减法操作,这会导致数据分布偏离训练时的状态,直接影响模型输出。
  • 验证方差计算方式:PyTorch中torch.var(dim=1)默认是有偏估计(除以N),而标准nn.BatchNorm2d使用的是无偏估计(除以N-1)。可修改为sigma2 = y.var(dim=1, unbiased=True),确保训练时的方差统计和标准BatchNorm一致。
  • 确认running stats的加载状态:打印自定义BatchNorm层的running_mean和running_var,检查是否与训练结束时的数值一致(而非初始的0和1)。若加载后仍是初始值,说明checkpoint未正确保存或加载这些统计量。
  • 对比标准BatchNorm的基准测试:临时将自定义BFBatchNorm2d替换为标准nn.BatchNorm2d,重新训练一个小型模型并测试加载后的PSNR,若恢复正常,即可确认自定义BatchNorm是问题根源。
  • 检查forward方法的train_mode参数:DnCNN_lightning的forward方法定义了train_mode=True参数但未实际使用,需确认该参数是否会影响模型内部的训练状态切换,避免与model.eval()的设置冲突。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 18:55:14