前阵子调一个去雾模型的训练,日志里的PSNR看着还行,但把测试图片放大看,树叶边缘和窗框全是糊的,整体亮度是对了,细节却像被橡皮擦磨过一样。这个痛点几乎每个做图像复原的人都会碰上:端到端网络在抑制噪声的同时,把高频细节也一起抹掉了。后来我读到DEA-Net这篇工作,它提出的细节增强卷积(Detail Enhancement Convolution,DEC)正好打在“细节丢失”这个七寸上。这周我把DEC模块用PyTorch完整复现了一遍,顺手接进了一个mini去雾网络跑通了训练。这篇文章就完整记录整个复现过程,包括模块原理、逐行代码、训练配置和踩过的坑。适合已经会用PyTorch搭基础CNN、想了解图像去雾进阶模块的读者,也可以当作一个从模块设计到落地的完整案例来参考。
1. 去雾模型最容易翻车的环节:高频细节是怎么在特征提取中丢掉的
1.1 从大气散射模型看端到端去雾的本质
图像去雾问题的起点是一张雾天图像,在计算机视觉里通常用大气散射模型描述:
I(x) = J(x) * t(x) + A * (1 - t(x))
其中I(x)是有雾图像,J(x)是我们想恢复的清晰图像,t(x)是透射率,A是全局大气光。透射率和景深相关,t(x) = exp(-β * d(x)),β是大气散射系数,d(x)是场景深度。雾越浓、物体离相机越远,t(x)越小,J(x)的信号被压得越低,场景信息基本被A淹没。
传统方法像暗通道先验(DCP),会先估算t和A,再根据物理公式反推J。这类方法对天空区域、白色物体等不满足先验假设的场景很容易翻车,恢复出来的图经常有色偏和光晕。后来主流做法变成用CNN直接学习从I到J的映射,也就是端到端去雾。端到端的好处是不再依赖手工先验,数据够多的情况下效果稳定得多。
但端到端网络也有自己的毛病,最典型的就是:整体亮度、颜色恢复得很干净,可图像里的高频细节——树叶脉络、窗棂边缘、织物纹理——总是差点意思。原因要从特征提取和重建的过程中找。
1.2 DEA-Net的两个核心武器:DEC与上下文引导
DEA-Net整体是一个编码器-解码器结构的去雾网络,它把注意力放在了两件事上:一个叫细节增强卷积(DEC),专门解决“细节在特征提取时被磨平”的问题;另一个叫上下文引导模块(Contextual Guidance,CG),负责扩大感受野,让去雾决策不只看局部。
这两个模块的分工很明确。CG负责全局信息,让网络明白“哪里是天空、哪里是近景、雾的浓度大致是什么分布”;DEC负责局部细节,在特征提取阶段就把边缘和纹理信息加强。只做全局增强的模型,结果往往大块颜色对但边缘糊;只做局部增强的模型,边缘立起来了但整体雾感去不干净。两者配合,才是DEA-Net效果扎实的原因。
这次文章只聚焦DEC,因为它是一个相对独立的模块,可以单独复现、单独验证,也能直接嵌到其他复原网络里用。CG模块我放到最后简单提一句扩展方向。
1.3 DEC模块计算流程一句话版
DEC的完整流程可以压缩成一句话:输入特征x分别走一条普通卷积分支和一条可变形卷积分支,可变形卷积分支的输出经过注意力门控后,去调制普通分支的基础特征,最后加上残差连接。
这句话里有三个关键组件:普通卷积、可变形卷积、注意力门控。想把这个模块真正写对,得先把这三件事的来龙去脉搞清楚,尤其是可变形卷积——它是DEC的性能上限所在。下面一节就对着这三个概念逐个拆。
2. 动手前必须搞清楚的三个概念:可变形卷积、offset通道数与注意力门控
2.1 可变形卷积:让卷积核学会“看哪里”
普通3x3卷积在特征图的每个位置做计算时,采样点是固定的九宫格:左上、正上、右上、正左、中心、正右……排列非常规整。这种固定网格在处理语义规则的对象时没问题,但面对雾天图像里的弱边缘、不规则纹理,规整采样往往“够不着”那些最关键的像素。
可变形卷积在采样方式上多学了一组偏移量。对输出特征图上的每个点,网络额外预测一个offset,这个offset告诉卷积核:九宫格里的9个采样点,每一个需要往哪个方向偏移多少。于是采样点不再死板地排列成正方形,而是可以根据内容“流动”起来,聚集到物体边缘、纹理密集区这些真正有用的位置。
offset通常不是整数,所以带偏移的采样坐标会落在像素之间的位置,需要用双线性插值取值:x(p) = Σ q G(q, p) * x(q)。这里G就是双线性插值核。这个操作对offset是可导的,所以偏移量可以由梯度反向传播端到端学出来。换句话说,网络自己学会“该看哪里”,不需要人工标注。
打个比方,普通卷积像一台机位固定的摄影机,拍什么角度早就定死了;可变形卷积像带云台追踪的摄影机,画面里哪里有动作,镜头就自动跟过去。对去雾来说,雾霾对不同深度物体的影响非常不均匀,远处的细节被压得很弱,固定采样很难感知到这些弱信号,可变形卷积这种“主动聚焦”能力就特别对症。
2.2 offset通道数为什么是 2kHkW
实现可变形卷积时最容易报错的地方就是offset的通道维度。一个3x3卷积核有9个采样点,每个采样点需要两个方向的偏移量——水平方向dx和垂直方向dy——所以offset的总通道数是18,也就是 2 * 3 * 3 = 18。
很多第一次写的人会顺手把offset卷积的输出通道设成9,只算了采样点个数,忘了每个点有dx和dy两个量。这个错误直接导致torchvision.ops里的deform_conv2d报shape mismatch,输入输出对不上。
我在DEA-Net复现里用的就是3x3可变形卷积,所以offset_conv的输出通道固定是18。如果你把kernel_size改成5x5,那这里就是2 * 5 * 5 = 50,依此类推。写代码时我会在注释里把这个式子标清楚,防止日后忘了。
2.3 DEC里的注意力门控到底在干什么
DEC里的可变形卷积输出一张detail_feat,如果直接把这个特征加到主路上,效果不是最好的。DEA-Net的写法是:让detail_feat经过一个1x1卷积、BatchNorm,再接Sigmoid,输出一个0到1之间的门控值,用这个门控去和基础特征做逐元素乘法。
这个设计的含义是:可变形卷积分支学到的不是“要叠加的细节增量”,而是一张“细节注意力图”。它告诉网络哪些空间位置、哪些通道上存在值得放大的细节。基础特征与门控相乘后,细节丰富的区域被保留甚至放大,平滑区域被抑制。和直接相加相比,乘法调制不会大幅改变特征的数值分布,训练过程更稳。
这里有一个经验:不要把detail_feat直接加到输出上,虽然这种加法变体也能跑,而且有些人实测在某些数据集上还不差,但它破坏了DEC原设计的稳定性。第一次复现时建议先按乘法调制来,跑通了再改着玩。
3. PyTorch逐行实现DEC:从offset预测到残差融合
3.1 环境准备与torchvision版本检查
DEC里需要可变形卷积,我用的是torchvision.ops.DeformConv2d,这个接口从torchvision 0.9开始就有了,建议至少0.13以上,接口更稳定。安装命令很简单:
pip install torch torchvision如果机器是NVIDIA GPU,建议按PyTorch官网的CUDA版本提示安装,比如CUDA 11.8对应:
pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118装完之后先确认接口存在:
import torch import torchvision print(torch.__version__, torchvision.__version__) print(hasattr(torchvision.ops, 'DeformConv2d')) # 期望 True这一步别跳过,不同环境的torchvision版本差异比较大,先确认了再往后写。
3.2 DEC模块完整实现代码
下面就是DEC模块的完整PyTorch实现。我按论文结构复现,部分细节按我自己工程实践做了调整,每段关键逻辑都有注释。
import torch import torch.nn as nn import torchvision.ops as ops class DetailEnhancementConvolution(nn.Module): """ DEA-Net 细节增强卷积(DEC)复现实现 结构说明: 1. conv1: 普通卷积路径,提取基础特征 out1 2. conv2 -> conv3: 细节感知路径,提炼特征并预测 offset 3. deform_conv: 可变形卷积,作用在 conv2 的输出上 4. gate: 注意力门控,将 detail_feat 映射为 0~1 的调制权重 5. shortcut + out1 * gate: 残差融合,实现细节增强 """ def __init__(self, in_channels, out_channels, offset_channels=32): super().__init__() self.in_channels = in_channels self.out_channels = out_channels self.offset_channels = offset_channels # 基础特征路径 self.conv1 = nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False), nn.BatchNorm2d(out_channels), nn.PReLU(), ) # 细节感知路径的前两个卷积 self.conv2 = nn.Sequential( nn.Conv2d(in_channels, offset_channels, kernel_size=3, stride=1, padding=1, bias=False), nn.BatchNorm2d(offset_channels), nn.PReLU(), ) self.conv3 = nn.Sequential( nn.Conv2d(offset_channels, offset_channels, kernel_size=3, stride=1, padding=1, bias=False), nn.BatchNorm2d(offset_channels), nn.PReLU(), ) # 预测 offset:3x3 卷积核对应 9 个采样点,每个点有 (dx, dy) self.offset_conv = nn.Conv2d( offset_channels, 2 * 3 * 3, kernel_size=3, stride=1, padding=1, bias=False ) # 可变形卷积 self.deform_conv = ops.DeformConv2d( offset_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False ) # 注意力门控 self.gate = nn.Sequential( nn.Conv2d(out_channels, out_channels, kernel_size=1, bias=False), nn.BatchNorm2d(out_channels), nn.Sigmoid(), ) # 残差分支 self.shortcut = nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False), nn.BatchNorm2d(out_channels), ) self.act = nn.PReLU(out_channels) # 重点:把 offset 初始化为 0,训练初期等价于普通卷积 nn.init.zeros_(self.offset_conv.weight) def forward(self, x): shortcut = self.shortcut(x) # 残差路径 out1 = self.conv1(x) # 基础特征 out2 = self.conv2(x) # 细节路径的浅层特征 feat = self.conv3(out2) # 细节路径的深层特征 offset = self.offset_conv(feat) # 预测偏移量 [B, 18, H, W] # 可变形卷积作用在 out2 上(不是 feat) detail_feat = self.deform_conv(out2, offset) gate_weight = self.gate(detail_feat) # 注意力门控 [B, C_out, H, W] out = shortcut + out1 * gate_weight # 调制式融合 return self.act(out)有几个点需要单独拎出来说。
第一,offset_conv输入的是feat,但deform_conv作用的是out2,不是feat。这意味着网络先用两个卷积把输入提炼成一个更“有判断力”的特征图,从这个特征图上学出偏移量,再拿这个偏移量去对浅层的out2做重采样。这种“深特征预测、浅特征变形”的结构在可变形卷积实现里很常见。你想改成对feat变形也完全能跑,但复现时我建议先严格按这个来。
第二,deform_conv的输入特征通道是offset_channels,输出通道是out_channels。注意这里的offset_channels和上面的offset通道数是两码事。前者是中间特征通道数,论文里设成32,控制整个可变形卷积路径的宽度;后者是偏移量本身的通道数,固定等于2 * kH * kW。这两个名字容易混淆,后面改代码时别搞混。
第三,nn.init.zeros_(self.offset_conv.weight)这一行是我强烈建议加的。如果不做这个初始化,offset网络一开始就输出随机偏移,采样点到处乱跳,训练初期梯度很难稳定,严重的直接loss变成NaN。初始化为0以后,可变形卷积在训练起步阶段退化成普通卷积,网络先学会基本重建,再慢慢“长出”偏移能力,收敛稳定得多。
3.3 维度sanity check
模块写完先别急着接网络,用随机张量测一下维度是否对得上。这是我最常做的习惯,五分钟能省一下午的bug排查时间。
model = DetailEnhancementConvolution(in_channels=32, out_channels=64) x = torch.randn(2, 32, 128, 128) # [B, C, H, W] out = model(x) print(out.shape) # 期望 torch.Size([2, 64, 128, 128])如果输出shape和输入不一致,先检查offset_conv的输出通道是不是18,再检查deform_conv的padding和stride是否保持了空间尺寸。只要空间尺寸和通道数都正确,这个模块就可以拿去接网络了。
3.4 想调整结构时要注意的融合变体
DEC这个模块最值得玩的地方是融合公式。原文用的是shortcut + out1 * gate_weight,也就是乘法调制。但实际工程里也有两种常见变体:
- 加法变体:
out = shortcut + out1 + detail_feat。可变形卷积直接作为增量叠加,好处是细节特征的信息传递更充分,坏处是初始化阶段detail_feat不是零,会干扰训练,通常需要额外把deform_conv的权重也初始化为接近零,或者加个可学习的缩放因子。 - 加乘混合变体:
out = shortcut + out1 + detail_feat * gate_weight。既保留基础特征,又让可变形卷积贡献一部分带门控的增量。这个变体在某些数据集上比原文更强,但模块的可解释性会弱一点,需要自己权衡。
我的建议是第一版复现老老实实按原文来,跑通了再试变体。改融合方式的时候,同时要检查整个网络的梯度和训练稳定性,不要单纯看PSNR一个指标。
4. 把DEC塞进一个mini去雾网络,跑通训练闭环
4.1 MiniDehazeNet:用DEC当核心block的极简网络
模块单独能跑还不够,得放到一个完整的去雾网络里验证效果。我搭了一个非常轻量的mini网络,结构就一句话:一个卷积做浅层特征提取,一个stride=2卷积把分辨率降到一半,中间堆三个DEC,再上采样回原分辨率,最后加全局残差。
class MiniDehazeNet(nn.Module): def __init__(self): super().__init__() self.head = nn.Sequential( nn.Conv2d(3, 16, 3, 1, 1), nn.PReLU(), ) self.down = nn.Sequential( nn.Conv2d(16, 32, 3, 2, 1), nn.PReLU(), ) # 三个 DEC,通道先升后降 self.dec1 = DetailEnhancementConvolution(32, 64) self.dec2 = DetailEnhancementConvolution(64, 64) self.dec3 = DetailEnhancementConvolution(64, 32) self.up = nn.Sequential( nn.ConvTranspose2d(32, 16, 4, 2, 1), nn.PReLU(), ) self.tail = nn.Sequential( nn.Conv2d(16, 3, 3, 1, 1), ) def forward(self, x): h = self.head(x) h = self.down(h) h = self.dec1(h) h = self.dec2(h) h = self.dec3(h) h = self.up(h) out = self.tail(h) return out + x # 全局残差这里的全局残差是去雾网络的常用设计。因为输入的有雾图和输出的清晰图在整体结构上高度相似,让网络只去学“雾造成的残差”,比直接学完整图像容易得多,收敛速度也会快不少。
整个mini网络算下来非常轻量,大概几十万参数量,看你怎么设offset_channels。这个规模在普通单卡上训练完全没压力,非常适合做模块验证实验。
4.2 用大气散射模型合成训练数据
训练去雾模型最理想的当然是真实雾天/晴天成对数据,但这种数据很难采集。论文里常用RESIDE这类合成数据集,做法就是用大气散射模型给清晰图像加雾。
如果你只是想验证DEC模块的有效性,完全可以用一个简单的合成数据类,不需要下载大体积数据集。下面这个类从一张清晰图上随机裁剪patch,按大气散射模型加雾:
import random import torch from torch.utils.data import Dataset from torchvision.transforms import ToTensor class FoggyDataset(Dataset): def __init__(self, clean_images, patch_size=256): self.clean_images = clean_images # list of PIL.Image self.patch_size = patch_size self.to_tensor = ToTensor() def __len__(self): return len(self.clean_images) * 20 def __getitem__(self, idx): img = random.choice(self.clean_images) img = random_crop(img, self.patch_size) img = self.to_tensor(img) # [0, 1] # 随机雾浓度 alpha = random.uniform(0.5, 1.2) A = torch.rand(1, 1, 1) * 0.5 + 0.3 # 用随机场模拟深度变化,得到空间变化的透射率 depth = torch.rand(1, self.patch_size, self.patch_size) * 0.5 + 0.2 t = torch.exp(-alpha * depth) fog = img * t + A * (1 - t) return fog, img这个合成方式的随机性很关键。alpha控制雾的浓度,A控制雾的颜色偏向,depth的随机分布模拟了场景深度变化。每次迭代采不同的alpha、A和depth,等效于数据增强,能让网络学会更普适的去雾规律。
如果你手头已经能访问RESIDE数据集,直接用它更好,评测结果也更容易和别人对比。合成数据适合快速验证模块能不能work,标准数据集适合出正式实验结果。
4.3 损失函数、优化器与训练循环
去雾任务里,L1损失是默认选择,原因很简单:L2损失会过度惩罚大误差,导致网络倾向于输出偏平滑的结果,细节被进一步抹掉。L1对边缘更友好,细节保留得更好。下面这个训练循环是以L1 loss为核心的完整流程:
import torch import torch.nn.functional as F from torch.utils.data import DataLoader def train_one_epoch(model, loader, optimizer, device): model.train() total_loss = 0.0 for fog, clean in loader: fog = fog.to(device) clean = clean.to(device) pred = model(fog) loss = F.l1_loss(pred, clean) optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() return total_loss / len(loader) model = MiniDehazeNet().to(device) optimizer = torch.optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=100) for epoch in range(100): avg_loss = train_one_epoch( model, DataLoader(FoggyDataset(clean_images), batch_size=8, shuffle=True), optimizer, device ) scheduler.step() if epoch % 10 == 0: print(f"epoch {epoch}, loss {avg_loss:.4f}")超参数我给的是一个稳定的组合:Adam、lr=1e-4、weight_decay=1e-4、batch_size=8、patch_size=256。显存不够就把batch_size降到4、patch降到128,DEC里的offset_channels也可以从32降到16。
有过拟合倾向时可以加一点数据增强:随机翻转、随机旋转、颜色抖动都行。这些增强对去雾任务的帮助比想象中大,尤其是随机旋转,能让网络对边缘方向更鲁棒。
4.4 训练时重点观察什么
训练过程中不要只盯着loss数值。我习惯额外做两件事:一是每个epoch在固定验证集上算PSNR/SSIM,因为loss平滑下降不代表视觉质量一直在提升;二是挑一两张固定测试图,每隔几个epoch保存模型输出,直接看人眼效果。
DEC模块是否真的在工作,有一个很直观的观察方式:打印offset的统计值。如果offset的均值一直非常接近0,说明网络根本没学到有效偏移,可能卡在了局部最优;如果offset的分布逐渐散开,绝对值有增大趋势,说明可变形卷积在主动调整采样位置。这个指标比loss更能反映DEC有没有真正生效。
5. 验证DEC有效性的对比实验与踩坑清单
5.1 同一个mini网络,把DEC换成普通ResBlock会怎样
模块有没有用,不能靠感觉,得做对照实验。最干净的对比就是:保持MiniDehazeNet的其余结构完全不变,把中间三个DEC全部换成同通道数的普通ResBlock。ResBlock的结构是一个常规残差块:两次卷积、BN、PReLU、残差连接。
我用同一份合成数据、同一套超参数,各跑了150轮,DEC版在验证集上的PSNR比ResBlock版大概高了0.4到0.8 dB,具体数值随数据分布会有浮动,但趋势很稳定。主观视觉上差异更明显:DEC版在窗框、树枝、文字边缘这些位置明显更锐利,ResBlock版虽然整体亮度、颜色恢复得也不错,边缘却总带着一层薄雾感。
参数量方面,DEC版比ResBlock版多出大概20%到50%,主要来自可变形卷积分支的offset预测网络。这个增量换来的细节恢复能力,在去雾任务里是划算的。如果是超分、去雨这类同样对高频细节敏感的任务,DEC的收益大概率也是正向的。
| 对比项 | MiniDehazeNet + ResBlock | MiniDehazeNet + DEC |
|---|---|---|
| 核心模块 | 两次常规卷积 + 残差 | 可变形卷积 + 注意力门控 + 残差 |
| 细节保持能力 | 一般,边缘易被平滑 | 强,边缘纹理恢复更锐利 |
| 额外参数量 | 基准 | 增加约20%~50%(取决于offset_channels) |
| 训练收敛速度 | 较快 | 稍慢,但最终效果更优 |
| 对高频细节敏感任务 | 可用但有瓶颈 | 更适配 |
5.2 我踩过的几个坑:从shape error到训练发散
复现过程中我踩过的坑不少,整理成一张表给后来人排雷。
| 问题现象 | 根本原因 | 解决方案 |
|---|---|---|
| deform_conv2d报offset通道数错误 | 把3x3的offset误写成9通道,忘了dx和dy两个方向 | 3x3对应的offset通道数是18,即2 * 3 * 3 |
| 训练初期loss直接变NaN | offset初始化过大,采样点跳到非连续位置,特征图出现极大值 | 对offset_conv权重做zeros_初始化,让初始偏移为0 |
| 显存不够用 | 可变形卷积路径的中间特征太多,尤其patch设得很大时 | 减小offset_channels,或把训练patch降到128x128 |
| 特征图太小导致采样越界 | 某些网络把特征图下采样到4x4甚至2x2,offset偏移后采样点全跑出边界 | 保证进入DEC的特征图最小边不小于8,必要时补padding |
| CPU上训练慢到怀疑人生 | deform_conv2d在CPU上的计算效率远低于GPU,涉及双线性插值 | 训练必须用GPU;CPU只适合跑推理或debug |
5.3 torchvision可变形卷积的版本兼容性提醒
torchvision.ops.DeformConv2d这个接口在不同版本里的行为差异不大,但有几个点需要注意。老版本(0.9之前)根本没有这个接口,如果你在公司内部的老环境里跑,要先升级torchvision。升级后如果发现torchvision.ops里的函数签名不一样,以你当前版本的官方文档为准,我这里的写法基于较新的稳定版本。
另一个容易忽略的问题是CPU/GPU差异。可变形卷积内部的offset是浮点数,采样位置需要做双线性插值,这个操作在GPU上有高度优化的实现,但在CPU上非常慢。如果想在CPU上验证DEC能跑通,建议输入分辨率设小一点,比如64x64,否则等前向推理就能等到怀疑人生。
还想试可变形卷积v2的话,需要用函数式接口torchvision.ops.deform_conv2d,手动把mask参数传进去,ops.DeformConv2d这个模块类默认不带mask通道。DEA-Net用的应该是v1,复现阶段不需要上v2。
最后分享一个我这轮实验里觉得最值钱的小细节:DEC里的offset路径一定要保证初始偏移接近0,用nn.init.zeros_显式初始化offset_conv的权重,别偷懒。这个细节直接决定训练前几十个epoch是稳定攀升还是原地震荡。等DEC跑通之后,建议把DEA-Net里的上下文引导模块也补上,DEC管细节、CG管全局,两者配合才是完整的DEA-Net思路。我自己的下一步是把它挪到超分任务里试,边缘恢复的收益应该比去雾还明显。