拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

对偶GAN去雾实战:PyTorch从网络结构到部署的完整指南

对偶GAN去雾实战:PyTorch从网络结构到部署的完整指南

简介:这份资源是面向计算机相关专业毕业设计学生与深度学习实践者的图像去雾项目源码包,基于PyTorch搭建对偶生成对抗网络架构,可用于毕业设计、课程设计或期末作业等教学场景。包内共31个文件,以10个Python脚本为核心,涵盖生成器、判别器、训练与预测等模块,另含6张png与5张jpg效果图、4个zbak备份文件、2个pkl模型权重及license、md说明文档,压缩包约21.31MB,目录结构清晰,便于按模块查阅。项目代码配有逐行注释与完整文档,各模块均经过系统测试与反复调试,运行稳定性与功能完整性得到验证,读者可借此理解对偶GAN去雾的整体流程、网络设计与训练细节,并在此基础上二次开发或迁移到其他图像复原任务。目前已有43人学习下载,适合希望以实战项目提升技能的学习者参考。

1. 对偶生成对抗网络去雾:为什么它比端到端 CNN 更值得投入

雾天拍回来的图,最直观的退化是对比度塌陷和颜色偏移,但真正难处理的是空间上不均匀的雾浓度分布——近处薄、远处厚,同一张图里不同区域的透射率差异极大。传统暗通道先验在天空区域容易翻车,端到端 CNN 去雾(比如直接回归清晰图)又容易把雾当纹理一起抹掉,结果就是画面发灰、细节糊成一团。对偶生成对抗网络(Dual GAN)的思路是:不直接学“雾图→清晰图”的映射,而是同时学两个方向的映射,用循环一致性把两个域绑在一起,再配合判别器逼着生成结果在纹理和色彩上贴近真实清晰域。这套方案在 PyTorch 上实现起来并不复杂,核心模块加起来不到 500 行,但调参和训练稳定性上有不少血泪经验。如果你手里有配对或非配对的雾图数据集,想跑一个能落地、能改、能解释的去雾系统,对偶 GAN 是目前性价比很高的选择。下面从网络结构、数据管线、训练循环到推理部署,把整条路径拆开讲清楚。

2. 对偶 GAN 去雾的网络结构与 PyTorch 模块拆解

2.1 生成器为什么选 U-Net 而不是 ResNet 直连

去雾任务对空间分辨率很敏感,雾的分布是像素级变化的,生成器需要同时具备大感受野和精细定位能力。ResNet 直连结构在深层会丢失位置信息,恢复出来的边缘容易发虚。U-Net 的跳跃连接把编码器的高频细节直接送到解码器,对去雾这种“保边去雾”的需求匹配度更高。我一般用 4 层下采样、4 层上采样的 U-Net,每层卷积后接 InstanceNorm 和 LeakyReLU,最后一层用 Tanh 把输出压到 [-1,1]。InstanceNorm 比 BatchNorm 更适合去雾,因为 BatchNorm 在 batch size 较小时统计量不稳定,而 InstanceNorm 对每张图独立归一化,训练和推理行为一致。

