Following the close of #1318 for size, this is the plan @Tmodrzyk asked for: what the feature is, which existing abstractions it reuses, and how it splits into PRs under the 3000 line limit. I would rather agree the shape here than open code that has to be redone.
What it adds
Bilevel learning of prior parameters: choose theta so that reconstructions x*(theta) = argmin_x 0.5||Ax - y||^2 + R_theta(x) match ground truth on a training set. Today a user picks a regularisation weight by hand or grid search; anything beyond a couple of parameters is out of reach.
The method is MAID, Salehi, Mukherjee, Roberts, Ehrhardt, SIAM J. Math. Data Sci. 2025. Its contribution is not a faster hypergradient but a rule for how accurately the lower level needs solving at each outer step, with a certificate: for a mu-strongly convex lower level, ||x* - x|| <= ||grad_x h(x)|| / mu, computable from the returned iterate. Inexactness is quantified rather than hoped away.
One honest note from the closed PR, since it changes the pitch. On a matched comparison, adaptive accuracy alone does not beat a well-chosen fixed accuracy on PSNR; the two land within 0.01 dB. What it buys is not needing to know that accuracy in advance: the fixed arm only matched after its tolerance was picked by watching worse values fail. The user guide already says fixed accuracy wins on cheap lower levels, and the feature should be framed that way.
What it reuses
The closed PR reimplemented three things DeepInverse already has. The rewrite will not.
| need |
existing abstraction |
what #1318 did instead |
| hypergradient, an implicit linear solve |
deepinv.optim.linear.least_squares_implicit_backward |
its own CG in cg_utils.py |
| lower-level solver |
optim_builder / FixedPoint |
its own Newton, FISTA and GD |
| a prior with parameters |
deepinv.optim.Prior, whose fn(x, *args, **kwargs) already returns (B,) and accepts theta |
a parallel ParametricPrior hierarchy |
Two things I would like a steer on before writing them:
- Learnable priors. Subclass
deepinv.optim.Prior and pass theta through *args, or add a small init_theta mixin so a prior can say how to initialise its parameters? I lean to the mixin, since a Prior alone cannot tell the outer loop what shape theta has.
- Stopping criterion. The certificate needs the lower level to stop on a gradient-residual tolerance, not an iteration count. Is adding a residual-based
check_conv_fn to FixedPoint acceptable, or should that live in the bilevel module?
Split
Each under 3000 lines including tests. Each usable on its own.
- MAID core.
MAID, MAIDConfig, the oracle interface, one quadratic oracle for tests, hypergradient via least_squares_implicit_backward. About 1200 lines. Docs: docstrings plus a short section in the optimisation guide.
- Learnable
Prior. The oracle for any parametric deepinv.optim.Prior, lower level via optim_builder, with a convexity check and a finite-difference check on the hypergradient as tests. About 900.
- Batched oracle. Samples solved together, hypergradients accumulated exactly. About 700.
- Convex ridge regulariser. As a prior in
deepinv.optim, one example on CBSD68 with three arms, short docs. About 1300.
Dropped
The 1086 line guide, two of three examples, the saddle-point, smooth and stochastic variants, the TV baseline, and the accelerated ablation. Any of these can come later as its own small PR if wanted.
Total is roughly 4000 lines across four PRs, against 12700 in one. I will not open the first until the two questions above are settled.
Following the close of #1318 for size, this is the plan @Tmodrzyk asked for: what the feature is, which existing abstractions it reuses, and how it splits into PRs under the 3000 line limit. I would rather agree the shape here than open code that has to be redone.
What it adds
Bilevel learning of prior parameters: choose
thetaso that reconstructionsx*(theta) = argmin_x 0.5||Ax - y||^2 + R_theta(x)match ground truth on a training set. Today a user picks a regularisation weight by hand or grid search; anything beyond a couple of parameters is out of reach.The method is MAID,
Salehi, Mukherjee, Roberts, Ehrhardt, SIAM J. Math. Data Sci. 2025. Its contribution is not a faster hypergradient but a rule for how accurately the lower level needs solving at each outer step, with a certificate: for amu-strongly convex lower level,||x* - x|| <= ||grad_x h(x)|| / mu, computable from the returned iterate. Inexactness is quantified rather than hoped away.One honest note from the closed PR, since it changes the pitch. On a matched comparison, adaptive accuracy alone does not beat a well-chosen fixed accuracy on PSNR; the two land within 0.01 dB. What it buys is not needing to know that accuracy in advance: the fixed arm only matched after its tolerance was picked by watching worse values fail. The user guide already says fixed accuracy wins on cheap lower levels, and the feature should be framed that way.
What it reuses
The closed PR reimplemented three things DeepInverse already has. The rewrite will not.
deepinv.optim.linear.least_squares_implicit_backwardcg_utils.pyoptim_builder/FixedPointdeepinv.optim.Prior, whosefn(x, *args, **kwargs)already returns(B,)and acceptsthetaParametricPriorhierarchyTwo things I would like a steer on before writing them:
deepinv.optim.Priorand passthetathrough*args, or add a smallinit_thetamixin so a prior can say how to initialise its parameters? I lean to the mixin, since aPrioralone cannot tell the outer loop what shapethetahas.check_conv_fntoFixedPointacceptable, or should that live in the bilevel module?Split
Each under 3000 lines including tests. Each usable on its own.
MAID,MAIDConfig, the oracle interface, one quadratic oracle for tests, hypergradient vialeast_squares_implicit_backward. About 1200 lines. Docs: docstrings plus a short section in the optimisation guide.Prior. The oracle for any parametricdeepinv.optim.Prior, lower level viaoptim_builder, with a convexity check and a finite-difference check on the hypergradient as tests. About 900.deepinv.optim, one example on CBSD68 with three arms, short docs. About 1300.Dropped
The 1086 line guide, two of three examples, the saddle-point, smooth and stochastic variants, the TV baseline, and the accelerated ablation. Any of these can come later as its own small PR if wanted.
Total is roughly 4000 lines across four PRs, against 12700 in one. I will not open the first until the two questions above are settled.