简介:基于Pytorch实现的对偶生成对抗网络(Dual GAN)图像去雾项目,面向计算机相关专业学生、课程设计及毕业设计人群,可解决真实场景图像复原与生成对抗网络实战训练问题。项目经导师指导并获评审98分,整体完成度高,适合作为课程设计或期末大作业的高分范例。压缩包共54个文件,大小42.5MB,主要包含10个Python源代码文件,覆盖生成器与判别器网络定义、训练与预测主脚本、数据加载器、参数解析、日志展示、模型保存与加载等完整工程模块;同时提供14张测试样例图片与5张预测效果对比图,2个预训练模型pkl文件,以及README说明文档、git配置等辅助内容,目录结构清晰,便于学习者直接运行复现和二次开发。资源已有60人学习下载,适合希望掌握GAN在图像去雾中的应用、快速搭建深度学习项目并提升实践能力的读者,也可作为后续扩展图像增强、风格迁移等方向的起点。
1. 用 PyTorch 实现对偶生成对抗网络做图像去雾,为什么不是“换个网络”那么简单
图像去雾这个任务,表面看是把一张灰蒙蒙的照片变清晰,但真正动手才会发现难点不在“去雾”本身,而在“没雾的图从哪来”。真实世界里你很难拍到同一场景“有雾”和“无雾”的成对照片,所以监督学习所需的 paired data 几乎不存在。Dual GAN(对偶生成对抗网络)恰好是冲着这个约束去的:它不需要成对样本,只需要两个域的图像集合——有雾图集合和清晰图集合——就能通过循环一致性约束学出二者之间的映射。配合 PyTorch 的动态图机制,这套方案在工程上非常容易落地,而且显存占用比想象中温和。
这篇文章以“PyTorch 实现对偶生成对抗网络来实现图像去雾”为线索,从一个可运行的最小框架讲起,覆盖数据组织、生成器与判别器的选型、循环一致性损失和身份损失的配比,再到训练稳定性和推理阶段的坑。源码部分不假装是某个现成项目的解读,而是按从业者最常见的做法拆解——你可以直接照着搭,也能把里面的模块替换成自己的结构。整个思路对做过 GAN 但没碰过无配对图像翻译的人同样适用,读完你至少知道一张有雾图从输入到输出,中间经历了哪些张量变换和损失回传。
2. 数据组织与预处理:Dual GAN 去雾的输入输出到底该怎么对齐
2.1 无配对数据集的目录结构与 Dataset 实现
Dual GAN 的训练不需要成对样本,但目录结构最好还是分成两个大域,方便 DataLoader 按域独立采样。通常的做法是:
data/ haze/ # 有雾图像,全部放这里 clear/ # 清晰图像,全部放这里这两个目录下的文件名完全不需要对应。你需要关心的是每张图的尺寸是否接近,因为后续生成器通常下采样到 256×256 或 512×512 分辨率,尺寸差异太大会导致缩放后内容失真。常见做法是统一短边缩放到 256,再做随机裁剪,这样既保留纹理细节,又给训练增加随机性。
PyTorch 的 Dataset 写法比较直接,核心是在__getitem__里分别从两个目录读图并返回:
import torch from torch.utils.data import Dataset from PIL import Image import os import torchvision.transforms as T class UnpairedHazeDataset(Dataset): def __init__(self, haze_dir, clear_dir, size=256): self.haze_paths = sorted( [os.path.join(haze_dir, f) for f in os.listdir(haze_dir)] ) self.clear_paths = sorted( [os.path.join(clear_dir, f) for f in os.listdir(clear_dir)] ) self.transform = T.Compose([ T.Resize((size, size)), T.ToTensor(), T.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]), ]) def __len__(self): # 以较大的域为准,小的域做循环采样 return max(len(self.haze_paths), len(self.clear_paths)) def __getitem__(self, idx): haze_path = self.haze_paths[idx % len(self.haze_paths)] clear_path = self.clear_paths[torch.randint(0, len(self.clear_paths), (1,)).item()] haze = Image.open(haze_path).convert("RGB") clear = Image.open(clear_path).convert("RGB") return self.transform(haze), self.transform(clear)这段代码有几个值得留意的点。Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5])把像素从 [0, 1] 映射到 [-1, 1],这是大多数 GAN 生成器输出层用 Tanh 的前提,如果你改成默认的均值 0 方差 1 归一化,生成器输出会被限制在一个非常窄的范围内,很难拟合真实图像的分布。长度取两个域的最大值,是为了让训练过程交替看到不同域的样本,而不是某个域先被耗尽。
2.2 批大小与分辨率的权衡
Dual GAN 和 CycleGAN 类似,显存消耗集中在生成器和两个判别器上。256×256 分辨率下,批大小设为 1 是安全的,视觉结果也不错;如果想批量训练,通常只能降到 128×128。从业者的经验是:去雾任务对细节敏感,宁可 batch size 小一点,也要保住分辨率。
训练时用torch.utils.data.DataLoader加载,注意drop_last=True避免最后一个 batch 形状不一致:
dataset = UnpairedHazeDataset("data/haze", "data/clear", size=256) loader = DataLoader(dataset, batch_size=1, shuffle=True, num_workers=4, drop_last=True)num_workers在 Linux 下可以开到 4 或 8,Windows 下建议 2,否则数据加载经常成为瓶颈。这里的图像增强没有做随机翻转,实际训练时可以在__getitem__里加T.RandomHorizontalFlip(),代价是多一次张量操作,但对生成器的泛化性有不少帮助。
2.3 去雾任务特有的预处理技巧:大气光归一化
做去雾的人都知道大气散射模型,简单说一张有雾图可以看成清晰图衰减后叠加了大气光。CycleGAN 这类方法不显式建模这个物理过程,但如果把输入图像直接丢给网络,生成器需要自己学到“雾的密度”和“场景深度”的隐含关系,这比普通风格迁移更难收敛。
一个常见技巧是:训练前对每一张有雾图计算暗通道,粗略估计大气光值,然后把图像减去大气光再做归一化。这个操作不需要精确的物理估计,只需要把输入的分布拉到一个更“平”的位置,让生成器更容易学到残差。
def estimate_atmospheric_light(img_tensor, percent=0.001): # img_tensor: [C, H, W], 值范围 [0, 1] dark = img_tensor.min(dim=0, keepdim=True)[0] flat = dark.flatten() k = max(1, int(flat.shape[0] * percent)) indices = torch.topk(flat, k, largest=True).indices # 取原图中这些最亮暗通道位置的像素均值 flat_img = img_tensor.reshape(img_tensor.shape[0], -1) atmosphere = flat_img[:, indices].mean(dim=1).view(-1, 1, 1) return atmosphere # 使用示例 haze = torch.rand(3, 256, 256) * 0.5 + 0.1 A = estimate_atmospheric_light(haze) haze_norm = (haze - A) / (1 - A + 1e-8)这段预处理在很多去雾网络里被视为“去雾前处理”的标准操作。需要提醒的是,Dual GAN 的生成器输出范围是 [-1, 1](Tanh),而这里的计算假设输入是 [0, 1],所以预处理要在 ToTensor 之后、Normalize 之前做。如果你在往自己的数据集上套这套流程,务必把顺序理清楚,否则输入分布和生成器输出域不匹配,训练会一直震荡。
3. 生成器与判别器架构:在 PyTorch 里搭出适合去雾的 Dual GAN 主体
3.1 生成器选型:ResNet 块堆叠为什么比 U-Net 更合适
图像去雾本质上是“输入输出同尺寸的图像翻译”,U-Net 通过跳跃连接能保留边缘细节,但 Dual GAN 的场景里生成器需要处理的是“雾”这种全局低频干扰,而不是局部结构缺失。ResNet 风格生成器——先下采样再上采样,中间夹若干残差块——更擅长建模这种全局变换。PyTorch 里实现一个残差块非常简洁:
import torch.nn as nn class ResidualBlock(nn.Module): def __init__(self, in_channels): super().__init__() self.block = nn.Sequential( nn.ReflectionPad2d(1), nn.Conv2d(in_channels, in_channels, 3), nn.InstanceNorm2d(in_channels), nn.ReLU(inplace=True), nn.ReflectionPad2d(1), nn.Conv2d(in_channels, in_channels, 3), nn.InstanceNorm2d(in_channels), ) def forward(self, x): return x + self.block(x)这里刻意选了ReflectionPad2d而不是ZeroPad2d,原因是反射填充不会在图像边缘产生突兀的暗边,这对去雾结果影响很大。InstanceNorm2d是关键中的关键:图像去雾任务里,不同图的雾浓度差异很大,BatchNorm 会把整个 batch 的均值和方差拉平,导致雾浓的图生成结果发白。InstanceNorm 逐样本归一化,能保住每张图自身的对比度。
3.2 完整生成器:从下采样到残差再到上采样
一个标准的 256×256 输入的生成器结构是:三个下采样卷积 + 九个残差块 + 三个上采样转置卷积。转置卷积容易产生棋盘伪影,常见做法是换成最近邻上采样加普通卷积:
class DualGANGenerator(nn.Module): def __init__(self, in_channels=3, out_channels=3, n_res=9): super().__init__() # 下采样 down_layers = [ nn.ReflectionPad2d(3), nn.Conv2d(in_channels, 64, 7), nn.InstanceNorm2d(64), nn.ReLU(inplace=True), nn.Conv2d(64, 128, 3, stride=2, padding=1), nn.InstanceNorm2d(128), nn.ReLU(inplace=True), nn.Conv2d(128, 256, 3, stride=2, padding=1), nn.InstanceNorm2d(256), nn.ReLU(inplace=True), ] res_blocks = [ResidualBlock(256) for _ in range(n_res)] # 上采样:先用最近邻放大,再卷积 up_layers = [ nn.Upsample(scale_factor=2, mode="nearest"), nn.ReflectionPad2d(1), nn.Conv2d(256, 128, 3), nn.InstanceNorm2d(128), nn.ReLU(inplace=True), nn.Upsample(scale_factor=2, mode="nearest"), nn.ReflectionPad2d(1), nn.Conv2d(128, 64, 3), nn.InstanceNorm2d(64), nn.ReLU(inplace=True), nn.ReflectionPad2d(3), nn.Conv2d(64, out_channels, 7), nn.Tanh(), ] self.model = nn.Sequential( *down_layers, *res_blocks, *up_layers ) def forward(self, x): return self.model(x)为什么用 9 个残差块而不是 6 个或 12 个?这个数字来自 CycleGAN 原论文对 256×256 输入的经验配置,残差块越多感受野越大,对全局雾霾分布建模能力越强,但训练代价线性增长。实际使用中如果你的图里有大面积浓雾区域,9 个块是底线;如果是薄雾场景,6 个块也能出结果,收敛更快。上采样选nearest + conv,这是避免棋盘伪影最稳妥的做法。
3.3 判别器:70×70 PatchGAN 的 PyTorch 实现
Dual GAN 的判别器不需要看整张图来决定真伪,只需要对局部 patch 输出真伪概率。PatchGAN 的做法是输出一个 N×N 的特征图,每个位置对应输入图像的一个感受野区域,这样既能捕捉高频纹理,参数还少。PyTorch 实现:
class PatchDiscriminator(nn.Module): def __init__(self, in_channels=3, base=64): super().__init__() layers = [ nn.Conv2d(in_channels, base, 4, stride=2, padding=1), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base, base * 2, 4, stride=2, padding=1), nn.InstanceNorm2d(base * 2), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base * 2, base * 4, 4, stride=2, padding=1), nn.InstanceNorm2d(base * 4), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base * 4, 1, 4, padding=1), ] self.model = nn.Sequential(*layers) def forward(self, x): return self.model(x)注意这里没有在最后一层加 Sigmoid。实际训练中,PyTorch 的BCEWithLogitsLoss会自己完成 Sigmoid 和交叉熵的计算,如果网络层里已经加了 Sigmoid,训练时会因为数值不稳定导致梯度消失。很多初学者在这里栽跟头:判别器 loss 一直在 0.693 附近不动,就是因为输出层和损失函数不匹配。
这个判别器对 256×256 输入会输出 30×30 的特征图,每个点对应原图 70×70 的感受野。对去雾任务来说,70×70 已经足够覆盖局部纹理和边缘信息,更大感受野反而会过分关注全局亮度一致性,干扰生成器学到正确的物理去雾方向。
3.4 双生成器结构:为什么要两个 G 而不是一个
Dual GAN 的核心是对偶映射:一个生成器负责“有雾→清晰”,另一个负责“清晰→有雾”。这个设计的物理意义在于,去雾映射如果没有反向映射的约束,解空间会非常大——一张有雾图可以对应无数张“清晰”图。有了反向生成器,正向结果必须能被反向生成器还原成原来的有雾图,这就大大压缩了可行解的范围。
在代码层面,两个生成器结构完全相同,只是学习目标不同。初始化时不要用默认的均匀分布,常见做法是nn.init.normal_(weight, 0.0, 0.02),这个初始化的标准差 0.02 是 DCGAN 系列论文验证过的经验值,能有效避免训练初期判别器瞬间碾压生成器。
4. 训练循环与损失函数:对抗损失、循环一致性、身份损失的权重怎么配
4.1 三个损失函数的权重设定逻辑
Dual GAN 去雾和 CycleGAN 在损失设计上几乎一致:
- 对抗损失(adversarial loss):让生成器的输出看起来属于目标域
- 循环一致性损失(cycle consistency loss):正向去雾再反向加雾后,必须能还原回原图
- 身份损失(identity loss):输入本来就是清晰图时,去雾生成器应尽量保持不变
权重配比上,常见做法是给循环一致性损失一个较大的系数,通常是 10,身份损失系数为 5,对抗损失系数为 1。为什么循环一致性权重这么大?因为对抗损失只负责“像不像”,循环一致性负责“内容保真”。如果只有一个对抗损失,生成器可以牺牲内容细节来骗过判别器,比如把图像变成全灰,判别器无法识破。循环一致性损失会惩罚这种投机行为。
训练时总损失可以组织成字典,方便后续断点续训和调参:
criterion_idt = nn.L1Loss() criterion_cycle = nn.L1Loss() criterion_gan = nn.MSELoss() # LSGAN 形式,训练更稳定 lambda_cycle = 10.0 lambda_idt = 5.0 # 前向计算 fake_clear = gen_h2c(haze) # 有雾 -> 清晰 fake_haze = gen_c2h(clear) # 清晰 -> 有雾 rec_haze = gen_c2h(fake_clear) # 还原有雾 rec_clear = gen_h2c(fake_haze) # 还原清晰 # 循环一致性损失 loss_cycle = criterion_cycle(rec_haze, haze) + criterion_cycle(rec_clear, clear) # 身份损失 loss_idt = criterion_idt(fake_clear, clear) + criterion_idt(fake_haze, haze)身份损失这一段有几个容易踩的坑。输入清晰图给gen_h2c,期望输出仍然接近清晰图;同理,输入有雾图给gen_c2h,输出要接近有雾图。这保证了生成器不会过度调整输入的色彩分布。去雾任务里如果不加身份损失,生成结果经常出现偏绿或偏蓝的情况,因为对抗损失只要求“看起来像清晰图”,而清晰图数据集中若有色偏样本,生成器会主动去模仿。
4.2 优化器分组与学习率调度
两个生成器共享一个优化器还是分开用?实际工程中更常见的是四个网络——两个生成器和两个判别器——分别建优化器,这样能独立控制学习率。生成器学习率设为 0.0002,判别器也是 0.0002,但如果有判别器 loss 掉得太快的情况,可以单独把判别器学习率降到 0.0001。
PyTorch 中最常用的做法是用两个 Adam 优化器分别管理生成器参数组和判别器参数组:
gen_params = list(gen_h2c.parameters()) + list(gen_c2h.parameters()) dis_params = list(dis_haze.parameters()) + list(dis_clear.parameters()) opt_gen = torch.optim.Adam(gen_params, lr=0.0002, betas=(0.5, 0.999)) opt_dis = torch.optim.Adam(dis_params, lr=0.0002, betas=(0.5, 0.999))betas=(0.5, 0.999)里的 0.5 是关键参数。默认的 Adam 第一动量系数是 0.9,在 GAN 训练里容易产生震荡;0.5 是很多 GAN 从业者验证过的选择,能让梯度变化更平缓。
学习率调度上,常见做法有两种:一是每 N 个 epoch 乘以 0.5,二是前一半 epoch 保持常数,后一半线性衰减到 0。对 Dual GAN 去雾来说,线性衰减更常见,因为训练后期需要更小的更新步长来微调生成器的纹理细节。PyTorch 里可以直接用torch.optim.lr_scheduler.LambdaLR实现,不必额外引入外部库。
4.3 判别器的训练策略:先更新谁,梯度怎么停
每个训练 step 里,参数的更新顺序会影响收敛效果。常见做法是:
- 前向计算所有生成和重建结果
- 计算判别器损失,反向传播后更新判别器参数
- 计算生成器损失,反向传播后更新生成器参数
但这里有一个容易忽视的点:在计算判别器损失时,生成器的输出是带梯度的,如果直接用这些张量计算判别器损失并回传,梯度流会同时更新生成器。PyTorch 里需要显式切断:
# 更新判别器 opt_dis.zero_grad() loss_dis = criterion_gan(dis_clear(fake_clear.detach()), torch.ones_like(dis_clear(fake_clear.detach()))) loss_dis += criterion_gan(dis_clear(clear), torch.zeros_like(dis_clear(clear))) loss_dis.backward() opt_dis.step()fake_clear.detach()是这里的主角。它的作用是切断梯度回传到生成器的路径,让判别器可以专注于把自己训练得更强,而不会顺带把生成器带偏。如果你忘记加 detach,生成器和判别器的梯度混在一起,训练几乎必然发散。
4.4 梯度惩罚与谱归一化要不要加
标准 Dual GAN 用 LSGAN 或 BCE 就够用了,但训练不稳定时可以考虑加谱归一化,也就是给判别器的每一层卷积权重做奇异值约束。PyTorch 2.x 里可以直接用torch.nn.utils.spectral_norm包装卷积层:
from torch.nn.utils import spectral_norm conv = spectral_norm(nn.Conv2d(64, 128, 4, stride=2, padding=1))谱归一化能限制判别器的水印能力(即函数对输入的敏感度上限),防止判别器在训练初期学得太快。但要注意:谱归一化会增加约 10% 的显存开销,而且由于它约束了权重的谱范数,收敛后的生成结果可能会略显平滑。
对去雾这种低层次视觉任务,更推荐的稳定手段是“标签平滑”。把真实标签从 1 改成 0.9,伪造标签从 0 改成 0.1,这能防止判别器输出极端值,从而避免生成器因为过大的梯度而震荡。代码上和普通 MSE 损失几乎一致,只是目标值变了:
real_label = torch.ones(batch_size, device=device) * 0.9 fake_label = torch.zeros(batch_size, device=device) * 0.1标签平滑的效果在去雾任务里尤其明显,因为有雾图像中本来就存在天然的模糊区域,判别器如果“过于自信”地判断某块区域是假的,会给生成器传递错误的梯度方向。
5. 训练稳定性调优与推理验证:让 Dual GAN 去雾结果从“能看”到“干净”
5.1 训练过程中必须盯的三个指标
Dual GAN 去雾训练和图像分类不一样,没有明确的准确率指标,训练过程盯着 loss 曲线是不够的。从业者一般会同时关注三个指标:
- 生成器的对抗损失:如果持续上升,说明判别器太强,需要降低判别器学习率或增强生成器容量
- 循环一致性损失:如果下降缓慢,说明两个域的映射还没有建立起来,需要检查数据预处理
- 生成图像的平均亮度:这个指标很朴素但有效,去雾后的图像亮度应该落在一个合理区间内
固定一个频率把生成结果保存成图片是最可靠的验证方式。PyTorch 里可以用torchvision.utils.save_image:
from torchvision.utils import save_image if step % 500 == 0: with torch.no_grad(): fake_clear = gen_h2c(haze) save_image(fake_clear, f"outputs/step_{step}.png", normalize=True, range=(-1, 1))normalize=True配合range=(-1, 1)会先把张量值从 [-1, 1] 映射到 [0, 1] 再保存,这样保存出来的图片不会发灰。很多人在这一步忘记设定range参数,导致保存的图片整体偏暗,误以为是网络没收敛。
5.2 常见的不收敛症状和调参方向
Dual GAN 去雾训练中常见的失败模式有这么几种,各有应对方法:
第一种是生成器输出全黑或全白。这通常是判别器直接输出了极端值,导致生成器梯度爆炸。解决办法是把判别器换用 PatchGAN 后加入谱归一化,或者把对抗损失从 BCE 换成 LSGAN,也就是上面代码里使用的 MSE 形式。
第二种是生成结果有严重棋盘伪影。这个多发生在上采样层,检查一下你的转置卷积是否被普通卷积 + 最近邻上采样替代。如果用的是nn.ConvTranspose2d,把kernel_size设为stride的整数倍会缓解问题,但根治还是要换结构。
第三种是雾去除得不彻底,图像整体还是发白。这往往不是网络结构问题,而是身份损失的权重过大。当lambda_idt设置过高时,生成器会小心翼翼地不改变输入,导致去雾力度不够。把身份损失系数从 5 降到 2 或 3 通常会有效果。
5.3 推理阶段:加载权重后如何保持结果稳定
训练完成后,推理代码要单独写,和训练逻辑解耦:
import torch def dehaze(model, img_tensor, device="cuda"): model.eval() img_tensor = img_tensor.to(device) with torch.no_grad(): result = model(img_tensor) return resultmodel.eval()是必要的,虽然生成器里没有 Dropout 和 BatchNorm,但如果你的实现里用了 BatchNorm,不切到 eval 会导致推理时仍然按 batch 统计量归一化,结果会与训练时不一致。对去雾任务,输入图片的尺寸可能与训练尺寸不同,这时生成器中的下采样和上采样倍数要能整除输入尺寸,否则会出现尺寸错位。保险做法是推理前把图片 resize 到 256 的整数倍。
5.4 量化去雾效果:PSNR 和 SSIM 怎么才算合格
去雾效果的主观判断容易产生偏差,量化指标仍然是必要的。pytorch 里可以用torvision.transforms.functional配合skimage.metrics计算 PSNR 和 SSIM:
from skimage.metrics import peak_signal_noise_ratio, structural_similarity import numpy as np def evaluate_dehaze(fake, real): # fake 和 real 均为 [0,1] 范围的 numpy 数组,HWC psnr = peak_signal_noise_ratio(real, fake, data_range=1.0) ssim = structural_similarity(real, fake, channel_axis=-1, data_range=1.0) return psnr, ssim在合成数据集上,如果 PSNR 能达到 20 以上、SSIM 能达到 0.85 以上,视觉上基本是干净的了。真实雾图没有 ground truth,只能靠主观视觉,这时可以从“边缘锐度、颜色保真度、高光区是否过曝”三个角度评价。
6. 从 Dual GAN 到工程落地的最后几步:权重固化、批处理推理与常见报错排查
6.1 权重固化与模型导出
训练完以后,实际工程中往往需要把模型部署到没有训练环境的机器上。PyTorch 里最直接的方式是保存state_dict而不是整个模型对象,前者只包含参数张量,跨版本兼容性更好:
torch.save({ 'gen_h2c': gen_h2c.state_dict(), 'gen_c2h': gen_c2h.state_dict(), 'opt_gen': opt_gen.state_dict(), 'epoch': epoch, }, "dual_gan_dehaze.pth")加载时要注意先初始化模型,再加载权重。这里有一个常见报错:Missing key(s) in state_dict,通常是模型定义里的层名与保存时的层名不一致,比如保存时用了nn.DataParallel,加载时却是单卡模型,键名多了module.前缀。解决办法是在load_state_dict时设置strict=False,或者在保存前先.module取回原始模型。
6.2 批处理推理脚本的写法
真实使用场景往往不止处理一张图,而是一个文件夹下的所有图片。批处理推理的关键是控制内存,不要一次性把所有图片读入。一个实用的模式是边读边推理边保存:
from PIL import Image import torchvision.transforms as T transform = T.Compose([ T.Resize((256, 256)), T.ToTensor(), T.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]), ]) for img_path in sorted(Path("test/haze").glob("*.png")): img = Image.open(img_path).convert("RGB") input_tensor = transform(img).unsqueeze(0).to(device) with torch.no_grad(): out = gen_h2c(input_tensor) out_img = (out.squeeze(0).cpu().permute(1, 2, 0) + 1) / 2 out_img = (out_img * 255).numpy().astype(np.uint8) Image.fromarray(out_img).save(f"outputs/{img_path.name}")这段代码里(out + 1) / 2是把生成器的 Tanh 输出从 [-1, 1] 映射回 [0, 1],然后转换成 8 位整数存储。换一个思路,如果你在训练时记录过归一化时用的均值和方差,推理时反归一化更精准,但大多数去雾场景用 Tanh 的对称输出域就够了,不需要额外反归一化。
6.3 推理阶段常见的三个报错与含义
CUDA out of memory是推理时最常碰到的错误,原因是输入分辨率过大导致中间特征图膨胀。解决办法不是换更大的显卡,而是把torch.no_grad()写在代码外层,确保没有建立计算图。这一点对已经在训练代码里用了model.eval()的人可能觉得无所谓,但对于从训练代码直接改推理的人来说非常关键。
另一个常见报错是输入通道数不符合预期。去雾模型默认输入是三通道 RGB,如果你传入了带透明通道的 RGBA 图片,会直接报Expected 3 channels, got 4。打开图像后加一句.convert("RGB")能一劳永逸地规避这个问题。
还有一个隐蔽的问题是图片方向。PIL 读取时不会自动应用 EXIF 旋转信息,手机拍摄的照片可能会出现旋转 90 度的情况。推理脚本里最好加上ImageOps.exif_transpose(img),否则去雾结果的构图方向是错的,但算法本身没有错误,你还会误以为是模型的问题。
6.4 计算效率优化:半精度推理与批处理
去雾模型推理速度在 GPU 上通常不是瓶颈,真正慢的是数据读取和图像预处理。如果想进一步加速,可以考虑 PyTorch 的自动混合精度推理:
model.half() input_tensor = input_tensor.half() with torch.no_grad(): out = model(input_tensor)半精度对生成器的输出影响很小,因为 Tanh 的数值范围本身有限,但带来的加速在批量推理时非常明显。需要注意的是,InstanceNorm 在半精度下偶尔会出现数值不稳定,遇到结果出现极亮或极暗像素时,把对应层强制回float32即可。
批处理推理时,手动把多张图打包成一个 batch 可以减少 kernel launch 的开销,但要注意显存限制。一个折中的做法是每次打包 4 张,循环处理,这比逐张处理快 2 到 3 倍,代码改动也非常小。
本文还有配套的精品资源,点击获取