尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

EnlightenGAN复现实践:无监督低光图像增强与训练调优全攻略

EnlightenGAN复现实践:无监督低光图像增强与训练调优全攻略 1. 项目概述为什么选择EnlightenGAN作为复现目标1.1 核心需求解析说到低光图像增强很多人第一时间想到的可能是RetinexNet、Zero-DCE或者最近大火的SCI等算法。但我这次选择复现的是EnlightenGAN原因有三一是它在无监督低光增强这条技术路线里属于开山之作很多后续方法都在它的基础上做改进二是它的核心思路——全局-局部判别器加自特征保留损失——即便放到现在也有很强的参考价值三是官方代码是PyTorch实现的结构清晰非常适合用来做模型复现和训练调优的练手项目。如果你也是第一次接触这种基于GAN的低光增强模型我的建议是不要一上来就扎进论文公式里先把代码跑通、把训练流程走一遍再回头看论文你对那些模块的理解会快得多。这也是这篇文章存在的意义——把我从环境搭建、数据准备、参数调优到踩坑排错的全过程记录下来给你一条可以直接照走的路径。简单说一下这个项目是干什么的。EnlightenGAN的核心目标很明确给定一张暗光下的照片模型能够把它增强成正常光照下的效果而且不需要成对的低光/正常光训练数据。这意味着你不需要费尽心思去采集同一场景不同曝光程度的图片对只需要准备两类图像——暗光图和无所谓内容是否对应的正常光图就可以完成训练。这个特性在实际工程落地中太重要了。1.2 复现前的技术准备复现一个深度学习模型最忌讳的就是什么都不想直接git clone然后开始python train.py。我在动手之前先花了一整天梳理整个项目包括论文的核心创新点、代码仓库的结构、训练机制和损失函数构成。下面是我的准备清单你也可以照着准备硬件环境方面EnlightenGAN的训练对显存有一定要求。官方默认的训练分辨率是--fineSize 400在400×400的输入尺度下批量大小设为8比较稳妥。我自己用的是RTX 309024GB显存训练时占用了大概15GB左右。如果你用的是8GB或12GB显存的中端卡可以把批量大小调整为4到6或者把fineSize降到360训练效果不会有明显下降。软件环境方面官方代码是在PyTorch 0.4.1时代写的直接在新版本PyTorch上运行会有一堆兼容性问题。我这里实测后的推荐组合是组件版本说明Python3.8太新的版本会报typing相关错误PyTorch1.8.1兼容性最好再新版也能跑但需要改几处代码torchvision0.9.1与PyTorch版本严格对应CUDA11.1根据显卡驱动版本选择即可其他依赖numpy, scipy, opencv-python, pillow按需安装提示千万别直接用pip install torch安装最新的PyTorch版本我试过在PyTorch 2.0上跑torchvision.transforms的接口变化会导致一堆报错。建议用conda create -n enlighten python3.8先创建独立的虚拟环境再在环境内安装对应版本。2. 环境搭建与数据预处理细节2.1 虚拟环境与依赖安装我从踩坑中得出的经验是复现这类老项目最关键的一步就是环境隔离。下面是完整的安装命令你直接复制到终端里执行就行# 创建独立虚拟环境 conda create -n enlighten python3.8 conda activate enlighten # 安装CUDA版本的PyTorch pip install torch1.8.1cu111 torchvision0.9.1cu111 -f https://download.pytorch.org/whl/torch_stable.html # 安装其他依赖 pip install numpy scipy opencv-python pillow tqdm tensorboard装完之后建议立刻测试一下CUDA是否可用import torch print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0))如果输出结果为True和显卡型号就说明环境没问题。如果打印False大概率是PyTorch版本与CUDA驱动不匹配这个时候先检查驱动版本再考虑重装PyTorch。2.2 数据集获取与目录结构规划EnlightenGAN官方提供了两个主要数据集LOL数据集包含成对的低光/正常光图像用于有监督评估以及Unpaired数据集低光和正常光图像不配对用于无监督训练。在复现时我建议你用LOL数据集来训练用Unpaired数据集也可以但LOL数据的质量更高训练出来的效果更好。下载好数据后需要按照官方代码要求的方式组织目录结构。官方代码读取数据的方式是通过一个CSV文件记录图像路径因此在训练之前需要先生成一个索引文件。我的做法是写一个简单的Python脚本来完成import os, random low_light_dir LOLdataset/our485/low normal_light_dir LOLdataset/our485/high # 获取文件列表 low_imgs [os.path.join(low_light_dir, f) for f in os.listdir(low_light_dir) if f.endswith(.png) or f.endswith(.jpg)] normal_imgs [os.path.join(normal_light_dir, f) for f in os.listdir(normal_light_dir) if f.endswith(.png) or f.endswith(.jpg)] # 确保有足够多的正常光图像 assert len(normal_imgs) len(low_imgs), 正常光图像数量应大于低光图像数量 # 随机抽样配对注意这是无监督学习其实不要求严格对应同一场景 random.shuffle(normal_imgs) pairs list(zip(low_imgs, normal_imgs)) # 写入CSV with open(train.csv, w) as f: f.write(low_light,normal_light\n) for low, normal in pairs: f.write(f{low},{normal}\n)2.3 图像预处理的坑点与对策数据处理这一块有几个容易被忽视的细节我在这里专门说下图像归一化范围。官方代码里对图像的归一化方式是transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))这意味着像素值从[0, 1]映射到[-1, 1]。很多人在复现时习惯用ImageNet的均值和方差来做归一化这其实是错的。在GAN的训练中使用[-1, 1]范围是主流做法因为生成器的最后一层激活函数通常用的是Tanh输出范围就是[-1, 1]如果输入输出范围不一致模型很难收敛。数据加载方式。官方代码中提供了一个自定义的Dataset类内部会做随机裁剪和翻转增广。裁剪大小默认是--fineSize 400我个人实测400是一个比较平衡的值太大显存吃不消太小影响恢复效果。如果你想用更大尺度训练建议优先考虑增大裁剪尺寸而不是增大批量大小这样效果更好。验证集的设置。建议从训练集中每个场景挑选1到2张图像单独保存到一个文件夹里用于训练过程中定期生成增强前后的对比图。这样你可以直观地看到模型在训练不同阶段的恢复效果比只看loss曲线有用得多。我在实际训练时每一代epoch结束都会跑一次验证将模型的输出保存下来最后合成一个GIF来查看变化趋势。3. 模型结构与损失函数拆解3.1 生成器、全局判别器、局部判别器的工作分工EnlightenGAN的生成器是基于Attention U-Net结构的。和普通U-Net不同的是它在跳跃连接处引入了注意力门控机制让模型自动学习哪些特征需要被传递、哪些需要被抑制。这种设计特别适合低光增强任务因为暗光图像中不同区域的增强需求不一样比如窗户区域和阴影区域的处理方式就完全不同。生成器的输入是低光图像经过归一化后的张量形状为[B, 3, H, W]输出是增强后图像形状相同的张量。在编码器部分图像经过4次下采样特征通道逐渐增加到64、128、256、512而在解码器部分则对称地上采样。这里的注意力模块会计算一个权重图告诉网络哪些位置的特征更重要从而更精确地保留细节。而判别器采用的是全局-局部双判别器架构这是EnlightenGAN的一个核心创新点。全局判别器接收完整的增强图像和正常光图像判断两者是否为同一个分布局部判别器则随机裁剪增强图像中的几个小区域块--num_patch参数默认是16个50×50大小的patch判断这些局部区域是否也能骗过判别器。这样设计的目的很直接全局判别器保证整张图的色调和光照风格接近真实正常光图局部判别器则强制模型在每个细节区域都要有真实感防止出现局部过曝或者色彩失真的问题。如果你之前接触过普通GAN可以这样理解把生成器想象成一个考生它负责把低光图修成正常光图判别器是阅卷老师它负责判断考生交上来的图是真图还是伪图。全局判别器负责看整体印象分局部判别器负责看细节对不对两个老师同时打分学生就得两边都照顾到才行。3.2 损失函数构成与权重分配逻辑EnlightenGAN的损失由三部分构成生成器对抗损失使用最小二乘GANLSGAN的形式公式表达为0.5 * mean((D(G(x)) - 1)^2)。这里用LSGAN而不是传统GAN的交叉熵损失是因为LSGAN能提供更平滑的梯度训练更稳定。全局-局部判别器损失全局判别器和局部判别器各自独立计算对抗损失。判别器的目标是区分“真图”和“生成图”其损失为0.5 * mean((D(y) - 1)^2) 0.5 * mean((D(G(x)) - 0)^2)。自特征保留损失这是EnlightenGAN区别于其他GAN的另一个亮点。将输入的低光图和生成的增强图分别送入一个预训练的VGG16网络提取它们的深层特征然后计算这些特征之间的L1距离。这个损失的作用是保证增强后的图像在内容和语义上与原始低光图保持一致防止生成器“自由发挥”过度导致内容漂移。各损失的权重分配如下损失名称权重系数作用GAN对抗损失全局1.0保证整体光照风格一致GAN对抗损失局部1.0保证局部细节真实自特征保留损失10.0保持内容不变性自特征保留损失的权重为什么设这么大我在实际训练中发现如果没有这个损失生成器很容易把暗部提得很亮但画面中的文字、纹理等内容信息会被破坏掉。VGG特征的L1距离能够起到一个“锚点”的作用拉住生成器不让它跑偏。3.3 优化器与学习率调度细节官方代码中生成器和判别器分别使用两个Adam优化器学习率都是0.0001beta1设为0.5。这个beta1 0.5是GAN训练中一个经典的设置它能让优化器对梯度的一阶矩估计衰减更快从而避免训练震荡。学习率方面官方默认是固定学习率不设置衰减。我的实测建议是如果你训练超过100个epoch可以考虑在第80个epoch之后把学习率线性衰减到原来的0.1倍这样收敛得更平稳。具体实现可以用PyTorch的lr_scheduler.LambdaLRdef lambda_rule(epoch): if epoch 80: return 1.0 else: return 0.1 scheduler_G torch.optim.lr_scheduler.LambdaLR(optimizer_G, lr_lambdalambda_rule) scheduler_D torch.optim.lr_scheduler.LambdaLR(optimizer_D, lr_lambdalambda_rule)4. 训练过程的实操记录与参数调优4.1 训练启动命令与核心参数说明把环境、数据、代码都准备好之后接下来就是最核心的训练环节。官方代码的入口是train.py启动之前先确认训练参数。以下是我在复现过程中常用的一套参数组合python train.py \ --dataset unpaired \ --dataroot ./datasets/LOLdataset \ --fineSize 400 \ --num_patch 16 \ --batch_size 8 \ --n_epochs 120 \ --decay_epoch 80 \ --lr 0.0001 \ --gpu_ids 0 \ --display_freq 200 \ --print_freq 100 \ --save_epoch_freq 10这里重点解释几个容易被忽视的参数--num_patch 16局部判别器每次从生成图像中随机裁剪的补丁数量。这个值太小会导致局部判别效果不佳太大则会增加显存占用并拖慢训练速度。16是一个经过很多实验验证的折中值。--decay_epoch 80从第80个epoch开始学习率线性衰减。配合n_epochs 120意味着最后40个epoch的学习率从0.0001逐渐降到0。这个策略能帮助模型在训练后期稳定收敛避免在最优解附近反复震荡。--display_freq 200每200次迭代在TensorBoard中输出一次当前生成器和判别器的损失值。通过观察这些损失曲线的变化趋势你可以实时判断训练是否正常。4.2 训练中的观察指标与判断依据训练开始后你要养成定期看两个东西的习惯loss曲线和验证集输出图。Loss曲线的正常表现训练初期前10个epoch生成器的损失会比较高判别器的损失相对较低这说明判别器很容易分辨出增强图和真实图。随着训练进行生成器的损失逐渐下降判别器的损失会上下波动这是正常的对抗现象。如果生成器损失降得很快而判别器损失一直在低位徘徊说明生成器找到了骗过判别器的方法但增强效果可能并不好这时候需要检查是不是自特征保留损失的权重太小了。验证集输出图的判断标准训练到第10个epoch左右你应该能在验证输出图中看到明显的增强效果比如暗部亮度提升、色彩饱和度有所恢复。如果你的输出图像到第20个epoch还是一团黑或者一片白那就需要停下来检查问题了。我自己根据实验经验整理了这样一张对照表观察现象可能原因解决方案输出图像偏黑几乎无增强效果生成器学习率过小或模型未收敛调大学习率至0.0003或检查输入数据是否归一化正确输出图像过曝暗部细节丢失自特征保留损失权重太小将VGG损失权重从10增大至20或降低对抗损失权重Loss值出现NaN学习率过大导致梯度爆炸或输入数据有问题降低学习率至0.00005检查数据是否含异常值判别器loss恒为0判别器太强生成器还没学会降低判别器学习率或增大生成器的训练频率4.3 从零到收敛的训练时间线参考我以一个具体实验为例给出训练时间线的参考。使用LOL数据集的485张低光图和485张正常光图batch size为8时每个epoch大约有60个迭代。结合10个epoch保存一次模型单卡RTX 3090从零训练到120个epoch大约需要11到13小时。前10个epoch模型快速学习增强效果从无到有图像亮度和色彩逐渐恢复。此阶段生成器损失从初始值快速下降判别器损失也随之出现波动。10到50个epoch增强效果逐步精细局部细节如边缘锐度、纹理清晰度得到改善。局部判别器开始发挥更明显的作用输出图像的局部区域不再出现奇怪的颜色斑块。50到100个epoch模型进入微调阶段loss曲线趋于平稳增强效果变化不大但更稳定。此时要注意不要过拟合建议在验证集上选择最佳epoch的模型权重进行最终评估。100到120个epoch学习率衰减阶段模型收敛到最优区域。5. 常见问题与排查技巧实录5.1 显存溢出与训练崩溃的应对这是我在训练过程中遇到的第一类问题也是初学者最容易踩的坑。有几次我试图把batch size从8提高到16结果瞬间就报错CUDA out of memory。如果你也遇到类似的情况按以下顺序排查查看GPU占用情况用nvidia-smi查看其他进程是否占用了显存必要时kill掉无关进程。降低batch_size从8降到4或2观察显存占用变化。降低批量大小对训练效果影响并不大尤其是使用Adam优化器的情况下。降低--fineSize400改为320或256这是最直接有效的办法。使用torch.cuda.empty_cache()在代码中每隔几个迭代手动释放缓存可以缓解显存碎片化问题。如果以上方法都尝试后仍然溢出建议检查一下你的PyTorch版本是否与CUDA驱动匹配。我遇到过因为CUDA驱动版本过旧导致PyTorch无法利用全部显存的情况升级驱动后问题解决。5.2 训练不收敛或者模型崩溃的排查思路训练不收敛是GAN复现中常见的现象具体表现为loss值出现NaN、生成器输出全黑或全白图像等。排查思路如下数据问题检查首先确认数据集路径是否正确图像是否能够正常加载。我遇到过因为数据集文件名包含中文字符导致读取失败的情况这时候把文件名改为纯英文的编号即可。归一化检查输入图像是否被正确映射到了[-1, 1]区间。如果原始图像是0到255的uint8类型而你没有做归一化训练必定崩溃。学习率检查GAN训练对学习率非常敏感。如果使用默认的0.0001训练时loss一直震荡尝试将学习率降低到0.00003甚至更低看loss是否能够稳定下降。梯度裁剪在生成器和判别器的反向传播之前添加一个梯度裁剪操作范数限制在5以内。这一步能有效防止梯度爆炸torch.nn.utils.clip_grad_norm_(model_G.parameters(), 5.0)5.3 代码版本兼容性问题与解决记录官方代码中使用了torchvision.models.vgg19_bn(pretrainedTrue)来构建VGG特征提取器。在PyTorch 1.8.1中这个模型仍然可以正常加载预训练权重但在PyTorch 2.x的版本中pretrainedTrue的写法已经被官方弃用需要改为weightstorchvision.models.VGG19_BN_Weights.DEFAULT。如果你坚持用新版PyTorch记得做以下修改# 旧代码PyTorch 1.x from torchvision.models import vgg19_bn vgg vgg19_bn(pretrainedTrue) # 新代码PyTorch 2.x from torchvision.models import vgg19_bn, VGG19_BN_Weights vgg vgg19_bn(weightsVGG19_BN_Weights.DEFAULT)另外官方代码中的data.py文件里自定义的Dataset类在PyTorch 1.8.1之后版本中可能因为torchvision.transforms接口变化而报错。我的建议是保持PyTorch 1.8.1不变这样最省心。如果你必须使用新版本需要将data.py中的transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))用transforms.Normalize(mean[0.5, 0.5, 0.5], std[0.5, 0.5, 0.5])替代并将所有torchvision.transforms中的lambda函数改写为显式函数定义。5.4 预训练权重加载的几个易错点训练过程中需要加载预训练的VGG19权重来计算特征保留损失很多人在这一步会遇到下载超时的问题。我建议使用torchvision自带的权重缓存机制首次使用之前手动下载权重文件将其放到~/.cache/torch/hub/checkpoints/目录下。具体操作为# 手动下载vgg19_bn权重文件 wget -P ~/.cache/torch/hub/checkpoints/ https://download.pytorch.org/models/vgg19_bn-c79401a0.pth下载完成后再运行训练脚本时就不会因为网络问题而卡住。另一个常见的错误是加载权重时把模型结构搞混。VGG19_bn和VGG19的结构是不一样的一个是带BatchNorm的版本一个不带权重文件不能混用。使用官方代码时一定注意用的是vgg19_bn还是vgg19保持一致。5.5 训练后模型评估的实操要点训练完成后通常会在测试集上计算PSNR峰值信噪比和SSIM结构相似度作为客观评价指标。我这里给出评估脚本的大致流程import numpy as np from skimage.metrics import structural_similarity as ssim from skimage.metrics import peak_signal_noise_ratio as psnr def evaluate(model, test_loader, device): model.eval() psnr_list, ssim_list [], [] with torch.no_grad(): for low_img, high_img in test_loader: low_img low_img.to(device) enhanced model(low_img) # 将输出从[-1,1]映射回[0,1] enhanced (enhanced 1) / 2 high_img (high_img 1) / 2 # 转为numpy数组计算指标 enh_np enhanced.cpu().numpy().transpose(0, 2, 3, 1) high_np high_img.cpu().numpy().transpose(0, 2, 3, 1) for i in range(enh_np.shape[0]): psnr_list.append(psnr(high_np[i], enh_np[i], data_range1.0)) ssim_list.append(ssim(high_np[i], enh_np[i], multichannelTrue, data_range1.0)) print(fPSNR: {np.mean(psnr_list):.4f}, SSIM: {np.mean(ssim_list):.4f})需要注意的是PSNR和SSIM是参考图像与增强图像之间的指标对于无监督增强任务来说它们只能部分反映模型效果。建议同时保存一组可视化对比图人工观察色彩、纹理、边缘等维度的表现。如果一张低光图的暗部细节恢复得很好但整体色调偏蓝或偏绿PSNR指标可能不会太高但视觉效果却是可用的。6. 延伸思考与后续扩展建议6.1 在自有数据集上微调的方案复现官方代码只是第一步在实际项目中如果我们有自己拍摄的低光图像就可以在EnlightenGAN基础上做微调。操作方法很简单将自有数据集按照相同目录格式整理好使用预训练好的权重作为初始化参数用较小的学习率继续训练20到30个epoch。在实践中我发现这种微调方式能够快速适应特定场景的光照分布效果比完全从零训练更好收敛速度也更快。举个例子如果你要处理的是夜间监控视频那么从官方权重开始微调只需要很少量的样本就能获得不错的增强效果。6.2 把EnlightenGAN嵌入到自己的实际项目中复现最终还是要服务于实际项目需求EnlightenGAN这样的无监督低光增强模型可以方便地集成到图像质量和视频质量相关的业务中。在实际工程中model_G训练好之后会导出为ONNX格式或TorchScript格式然后部署到服务端或前端。以下是一个简单的ONNX导出例子import torch from models import Generator model_G Generator() model_G.load_state_dict(torch.load(outputs/checkpoints/best_G.pth)) model_G.eval() dummy_input torch.randn(1, 3, 400, 400) torch.onnx.export(model_G, dummy_input, enlighten_gan.onnx, opset_version11, input_names[input], output_names[output])导出ONNX时需要特别注意Attention U-Net中的一些自定义操作是否能够被ONNX算子集支持。如果遇到不支持的算子可以将部分操作改写为PyTorch基础函数后再导出或者直接用TorchScript格式替代ONNX兼容性更好。6.3 从复现到创新的进阶思路复现文献代码不是终点理解它之后的改进空间在哪里才是关键。EnlightenGAN的生成器是基于Attention U-Net的这是一种轻量且高效的结构后续的很多低光增强方法都在此基础上进行改进例如将U-Net替换为具有更强特征表达能力的Transformer块或是在特征保留损失中引入感知损失。一个可行的改进方向是引入频率域的损失惩罚增强结果在频域上的失真从而保留更多高频细节。我自己也在尝试将EnlightenGAN的训练框架与最新的扩散模型思路结合让增强后的图像在细节真实性和自然度上更上一层楼。如果你也是正在复现EnlightenGAN的开发者希望这篇记录能帮你少走一些弯路。当然每个人的硬件环境、数据情况和具体需求都不一样参数上不可能一套方案走到底。根据我的经验最有效的做法是先在完整数据集中各取20张图像跑一个快速实验确认代码能够跑通、损失值在正常范围再启动完整训练。这样即使遇到问题排查起来也不会太痛苦。等你把整个流程跑通对GAN的训练机制、判别器的设计逻辑、特征损失的作用都会有更深的理解那时候再去看其他类似的工作会轻松很多。
返回列表