Skip to content

Bilevel learning of prior parameters with MAID: proposed incremental plan #1390

Description

@MohammadSadeghSalehi

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:

  1. 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.
  2. 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.

  1. 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.
  2. 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.
  3. Batched oracle. Samples solved together, hypergradients accumulated exactly. About 700.
  4. 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.

No activity

Activity on this issue will appear here.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    priority: lowNice to have, non-urgent issue or PR.type: featureNew feature, enhancement or request

    Projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions