损失函数
L1 与 SSIM
在评估预测图像(Prediction)与目标图像(Ground Truth)之间的差距时,单一的损失函数往往是不够的。在此处,我们采用了 L1 Loss 和 SSIM Loss 的组合策略。根据论文设定,总损失函数表达为:
为什么要这么组合?这得从它们各自的“性格”说起。
L1 计算的是像素到像素的绝对差异。它能很好地保证整体的亮度一致性,是最基础的保底指标。但它的缺点也很明显:为了追求整体误差最小,它倾向于把不确定的高频细节(比如精细的纹理、头发丝)进行“平均化”处理,导致画面边缘模糊。
为了弥补 L1 的不足,我们需要引入 SSIM。它不逐个死抠像素,而是模拟人类视觉系统的感知方式,将图像切分成一个个局部窗口,综合考量三个维度:
1. 亮度 (Luminance,
2. 对比度 (Contrast,
3. 结构 (Structure,
在 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 格式。我们需要进行维度重排。
- # 5. 使用 TorchMetrics 定义 SSIM 与混合损失函数
- from torchmetrics.functional.image import structural_similarity_index_measure as ssim
- from torchmetrics.functional.image import peak_signal_noise_ratio as psnr
- def compute_loss(pred, target):
- # 将形状从 (H, W, C) 转换为 (1, C, H, W) 以符合 TorchMetrics 的要求
- pred_trans = pred.permute(2, 0, 1).unsqueeze(0)
- target_trans = target.permute(2, 0, 1).unsqueeze(0)
- l1 = F.l1_loss(pred, target)
- ssim_val = ssim(pred_trans, target_trans)
- return 0.8 * l1 + 0.2 * (1.0 - ssim_val)