Flare Removal 训练全指南:基于 ICCV 2021 论文复现镜头光晕去除的神经网络训练、评估与推理
2026/9/20 13:49:53 网站建设 项目流程

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.pyu_net.pyvgg.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)。

主流程分三步:

  1. 生成缺陷光圈:调用RandomDirtyAperture.m,在光圈上随机添加"尘点"(dots)与"划痕"(polylines);
  2. 计算 PSF:在 380nm-740nm 之间采样 73 个波长,通过RandomSpectralResponse.m生成随机的 RGB 光谱响应,再结合随机散焦量(GetDefocusPhase.mGetPsf.m)得到 RGB 点扩散函数;
  3. 随机裁剪与畸变:对 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_dirNone所有 flare-free 图像的父目录,任意 RGB 且等尺寸的自然图像数据集均可
--flare_dirNone所有 flare-only 图像的父目录;若按上文下载官方数据,传--flare_dir=/your/local/path/lens-flare
--data_sourcejpg数据来源枚举:jpg(单张 JPG/PNG 文件)或tfrecord(预烘焙的分片 TFRecord 文件)
--modelunet模型名:unetcan
--losspercep损失函数名:percep(感知损失)或l1/l2
--batch_size2训练 batch 大小
--epochs100训练轮数
--ckpt_period1000每隔多少步写一次 checkpoint 与 summary
--learning_rate1e-4初始学习率(Adam 优化器)
--scene_noise0.01合成数据中加到场景上的高斯噪声 sigma;每张图的实际方差从以scene_noise为尺度的卡方(Chi-squared)分布中抽取
--flare_max_gain10.0合成时施加到 flare 图案上的最大数字增益(线性域内,RGB 三通道各自随机独立、不超过该上限)
--flare_loss_weight1.0flare 损失的权重(场景损失权重固定为 1)
--training_res512训练分辨率(方形图边长)

从源码看,train.py的主流程是:通过 data_provider.py 加载场景与 flare 两个数据集并zip配对;构建模型(models.py);用 Adam 优化器在train_step中执行梯度下降,并做全局梯度裁剪(tf.clip_by_global_norm(grads, 5.0));通过tf.train.CheckpointManagerckpt_period步保存一次权重与全量 SavedModel,同时向 TensorBoard 写入prediction图像、lossstep_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 的前缀而非目录。
  • --modelunetcan,仅当--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.6block2_conv2: 1/4.8block3_conv2: 1/3.7block4_conv2: 1/5.6block5_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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询