ARTICLE DETAIL

资讯详情

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

PyTorch实现对偶生成对抗网络图像去雾:从数据到部署

PyTorch实现对偶生成对抗网络图像去雾:从数据到部署 简介基于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, size256): 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_lastTrue避免最后一个 batch 形状不一致dataset UnpairedHazeDataset(data/haze, data/clear, size256) loader DataLoader(dataset, batch_size1, shuffleTrue, num_workers4, drop_lastTrue)num_workers在 Linux 下可以开到 4 或 8Windows 下建议 2否则数据加载经常成为瓶颈。这里的图像增强没有做随机翻转实际训练时可以在__getitem__里加T.RandomHorizontalFlip()代价是多一次张量操作但对生成器的泛化性有不少帮助。2.3 去雾任务特有的预处理技巧大气光归一化做去雾的人都知道大气散射模型简单说一张有雾图可以看成清晰图衰减后叠加了大气光。CycleGAN 这类方法不显式建模这个物理过程但如果把输入图像直接丢给网络生成器需要自己学到“雾的密度”和“场景深度”的隐含关系这比普通风格迁移更难收敛。一个常见技巧是训练前对每一张有雾图计算暗通道粗略估计大气光值然后把图像减去大气光再做归一化。这个操作不需要精确的物理估计只需要把输入的分布拉到一个更“平”的位置让生成器更容易学到残差。def estimate_atmospheric_light(img_tensor, percent0.001): # img_tensor: [C, H, W], 值范围 [0, 1] dark img_tensor.min(dim0, keepdimTrue)[0] flat dark.flatten() k max(1, int(flat.shape[0] * percent)) indices torch.topk(flat, k, largestTrue).indices # 取原图中这些最亮暗通道位置的像素均值 flat_img img_tensor.reshape(img_tensor.shape[0], -1) atmosphere flat_img[:, indices].mean(dim1).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(inplaceTrue), 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_channels3, out_channels3, n_res9): super().__init__() # 下采样 down_layers [ nn.ReflectionPad2d(3), nn.Conv2d(in_channels, 64, 7), nn.InstanceNorm2d(64), nn.ReLU(inplaceTrue), nn.Conv2d(64, 128, 3, stride2, padding1), nn.InstanceNorm2d(128), nn.ReLU(inplaceTrue), nn.Conv2d(128, 256, 3, stride2, padding1), nn.InstanceNorm2d(256), nn.ReLU(inplaceTrue), ] res_blocks [ResidualBlock(256) for _ in range(n_res)] # 上采样先用最近邻放大再卷积 up_layers [ nn.Upsample(scale_factor2, modenearest), nn.ReflectionPad2d(1), nn.Conv2d(256, 128, 3), nn.InstanceNorm2d(128), nn.ReLU(inplaceTrue), nn.Upsample(scale_factor2, modenearest), nn.ReflectionPad2d(1), nn.Conv2d(128, 64, 3), nn.InstanceNorm2d(64), nn.ReLU(inplaceTrue), 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_channels3, base64): super().__init__() layers [ nn.Conv2d(in_channels, base, 4, stride2, padding1), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(base, base * 2, 4, stride2, padding1), nn.InstanceNorm2d(base * 2), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(base * 2, base * 4, 4, stride2, padding1), nn.InstanceNorm2d(base * 4), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(base * 4, 1, 4, padding1), ] 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, lr0.0002, betas(0.5, 0.999)) opt_dis torch.optim.Adam(dis_params, lr0.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, stride2, padding1))谱归一化能限制判别器的水印能力即函数对输入的敏感度上限防止判别器在训练初期学得太快。但要注意谱归一化会增加约 10% 的显存开销而且由于它约束了权重的谱范数收敛后的生成结果可能会略显平滑。对去雾这种低层次视觉任务更推荐的稳定手段是“标签平滑”。把真实标签从 1 改成 0.9伪造标签从 0 改成 0.1这能防止判别器输出极端值从而避免生成器因为过大的梯度而震荡。代码上和普通 MSE 损失几乎一致只是目标值变了real_label torch.ones(batch_size, devicedevice) * 0.9 fake_label torch.zeros(batch_size, devicedevice) * 0.1标签平滑的效果在去雾任务里尤其明显因为有雾图像中本来就存在天然的模糊区域判别器如果“过于自信”地判断某块区域是假的会给生成器传递错误的梯度方向。5. 训练稳定性调优与推理验证让 Dual GAN 去雾结果从“能看”到“干净”5.1 训练过程中必须盯的三个指标Dual GAN 去雾训练和图像分类不一样没有明确的准确率指标训练过程盯着 loss 曲线是不够的。从业者一般会同时关注三个指标生成器的对抗损失如果持续上升说明判别器太强需要降低判别器学习率或增强生成器容量循环一致性损失如果下降缓慢说明两个域的映射还没有建立起来需要检查数据预处理生成图像的平均亮度这个指标很朴素但有效去雾后的图像亮度应该落在一个合理区间内固定一个频率把生成结果保存成图片是最可靠的验证方式。PyTorch 里可以用torchvision.utils.save_imagefrom torchvision.utils import save_image if step % 500 0: with torch.no_grad(): fake_clear gen_h2c(haze) save_image(fake_clear, foutputs/step_{step}.png, normalizeTrue, range(-1, 1))normalizeTrue配合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, devicecuda): 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 和 SSIMfrom 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_range1.0) ssim structural_similarity(real, fake, channel_axis-1, data_range1.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时设置strictFalse或者在保存前先.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(foutputs/{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 倍代码改动也非常小。本文还有配套的精品资源点击获取
返回列表