夜间监控的画面一直是个老大难:光线不足时噪点连成一片,暗部细节被按死,强行拉高亮度,天空和墙面又泛出奇怪的紫红色。我之前做低光图像增强项目时被这个问题卡了很久,后来干脆从实际需求出发,打磨了一个轻量级低光增强网络 StarNet。它做的事情不花哨:把暗部亮度拉起来,同时把噪声压住,尽量保留原有细节和色彩。模型参数控制在 1M 上下,在端侧设备上可以实时跑,既适合刚入门的同学研究完整的训练和部署链路,也适合工程背景的开发者直接拿来做夜间图像增强的基线方案。这篇文章把这套东西从设计思路、数据准备、网络结构、训练参数到真实场景中的踩坑经验全部讲清楚。
1. 自研 StarNet 之前,我为什么放弃“直接调亮度”
1.1 传统低光增强方法的问题在哪
很多人拿到暗图第一反应是调 gamma 曲线,或者做直方图均衡化。这两种办法在亮度分布均衡的图上确实有效,但放到真实夜拍场景里就露馅了。gamma 校正是逐像素的非线性映射,它把暗部整体提亮的同时,也会把传感器噪声一起放大,结果图像是亮了,噪点也更明显了。直方图均衡化更激进,它把暗部灰度范围拉伸到整个动态范围,经常导致对比度过猛,人脸皮肤变成蜡黄色,暗区的色调完全失真。
Retinex 理论看起来更科学,它把图像分解成反射分量和光照分量,然后只对光照做调整。思路没毛病,但光照估计本身就是病态问题,在亮度突变的边缘很容易产生光晕。而且传统 Retinex 算法需要多尺度高斯滤波,计算量不小,做出来的效果也未必比一张精心调过的 gamma 图好多少。
所以低光增强本质上不是一个逐像素的映射问题,它需要模型理解图像的空间上下文,知道哪里是噪声、哪里是纹理、哪里有边缘,然后才能决定对该提亮多少。这些先验靠人工规则很难写全,反而是让网络从成对数据里自己学更靠谱。
1.2 现有深度学习方案的痛点与 StarNet 的定位
到这一步目光自然转向深度学习。当时我对比了若干主流方案:Zero-DCE 不做成对数据监督,通过估计一系列亮度曲线来增强图像,思路很简洁,但缺少强监督约束时容易过度增强,天空和皮肤经常出现过饱和;RetinexNet 按 Retinex 理论把网络拆成分解和增强两个模块,效果好一些,但模块之间需要分别调优,端到端训练不友好,模型体积也比较大;MIRNet 这类大模型效果更强,但动辄几十兆参数,边缘设备根本跑不动。
我当时的实际场景是给夜间监控画面做实时增强,输入分辨率 1080P,推理时间要求控制在几十毫秒内。这样一来选型标准就很清晰:单阶段端到端、参数量尽量小、推理延迟低、部署友好。所以我最后既没有照搬哪一个现成模型,而是针对这个约束设计了一个很朴素的编码器-解码器结构,并在后续迭代里不断压缩和优化它的瓶颈点。StarNet 这个名字也是那时候定的,它更像是一个轻量基线的代号,核心在“够用、好部署、容易调”。后面所有设计决策,都围绕这三个关键词展开。
2. 核心细节解析与实现要点
2.1 数据集的选择与预处理思路
训练低光增强模型,我优先使用的是公开的低光/正常光成对数据集。LOL 数据集包含真实场景下不同曝光等级的照片对,覆盖室内外多种环境,比较适合做基线验证;LSRW 也有大量真实拍摄的成对图像,风格偏夜间街景和建筑,和监控场景比较接近。如果条件允许,可以再混入一小部分自己拍摄的夜间长曝光照片,用来提升模型的泛化能力。
预处理这里有一个关键选择,我把它直接写在前面:我没有在 RGB 三通道上直接做增强,而是先把图像从 RGB 转换到 YCbCr 色彩空间,只对 Y(亮度)通道做增强,然后和原始 CbCr 通道合并再转回 RGB。
这样设计有实际依据。低光图像的主要问题集中在亮度不足和亮度通道上的噪声,色度通道相对稳定。如果让网络同时处理 RGB 三通道,模型很容易学到错误的色相偏差,比如把暗部拉亮后整个画面偏绿或偏紫。只处理 Y 通道相当于把“提亮去噪”和“色彩恢复”强行解耦,网络专注学亮度结构,色彩偏移问题天然少很多。这也是后来很多实际项目里会用到的小技巧,但对新人来说不太容易想到。
预处理流程具体如下:
- 把成对图像调整到长边不超过 1024,避免直接缩放太狠导致细节丢失。
- 读取图像后转换为 YCbCr 色彩空间,只取 Y 通道作为网络的输入和目标。
- 随机裁剪到 256×256,同时给低光和正常光两张图用相同的裁剪位置。
- 随机水平翻转,翻转操作同样要保持一致。
- 对输入 Y 通道做轻微的高斯模糊增强,模拟大光圈镜头的夜间成像模糊。
这里最需要注意的是第 3 步和第 4 步的“一致性”:低光图和目标图必须做一模一样的几何变换,否则配对被破坏,训练出来的模型会学出重影。另外也不要为了增加数据多样性而对输入加太强的高斯噪声,模型会误以为噪声是有效纹理,最后把细节一起抹掉。
2.2 网络结构设计:轻量化不是单纯堆小通道数
StarNet 的结构走的是简单有效的路线:一个三层的卷积编码器,一个三层的解码器,中间有跳连,加上一个全局光照统计分支。整体结构很接近轻量化 U-Net,但细节上有几个改动值得单独拿出来说。
编码器的通道数从 32 开始,逐层翻倍到 64、128。每层由一个 3×3 卷积、BatchNorm 和 LeakyReLU 组成。第二层和第三层用 stride=2 做下采样,扩大感受野。这样设计后,网络能获取更大范围的光照信息,不会只盯着局部像素疯狂提亮度。解码器则从 128 通道逐步降回 32,最后用一个 3×3 卷积输出增强后的 Y 通道。
跳连是必须保留的。没有跳连,编码器下采样会把高频边缘细节丢掉,解码器上采样时又补不回来,最终生成的图像会明显偏软,边缘像蒙了一层雾。有了跳连,模型可以直接把编码器提取到的边缘特征传给解码器,细节保真度会上升一个档次。
全局光照统计分支是我后来加上的。具体做法是对编码器最深层的特征做全局平均池化,得到一个很小的向量,然后通过一层全连接把它变成和该层特征相同通道数的权重,再拼回原来的特征图。这个分支让网络能感知整张图的平均亮度和光照分布。它的价值体现在拍摄光照极不均匀的场景:前景亮、背景全黑,如果没有全局信息,模型容易顾此失彼。
模型整体参数量大概在 0.9M 到 1.1M 之间,取决于是否有额外的分支。用 256×256 输入在 1080Ti 上单帧推理大约是 5ms 左右,在端侧设备上经过量化后也能控制在几十毫秒内。为什么不直接用 Transformer 或者更复杂的结构?因为低光增强任务并不需要长距离建模到那种程度,3 层下采样的感受野已经足够覆盖常规夜景的空间结构,复杂结构只会白白增加延迟。
2.3 损失函数与训练策略的设计取舍
损失函数是模型效果的上限约束,我最终的损失组合是:
L = L1 + 0.2 * TV + 0.05 * perceptualL1 损失是主监督项,它要求预测的 Y 通道和目标 Y 通道逐像素尽量接近。L1 比 L2 更不容易产生过于平滑的结果,在图像增强任务里更常用,因为 L2 会把小的残差重罚过大,导致模型倾向于输出保守的均值。
TV(Total Variation)损失用来约束输出图像的平滑度。低光图像的噪声非常明显,模型如果只优化 L1,极容易把噪声当作细节重建出来。TV 损失鼓励相邻像素亮度变化尽量平滑,能有效抑制噪声。但是权重不能太高,之前我试过 0.5 以上的权重,结果噪声没了,纹理也变得像塑料一样。调到 0.2 左右在去噪和保纹理之间比较平衡。
感知损失用的是 VGG16 的 relu3_3 层特征。它衡量两个图像在高层语义特征上的距离,虽然不像像素级损失那样精确,但能帮助模型保留结构和边缘的自然感。感知损失的权重不必太大,0.05 就够了,它更多是辅助作用。如果训练资源有限,去掉这一项也能得到一个可用的模型,只是会产生细节偏少的问题。
训练策略上我选 Adam 优化器,初始学习率 1e-4,batch size 16,开混合精度训练以节省显存。学习率采用余弦退火方式在 100 个 epoch 内衰减到 1e-6。每个 epoch 大概迭代几百次,跑完整个训练过程在单张 1080Ti 上约需要几个小时,完全可接受。训练时还会做梯度裁剪,最大范数设为 1.0,防止个别极端样本把参数一下拉偏。
3. 从零实现:数据准备与训练全流程
3.1 环境配置与工程目录
先交代一下运行环境。我用的是 PyTorch 2.0 以上的版本,配合 torchvision 提供预训练 VGG16 来提取感知特征。图像处理用 OpenCV 和 NumPy,数据增强用 Albumentations 来做,关键是它能统一处理成对图像的变换。
如果你要复现,建议的工程结构如下:
starnet/ ├── dataset.py ├── model.py ├── train.py ├── export.py └── config.pyconfig.py 里集中管理数据路径、输入尺寸、学习率、batch size 等超参数。把所有可调参数集中在一起,后续做实验调节版本会非常方便,不需要在代码里到处翻。
3.2 数据读取与成对增强代码
数据读取的核心是保证低光图和目标图经过完全一致的几何变换。我直接基于 OpenCV 写了一个简单的 Dataset:
import cv2 import numpy as np from torch.utils.data import Dataset class LowLightDataset(Dataset): def __init__(self, low_paths, high_paths, img_size=256): self.low_paths = low_paths self.high_paths = high_paths self.img_size = img_size def _read_y_channel(self, path): img = cv2.imread(path) ycrcb = cv2.cvtColor(img, cv2.COLOR_BGR2YCrCb) return ycrcb[:, :, 0] def __getitem__(self, idx): low_y = self._read_y_channel(self.low_paths[idx]).astype(np.float32) / 255.0 high_y = self._read_y_channel(self.high_paths[idx]).astype(np.float32) / 255.0 # 随机裁剪到固定尺寸,位置保持一致 h, w = low_y.shape y = np.random.randint(0, h - self.img_size + 1) x = np.random.randint(0, w - self.img_size + 1) low_y = low_y[y:y + self.img_size, x:x + self.img_size] high_y = high_y[y:y + self.img_size, x:x + self.img_size] # 随机水平翻转,保持配对 if np.random.rand() > 0.5: low_y = low_y[:, ::-1] high_y = high_y[:, ::-1] low_y = np.expand_dims(low_y, 0) high_y = np.expand_dims(high_y, 0) return low_y, high_y这里把图像转成 Y 通道后,用一个(1, H, W)的维度送入网络。实际上 CbCr 通道在推理阶段从原图中取,不需要进网络。
3.3 模型定义与训练循环核心代码
模型定义直接写成 PyTorch 的nn.Module:
import torch import torch.nn as nn class StarNet(nn.Module): def __init__(self, base_channels=32): super().__init__() self.enc1 = self._block(1, base_channels) self.enc2 = self._block(base_channels, base_channels * 2, stride=2) self.enc3 = self._block(base_channels * 2, base_channels * 4, stride=2) self.global_pool = nn.AdaptiveAvgPool2d(1) self.gp_fc = nn.Linear(base_channels * 4, base_channels * 4) self.dec3 = self._block(base_channels * 4, base_channels * 2) self.dec2 = self._block(base_channels * 2, base_channels) self.dec1 = nn.Conv2d(base_channels, 1, 3, 1, 1) def _block(self, in_c, out_c, stride=1): return nn.Sequential( nn.Conv2d(in_c, out_c, 3, stride, 1), nn.BatchNorm2d(out_c), nn.LeakyReLU(0.2) ) def forward(self, x): e1 = self.enc1(x) e2 = self.enc2(e1) e3 = self.enc3(e2) gp = self.global_pool(e3).view(e3.size(0), -1) gp = torch.sigmoid(self.gp_fc(gp)).view(e3.size(0), -1, 1, 1) e3 = e3 * gp d3 = self.dec3(torch.cat([e3, e2], dim=1)) d2 = self.dec2(torch.cat([d3, e1], dim=1)) out = self.dec1(d2) return out这里跳连用了torch.cat做通道拼接,而不是常见的加法拼接。通道拼接能让解码器同时看到来自不同层的特征,虽然会多占用一些内存和计算量,但效果更好。
训练循环核心部分如下:
def train_one_epoch(model, loader, opt, vgg, epoch): model.train() total_loss = 0.0 for low_y, high_y in loader: low_y, high_y = low_y.cuda(), high_y.cuda() pred = model(low_y) loss_l1 = F.l1_loss(pred, high_y) loss_tv = tv_loss(pred) loss_perc = perceptual_loss(pred, high_y, vgg) loss = loss_l1 + 0.2 * loss_tv + 0.05 * loss_perc opt.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) opt.step() total_loss += loss.item() print(f"Epoch {epoch} loss: {total_loss / len(loader):.6f}")TV 损失的计算就是对输出做一次水平和垂直方向的差分:
def tv_loss(img): h_tv = torch.mean(torch.abs(img[:, :, 1:, :] - img[:, :, :-1, :])) w_tv = torch.mean(torch.abs(img[:, :, :, 1:] - img[:, :, :, :-1])) return h_tv + w_tv感知损失直接用预训练 VGG16 的第 3 层 ReLU 特征做 MSE:
def perceptual_loss(pred, target, vgg): pred_feat = vgg(torch.cat([pred] * 3, dim=1)) target_feat = vgg(torch.cat([target] * 3, dim=1)) return F.mse_loss(pred_feat, target_feat)因为模型只输出单通道 Y,送入 VGG 前先复制成三通道,这一点很简单也容易漏。
3.4 导出 ONNX 与端侧部署路线
训练完的模型要落地,我首先导出 ONNX。PyTorch 自带的导出接口足够用:
model.eval() dummy = torch.randn(1, 1, 256, 256) torch.onnx.export( model, dummy, "starnet.onnx", input_names=["low_y"], output_names=["enhanced_y"], dynamic_axes={"low_y": {2: "height", 3: "width"}, "enhanced_y": {2: "height", 3: "width"}}, opset_version=11 )动态轴是必须加的,因为实际推理时输入可能是 512×512、768×768 甚至更大,不可能固定成 256×256。导出后用 onnxruntime 再跑一遍,和 PyTorch 的输出做对比,确认误差在可接受范围。
从 ONNX 再转端侧格式就看平台了。Android 可以转 NCNN 或 TFLite,iOS 可以转 Core ML。我实测下来 int8 量化对低光图像的影响比较明显,因为暗部动态范围窄,量化误差会被衬托得更清楚。如果目标设备支持 FP16,优先用半精度;只有必须压缩体积时才做 int8,并建议用 QAT 而不是简单 PTQ。
4. 真实场景踩坑实录
4.1 过度增强导致的色偏问题
模型刚训练完在 LOL 测试集上指标很漂亮,PSNR 有 24 以上,SSIM 也不错,但一放到真实夜拍照片上就露馅:暗部被提亮后泛着一层淡青色,高光区域变得苍白,整个画面像蒙了层雾。
排查后发现,问题出在我后续处理时,直接把网络输出的 Y 通道替换原始 Y 通道,然后和原始 CbCr 通道合并。尽管网络学的是正常光照下的 Y,但普通照片的 CbCr 采集环境与训练数据存在差异,亮度变化后色度没有同步调整,就会产生色偏。
我的解决办法是做一个“亮度混合”:把增强后的 Y 通道与原始 Y 通道按照一定比例混合,enhanced_Y = 0.85 * enhanced_Y + 0.15 * original_Y。这样既保留了模型提亮的力度,又让图像的亮度变化不会过于夸张,色彩偏移的观感显著下降。另一个思路是直接让网络同时输出色度通道的修正量,但复杂度会上升,我在当前版本里没有采用。
4.2 边缘光晕和分块伪影
使用过程中另一个高频问题是在强边缘附近出现一圈光晕。比如夜空下建筑物的轮廓,原本清晰的边界被提亮后反而变得模模糊糊,边缘外侧有一圈亮度异常的过渡带。
造成这个现象的原因有两方面:一是 TV 损失权重过高时,模型为了追求平滑,会牺牲边缘的锐度;二是下采样层数过多后,高频信息在跳连恢复过程中没有被充分传递,解码器只能根据粗略的光照信息猜测边缘位置。
针对这个问题我做了两个改动。第一,把 TV 损失权重从 0.2 适度下调,给 L1 损失更多主导权,让模型优先还原结构;第二,在推理大图时不用整图直接输入,而是切成带重叠区域的小块分别推理,再把结果拼回去,重叠区域用线性羽化消除接缝。两个改动叠加后,边缘光晕基本看不到了。
4.3 训练 loss 震荡与不收敛
还有一个很让人头疼的情况是训练前期 loss 下降正常,但到了中后期开始震荡,偶尔还会突然飙升几个数量级。我逐项排查后发现,根源在于数据集中低光图像的亮度分布差异巨大:有的图整体很暗,有的图只有局部暗,极端样本会让梯度出现异常。单纯调低学习率只能缓解,无法根除。
我的解决方案有三步:第一,对所有输入做一次全局亮度归一化,把每张图的 Y 通道除以该图 Y 均值,让网络学习的输入更规范;第二,启用梯度裁剪,把梯度最大范数限制在 1.0;第三,用指数移动平均(EMA)维护一份模型参数副本用于最终推理,训练稳定性因此提升很明显,最后几个 epoch 的性能也更好。
4.4 推理速度不达标与模型裁剪
最初版本的 StarNet 通道数设得偏大,在 PC 上速度没问题,但在移动端 GPU 上依然达不到 30fps。于是我开始裁剪模型。先把编码器的最大通道从 128 降到 96,参数量下降了约 30%,PSNR 只掉了 0.3 左右,肉眼几乎没区别;再把解码器里的拼接改成加法,虽然效果略降,但推理速度又上了一个台阶。如果你的部署目标性能实在有限,这个裁剪思路可以参考。与其一上来就设计超大网络,不如先跑通一个能用的版本,再根据实际延迟倒推容量上限。
下面是我整理的一套常见问题速查表,方便后续排查:
| 问题现象 | 可能原因 | 解决方案 |
|---|---|---|
| 暗部发青或偏紫 | 色度通道与亮度不匹配 | 使用 YCbCr 只增强 Y 通道,并做 Y 与原始亮度混合 |
| 边缘出现光晕 | TV 权重过高或下采样过多 | 降低 TV 权重,推理时分块重叠并羽化 |
| 输出图像过亮发白 | 增强强度过大 | 降低输出混合比例,或对输出做亮度上限截断 |
| 细节被抹平 | L1 权重不足或输入加噪过强 | 加大 L1 权重,去除额外噪声增强 |
| 分块推理有接缝 | 分块之间没有重叠 | 使用 16~32px 重叠区域并做线性混合 |
| 训练 loss 震荡 | 极端样本干扰梯度 | 光照归一化、梯度裁剪、EMA |
我在实际做这个项目的过程中体会最深的一点是,指标好看不如观感可靠,PSNR 和 SSIM 只能说明模型学到了数据集内部的分布规律,真实世界里的夜拍环境永远比数据集更杂乱。所以无论模型设计得多巧妙,最后一定要留出充足的调优空间,在真实照片上反复看效果。如果你也想做一个类似的轻量级低光增强模型,建议先在我说的这个流程上跑通,再根据你的部署场景去裁剪和调整。后续这个方向还可以扩展做视频增强,把当前单帧模型加一个轻量的时序对齐模块,就能在低光视频流上实现更稳定的亮度输出,这也是我目前在探索的一条路。