Flare Removal 训练全指南:基于 ICCV 2021 论文复现镜头光晕去除的神经网络训练、评估与推理
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 removalICCV 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 LossesZhang 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_dirNone所有 flare-free 图像的父目录任意 RGB 且等尺寸的自然图像数据集均可--flare_dirNone所有 flare-only 图像的父目录若按上文下载官方数据传--flare_dir/your/local/path/lens-flare--data_sourcejpg数据来源枚举jpg单张 JPG/PNG 文件或tfrecord预烘焙的分片 TFRecord 文件--modelunet模型名unet或can--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.CheckpointManager每ckpt_period步保存一次权重与全量 SavedModel同时向 TensorBoard 写入prediction图像、loss与step_time标量。训练结束时将training_finished标记置为True并做最后一次保存——该标记正是评估脚本判定训练结束的信号。并行监控评估evaluate.pypython/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, timeout30, timeout_fn...)持续轮询训练目录中的最新 checkpoint30 秒超时直到训练脚本写出的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_formattf)保存的 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 的前缀而非目录。--modelunet或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像素级 MAEMeanAbsoluteErrorl2像素级 MSEMeanSquaredErrorpercep默认/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.pymodels.py 暴露两个模型名--modelunet自定义 U-Netu_net.py参考 Ronneberger et al. 2015 的结构输入 512×512×3scales44 级下采样/上采样、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.pyMatlab 仿真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),仅供参考
上一篇/下一篇内容由系统自动关联
返回资讯列表 →