损失函数

L1 与 SSIM

在评估预测图像(Prediction)与目标图像(Ground Truth)之间的差距时,单一的损失函数往往是不够的。在此处,我们采用了 L1 Loss 和 SSIM Loss 的组合策略。根据论文设定,总损失函数表达为:

Loss = 0.8 \times L_{1} + 0.2 \times L_{SSIM}

为什么要这么组合?这得从它们各自的“性格”说起。

L1 计算的是像素到像素的绝对差异。它能很好地保证整体的亮度一致性,是最基础的保底指标。但它的缺点也很明显:为了追求整体误差最小,它倾向于把不确定的高频细节(比如精细的纹理、头发丝)进行“平均化”处理,导致画面边缘模糊。

为了弥补 L1 的不足,我们需要引入 SSIM。它不逐个死抠像素,而是模拟人类视觉系统的感知方式,将图像切分成一个个局部窗口,综合考量三个维度:

1. 亮度 (Luminance, l):比较两个窗口的平均像素亮度(均值 \mu)。

2. 对比度 (Contrast, c):比较两个窗口内像素的波动程度(方差 \sigma^2)。

3. 结构 (Structure, s):比较两个窗口内像素分布的相似性(协方差 \sigma_{xy})。

在 PyTorch 中,SSIM 本质上是一系列卷积操作。为了计算局部的均值和方差,底层会利用一个 1D 或 2D 的高斯核(Gaussian Kernel)在图像上滑动。这也是为什么 SSIM 的计算开销远大于 L1。


PSNR (峰值信噪比)

如果说 SSIM 是看重构图和光影的艺术评委,那么 PSNR 就是一个手里拿着游标卡尺、极其严谨的理工男。虽然我们在训练时不用它做 Loss,但它是评估模型最终效果不可或缺的客观指标。

PSNR 衡量的是“图像中有用信号的峰值强度”与“误差(噪声)强度”之间的比值。比值越大,说明图像失真越小。

一般来说,PSNR > 30 dB 就已经是高质量图像了;如果 > 40 dB,肉眼几乎无法分辨出与原图的区别。


代码实现

在实现方面,我们无需手写上述复杂的功能,直接站在巨人的肩膀上,利用 torchmetrics 库即可。

但这里有一个需要注意的地方。PyTorch 对图像数据有着严格的要求:必须是 (Batch_size, Channels, Height, Width),简称 BCHW。而我们外部读取的图像是 HWC 格式。我们需要进行维度重排。

  1. # 5. 使用 TorchMetrics 定义 SSIM 与混合损失函数
  2. from torchmetrics.functional.image import structural_similarity_index_measure as ssim
  3. from torchmetrics.functional.image import peak_signal_noise_ratio as psnr
  4.  
  5. def compute_loss(pred, target):
  6.     # 将形状从 (H, W, C) 转换为 (1, C, H, W) 以符合 TorchMetrics 的要求
  7.     pred_trans = pred.permute(2, 0, 1).unsqueeze(0)
  8.     target_trans = target.permute(2, 0, 1).unsqueeze(0)
  9.  
  10.     l1 = F.l1_loss(pred, target)
  11.     ssim_val = ssim(pred_trans, target_trans)
  12.     return 0.8 * l1 + 0.2 * (1.0 - ssim_val)