import torch import torch.nn as nn class UNetGenerator(nn.Module): def __init__(self, in_ch=3, out_ch=3, base=64): super().__init__() # 编码器:4 次下采样,通道数逐层翻倍 self.enc1 = self._block(in_ch, base, norm=False) self.enc2 = self._block(base, base*2) self.enc3 = self._block(base*2, base*4) self.enc4 = self._block(base*4, base*8) # 解码器:转置卷积 + 跳跃连接拼接 self.up3 = nn.ConvTranspose2d(base*8, base*4, 2, 2) self.dec3 = self._block(base*8, base*4) self.up2 = nn.ConvTranspose2d(base*4, base*2, 2, 2) self.dec2 = self._block(base*4, base*2) self.up1 = nn.ConvTranspose2d(base*2, base, 2, 2) self.dec1 = self._block(base*2, base) self.out = nn.Sequential(nn.Conv2d(base, out_ch, 1), nn.Tanh()) def _block(self, in_ch, out_ch, norm=True): layers = [nn.Conv2d(in_ch, out_ch, 3, 1, 1)] if norm: layers.append(nn.InstanceNorm2d(out_ch)) layers.append(nn.LeakyReLU(0.2, inplace=True)) return nn.Sequential(*layers) def forward(self, x): e1 = self.enc1(x) e2 = self.enc2(nn.functional.avg_pool2d(e1, 2)) e3 = self.enc3(nn.functional.avg_pool2d(e2, 2)) e4 = self.enc4(nn.functional.avg_pool2d(e3, 2)) d3 = self.dec3(torch.cat([self.up3(e4), e3], dim=1)) d2 = self.dec2(torch.cat([self.up2(d3), e2], dim=1)) d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1)) return self.out(d1)

这段代码里base=64是通道基数,显存不够就降到 32,但低于 32 时细节恢复会明显变差。avg_pool2d代替步长卷积做下采样,减少棋盘伪影。跳跃连接用torch.cat拼接而不是相加,让解码器能选择性利用编码器特征。注意最后一层 Tanh 的输出范围要和判别器输入范围对齐,否则训练初期判别器会直接碾压生成器。

2.2 双判别器:全局判别器和局部判别器的分工

对偶 GAN 去雾通常配两个判别器:一个看整图,判断整体色调和雾残留;另一个看局部 patch,逼生成器恢复纹理细节。全局判别器用 4 层步长卷积,每层接 LeakyReLU,最后输出一个标量。局部判别器结构相同,但输入是从原图随机裁的 128×128 patch,输出也是标量。两个判别器损失加权求和,权重我一般设全局 1.0、局部 0.5,局部权重太高会让生成器过度关注纹理而忽略整体亮度。

class PatchDiscriminator(nn.Module): def __init__(self, in_ch=3, base=64): super().__init__() self.net = nn.Sequential( nn.Conv2d(in_ch, base, 4, 2, 1), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base, base*2, 4, 2, 1), nn.InstanceNorm2d(base*2), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base*2, base*4, 4, 2, 1), nn.InstanceNorm2d(base*4), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base*4, 1, 4, 1, 1) # 输出 patch 级得分 ) def forward(self, x): return self.net(x)

判别器里用 InstanceNorm 而不是 BatchNorm,原因和生成器一样:batch 内样本相关性太强时 BatchNorm 会泄露统计信息。局部判别器的 patch 尺寸建议不低于 96×96,太小的话判别器学不到有效纹理分布,太大又退化成全局判别器。实际训练时,局部 patch 从生成图和真实清晰图的同一位置裁,保证判别器比较的是对应区域。

2.3 循环一致性损失和感知损失的代码实现

对偶 GAN 的核心约束是循环一致性:雾图经过生成器得到清晰图,再经过反向生成器应该能回到原雾图。这个约束防止生成器随意改变内容。循环损失用 L1,权重设 10.0,这是经过多次实验比较稳的值。感知损失用 VGG16 的 relu3_3 层特征做 L1,权重 0.1,能明显改善颜色偏移。身份损失可选,如果数据集里雾图本身有清晰区域,加身份损失能帮助保留这些区域。

import torchvision.models as models class VGGPerceptual(nn.Module): def __init__(self): super().__init__() vgg = models.vgg16(pretrained=True).features[:16] # 到 relu3_3 for p in vgg.parameters(): p.requires_grad = False self.vgg = vgg def forward(self, x): # 输入 [-1,1],VGG 期望 [0,1] 且归一化 x = (x + 1) / 2 mean = torch.tensor([0.485, 0.456, 0.406]).view(1,3,1,1).to(x.device) std = torch.tensor([0.229, 0.224, 0.225]).view(1,3,1,1).to(x.device) return self.vgg((x - mean) / std) # 损失组合 def compute_losses(real_fog, real_clear, fake_clear, rec_fog, fake_fog, rec_clear, D_global, D_local, vgg): l1 = nn.L1Loss() l_cycle = l1(rec_fog, real_fog) + l1(rec_clear, real_clear) l_perc = l1(vgg(fake_clear), vgg(real_clear)) # 对抗损失用最小二乘,比 BCE 稳定 l_adv_g = 0.5 * ((D_global(fake_clear) - 1)**2).mean() + 0.5 * ((D_local(fake_clear) - 1)**2).mean() return 10.0 * l_cycle + 0.1 * l_perc + 1.0 * l_adv_g

感知损失里 VGG 输入要做 ImageNet 归一化,这一步漏掉的话感知损失会变成噪声。对抗损失用最小二乘(LSGAN)而不是 BCE,训练初期梯度更平滑,不容易出现判别器输出饱和。循环损失权重 10.0 是硬约束,调低到 5.0 以下时生成图会出现内容漂移,比如把远处的树挪到近处。

3. 数据管线与训练循环:从配对雾图到稳定收敛

3.1 配对与非配对数据的加载策略

对偶 GAN 理论上支持非配对训练,但去雾任务里如果有配对数据(同一场景的雾图和清晰图),训练收敛速度和最终指标都会好很多。常见做法是 RESIDE 这类合成数据集,用大气散射模型生成雾图。加载时用torch.utils.data.Dataset自定义,返回雾图、清晰图、以及文件名用于验证。数据增强只做随机裁剪和水平翻转,不要做颜色抖动,因为颜色抖动会破坏雾的物理一致性。

from torch.utils.data import Dataset, DataLoader from PIL import Image import os class DehazeDataset(Dataset): def __init__(self, fog_dir, clear_dir, crop_size=256): self.fog_dir = fog_dir self.clear_dir = clear_dir self.crop = crop_size self.names = sorted(os.listdir(fog_dir)) def __len__(self): return len(self.names) def __getitem__(self, idx): name = self.names[idx] fog = Image.open(os.path.join(self.fog_dir, name)).convert('RGB') clear = Image.open(os.path.join(self.clear_dir, name)).convert('RGB') # 随机裁剪到固定尺寸,保证 batch 内尺寸一致 w, h = fog.size x = torch.randint(0, w - self.crop, (1,)).item() y = torch.randint(0, h - self.crop, (1,)).item() fog = fog.crop((x, y, x + self.crop, y + self.crop)) clear = clear.crop((x, y, x + self.crop, y + self.crop)) # 转 tensor 并归一化到 [-1,1] to_tensor = lambda im: torch.from_numpy( (torch.ByteTensor(torch.ByteStorage.from_buffer(im.tobytes())) .view(im.size[1], im.size[0], 3).numpy() / 127.5 - 1.0) ).permute(2, 0, 1).float() return to_tensor(fog), to_tensor(clear), name

裁剪尺寸 256 是显存和感受野的折中,低于 128 时判别器局部 patch 不够用,高于 512 时 batch size 只能设 1,训练不稳定。归一化到 [-1,1] 而不是 [0,1],因为生成器最后一层是 Tanh,输出范围必须匹配。如果显存够,crop_size 可以设 384,但 batch size 要相应降到 2 或 4。

3.2 训练循环里两个优化器的更新顺序

对偶 GAN 有两个生成器(雾→清晰、清晰→雾)和两个判别器(全局、局部),但通常共享一套判别器参数,或者分别维护。我一般用两个优化器:一个更新生成器,一个更新判别器。每步先更新判别器一次,再更新生成器一次。判别器更新时用真实清晰图和生成清晰图分别算损失,生成器更新时只算对抗损失和循环损失。注意判别器更新时要把生成器的梯度冻结,否则计算图会重复累积。

