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
相关产品推荐
相关产品推荐

