ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

DEA-Net细节增强卷积DEC的PyTorch复现:可变形卷积与注意力门控提升去雾细节

DEA-Net细节增强卷积DEC的PyTorch复现:可变形卷积与注意力门控提升去雾细节 前阵子调一个去雾模型的训练日志里的PSNR看着还行但把测试图片放大看树叶边缘和窗框全是糊的整体亮度是对了细节却像被橡皮擦磨过一样。这个痛点几乎每个做图像复原的人都会碰上端到端网络在抑制噪声的同时把高频细节也一起抹掉了。后来我读到DEA-Net这篇工作它提出的细节增强卷积Detail Enhancement ConvolutionDEC正好打在“细节丢失”这个七寸上。这周我把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 GuidanceCG负责扩大感受野让去雾决策不只看局部。这两个模块的分工很明确。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_channels32): 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_size3, stride1, padding1, biasFalse), nn.BatchNorm2d(out_channels), nn.PReLU(), ) # 细节感知路径的前两个卷积 self.conv2 nn.Sequential( nn.Conv2d(in_channels, offset_channels, kernel_size3, stride1, padding1, biasFalse), nn.BatchNorm2d(offset_channels), nn.PReLU(), ) self.conv3 nn.Sequential( nn.Conv2d(offset_channels, offset_channels, kernel_size3, stride1, padding1, biasFalse), nn.BatchNorm2d(offset_channels), nn.PReLU(), ) # 预测 offset3x3 卷积核对应 9 个采样点每个点有 (dx, dy) self.offset_conv nn.Conv2d( offset_channels, 2 * 3 * 3, kernel_size3, stride1, padding1, biasFalse ) # 可变形卷积 self.deform_conv ops.DeformConv2d( offset_channels, out_channels, kernel_size3, stride1, padding1, biasFalse ) # 注意力门控 self.gate nn.Sequential( nn.Conv2d(out_channels, out_channels, kernel_size1, biasFalse), nn.BatchNorm2d(out_channels), nn.Sigmoid(), ) # 残差分支 self.shortcut nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size1, biasFalse), 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_channels32, out_channels64) 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网络结构就一句话一个卷积做浅层特征提取一个stride2卷积把分辨率降到一半中间堆三个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_size256): 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(), lr1e-4, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max100) for epoch in range(100): avg_loss train_one_epoch( model, DataLoader(FoggyDataset(clean_images), batch_size8, shuffleTrue), optimizer, device ) scheduler.step() if epoch % 10 0: print(fepoch {epoch}, loss {avg_loss:.4f})超参数我给的是一个稳定的组合Adam、lr1e-4、weight_decay1e-4、batch_size8、patch_size256。显存不够就把batch_size降到4、patch降到128DEC里的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 ResBlockMiniDehazeNet DEC核心模块两次常规卷积 残差可变形卷积 注意力门控 残差细节保持能力一般边缘易被平滑强边缘纹理恢复更锐利额外参数量基准增加约20%~50%取决于offset_channels训练收敛速度较快稍慢但最终效果更优对高频细节敏感任务可用但有瓶颈更适配5.2 我踩过的几个坑从shape error到训练发散复现过程中我踩过的坑不少整理成一张表给后来人排雷。问题现象根本原因解决方案deform_conv2d报offset通道数错误把3x3的offset误写成9通道忘了dx和dy两个方向3x3对应的offset通道数是18即2 * 3 * 3训练初期loss直接变NaNoffset初始化过大采样点跳到非连续位置特征图出现极大值对offset_conv权重做zeros_初始化让初始偏移为0显存不够用可变形卷积路径的中间特征太多尤其patch设得很大时减小offset_channels或把训练patch降到128x128特征图太小导致采样越界某些网络把特征图下采样到4x4甚至2x2offset偏移后采样点全跑出边界保证进入DEC的特征图最小边不小于8必要时补paddingCPU上训练慢到怀疑人生deform_conv2d在CPU上的计算效率远低于GPU涉及双线性插值训练必须用GPUCPU只适合跑推理或debug5.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思路。我自己的下一步是把它挪到超分任务里试边缘恢复的收益应该比去雾还明显。
返回列表