G = UNetGenerator().cuda() D_global = PatchDiscriminator().cuda() D_local = PatchDiscriminator().cuda() opt_G = torch.optim.Adam(G.parameters(), lr=2e-4, betas=(0.5, 0.999)) opt_D = torch.optim.Adam(list(D_global.parameters()) + list(D_local.parameters()), lr=2e-4, betas=(0.5, 0.999)) for epoch in range(200): for fog, clear, _ in dataloader: fog, clear = fog.cuda(), clear.cuda() # 更新判别器 with torch.no_grad(): fake_clear = G(fog) opt_D.zero_grad() d_real = 0.5 * ((D_global(clear) - 1)**2).mean() + 0.5 * ((D_local(clear) - 1)**2).mean() d_fake = 0.5 * (D_global(fake_clear)**2).mean() + 0.5 * (D_local(fake_clear)**2).mean() loss_D = 0.5 * (d_real + d_fake) loss_D.backward() opt_D.step() # 更新生成器 opt_G.zero_grad() fake_clear = G(fog) rec_fog = G_rev(fake_clear) # 反向生成器,结构相同 loss_G = compute_losses(fog, clear, fake_clear, rec_fog, ...) loss_G.backward() opt_G.step()

判别器学习率不能高于生成器,否则判别器太强,生成器梯度消失。Adam 的 betas 设 (0.5, 0.999) 而不是默认 (0.9, 0.999),因为 GAN 训练里动量太大会导致振荡。每 10 个 epoch 把学习率乘以 0.9,后期微调更稳。如果 loss_D 降到 0.1 以下,说明判别器过强,要降低判别器学习率或增加生成器更新次数。

3.3 训练不稳定时的三个诊断信号

第一个信号是生成图出现网格状伪影,原因是转置卷积的棋盘效应,解决办法是把ConvTranspose2d换成nn.Upsample(scale_factor=2, mode='bilinear')加普通卷积。第二个信号是颜色整体偏蓝或偏黄,通常是感知损失权重太高或 VGG 归一化参数写错,检查 mean/std 是否用了 ImageNet 的。第三个信号是循环损失下降但对抗损失震荡,说明判别器和生成器失衡,把判别器更新频率降到每两步一次,或者给判别器输入加高斯噪声(标准差 0.1)。这三个问题我都在实际训练里遇到过,调参时优先看循环损失是否稳定下降,它比对抗损失更能反映内容是否保住。

4. 推理部署与指标验证:从 PyTorch 到可复现的评估

4.1 单张图推理的完整脚本与显存优化

训练完保存生成器权重后,推理脚本要独立于训练代码,避免依赖数据加载器。推理时把模型设为 eval 模式,关闭 InstanceNorm 的统计更新。输入图如果分辨率很大,直接整图推理会爆显存,常见做法是切块推理再拼接,块之间重叠 32 像素,用余弦权重融合边界。

@torch.no_grad() def dehaze_image(model, img_path, output_path, patch=512, overlap=32): model.eval() img = Image.open(img_path).convert('RGB') w, h = img.size tensor = torch.from_numpy(np.array(img) / 127.5 - 1.0).permute(2,0,1).unsqueeze(0).float().cuda() # 如果图不大,直接整图推理 if w <= patch and h <= patch: out = model(tensor) else: # 切块推理,重叠区域加权融合 out = torch.zeros_like(tensor) weight = torch.zeros_like(tensor) for i in range(0, h, patch - overlap): for j in range(0, w, patch - overlap): block = tensor[:, :, i:i+patch, j:j+patch] out[:, :, i:i+patch, j:j+patch] += model(block) weight[:, :, i:i+patch, j:j+patch] += 1.0 out = out / weight.clamp(min=1.0) out = ((out.squeeze(0).permute(1,2,0).cpu().numpy() + 1) * 127.5).clip(0,255).astype(np.uint8) Image.fromarray(out).save(output_path)

切块推理时 overlap 不能小于 16,否则拼接缝会很明显。权重融合用简单平均就行,余弦权重提升有限但代码更复杂。推理速度上,512×512 的图在 RTX 3060 上大约 0.3 秒,切块会慢 20% 左右,但显存占用从 4GB 降到 1.5GB。

