ARTICLE DETAIL

资讯详情

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

RetinexNet低光图像增强PyTorch复现:原理、代码与实战

RetinexNet低光图像增强PyTorch复现:原理、代码与实战 简介RetinexNet模型基于Retinex理论设计通过深度学习拆分光照与反射信息实现低光照增强、噪声抑制与色彩修复有效改善图像的整体视觉质量。下载包提供完整的PyTorch工程实现包括网络结构定义、数据加载与预处理、训练主脚本、参数配置和辅助工具函数并附带可直接用于训练和验证的图像数据集非常适合图像处理初学者、算法工程师和科研人员复现算法或进行二次开发。压缩包内共约2000个文件其中大多数为png格式样本图像还包含tar打包的数据集、若干JPG和BMP示例图片、Python源码文件以及LICENSE说明文档整体包体约817.73MB目录层次清楚便于按模块查阅。这一资源已有1619人学习或下载使用。启动训练脚本后即可利用自带数据快速训练模型也可把网络结构和训练流程扩展到去雾、低光增强、去噪等图像质量优化任务中兼顾理论理解与工程落地。1. 项目核心思路Retinex理论在低光增强中的落地低光图像增强一直是计算机视觉里一个很实际的需求——夜间监控、手机夜景拍摄、暗光环境下的目标检测都离不开这一环。传统方法直方图均衡化、Gamma校正虽然简单快速但容易把噪声一并放大色彩也容易失真。近几年的主流方案基本都转向了深度学习而RetinexNet就是其中一个非常有代表性的工作。先说清楚Retinex理论本身。Retinex这个词是“Retina”和“Cortex”的合成词它的核心假设是一张图像可以分解成反射分量Reflectance和光照分量Illumination两者逐像素相乘得到原始图像。用公式表达就是S R * I其中S是观测到的图像R是物体本身的反射属性理论上不受光照影响代表物体的固有颜色I是环境光照。低光图像的问题就在于I的亮度太低导致整体画面偏暗。所以低光增强的自然思路就是把R和I分开调整I到正常光照水平再和R重新合成。RetinexNet整篇工作的思路其实很清晰就是用一个神经网络去拟合这个物理模型。它设计了三个子网络分别解决“分解”、“调光”、“重建”三个问题Decom-Net将输入图像分解为R和IEnhance-Net对光照分量I做亮度调整Recon-Net用调整后的光照和反射分量重建增强图像作为Pytorch复现的项目最重要的就是把这套网络结构用代码准确表达出来同时配套好训练数据和评估流程让研究者能直接上手训练、测试、对比。2. 代码工程结构与环境准备拿到这个项目第一步建议先把目录结构理清楚。一个规范的RetinexNet工程通常长这样RetinexNet-PyTorch/ ├── data/ │ ├── train/ │ │ ├── low/ # 低光训练图像 │ │ └── high/ # 正常光照训练图像GT │ └── test/ │ ├── low/ │ └── high/ ├── models/ │ ├── decompose.py # Decom-Net 定义 │ ├── enhance.py # Enhance-Net 定义 │ └── retinexnet.py # 整体模型封装 ├── utils/ │ ├── dataset.py # 数据加载与预处理 │ ├── loss.py # 损失函数 │ └── metrics.py # PSNR/SSIM 计算 ├── train.py # 训练入口 ├── test.py # 测试/推理入口 └── checkpoints/ # 模型权重保存目录环境方面我建议直接用Anaconda管理避免把系统Python环境搞乱。创建一个干净的虚拟环境conda create -n retinex python3.8 conda activate retinex pip install torch1.13.0 torchvision0.14.0 --index-url https://download.pytorch.org/whl/cu116 pip install opencv-python pillow numpy tqdm scikit-image版本对齐是个容易踩坑的地方。Pytorch版本和CUDA版本必须配套否则GPU根本调用不起来。上面的命令装的是CUDA 11.6对应的Pytorch 1.13稳定性和兼容性都经过了大量项目验证。如果你用的是更新版本的Pytorch比如2.x代码层面的兼容性基本没问题但要注意有些旧API比如torch.nn.functional.grid_sample的某些参数在新版本里的行为变化建议跑之前先做一次小规模冒烟测试。3. 核心模块实现详解3.1 Decom-Net图像分解网络Decom-Net的作用是学习图像中的反射分量和光照分量。它的网络结构并不复杂核心是一个5层的全卷积网络前4层是“卷积ReLU”最后一层是“卷积Sigmoid”分别输出R和I两个分支。import torch import torch.nn as nn import torch.nn.functional as F class DecomNet(nn.Module): def __init__(self, channel64, kernel_size3): super(DecomNet, self).__init__() self.conv1 self._conv_block(3, channel, kernel_size) self.conv2 self._conv_block(channel, channel, kernel_size) self.conv3 self._conv_block(channel, channel, kernel_size) self.conv4 self._conv_block(channel, channel, kernel_size) # 反射分量和光照分量各自一个卷积头 self.reflectance nn.Conv2d(channel, 3, kernel_size, paddingkernel_size//2) self.illumination nn.Conv2d(channel, 3, kernel_size, paddingkernel_size//2) def _conv_block(self, in_ch, out_ch, kernel_size): return nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size, paddingkernel_size//2), nn.ReLU(inplaceTrue) ) def forward(self, x): feat self.conv1(x) feat self.conv2(feat) feat self.conv3(feat) feat self.conv4(feat) R torch.sigmoid(self.reflectance(feat)) I torch.sigmoid(self.illumination(feat)) return R, I这里有几个设计细节值得注意。第一最后一层用Sigmoid而不是ReLU是因为R和I的数值范围都在[0,1]之间Sigmoid天然满足这个约束。R代表反射率取值在0到1之间物理上是合理的I代表光照强度归一化到0到1也方便后续处理。第二输入图像本身也需要归一化到[0,1]我一般习惯在数据加载的时候就完成这个操作而不是在模型里做。第三这个网络是全卷积结构没有全连接层所以理论上可以处理任意尺寸的输入。但在实际训练时LOL数据集的图像分辨率基本固定400x600测试时如果遇到尺寸差异过大的图建议做resize保持一致性。3.2 Enhance-Net光照分量调整分解出I之后需要对它做亮度增强。RetinexNet原文用的是Enhance-Net本质上是一个多尺度特征提取调整的网络。我这里推荐一个简化但效果稳定的实现class EnhanceNet(nn.Module): def __init__(self, scale_factor2): super(EnhanceNet, self).__init__() self.scale_factor scale_factor # 多尺度特征提取 self.conv1 nn.Conv2d(3, 32, 3, padding1) self.conv2 nn.Conv2d(32, 32, 3, padding1) self.conv3 nn.Conv2d(32, 32, 5, padding2) self.conv4 nn.Conv2d(32, 32, 7, padding3) # 特征融合 self.fusion nn.Conv2d(32*3, 32, 3, padding1) # 上采样增强 self.up nn.Upsample(scale_factorscale_factor, modebilinear, align_cornersFalse) self.conv5 nn.Conv2d(32, 3, 3, padding1) def forward(self, x): f1 F.relu(self.conv1(x)) f2 F.relu(self.conv2(f1)) f3 F.relu(self.conv3(f1)) f4 F.relu(self.conv4(f1)) # 猫眼聚合 fusion_feat torch.cat([f2, f3, f4], dim1) fused F.relu(self.fusion(fusion_feat)) # 上采样到更高分辨率 up_feat self.up(fused) enhanced torch.sigmoid(self.conv5(up_feat)) return enhanced值得强调的是这里的Enhance-Net并不是直接生成最终的增强图像它的输入输出都是光照分量。输入是Decom-Net分解出来的低光I输出是增强后的I_enhanced。原始图像S R * I我们用I_enhanced替换I得到增强结果 S_enhanced R * I_enhanced。这个设计很巧妙——它把“亮度调整”和“色彩保持”解耦了。因为R不参与光照调整所以图像的颜色信息理论上不会因为提亮而发生偏移这比直接做端到端图像映射要更可控。3.3 重建与整体网络封装重建部分就是简单的逐像素乘法不需要额外学习参数。整体网络封装成一个类方便训练时统一调用class RetinexNet(nn.Module): def __init__(self): super(RetinexNet, self).__init__() self.decom_net DecomNet() self.enhance_net EnhanceNet(scale_factor1) # 这里scale_factor设为1保持分辨率 def forward(self, x): # 分解 R, I self.decom_net(x) # 光照增强 I_enhanced self.enhance_net(I) # 重建 S_enhanced R * I_enhanced return R, I, I_enhanced, S_enhanced前面的EnhanceNet定义里预设了上采样操作实际调用时要注意保持输入输出尺寸一致避免分辨率不匹配导致尺寸报错。我在跑通这个项目的时候发现最稳妥的做法是让EnhanceNet保持输入输出尺寸不变后续需要scale再单独处理。3.4 损失函数设计损失函数是RetinexNet训练中的关键。原文用了三个损失的组合重建损失Reconstruction Loss衡量分解后的R*I能否重建回原图光照平滑损失Illumination Smoothness Loss让光照分量保持空间平滑反射一致性损失Reflectance Consistency Loss同一场景的低光/正常光图像应共享相同的反射分量class RetinexLoss(nn.Module): def __init__(self): super(RetinexLoss, self).__init__() self.l1_loss nn.L1Loss() self.mse_loss nn.MSELoss() def reconstruction_loss(self, R, I, S): S_hat R * I return self.l1_loss(S_hat, S) def illumination_smoothness_loss(self, I): # 光照平滑计算光照分量的梯度的L1范数 dx torch.abs(I[:, :, :, :-1] - I[:, :, :, 1:]) dy torch.abs(I[:, :, :-1, :] - I[:, :, 1:, :]) return torch.mean(dx) torch.mean(dy) def reflectance_consistency_loss(self, R_low, R_high): return self.mse_loss(R_low, R_high)计算光照平滑损失的时候不能简单粗暴地对整个I做梯度惩罚因为光照分量在物体边缘处本来就允许有突变比如室内室外交界处如果强制全局平滑会把边缘糊掉。实际操作中我曾试过加一个边缘感知权重比如在梯度大的地方减弱惩罚但这是后话了先按原始设计跑通再说。4. 数据集准备与预处理RetinexNet最常用的训练数据集是LOLLo-Light数据集它提供了成对的低光/正常光图像非常适合做监督学习。这个项目配套的数据集也主要基于LOL。LOL数据集的train集包含485对图像test集包含15对图像分辨率大约400x600。数据集结构分为low和high两个文件夹文件名一一对应。我在准备数据时写了这样一个简单的加载器import os from torch.utils.data import Dataset from PIL import Image import torchvision.transforms as transforms class LOLDataset(Dataset): def __init__(self, root_dir, modetrain, crop_size384): self.low_dir os.path.join(root_dir, mode, low) self.high_dir os.path.join(root_dir, mode, high) self.low_images sorted(os.listdir(self.low_dir)) self.crop_size crop_size def __len__(self): return len(self.low_images) def __getitem__(self, idx): low_path os.path.join(self.low_dir, self.low_images[idx]) high_path os.path.join(self.high_dir, self.low_images[idx]) low_img Image.open(low_path).convert(RGB) high_img Image.open(high_path).convert(RGB) # 随机裁剪 if self.crop_size 0: w, h low_img.size cw min(w, self.crop_size) ch min(h, self.crop_size) x torch.randint(0, w - cw 1, (1,)).item() y torch.randint(0, h - ch 1, (1,)).item() low_img low_img.crop((x, y, x cw, y ch)) high_img high_img.crop((x, y, x cw, y ch)) # 转tensor并归一化到[0,1] to_tensor transforms.ToTensor() low_tensor to_tensor(low_img) high_tensor to_tensor(high_img) return low_tensor, high_tensor随机裁剪是我特别加的一个数据增强策略。原始LOL图像只有400x600如果整图输入训练一张图只贡献一个样本太浪费了。通过随机裁剪成384x384的块既相当于数据增广又能显著增加训练样本量。你也可以把这个尺寸调小到256以支持更大的batch size但过小的裁剪会让光照信息的全局一致性变差这一点需要权衡。5. 训练流程与参数设置训练是一个两阶段过程先单独训练Decom-Net再联合训练整个网络。但我在实际操作中发现如果数据量不大直接端到端联合训练也能收敛得不错只是需要把各个损失的权重调好。我的训练配置思路很直接用Adam优化器初始学习率0.001训练200个epoch学习率在第100个epoch时衰减到0.0001。batch size在12G显存条件下设为8比较稳妥图像裁剪尺寸384x384。如果显存不够可以把batch size降到4或裁剪尺寸降到256。训练主循环的核心代码框架def train_one_epoch(model, dataloader, optimizer, criterion, device): model.train() total_loss 0.0 for low_img, high_img in dataloader: low_img low_img.to(device) high_img high_img.to(device) # 前向 R_low, I_low, I_enhanced, S_enhanced model(low_img) R_high, I_high, _, _ model(high_img) # 损失计算 recon_low criterion.reconstruction_loss(R_low, I_low, low_img) recon_high criterion.reconstruction_loss(R_high, I_high, high_img) smooth_low criterion.illumination_smoothness_loss(I_low) smooth_high criterion.illumination_smoothness_loss(I_high) consis criterion.reflectance_consistency_loss(R_low, R_high) loss recon_low recon_high 0.1 * (smooth_low smooth_high) 0.01 * consis optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() return total_loss / len(dataloader)这里有个值得反复强调的点低光图像和正常光图像共享Decom-Net的权重但它们在同一个batch里同时输入网络。这意味着网络要同时学会两件事——从低光图里提取R和I以及从正常光图里提取R和I。反射一致性损失强制让同一场景、不同光照条件下提取出来的R尽量一致这正是训练Decom-Net参数的核心驱动力。我个人的调参经验是平滑损失的权重系数0.1不能设太大。如果设到0.5以上光照分量会被磨得非常平导致增强结果出现明显的块状伪影但如果设得太小0.01以下光照分量噪声很大增强后的图像会保留很多原图噪声。0.1是我试过比较平衡的值。6. 测试与效果评估训练完成后用test.py对测试集进行推理就非常简单了def test(model, low_dir, output_dir, device): model.eval() os.makedirs(output_dir, exist_okTrue) with torch.no_grad(): for img_name in sorted(os.listdir(low_dir)): img_path os.path.join(low_dir, img_name) img Image.open(img_path).convert(RGB) # 预处理 to_tensor transforms.ToTensor() img_tensor to_tensor(img).unsqueeze(0).to(device) # 推理 R, I, I_enhanced, S_enhanced model(img_tensor) # 保存结果 output S_enhanced.squeeze(0).cpu() to_pil transforms.ToPILImage() output_img to_pil(output) output_img.save(os.path.join(output_dir, img_name))评估指标通常用PSNR峰值信噪比和SSIM结构相似性。这两个指标都是把增强结果和ground truth对比计算。PSNR越高代表像素层面越接近SSIM越高代表结构层面越相似。我在LOL测试集上跑通这个项目PSNR一般在19-21dB左右SSIM在0.75-0.82之间不同初始化参数会有差异。如果你的目标是做人眼观察的视觉效果指标只能作为参考。我见过有些模型PSNR不低但看起来颜色发灰反而是PSNR稍低的模型看起来更自然。所以测试阶段建议把增强结果图保存下来人工过一遍比只盯指标要靠谱得多。7. 常见问题与坑点排查7.1 显存不足OOM这是跑深度模型最经典的问题。如果报CUDA out of memory我一般按顺序排查先调小batch size其次调小裁剪尺寸最后考虑用混合精度训练torch.cuda.amp。另外检查一下是不是模型参数里有什么地方意外地把中间特征图缓存太多比如多个分支叠加的tensor没有被及时释放。7.2 训练损失不下降如果loss在训练初期就卡住不动大概率是学习率设置问题或者数据没对。先做一次过拟合测试——拿一批数据比如8张图反复训练看loss能不能降到很低。如果过拟合测试能通说明模型没问题是训练策略或数据分布的问题如果连过拟合都做不到就要怀疑代码哪里有bug了。7.3 增强结果颜色偏灰/偏白这个现象很常见根源多半是光照分量I被过度增强导致R * I_enhanced的输出饱和。可以检查一下I_enhanced的数值分布如果大量接近1说明增强过度了。解决方案是调整Enhance-Net的输出范围或者在loss里加一个对I_enhanced的约束项让它不要太极端。7.4 图像出现棋盘格伪影如果网络里有转置卷积层棋盘格伪影是个经典问题。RetinexNet原文的Enhance-Net在原有版本中使用了类似上采样的结构在Pytorch复现代码里如果处理不当就容易出现。我建议统一用nn.Upsample(modebilinear)配合普通卷积来替代转置卷积视觉效果会干净很多。7.5 数据集路径和文件匹配很多人在加载自己的数据集时遇到文件名对不上的问题。LOL数据集的low和high文件夹命名风格是xxx.png这种数字编号但如果换成其他数据集命名规则可能完全不同。我在写dataset时建议添加一个__init__参数专门控制文件名对应逻辑或者直接按索引对应这样遇到命名不统一的数据集时不用重写代码。8. 进一步优化方向跑通RetinexNet只是第一步。如果想让效果更好或者推向实际应用下面几个方向值得探索模型轻量化把Decom-Net和Enhance-Net替换为更紧凑的网络结构比如MobileNet风格的深度可分离卷积推理速度会有明显提升加入感知损失在计算增强结果时除了L1/MSE损失还可以加入VGG特征空间的感知损失能有效提升视觉质量无监督训练如果你的应用场景没有配对数据可以尝试用Zero-Reference风格的训练策略不再依赖低光/正常光的成对图像视频增强把单张图像增强扩展到视频序列加入时间维度的一致性约束避免连续帧闪烁我在实际使用中发现RetinexNet这种“分解-调整-重建”的范式比直接端到端的黑盒映射更有可解释性调参时出了问题也更容易定位。它可能不是当前指标最高的方法但作为学习和理解“物理模型深度学习”结合的范本价值非常大。如果你正打算入门低光图像增强或者研究Retinex相关方法拿这个Pytorch工程当起点会少走很多弯路。本文还有配套的精品资源点击获取
返回列表