DeepInverse迭代优化算法完全指南:ADMM、HQS、FISTA与PGD一次搞懂
【免费下载链接】deepinvDeepInverse: a PyTorch library for solving imaging inverse problems using deep learning项目地址: https://gitcode.com/gh_mirrors/de/deepinv
DeepInverse(DeepInverse: a PyTorch library for solving imaging inverse problems using deep learning)是一个基于 PyTorch 的图像逆问题深度学习库。本指南带你一次搞懂它内置的 4 大迭代优化算法——ADMM、HQS、PGD、FISTA:它们如何求解"数据保真项 + 正则项"的优化问题、各自适用什么场景,以及如何用几行代码跑通 PnP(Plug-and-Play)图像重建。
一、为什么逆问题需要迭代优化?
医学成像(MRI、CT)、遥感、显微镜等场景中,我们拿到的从来不是"原图" x,而是它经过物理过程 A 之后的测量 y ≈ A(x)——模糊、失线、加噪、欠采样……
绝大多数重建方法都归结为同一个优化问题:
x̂ = argmin_x D(x, y) + λ·R(x)- D(x, y):数据保真项,衡量重建结果与观测数据的一致性
- R(x):正则/先验项,注入"图像应该平滑、稀疏或干净"的先验知识
- λ:平衡两者的正则化参数
由于这个最小化通常没有解析解,只能迭代求解——这正是 PGD、FISTA、HQS、ADMM 登场的地方。🚀
二、DeepInverse 的优化模块结构
所有迭代优化算法都集中在deepinv/optim/模块中,结构非常清晰:
deepinv/optim/optimizers.py—— 面向用户的高层接口,包含ADMM、HQS、PGD、FISTA等类(如 deepinv/optim/optimizers.py 中的ADMM类、L1596 的PGD、L1737 的FISTA)deepinv/optim/optim_iterators/—— 底层迭代子,每个算法的"每一步更新公式"都在这里,例如admm.py、hqs.py、pgd.pydeepinv/optim/data_fidelity.py—— 数据保真项(如 L2 保真项),定义在deepinv/optim/data_fidelity.pydeepinv/optim/prior.py—— 先验/正则项,包括显式正则(TV、L1)和PnP 去噪器先验
统一的抽象是:所有算法都继承自BaseOptim,你只需要提供数据保真项 + 先验项 + 物理模型,剩下的交给算法迭代。🔧
三、四大迭代优化算法详解
3.1 PGD:近端梯度下降,最易上手的起点
Proximal Gradient Descent是四大算法中最简单直观的一个,每步只做两件事:
- 沿数据保真项的梯度走一小步(步长 γ)
- 用先验的近端算子(或去噪器)"压"一下结果
u_k = x_k − γ·∇D(x_k, y) x_{k+1} = prox_{γλR}(u_k)- ✅ 优点:实现简单、每步计算量小
- ⚠️ 缺点:收敛速度为 O(1/k),可能需要较多迭代
- 源码位置:deepinv/optim/optimizers.py 的
PGD类
3.2 FISTA:给 PGD 装上"加速器"
FISTA在 PGD 的基础上引入动量外推项,把收敛速度从 O(1/k) 提升到 O(1/k²):
u_k = z_k − γ·∇f(z_k) x_{k+1} = prox_{γλR}(u_k) z_{k+1} = x_{k+1} + α_k·(x_{k+1} − x_k) # 动量外推- ✅ 同样的迭代次数下,FISTA 通常比 PGD 收敛更快
- ⚠️ 对步长 γ 更敏感:需要 γ ≤ 1/Lip(∇f)(L 为梯度 Lipschitz 常数)
- 源码位置:deepinv/optim/optimizers.py 的
FISTA类
一句话选择:先验近端算子容易计算、问题梯度平滑时,优先用 FISTA。
3.3 HQS:半二次分裂,PnP 重建的常客
Half-Quadratic Splitting把问题拆成两个"近端步"交替求解,不引入对偶变量:
u_k = prox_{γf}(x_k) x_{k+1} = prox_{σλR}(u_k)- ✅ 结构比 ADMM 更简洁,天然适合把
prox换成神经网络去噪器(PnP) - ✅ DeepInverse 中 HQS 还支持DEQ(深度平衡)展开与Anderson 加速等高级选项
- 源码位置:deepinv/optim/optimizers.py 的
HQS类
3.4 ADMM:最通用、最强大的"老大哥"
Alternating Direction Method of Multipliers引入一个对偶变量 z 来"解耦"数据保真项和正则项:
u_{k+1} = prox_{γf}(x_k − z_k) x_{k+1} = prox_{γλR}(u_{k+1} + z_k) z_{k+1} = z_k + β·(u_{k+1} − x_{k+1})- ✅ 对"两项都不好处理"的复杂问题最鲁棒,是稀疏编码、块匹配(BM3D 类)等经典方法的标准解法
- ⚠️ 多维护变量,每步开销略大于 HQS
- 关键参数:步长 γ 与松弛参数 β(默认均为 1.0)
- 源码位置:deepinv/optim/optimizers.py 的
ADMM类,底层迭代在deepinv/optim/optim_iterators/admm.py
四、四大算法快速对比与选型指南
| 算法 | 核心思想 | 迭代变量 | 收敛速度 | 最佳适用场景 |
|---|---|---|---|---|
| PGD | 梯度步 + 近端步 | 仅原始变量 x | O(1/k) | 入门、调试、步数预算紧张 |
| FISTA | PGD + 动量加速 | x + 动量变量 z | O(1/k²) | 光滑数据保真项 + 简单近端算子 |
| HQS | 双近端步交替 | 仅原始变量 x | 较快 | PnP 去噪器先验、轻量重建 |
| ADMM | 增广拉格朗日解耦 | 原始变量 + 对偶变量 | 最鲁棒 | 复杂正则项、稀疏/非光滑问题 |
💡选型建议:不确定时用ADMM(最稳);追求简洁和速度且近端算子容易算时用HQS / FISTA;刚学习框架时用PGD把流程跑通。
五、跑一个 PnP 迭代重建:只需几步
以经典的"缺失值填充 + 高斯噪声"逆问题为例,用 PGD + PnP 去噪器先验的完整流程是:
- 定义物理模型(如
dinv.physics.Inpainting,含 mask 与噪声模型) - 定义数据保真项:
data_fidelity = dinv.optim.data_fidelity.L2() - 定义 PnP 先验:
prior = dinv.optim.Prior(denoiser=...)(去噪器可换成 BM3D、DnCNN 等) - 组装算法并调用:
model = dinv.optim.PGD(prior=prior, data_fidelity=data_fidelity),然后x_hat = model(y, physics)
四个算法的调用方式完全一致——只换算法名,其他参数不变,方便你横向对比效果。
更完整的可运行示例可参考:
examples/plug-and-play/demo_vanilla_PnP.py—— 最基础的 PnP 迭代重建examples/plug-and-play/demo_PnP_custom_optim.py—— 自定义迭代子examples/optimization/demo_TV_minimisation.py—— ADMM + 全变差正则
六、新手必须知道的 5 个通用参数
四个算法共享一套参数体系(定义在BaseOptim中,见deepinv/optim/optimizers.py):
| 参数 | 含义 | 默认值 |
|---|---|---|
stepsize(γ) | 步长,过大会发散 | 1.0 |
lambda_reg(λ) | 正则强度,权衡细节与噪声 | 1.0 |
g_param/sigma_denoiser | 先验参数(如去噪器噪声水平) | None |
max_iter/early_stop | 最大迭代数 / 提前停止 | 100 / False |
crit_conv | 收敛判据:"residual"或"cost" | "residual" |
⚡进阶技巧:把unfold=True打开,算法的步长、λ、甚至每层的去噪器都可以反向传播训练——这就是"展开网络(Unfolded Network)",让迭代优化算法与深度学习无缝衔接。这是 DeepInverse 区别于传统优化库的最大亮点。
七、总结
- DeepInverse 把逆问题重建统一为
argmin D(x,y) + λR(x),四大算法PGD / FISTA / HQS / ADMM提供了从入门到鲁棒的完整工具链,全部位于deepinv/optim/模块 - 所有算法支持PnP(去噪器当先验)与RED(去噪器定义梯度)两种深度学习先验方式,并可一键展开(unfold)进行端到端训练
- 新手路径:先用 PGD 跑通
examples/plug-and-play/demo_vanilla_PnP.py,再切换 ADMM/FISTA 对比效果,最后开启 unfold 训练参数
掌握这四个算法,你就拥有了 DeepInverse 迭代重建的核心钥匙。🔑
【免费下载链接】deepinvDeepInverse: a PyTorch library for solving imaging inverse problems using deep learning项目地址: https://gitcode.com/gh_mirrors/de/deepinv
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考