Flare Removal 训练全指南:基于 ICCV 2021 论文复现镜头光晕去除的神经网络训练、评估与推理
【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research
导读
本文围绕google-research仓库中 flare_removal/README.md 所描述的"How to train neural networks for flare removal"(ICCV 2021)开源实现展开,系统讲解从镜头光晕(lens flare)数据集的获取、散射光斑(streaks)的 Matlab 仿真合成,到基于 TensorFlow 的 U-Net 模型训练、并行评估与单图推理的完整流程。读完本文,你将能够复现该论文的数据管线、理解train.py/evaluate.py/remove_flare.py三个入口的全部命令行参数,并掌握感知损失(perceptual loss)与线性域 flare 合成等核心技术细节。
项目背景与论文定位
该目录是论文"How to train neural networks for flare removal"的官方开源代码,论文发表于 ICCV 2021(第 2239-2247 页),作者包括 Yicheng Wu、Qiurui He、Tianfan Xue、Rahul Garg、Jiawen Chen、Ashok Veeraraghavan 与 Jonathan T. Barron。核心任务是训练一个神经网络,输入一张受镜头光晕污染的 RGB 图像,输出对应的无 flare 场景图。
仓库按照 flare_removal/requirements.txt 声明依赖(TensorFlow >= 2.6、tensorflow-addons、absl-py、numpy、opencv-python、scikit-image、scipy、tqdm),全部代码分为两块:
matlab/:基于物理仿真生成"散射 flare"(即镜头产生的条状光斑 streaks)的 Matlab 代码;python/:数据加载(data_provider.py)、在线合成(synthesis.py)、模型定义(models.py、u_net.py、vgg.py)、损失函数(losses.py)、训练(train.py)、评估(evaluate.py)与推理(remove_flare.py)的完整 TensorFlow 实现。
值得注意的是,README 中有一则与代码质量相关的公告:官方曾对 VGG 损失做过一次小修复,并披露过 2022 年 1 月发现的训练代码潜在问题(该问题疑似在开源前的代码整理阶段引入,不影响remove_flare.py推理脚本,官方论文中的定量与定性结果均基于旧版内部代码复现)。如果你复现出的训练效果与论文存在差异,可结合该公告与 losses.py 的实现进行排查。
数据集准备
训练与评估需要两类图像:flare-only 图像(纯光晕图)与flare-free 场景图像(自然图像)。
Flare-only 图像:5,001 张 RGB 光晕图
官方在 CC BY 4.0 许可下发布了 5,001 张 RGB 光晕图像,分为两类:
- 2,001 张实验室实拍图(1,001 次拍摄 + 帧间插值),位于下载目录的
captured子目录; - 3,000 张计算仿真图,位于
simulated子目录。
获取方式需要先安装 Google Cloud SDK(会自动安装gcloud storage工具),然后执行:
$ gcloud storage cp --recursive gs://gresearch/lens-flare /your/local/path下载完成后,/your/local/path/lens-flare即为 flare 数据集的父目录,可直接作为训练脚本的--flare_dir参数。
Flare-free(场景)图像
场景图复用论文Single Image Reflection Removal with Perceptual Losses(Zhang et al., CVPR 2018)的图像数据集。与反射去除任务不同的是,本项目不区分反射层与透射层,而是将整个数据集打乱后作为一个统一的自然图像集合使用,因此你需要自行划分训练集与测试集。README 还特别提醒:evaluate.py使用的场景图应当与train.py使用的不重叠。所有场景图像必须为 RGB 且尺寸一致,因为训练管线(见下文"数据加载")会按固定image_shape解析。
散射光斑(Streaks)的 Matlab 物理仿真
matlab/目录用于仿真由"缺陷光圈"产生的随机散射 flare。直接在 MATLAB 中执行 main.m 即可复现结果。脚本内置了一组典型智能手机相机参数:名义波长 550nm、焦距 2.2mm、像素间距 1μm、传感器尺寸 6mm×6mm,并据此在频域计算散焦相位(defocus phase)与光圈掩膜(aperture mask)。
主流程分三步:
- 生成缺陷光圈:调用
RandomDirtyAperture.m,在光圈上随机添加"尘点"(dots)与"划痕"(polylines); - 计算 PSF:在 380nm-740nm 之间采样 73 个波长,通过
RandomSpectralResponse.m生成随机的 RGB 光谱响应,再结合随机散焦量(GetDefocusPhase.m、GetPsf.m)得到 RGB 点扩散函数; - 随机裁剪与畸变:对 PSF 施加随机径向畸变、缩放、旋转与平移等相机内参扰动,生成多组不同的 flare 图案。
脚本默认把产物写入两个目录:
matlab/apertures:模拟的缺陷光圈图(带尘点与划痕);matlab/streaks:由上述缺陷光圈产生的 flare 图案,每个光圈对应多张图案,覆盖不同的光源位置、散焦与畸变组合,这些图将用于进一步合成 flare 污染的拍摄照片。
环境搭建:从仓库根目录运行
README 给出了一个重要约束:所有 Python 命令都必须在仓库根目录google_research/(即本仓库根目录)下执行,否则 Python 无法正确解析flare_removal.python.*模块路径。
run.sh 演示了标准的环境搭建流程——创建并激活虚拟环境后安装依赖:
python3 -m venv env source ./env/bin/activate pip install -r flare_removal/requirements.txt注意:脚本最后执行的python3 -m flare_removal.python.remove_flare由于缺少模型与数据路径参数预期会失败,它的作用是验证依赖安装完整;实际运行需要按下文补齐参数(或修改源码中的默认值)。
训练模型:train.py 与全部参数说明
训练入口是 python/train.py,基本调用方式如下:
$ python3 -m flare_removal.python.train \ --train_dir=/path/to/training/logs/dir \ --scene_dir=/path/to/flare-free/training/image/dir \ --flare_dir=/path/to/flare-only/image/dir核心参数
| 参数 | 默认值 | 说明 |
|---|---|---|
--train_dir | /tmp/train | 训练状态目录:保存指标、summary 图像与模型权重 checkpoint。训练重启时会自动从该目录恢复上次状态,因此每个新实验应使用全新(空)目录 |
--scene_dir | None | 所有 flare-free 图像的父目录,任意 RGB 且等尺寸的自然图像数据集均可 |
--flare_dir | None | 所有 flare-only 图像的父目录;若按上文下载官方数据,传--flare_dir=/your/local/path/lens-flare |
--data_source | jpg | 数据来源枚举:jpg(单张 JPG/PNG 文件)或tfrecord(预烘焙的分片 TFRecord 文件) |
--model | unet | 模型名:unet或can |
--loss | percep | 损失函数名:percep(感知损失)或l1/l2 |
--batch_size | 2 | 训练 batch 大小 |
--epochs | 100 | 训练轮数 |
--ckpt_period | 1000 | 每隔多少步写一次 checkpoint 与 summary |
--learning_rate | 1e-4 | 初始学习率(Adam 优化器) |
--scene_noise | 0.01 | 合成数据中加到场景上的高斯噪声 sigma;每张图的实际方差从以scene_noise为尺度的卡方(Chi-squared)分布中抽取 |
--flare_max_gain | 10.0 | 合成时施加到 flare 图案上的最大数字增益(线性域内,RGB 三通道各自随机独立、不超过该上限) |
--flare_loss_weight | 1.0 | flare 损失的权重(场景损失权重固定为 1) |
--training_res | 512 | 训练分辨率(方形图边长) |
从源码看,train.py的主流程是:通过 data_provider.py 加载场景与 flare 两个数据集并zip配对;构建模型(models.py);用 Adam 优化器在train_step中执行梯度下降,并做全局梯度裁剪(tf.clip_by_global_norm(grads, 5.0));通过tf.train.CheckpointManager每ckpt_period步保存一次权重与全量 SavedModel,同时向 TensorBoard 写入prediction图像、loss与step_time标量。训练结束时将training_finished标记置为True并做最后一次保存——该标记正是评估脚本判定训练结束的信号。
并行监控评估:evaluate.py
python/evaluate.py 是可选的并行评估脚本,用于边训练边监控模型表现:
$ python3 -m flare_removal.python.evaluate \ --eval_dir=/path/to/evaluation/logs/dir \ --train_dir=/path/to/training/logs/dir \ --scene_dir=/path/to/flare-free/evaluation/image/dir \ --flare_dir=/path/to/flare-only/image/dir评估脚本会通过tf.train.checkpoints_iterator(train_dir, timeout=30, timeout_fn=...)持续轮询训练目录中的最新 checkpoint(30 秒超时),直到训练脚本写出的training_finished标记为真。它复用与训练相同的模型、损失与在线合成流程(synthesis.run_step),在恢复权重后于评估集上计算损失并写入--eval_dir/summary。其--learning_rate参数仅为满足参数扫描需求而存在的占位符,实际不使用。
训练产物:checkpoint 目录结构
训练与评估状态统一写入--train_dir(训练)与--eval_dir(评估),目录内容如下:
model/:最新模型文件(每ckpt_period步通过tf.keras.models.save_model(model, model_dir, save_format='tf')保存的 SavedModel,包含架构与权重,可直接被推理脚本加载);summary/:训练指标与 summary 图像,用 TensorBoard 可视化;ckpt-*:模型权重 checkpoint(不包含网络结构),用于恢复之前的模型权重;训练重启时ckpt.restore(latest_ckpt).expect_partial()会尝试从该目录恢复(由于惰性初始化,完整恢复校验在第一步训练后通过assert_consumed()完成)。
测试模型:remove_flare.py 推理
python/remove_flare.py 用于对真实世界图像做 flare 去除推理:
$ python3 -m flare_removal.python.remove_flare \ --ckpt=/path/to/training/logs/dir/model \ --input_dir=/path/to/test/image/dir \ --out_dir=/path/to/output/dir参数说明:
--ckpt:模型位置。可以是 SavedModel 目录(同时加载架构与权重,此时忽略--model),也可以是 TF checkpoint 路径(仅加载最新权重,加载更快,此时必须提供--model);若想加载某个特定 checkpoint,可传该 checkpoint 的前缀而非目录。--model:unet或can,仅当--ckpt指向 TF checkpoint/checkpoint 目录时必需。--batch_size:默认 1。部分网络(如 rain removal 网络)只能接受预定义的 batch 大小。--input_dir:输入图像目录。--out_dir:输出目录,缺省时为input_dir/model_output。--separate_out_dirs:默认True,将输出写入out_dir下的input/、output/、output_flare/、output_blend/四个子目录;设为0时所有结果写在同一目录,文件名带不同后缀(_input.png、_output.png、_output_flare.png、_output_blend.png)。
输入尺寸处理规则
从process_one_image的源码实现可以看到推理脚本对输入尺寸有明确的自动处理策略(对应论文第 6.4 节):
- 大于 512×512 的图像:先中心裁剪到 512×512 再送入模型;
- 大于 2048×2048 的图像:先中心裁剪到 2048×2048,再用 AREA 插值降采样到 512×512 送入模型,推断出的 flare-free 结果放大回 2048×2048(放大 flare 后经线性域相减得到场景,避免直接放大场景带来的伪影);
- 小于 512×512 的图像:不支持(会抛出
ValueError)。
推理脚本除输出input、去 flare 后的output、分离出的output_flare外,还会输出output_blend——这是论文第 5.2 节提出的光晕去除后保留光源的合成结果:把预测的场景与原始输入中截取的高光区域重新融合,避免去除 flare 时把真实光源一并抹掉。
线性域相减:remove_flare 的实现原理
flare 分离的关键操作 utils.remove_flare 不是简单的像素相减,而是在伽马编码反转后的线性域中做减法:输入与预测 flare 先各自做pow(x, gamma)线性化,相减得到线性域场景,再取pow(scene, 1/gamma)还原回伽马编码。gamma默认 2.2;两侧都通过clip_by_value夹在极小值1e-7与 1.0 之间,以避免pow在接近 0 时梯度未定义的问题。这也解释了训练合成中随机化 gamma 的必要性(见下文)。
关键实现原理:在线合成、损失函数与网络结构
数据合成:synthesis.py
训练时并不直接使用"带 flare 的图像对",而是由 synthesis.py 的add_flare在线把 flare-only 图与场景图合成出污染图。由于真实拍摄的伽马编码未知,脚本随机抽取gamma ∈ [1.8, 2.2]做随机伽马调整,使模型泛化到合理的伽马范围;随后:
- 用
remove_background去掉 flare 的直流背景; - 对 flare 施加随机的仿射变换(旋转 ∈ [-π, π]、平移(均值 0、标准差 10 像素)、剪切 ∈ [-π/9, π/9]、缩放 ∈ [0.9, 1.2])以模拟光源位置变化;
- 在线性域给 flare 施加随机 RGB 增益(上限
flare_max_gain); - 场景侧叠加从卡方分布抽取方差的高斯噪声(
scene_noise); - 最终在 sRGB 域合成污染图,供模型学习。
损失函数:losses.py 与 VGG 感知损失
losses.py 提供三种损失(通过--loss选择):
l1:像素级 MAE(MeanAbsoluteError);l2:像素级 MSE(MeanSquaredError);percep(默认)/perceptual:基于预训练 VGG19 的感知损失 + L1 损失的加权组合。感知部分在 VGG19 的 5 个 tap-out 层上计算加权 L1 距离,默认系数为block1_conv2: 1/2.6、block2_conv2: 1/4.8、block3_conv2: 1/3.7、block4_conv2: 1/5.6、block5_conv2: 10/1.5;由于感知损失内部按 [0,255] 量纲计算而输入约定为 [0,1],L1 分量乘以权重 255 以实现真正的 1:1 配比。README 提到的 "VGG loss 修复" 即针对此类细节。
网络结构:models.py、u_net.py 与 vgg.py
models.py 暴露两个模型名(--model):
unet:自定义 U-Net(u_net.py),参考 Ronneberger et al. 2015 的结构,输入 512×512×3,scales=4(4 级下采样/上采样)、bottleneck 深度 1024、bottleneck 2 层;下采样块由两个 3×3 卷积(ReLU)+ MaxPool2D 组成,上采样块使用双线性插值并与 skip connection 拼接;can:基于 vgg.py 构建的上下文聚合网络(context aggregation network),输入 512×512×3,卷积通道 64,输出 3 通道。
两个网络在build_model中都固定使用 512×512×3 的输入形状,这也是 README 与推理脚本以 512 作为基准分辨率的原因。
预训练模型与复现说明
由于许可限制,官方未发布预训练模型。复现论文结果需要自行完成:下载 5,001 张 flare-only 图像 → 获取场景图数据集 → 运行main.m(如需自产仿真 flare)→ 按上文命令训练与评估。从 README 公告看,论文中展示的定量与定性结果可由旧版内部代码复现,开源版本在部分环节(如 VGG 损失)曾有过修复,若训练结果与论文存在差异,可在 losses.py 与 synthesis.py 中核对合成与损失实现细节。
引用
若本工作对你有帮助,请按 README 提供的 BibTeX 引用论文:
@InProceedings{flareremvoal2021, author = {Wu, Yicheng and He, Qiurui and Xue, Tianfan and Garg, Rahul and Chen, Jiawen and Veeraraghavan, Ashok and Barron, Jonathan T.}, title = {How To Train Neural Networks for Flare Removal}, booktitle = {Proceedings of the IEEE/CVF International Conference on Computer Vision (ICCV)}, month = {October}, year = {2021}, pages = {2239-2247} }参考资料速查
- 项目说明与数据集下载:flare_removal/README.md
- 训练入口:flare_removal/python/train.py
- 评估入口:flare_removal/python/evaluate.py
- 推理入口:flare_removal/python/remove_flare.py
- 在线合成:flare_removal/python/synthesis.py
- 损失函数:flare_removal/python/losses.py(含测试 losses_test.py)
- 网络结构:flare_removal/python/models.py、flare_removal/python/u_net.py(含测试 u_net_test.py)、flare_removal/python/vgg.py(含测试 vgg_test.py)
- 数据加载:flare_removal/python/data_provider.py
- 通用工具(线性域相减、仿射变换、图像读写):flare_removal/python/utils.py
- Matlab 仿真:flare_removal/matlab/main.m
- 依赖清单与环境脚本:flare_removal/requirements.txt、flare_removal/run.sh
【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考