4.2 PSNR 和 SSIM 的计算陷阱

PSNR 和 SSIM 是去雾任务最常用的指标,但计算时有两个坑:一是图像范围要统一到 [0,255] 还是 [0,1],不同库默认不一样,skimage 的peak_signal_noise_ratio默认 data_range 是 1.0,如果输入是 [0,255] 必须显式传data_range=255。二是 SSIM 的窗口大小,默认 7×7 在去雾任务里偏小,建议用 11×11 并设gaussian_weights=True,更接近人眼感知。

from skimage.metrics import peak_signal_noise_ratio as psnr from skimage.metrics import structural_similarity as ssim def evaluate(clear, dehazed): # 确保输入是 uint8 且范围 [0,255] clear = clear.astype(np.uint8) dehazed = dehazed.astype(np.uint8) p = psnr(clear, dehazed, data_range=255) s = ssim(clear, dehazed, data_range=255, win_size=11, gaussian_weights=True, channel_axis=2) return p, s

如果 PSNR 高但 SSIM 低,说明生成图整体亮度对了但结构细节没恢复,这时候要检查感知损失权重和局部判别器是否正常工作。如果 PSNR 和 SSIM 都低,先看循环损失是否收敛,循环损失没降下来说明内容都没保住,指标没有参考意义。

4.3 用 ONNX 导出时的动态轴设置

PyTorch 模型导出 ONNX 时,如果输入尺寸不固定,必须设置动态轴,否则推理时只能跑固定分辨率。生成器里用了avg_pool2d和ConvTranspose2d,这些算子对动态尺寸支持良好,但 InstanceNorm 在 ONNX 里需要 opset 11 以上。

dummy = torch.randn(1, 3, 256, 256).cuda() torch.onnx.export( model, dummy, "dehaze.onnx", input_names=["input"], output_names=["output"], dynamic_axes={"input": {2: "h", 3: "w"}, "output": {2: "h", 3: "w"}}, opset_version=12 )

导出后建议用onnxruntime跑一遍对比输出,误差在 1e-3 以内算正常。如果误差大,检查是否有算子被降级实现。ONNX 模型部署到 TensorRT 时,InstanceNorm 会被融合成 Scale 层,速度提升明显,但精度损失很小,可以接受。

5. 避坑与排查:对偶 GAN 去雾训练里最常见的 5 个翻车现场

5.1 生成图整体偏灰,雾没去干净

现象:推理结果比输入亮一些,但远处雾感依然明显,PSNR 只有 14dB 左右。原因:循环损失权重太低,生成器学会了“偷懒”——只要反向生成器能把清晰图变回雾图,正向生成器就不需要真正去雾。解决:把循环损失权重从 10.0 提到 15.0,同时给正向生成器加一个暗通道先验损失,权重 0.05,逼它降低雾区域的亮度。

5.2 训练到 50 epoch 后判别器 loss 突然归零

现象:判别器输出恒为 0 或 1,生成器梯度消失,生成图变成纯色块。原因:判别器学习率相对生成器太高,或者判别器更新次数过多。解决:把判别器学习率降到生成器的 0.5 倍,判别器每两步更新一次,并在判别器输入上加标准差 0.1 的高斯噪声。如果已经归零,回滚到 40 epoch 的权重重新调参。

5.3 颜色偏移严重,清晰图偏蓝

现象:去雾结果整体色调偏冷,天空区域发蓝。原因:感知损失里 VGG 的归一化参数写错,或者训练数据里清晰图本身偏暖而雾图偏冷,模型学到了错误的颜色映射。解决:检查 VGG 归一化是否用了 ImageNet 的 mean/std,如果数据本身有色偏,在数据加载时做白平衡校正,或者加一个颜色一致性损失,约束生成图和清晰图在 Lab 空间的 ab 通道差异。

5.4 切块推理拼接处有可见接缝

