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

资讯详情

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

HDRnet图像增强:基于双边网格的可学习仿射变换

HDRnet图像增强:基于双边网格的可学习仿射变换 1. 先拆解整体设计思路1.1 为什么是双边网格而不是卷积网络HDRnet全称是Deep Bilateral Learning for Real-Time Image Enhancement2017年发表在SIGGRAPH上。做图像增强这行的应该都知道传统方法里有一个绕不开的痛全局映射太粗糙局部映射算起来又太慢。你把一张照片亮度提起来天空容易过曝暗部细节还压着出不来你做一个3D LUT查表颜色统一了但边缘又容易发灰。HDRnet最狠的地方在于它用双边网格bilateral grid把“局部颜色映射”这件事做成了可学习、可实时计算的模块。我们先理解一下为什么要选双边网格。普通卷积网络做图像增强一路conv下去到最后一层要么做全局映射要么逐像素预测颜色调整参数。前者丢失了空间信息后者参数量爆炸且很容易出现颜色断裂。双边网格的妙处在于它把空间信息和颜色信息先“降维揉在一起”再用一个小的三维网格来存储局部仿射系数。这样空间分辨率不需要很高颜色分辨率也不需要很高就能表达足够复杂的局部映射关系。我最早看到这篇论文的标题时以为它只是一个工程优化技巧。真正自己动手复现之后才明白双边网格在这里不是用来磨皮降噪的它是一整个可微的查表结构。网络学习的不是“怎么滤波”而是“网格里每个节点该存一个什么样的仿射变换系数”。标题里的“摇身一变”指的就是这件事——传统双边滤波里那个靠人工设定或者启发式计算的引导图被一个神经网络端到端地学出来了。1.2 两条路径coarse网络和guide网络HDRnet的结构非常清晰分两条路径下路是coarse网络输入低分辨率的图像通常256x256甚至128x128通过一系列卷积与全连接层输出一个低分辨率的双边网格。这个网格的尺寸一般是16x16x16左右空间维度x颜色维度每个节点存着12个仿射系数一个3x3颜色矩阵加一个3维偏置。上路是guide网络输入原分辨率图像用于计算亮度与色度特征输出的是一组低分辨率引导图guidance map。这个引导图用来在slice操作中决定每个像素在双边网格里对应哪个节点以及如何做三线性插值。为什么不能只用一条路只用低分辨率会丢失边缘信息指导图的意义在于它把全分辨率图像上每一个像素的亮度位置映射到网格坐标然后通过三线性插值取出该点的仿射矩阵。这个过程和双边滤波完全同构输入图像决定系数再作用回输入图像自身。你在实际部署时不需要关心这些设计哲学但你需要理解一个关键点coarse网络的参数量远小于常规U-Net因为它不需要对每个像素预测数值只需要对一个16x16x16的网格做预测中间再插值。这也是HDRnet在1080p实时处理中跑得动的原因之一。2. 从双边滤波到可学习的双边网格2.1 传统双边滤波与它的问题传统双边滤波大家应该不陌生。它保边因为它同时考虑空间距离和像素值差异边缘两侧像素参与权重很小。但双边滤波在图像增强上的用处其实有限——你最多用它做细节分离配合其他算子做局部色调映射。HDRnet换了个思路它不直接滤波而是用双边网格来“组织”一个局部的颜色变换域。双边网格的形式化描述是把一个二维图像加颜色维度堆成三维坐标(i, j, r)在空间维度降采样在颜色维度量化。每一个格子里存的值可以看作是这个颜色亮度区间对应的局部特征。原来双边滤波的核权重要靠手算HDRnet直接让这些格子里的值变成网络参数或者网络输出。我自己用一句话总结这个设计传统双边滤波是在“用局部邻域的加权平均”去处理像素HDRnet是在“用局部邻域共享的仿射变换”去处理像素。前者做平滑后者做映射。2.2 slice操作的数学本质slice是整个HDRnet里我最欣赏的操作。它的输入是低分辨率双边网格输出是全分辨率变换系数图像。具体来说对于图像上某个像素p先计算它在双边网格中的归一化坐标[ x \frac{i}{W} \cdot (N_s - 1), \quad y \frac{j}{H} \cdot (N_s - 1), \quad z I(p) \cdot (N_c - 1) ]其中W、H是图像宽高N_s是网格空间分辨率论文里是16N_c是颜色分辨率论文里是8或16I(p)是像素的亮度或色度归一化值。然后在这个三维坐标周围做三线性插值把周边的8个网格节点按权重混合得到该像素对应的仿射系数矩阵。我在复现时最初犯过一个错误把slice当成双线性插值来写。双线性只对二维坐标插值slice是三维插值而且第三个维度颜色维度和空间维度不能混为一谈。如果你的输入是RGB图像颜色维度是3通道分别量化还是统一亮度量化会直接影响输出画质。论文的tensorflow实现里用的是亮度灰度做颜色坐标因为这样可以保持边缘一致性也避免RGB三个通道分别插值产生色偏。2.3 仿射变换如何作用在像素上HDRnet对每一个像素使用取出的3x3颜色变换矩阵加上一个偏置向量对RGB三通道做一次仿射变换。注意这是像素级逐点的操作不是区域性的。仅仅一个逐点仿射变换当然学不出太复杂的效果但如果这个仿射变换在空间上是平滑变化的由双边网格和slice保证了这一点叠加起来就能表达极其丰富的局部色调映射。说一下仿射变换在这里和传统图形学仿射的区别。图形学里的仿射变换无论是做旋转、缩放还是平移作用对象是坐标系里的几何体通常有一个明确的矩阵乘法公式对整张图或某个区域应用同一个变换。HDRnet的仿射变换是每一个像素一个矩阵而且这个矩阵是从双边网格里学出来的。它不再是“固定的几何变换”而是“像素值域空间中的局部线性映射”。如果你做过传统图像增强里的色彩校正对3x3颜色矩阵一定不陌生。常见的白平衡就常常简化成对角矩阵颜色分级会用到3x3线性变换加offset。HDRnet的仿射变换本质上就是把这些人工调出来的经验参数变成了一个由输入图像自适应的函数。2.4 光照域的统一视角很多人忽略的一件事是HDRnet的网格坐标z轴选的是亮度域。这意味着两个空间不相邻但亮度相近的像素会在同一个网格节点附近共享一组仿射系数。这和很多传统方法里对亮度进行分段映射的思路一致但比传统的分段映射更平滑——三线性插值保证了亮度方向的连续性。这也解释了为什么HDRnet能处理HDR到LDR的压缩。高动态范围图像里暗部区域和亮部区域差异极大全局映射无法兼顾。如果用亮度域划分网格暗部和亮部的像素各自落在颜色维度的不同区间学习到的仿射系数就完全不同暗部被提亮的同时亮部不会被压死。这比3D LUT查表颜色变换更本质因为LUT是固定的HDRnet的网格是根据输入动态生成的。3. 自己搭一个简化版HDRnet3.1 工具选型与环境配置我建议用PyTorch复现HDRnet因为它对自定义层的支持友好slice操作可以用grid_sample实现也可以用纯NumPy先写一版验证逻辑。我的硬件环境是GeForce RTX 3060显存12GB跑512x512的图批量训练没有压力。环境如下Python 3.9 PyTorch 1.12 OpenCV 4.6 NumPy 1.23训练数据方面HDRnet自带数据集的构造方式是把原始RAW图做成两对一对原始线性图一对经Adobe Lightroom人工调色的图。如果你手头没有RAW数据可以退而求其次用几张高清图自己做一版“暗部提亮局部对比度增强”作为伪目标图。我实测下来用伪标签也能训练出不错的效果但泛化性会差一些后面会讲到怎么改善。3.2 核心代码slice操作与仿射变换Slice操作是实现难点。我用PyTorch配合unfold实现了一个简化版思路是先将双边网格按坐标索引展开再对每个像素采样对应的affine系数。核心代码如下def slice_grid(grid, guidemap): grid: (B, D, C) 双边网格, D Ns*Ns*Nc guidemap: (B, 3, H, W) 指导图, 包含xy坐标和亮度归一化值 B, _, H, W guidemap.shape # 将guidemap中的x, y, z坐标映射到网格节点索引 gx guidemap[:, 0, :, :] # 空间x坐标, 范围0~Ns-1 gy guidemap[:, 1, :, :] # 空间y坐标 gz guidemap[:, 2, :, :] # 亮度坐标, 范围0~Nc-1 x0 torch.floor(gx).long() y0 torch.floor(gy).long() z0 torch.floor(gz).long() x1 torch.clamp(x0 1, maxNs - 1) y1 torch.clamp(y0 1, maxNs - 1) z1 torch.clamp(z0 1, maxNc - 1) def gather(x, y, z): idx x * Ns * Nc y * Nc z # flatten index return grid.gather(1, idx.unsqueeze(1).expand(B, 3, H, W)) wx (gx - x0.float()).unsqueeze(1).unsqueeze(-1) wy (gy - y0.float()).unsqueeze(1).unsqueeze(-1) wz (gz - z0.float()).unsqueeze(1).unsqueeze(-1) c000 gather(x0, y0, z0) c100 gather(x1, y0, z0) c010 gather(x0, y1, z0) c110 gather(x1, y1, z0) c001 gather(x0, y0, z1) c101 gather(x1, y0, z1) c011 gather(x0, y1, z1) c111 gather(x1, y1, z1) # 三线性插值 c00 c000 * (1 - wx) c100 * wx c01 c001 * (1 - wx) c101 * wx c10 c010 * (1 - wx) c110 * wx c11 c011 * (1 - wx) c111 * wx c0 c00 * (1 - wy) c10 * wy c1 c01 * (1 - wy) c11 * wy return c0 * (1 - wz) c1 * wz注意上面代码里的grid_gather是对每个像素直接在网格中取对应节点。实际论文里B是batch维度网格是三维的我这里的维度和示意图略有简化理解原理即可。有了slice出的仿射系数把它作用到原图上def apply_affine(input_img, affine_coeff): input_img: (B, 3, H, W) affine_coeff: (B, 12, H, W), 包含3x3矩阵3维偏置 B, _, H, W input_img.shape coeff affine_coeff.view(B, 3, 4, H, W) A coeff[:, :, :3, :, :] # (B, 3, 3, H, W) b coeff[:, :, 3, :, :].unsqueeze(2) # (B, 3, 1, H, W) input_t input_img.unsqueeze(1) # (B, 1, 3, H, W) out torch.matmul(A, input_t).squeeze(2) b.squeeze(2) return out3.3 网络主体与训练流程Coarse网络我用了一个简化版输入128x128的RGB图经过4层步长2的卷积降到8x8再全局池化得到特征向量最后经过两层全连接输出13824个数16x16x16网格每个节点存12个仿射系数。输出reshape成(1, 16, 16, 16, 12)。Guide网络我直接用原图下采样到128x128计算灰度图、色度特征、和梯度特征堆成6通道再经过三层卷积生成3通道guidemap分别对应空间x、空间y、亮度坐标。这里其实有一个细节guidemap的x,y坐标在输入时已经是归一化的网格坐标论文里用了一个齐次坐标技巧便于slice直接索引。训练时loss函数用L1损失加一个轻微的梯度平滑项def hdr_loss(pred, target): l1 torch.mean(torch.abs(pred - target)) # 计算输出图的梯度差使颜色过渡更平滑 gx torch.abs(pred[:, :, 1:, :] - pred[:, :, :-1, :]).mean() gy torch.abs(pred[:, :, :, 1:] - pred[:, :, :, :-1]).mean() gt_gx torch.abs(target[:, :, 1:, :] - target[:, :, :-1, :]).mean() gt_gy torch.abs(target[:, :, :, 1:] - target[:, :, :-1]).mean() return l1 0.1 * (torch.abs(gx - gt_gx) torch.abs(gy - gt_gy))我自己用Adam优化器initial lr 1e-4batch size 8每20个epoch衰减0.5。在单张RTX 3060上512x512的图迭代200个epoch大约需要15小时。如果你不耐烦跑这么久可以先在256x256上验证正确性再放大到原图。3.4 训练数据准备的两个思路如果手头有Raw格式的原始图像最好。你可以同一张Raw用不同预设导出两版Tiff一份作为输入一份作为训练目标。HDRnet原始数据集的构造就是这么来的。没有Raw的话用高分辨率JPEG也行但目标图最好做点风格化的局部编辑比如局部加深减淡、颜色分级、提亮暗部等逼着网络学出任性的局部变换。我在实验中发现数据量不需要太大200张图的pair就基本够用了。重点是覆盖率场景亮度分布要有高有低颜色分布要有冷暖对比。如果你的数据全是白天风景拿到夜景图上就会翻车。此外图像增强的回归任务对噪声比较敏感输入图最好不要有太强的人工噪声不然网络会学到把噪声也提亮的坏习惯。数据增强方面随机旋转90度、水平翻转、随机裁剪都有效。不建议做随机色抖——色调偏移对颜色网格的学习干扰太大。4. 训练与部署中的常见问题及排查4.1 slice层反向传播报错或结果不更新很多人复现HDRnet时卡在slice层的梯度上。PyTorch的grid_sample本身是可微的如果你用我上面的gather方式实现注意索引计算一定要用detach否则梯度会把坐标搞乱。我踩了一个很深的坑guidemap的坐标范围必须是0到Ns或Nc减1。我最初把coord归一化到0到1没乘以网格尺寸导致所有像素都聚集在网格的0号格子附近失去了空间区分性。排查方式是打印slice输出的统计值如果发现所有仿射系数几乎相同大概率是坐标范围不对。还有一个常见问题是梯度爆炸。HDRnet里仿射变换对输入图像的乘法会使梯度成倍放大尤其是颜色矩阵初始值设置太大时。解决办法是初始化时把颜色矩阵设为对角占优的近似单位阵初始偏移设成0。我在代码里是这样初始化的def init_affine_coeff(m): if isinstance(m, nn.Linear): eye torch.eye(3).flatten() bias torch.zeros(3) init_val torch.cat([eye, bias]).unsqueeze(0).unsqueeze(-1).unsqueeze(-1) # 在网络输出最后加一个偏置让初始输出接近原图这个方法带来的稳定效果远超调低学习率。4.2 输出图出现明显色块和网格感色块通常说明双边网格的空间分辨率或颜色分辨率不够或者slice时插值权重没有归一化。检查一下坐标范围内是不是存在超出[0, N-1]的越界情况越界后clamp会导致多个像素映射到同一个边界节点视觉上就是色块。把clamp改成对坐标做更平滑的截断或者调整guidemap输出加一个sigmoid后再乘以网格尺寸。网格感则和训练时的图像分辨率有关。coarse网络输入分辨率越低输出网格的空间量子化越明显。我的做法是coarse网络输入保持256x256不要太低。虽然论文用128但我实测256画质更稳速度也只是略微下降。4.3 推理速度不达预期的排查HDRnet宣称实时但你在自己机器上跑不一定快。关键瓶颈往往不是卷积网络而是slice操作。如果切片操作是用Python循环实现的速度会慢到怀疑人生。建议用我上面的gather向量化版本或者直接用PyTorch自带的grid_sample在GPU上跑一次1080p的slice大约只需要2-5ms。另一个提速技巧是把guidemap和网格输出都转成半精度。在验证阶段半精度FP16几乎不影响最终视觉效果但吞吐量提升明显。4.4 传统图形学中的仿射变换与HDRnet的对比文章前面提过“java仿射变换图形期末作品”这让我想到许多做图形学的同学对仿射变换的理解停留在“几何变换的矩阵表示”这个层面。图形学课程里旋转、缩放、平移是重点一般还会教你怎么用齐次坐标把平移塞进矩阵里最后写个小程序旋转一个三角形交个期末大作业。HDRnet和这完全是两回事但底层共享同一个数学直觉用线性变换矩阵加上偏置向量去描述“值域空间的局部变化”。只不过图形学的仿射作用在x-y坐标上HDRnet的仿射作用在R-G-B颜色通道上。你可以把它想象成每个像素的RGB向量乘上一个3x3矩阵再加一个偏移就得到新的RGB向量。如果这个矩阵随空间平滑变化那么整张图就能表达复杂的局部色调映射。4.5 一个快速验证仿射变换效果的小实验在真正跑HDRnet前我用NumPy做了一个很简单的验证花十分钟就能完成。对一张图人为构造一个随空间变化的颜色矩阵亮部区域用偏蓝色调矩阵暗部用偏暖色调矩阵中灰区间做线性插值然后逐像素做矩阵乘法。视觉效果非常接近一种风格化的color grading。这让我彻底理解了HDRnet里仿射变换的作用——它就是让这种手工构造的渐变颜色变换变得自适应起来。如果读到这里的你还在迷茫我建议先做这个实验不要一上来就啃整个网络。确实理解了“空间变化的仿射变换长什么样”HDRnet的架构就能看懂一半。5. 效果评估与调试心得5.1 定量评价指标怎么选HDRnet的论文里主要用了HDR-VDP-2和用户偏好测试。如果你做研究这些指标要留意。日常工作里我一般用PSNR和SSIM做基线对比再用色调映射质量那个维度评主观效果。但这里有个反直觉的地方HDRnet的输出目标是一张已经调好色的图比如Lightroom导出的和输入图之间并非“越接近原图越好”所以PSNR虚高反而是错的方向。你更应该关注视觉效果看暗部细节有没有被提出来亮部有没有过曝颜色有没有偏色。我的建议是做一个简单的语义描述在测试集上分别看“低光区域细节保留程度”、“高光区域是否溢出”、“整体色彩是否自然”再结合算分。别只看RMSE不然调参方向会跑偏。5.2 训练不收敛的排查清单我把自己训练中遇到的情况写成一个速查表供参考现象可能原因检查方法loss下降非常慢学习率过小或网络输出初始化偏离打印affine系数的均值和方差输出图像整体偏灰偏置项初始化过大把bias初始化为0颜色断层严重颜色分辨率Nc太小从8提到16或加大数据多样性训练与验证loss差距大数据过拟合增加数据增强或降低网络容量推理时GPU显存不足网格过大或batch过大检查网格尺寸和coarse网络输入分辨率5.3 训练数据的质量和数量如何平衡我用过两份数据集一份是工业相机拍的Raw图约500对另一份是从开源数据集截取的大约900对网络图片用Lightroom人工过一遍当目标。两条路线都能跑通但差别很大。Raw数据集生成的模型在亮部和暗部的过渡上极其平滑因为Raw图的动态范围大曝光调节余地大网络图片数据模型的主观色彩更讨喜因为Lightroom预设本身就是经过精修的色调。所以做项目前先想一想你最终要输出什么效果。如果只是做色调节网络图片数据就够了如果是做HDR重建一定要用Raw数据。5.4 从HDRnet拓展到其他任务这个架构不只是HDR图片增强能用它其实是一个通用范式低分辨率预测局部仿射变换参数再用全分辨率指导图做稠密重建。后来有很多工作用这种思路做去噪、超分辨率、图像修复效果都不错。比如你可以把coarse网络输出的仿射系数换成去噪滤波器的核权重这样就能用双边网格表达一个内容自适应的保边去噪器。这里有一点很关键HDRnet本身是个很灵活的框架核心卖点是“双边网格slice仿射变换”的组合。我最近在试一个扩展把网格从3维扩到4维——空间2维加颜色2维比如加入色度信息效果确实比只按亮度量化更好但网格体量从16x16x8升级到16x16x8x8之后训练速度慢了不少参数也膨胀了一大截需要权衡。6. 几个实操中的细节补充6.1 网格尺寸的选择论文中建议空间分辨率16x16颜色分辨率8或16。我自己的经验是如果处理的是人像颜色分辨率16比8要好很多如果是夜景空间分辨率可以适当降低到12反正暗部细节本来就少。训练前先用你计划的最小网格尺寸跑通流程再逐渐放大别一上来就用满分辨率不然排查bug的周期会被拉得很长。另外网格也不一定是三个维度都均匀量化。可以让z轴亮度使用gamma校正后的非线性分布亮部区间分的更细暗部区间分的更粗这样色彩过渡会更自然尤其在HDR场景中。我试验后感觉主观画质能提升一点代价是代码变复杂了。6.2 训练过程的可视化在训练过程中最好每5个epoch保存一次网格和输出图观察网格节点里学出来的仿射矩阵。我发现一个有趣的规律训练初期网格里大部分节点存的是接近单位矩阵的参数随着epoch增加亮部节点和暗部节点的仿射系数差异越来越大。这说明网络真的在沿着亮度轴逐步学习局部映射。你还可以把双边网格的每个节点里的颜色矩阵按亮度排列画成一个三维色块立方体这样能直观看到网格学到了一个什么样的颜色映射表。如果可以的话建议把这个可视化做成训练时的日志输出对判断过拟合很有帮助。6.3 关于“仿射变换”术语的兼容性最后说一个很多人容易混淆的细节。传统图形学里的仿射变换包含平移而HDRnet使用的“仿射变换”术语实际上指的是对每个像素RGB值的线性变换再加上一个偏置。本质上是3x3线性变换加平移这正是仿射变换的数学定义。所以期刊和论文里直接用“affine”这个词是准确的。如果你拿这个项目去面试或给领导汇报可以一句话讲清楚HDRnet是一个把图像增强变成“学习空间-颜色联合网格上的局部仿射变换”的模型它把传统的双边滤波思想变成了一个可微分的网络模块。这个说法基本上能把不懂的人讲懂也能把懂的人讲服。
返回列表