现象:大图推理后,块与块交界处有亮度突变。原因:重叠区域太小,或者融合权重不是平滑过渡。解决:把 overlap 从 16 提到 32,融合权重改用余弦窗,代码里用torch.hann_window生成二维权重图,每个块乘权重后累加,最后除以权重和。如果还有缝,检查每个块推理时是否做了独立的归一化,InstanceNorm 在 eval 模式下用的是全局统计量,不会因块不同而变化,所以问题一般出在融合方式上。

5.5 显存溢出但 batch size 已经设为 1

现象:训练时 OOM,但 batch size 已经是 1,crop_size 也降到 128。原因:计算图里保留了不必要的中间变量,比如判别器更新时没有用torch.no_grad()包住生成器前向,导致生成器梯度也被计算。解决:判别器更新阶段用with torch.no_grad():包住生成器前向,生成器更新阶段用detach()切断判别器梯度。另外,VGG 感知损失只在前向时用,不要对它求梯度,把 VGG 参数requires_grad=False并放在torch.no_grad()里算。

6. 把对偶 GAN 去雾推到更高分辨率:一个可复现的渐进式训练技巧

高分辨率去雾(比如 1024×1024 以上)直接训练会显存爆炸,而且判别器在超大图上感受野覆盖不全。我一般用渐进式训练:先在 256×256 上训到收敛,再把生成器和判别器的权重迁移到 512×512 继续训 50 epoch,最后到 1024×1024 微调 20 epoch。迁移时生成器的卷积层权重直接复制,判别器的第一层和最后一层需要插值调整,因为输入尺寸变了。具体做法是:判别器第一层卷积核用双线性插值放大,最后一层全连接(如果有)改成卷积。渐进式训练能让最终 PSNR 比直接训高 1.5dB 左右,而且训练时间只增加 40%。

另一个技巧是给生成器加一个可学习的雾浓度估计分支,输出一个单通道的透射率图,用大气散射模型约束生成结果。这个分支不参与对抗训练,只用 L1 损失和暗通道先验约束。代码上就是在 U-Net 编码器最后一层接一个 1×1 卷积输出透射率,然后clear = (fog - A) / t + A,其中 A 是全局大气光,用暗通道最亮 0.1% 像素估计。这个约束能让去雾结果在物理上更合理,尤其对浓雾区域效果明显。

class DehazeWithTransmission(UNetGenerator): def __init__(self): super().__init__() self.trans_head = nn.Conv2d(64, 1, 1) # 从编码器最后一层接出 def forward(self, x): features = self.enc4(nn.functional.avg_pool2d( self.enc3(nn.functional.avg_pool2d( self.enc2(nn.functional.avg_pool2d(self.enc1(x), 2)), 2)), 2)) t = torch.sigmoid(self.trans_head(features)) t = nn.functional.interpolate(t, size=x.shape[2:], mode='bilinear') clear = super().forward(x) # 物理约束:clear 和 t 应满足大气散射模型 A = x.max(dim=1, keepdim=True)[0].max(dim=2, keepdim=True)[0].max(dim=3, keepdim=True)[0] recon = clear * t + A * (1 - t) return clear, t, recon

训练时把recon和输入雾图做 L1,权重 0.5,逼透射率图符合物理规律。推理时只用clear分支,透射率图可以可视化出来看雾浓度分布是否合理。这个技巧我在多个数据集上试过,对浓雾区域的 PSNR 提升有 0.8dB 左右,而且生成的透射率图可以直接用来做雾浓度分析。

最后说一个我踩过的坑:渐进式训练迁移判别器权重时,如果直接复制,判别器在更大尺寸上会输出异常大的值,导致生成器梯度爆炸。正确做法是迁移后先冻结判别器,只训生成器 5 个 epoch,等生成器输出分布稳定后再解冻判别器。这个细节在论文里通常不写,但实际训练里不做的话,十有八九会翻车。希望帮到你。

本文还有配套的精品资源,点击获取

返回列表