Skip to content

Bilevel hyperparameter learning with MAID - #1318

Closed
MohammadSadeghSalehi wants to merge 61 commits into
deepinv:mainfrom
MohammadSadeghSalehi:maid-bilevel
Closed

MohammadSadeghSalehi wants to merge 61 commits into
deepinv:mainfrom
MohammadSadeghSalehi:maid-bilevel

Conversation

@MohammadSadeghSalehi

@MohammadSadeghSalehi MohammadSadeghSalehi commented Aug 10, 2026 •

Copy link
Copy Markdown

Bilevel hyperparameter learning with MAID

Adds the Method of Adaptive Inexact Descent to deepinv.optim.bilevel, with a
prior-agnostic interface so any convex, twice-differentiable regulariser can be
learned through it. Implements Salehi, Mukherjee, Roberts and Ehrhardt, SIAM
Journal on Mathematics of Data Science
2025 (arXiv:2308.10098), with the
saddle-point instantiation of Bogensperger, Ehrhardt, Pock, Salehi and Wong
(arXiv:2412.06436).

Why this is needed

DeepInverse already has strong reconstruction algorithms and a rich library of
priors. What it does not have is a principled way to choose the parameters
those priors depend on
. Today a user picks a regularisation weight by hand or
by grid search, and anything with more than two or three parameters is out of
reach: a grid over a 7533-parameter filter bank does not exist.

The alternatives each give something up.

Grid search does not scale beyond a handful of parameters, and gives no
gradient to follow.

Unrolling differentiates the algorithm rather than the solution, so the
learned parameters are tied to the iteration count and the solver used at
training time. Change either and the result is no longer the one that was
learned. It also stores every iterate, so memory grows with the number of
iterations.

Implicit differentiation at fixed accuracy solves the right problem, but
must pick a lower-level tolerance in advance. Too loose and the hypergradient
is wrong in a way nothing detects; too tight and most of the compute is spent
polishing solutions that a coarse step would have discarded anyway.

MAID resolves that trade-off by adapting the accuracy: it solves the lower
level loosely while far from a solution, tightens only when its descent test
demands it, and certifies each reconstruction with an a posteriori bound so the
inexactness is quantified rather than hoped away. That combination, adaptive
accuracy with a certificate, is what neither grid search nor unrolling
provides.

What this brings to DeepInverse

Any prior becomes learnable. Supply a per-sample energy and everything else
follows by autograd. That turns parameter selection from a per-prior
engineering exercise into a property of the library: the same code path learns
3 parameters or 7533, a hand-designed regulariser or a neural one.

Reconstructions come with a guarantee. For a lower level that is
mu-strongly convex, every reported result carries

||x* - x|| <= ||grad_x h(x)|| / mu,

computable from the returned iterate alone. A reconstruction is then a bounded
distance from the true minimiser rather than "whatever the solver returned
after N iterations". This matters most where it is hardest to check: on real
data with no ground truth.

It runs where users run. float32 on CUDA and MPS, verified unbiased against
a float64 reference, with batched solves and memory-aware batch sizing. Bilevel
learning is expensive, and an implementation that only works in float64 on CPU
is not usable for the problems people care about.

It is a bridge to the variational literature. Learned convex regularisers,
input-convex networks and adaptive regularisation all need exactly this
machinery. Having it in DeepInverse makes those methods reproducible here
rather than living in separate research codebases.

Learning any prior

A prior supplies one function, the energy of a batch of images. Everything else
follows by automatic differentiation: the lower-level gradient is a backward
pass in x, the Hessian-vector product a second one, and the mixed Jacobian a
backward pass in x then in theta. No derivative is written by hand, and the
model that is solved and the model that is differentiated are the same object.

class MyPrior(ParametricPrior):
    def __init__(self, channels=3):
        self.n_params = channels

    def energy(self, x, theta):            # -> (B,)
        w = torch.exp(theta).view(1, -1, 1, 1)
        return (w * x.abs()).flatten(1).sum(dim=1)

    def init_theta(self, *, dtype=torch.float64, device="cpu", seed=0):
        return torch.zeros(self.n_params, dtype=dtype, device=device)

Three priors ship with it, chosen to span the range rather than to compete:
a multi-convolution ridge regulariser (7533 parameters), total variation with
learned per-channel weights (3), and an input-convex network (5968).

Results

Denoising at sigma = 0.05, trained on one 32x32 crop from each of 16 distinct
CBSD68 images, evaluated on 8 unseen images. Same oracle, solver and accuracy
rules throughout; only the energy differs.

Prior Parameters PSNR (dB) SSIM
noisy input 26.05 0.5992
total variation, grid-tuned 1 31.28 0.8215
convex ridge regulariser 7533 31.64 0.8499
learned total variation 3 30.82 0.8154
input-convex network 5968 29.20 0.7452

This is evidence that the interface carries priors across a 2500x range in
parameter count, not a ranking: the budgets differ and none of the runs is
converged. The guide records why each stopped. The learned total variation
halts after 25 iterations at the accuracy floor because it is initialised at
the weight grid search finds. The input-convex network begins with an energy
far below the data term and must grow the prior before its hypergradient
carries much information.

A prior learned on 32x32 patches transfers to whole 256x256 images unseen at
any crop, where the ridge regulariser reaches 32.62 dB against 30.81 for
grid-tuned TV. The regulariser is convolutional, so it carries no notion of
image size.

Design decisions worth review

Hyperparameters are derived, not fixed. eps0, delta0 and alpha0
default to None, meaning derive at theta0: the accuracy from the
stationarity residual at A* y, the step from the hypergradient norm. No
absolute constant serves every prior, because both references are properties of
the regulariser. Measured at the same initialisation, eps0 = 1e-1 leaves a
denoising problem below its noisy input (24.70 dB against 26.01) while sitting
inside a plateau four orders wide for inpainting; and ||z0|| ranges from 3.7
for the ridge prior to 6.9e4 for the input-convex network, so a step suited to
one moves the other by roughly 6900 on its first trial. Explicit floats are
still honoured, and oracles that cannot supply the references fall back with a
warning.

Tolerances are dimensionless and dtype-aware. eps is a per-element rms
gradient tolerance, so it means the same accuracy at 32x32 and 256x256, and the
floor comes from torch.finfo(dtype).eps. Falling short of eps is not an
error: the certificate holds at any residual, so a solve that stalls at the
precision floor returns a valid, wider bound rather than raising. The floor
itself is found by detecting that the residual has stopped improving, since it
depends on conditioning and dtype and cannot be predicted in advance.

float32 is supported and verified. Against a float64 reference on identical
inputs the float32 hypergradient is inexact but not biased: cosine 0.99999713,
relative norm 2.4e-03, and MPS float32 matches CPU float32 to seven significant
figures. The angle is what matters, since inexactness shortens the
hypergradient while bias would rotate it.

Batched lower-level solves. Samples are independent given theta, so the
batched Hessian is block diagonal and every operator applies to all of them at
once; only the scalars stay per sample. Batch size is chosen from measured
memory. Accumulation is exact, so splitting a sample set changes peak memory
and nothing else (1.9e-15 in float64). Any DeepInverse physics works, including
operators with per-sample parameters, and construction refuses a physics that
couples samples across the batch.

Testing

114 tests across deepinv/tests/test_maid*.py, covering the lower-level
solvers, the certificate, the oracles, the accelerated path, minibatch
accumulation, the prior interface and the batched path.

The last two were added after codecov flagged batched.py and priors.py
at 17-18% patch coverage: they had no tests, and the coverage shown was
import-time execution alone. test_maid_priors.py and test_maid_batched.py
close that. Both were checked by mutation rather than assumed: flipping the
sign of the hypergradient fails all three finite-difference tests, and
dropping a batch from the accumulation fails both exactness tests.

The central check is a finite difference on the hypergradient, taken per
parameter block and per coordinate against
g(x*(theta)). It is the only test that can detect a
disagreement between the model that is solved and the model that is
differentiated, since automatic differentiation reproduces whatever the
forward code expresses, including a mistake. The per-block split matters: an
error confined to one block is invisible in a single whole-vector comparison.

Two checks are recommended for any new prior, and both are cheap:

  • assemble the Hessian on a small image and confirm its smallest eigenvalue is
    non-negative, establishing convexity in x
  • compare the hypergradient against a central finite difference along a random
    direction

Both are now run for all three supplied priors as part of the suite, in
test_maid_priors.py, rather than quoted from a one-off script.

Run locally on Apple silicon: 114 MAID tests pass, and test_optim.py gives 14
failed, 221 passed, 32 skipped. Every failure is on the MPS device, and the
same suite on unmodified main at 522bba9 gives 14 failed, 221 passed, 32
skipped with a test-for-test identical failure set, so none of them is
introduced here. Eight are test_least_squares_implicit_backward, the rest
test_tvprior_gradient, test_linear_system, test_MLEM,
test_CP_datafidsplit, test_CP_K and
test_least_squares_implicit_backward_nonleaf_buffer_grad. They look like
gaps in MPS operator coverage rather than anything specific to this branch.

test_datasets.py was not run locally: it downloads full public datasets
(Flickr2K alone exceeded 3 GB before being stopped). CI covers both.

Breaking change

MAIDConfig.eps0, delta0 and alpha0 change from float to float | None
and now default to None. Code passing explicit values is unaffected. The
previous default of eps0 = 1e-1 was harmful under the per-element tolerance
introduced here, so anyone relying on it was already getting poor behaviour.

Documentation

  • docs/source/user_guide/reconstruction/bilevel.rst, covering the oracle
    interface, certified and non-certified paths, the accelerated heuristics with
    an ablation, numerical precision and GPU backends, batching, and learning a
    new prior
  • three gallery examples under examples/optimization/
  • API entries and a changelog entry

Checks to be done before submitting your PR

  • python3 -m pytest deepinv/tests runs successfully.
  • black . and ruff check . run successfully.
  • make html runs successfully (in the docs/ directory).
  • Updated docstrings related to the changes (as applicable).
  • Added an entry to the changelog.rst.

LLM policy

LLM usage is ok, but not PRs generated 100% by AI. See our LLM policy Tick below as appropriate:

  • I did not use LLM tools to write the code
  • LLM tools helped me to write part of the code
  • An LLM tool wrote all of the code.
  • An agent submitted the PR and wrote the description.

Implement Algorithm 3.1 (MAID) and Algorithm 3.2 (INEXACT_GRADIENT) with
the Theorem 2.1 a posteriori hypergradient error bound, specialised to a
gradient-descent lower level on the section 4.1 quadratic least-squares
bilevel problem with known optimum.
MAID now depends only on a HypergradientOracle. The smooth IFT+CG path is
unchanged in behaviour behind SmoothHypergradientOracle (certified,
Theorem 2.1). The saddle-point path adds PDHG piggyback with the
arXiv 2412.06436 residual distances and Theorem 2 hypergradient bound
(certified when mu_g and mu_f* are known). Non-certified oracles require
an explicit allow_uncertified opt-in.
check_descent_direction defaults to False: Algorithm 3.2 no longer forms
omega unless the user restores the certified path. Lemma 3.5 still
guarantees decrease on every accepted step; Theorem 3.19 does not apply
as proven. Backtracking failures are counted as the a posteriori
detector. Section 4.1 measurement shows the skip is not faster end to
end on that instance.
…factor

GoalOrientedEstimator forms a dual-weighted residual estimate of
||z - grad f|| from the linear-solve residual and the lower-level
gradient, with default safety_factor=1.25 and cg_budget=5. On quadratic
instances the certified bound is tens of times loose; DWR recovers the
error to within about one percent, and 1.25 removes every under-estimate
in the measured sweeps. The same factor remains conservative on a smooth
nonquadratic log-cosh lower level (0/36 under-estimates). The estimator
is non-certified and requires allow_uncertified=True.
Remove bare print from the DWR tests. Parameterise the nonquadratic sweep
by data scale 1, 5 and 20 with sech^2 spread regime checks so the
nonlinearity coverage is real. Document the supervisor nonlinearity table
on the estimator. Add a gallery demo comparing MAID to fixed-accuracy
descent on total lower-level GD iterations, and a user-guide page for
bilevel oracles, certification and measured tables.
Wrap BaseOptim behind residual stopping: gradient residual for GD,
proximal residual for PGD and FISTA. Warm starts pass the previous
reconstruction through init. Learn a real Tikhonov prior weight on
inpainting via TikhonovWeightOracle. Correct docstrings that previously
claimed optimisers that were not wired. Unit tests cover residuals,
warm start, finite-difference hypergradients and MAID with all three
solvers.
The previous flagship used a well-conditioned Tikhonov inpainting problem
where a residual of 1e-4 costs about one GD step per solve. MAID then pays
for line-search trials and loses (246 vs 26 BaseOptim iterations). That is
the wrong regime to advertise the method.

Replace it with:
- a condition-number-30 quadratic least-squares bilevel (section 4.1) where
  MAID reaches a 0.5% gap to f_star in 60062 GD steps against 129383 for
  fixed accuracy 1e-4;
- a crossover table over condition numbers 2, 5, 10, 20, 30 so the claim is
  when MAID is faster, not that it always is;
- the inpainting case kept as a labelled counterpoint with the measured
  overhead numbers.

Document the same tables in the bilevel user guide under "When to use MAID".
Evaluate the mean upper level over m samples by walking the dataset in
fixed index order and accumulating z_i, omega_i and g_i with sequential
floating-point addition. Chunk size bounds concurrent working memory only;
it does not change the reduction order, so hypergradients and MAID
trajectories are bitwise invariant to chunk size.

Error bound for the mean is omega = (1/m) sum omega_i, from the triangle
inequality on (1/m) sum (z_i - grad f_i). Documented in the module
docstring and tested against the true mean hypergradient error.

Also:
- promote the outer-step column in the crossover write-up (at cond 30,
  MAID uses 6 outer steps against 23; each costs a hypergradient);
- goal-oriented estimator remains per sample (Krylov recycling does not
  span a chunk);
- gallery example and user guide report trajectory invariance, peak
  working memory vs chunk size, and the crossover with chunking enabled
  on the expensive end (cond >= 5, m=4, chunk=2).
Two optional MAIDConfig switches (both off by default):

- nonmonotone: Zhang-Hager window reference C_k with a monotone Lemma 3.5
  fallback so a sandwich-depressed C cannot block a valid step. C is
  updated with U_lower at the accepted trial accuracy eps_k. Proof sketch
  in the maid module docstring; author verification required.
- bb_init: Barzilai-Borwein initial step from (s, y), clamped, with
  fallback rho_bar * alpha_k. Changes only where backtracking starts.

accelerated_maid_config() enables both. Nesterov/heavy-ball on theta is
explicitly out of scope (would leave -z).

Three-way crossover (gap 5e-3, n=80, d=4): accelerated MAID is already
cheaper at condition number 5 (GD ratio 0.58 vs fixed) where vanilla is
even (1.01), and holds ratio about 0.20 above that. At condition number 2
it still loses to fixed (1.20) but improves on vanilla (1.45). Backtracking
failures drop at every condition number. On the inpainting counterpoint,
BT falls from 16 to 1 and BaseOptim iters from 246 to 24.
Memory: replace the 24/48/96-byte chunk table with an MB-scale measurement
that varies m at fixed chunk size. Peak working memory for float64 states of
length 200000 is flat at 6.10 MB for chunk 4 across m in {16,32,64,128};
warm-start storage grows as O(m). Against chunk size at m=64 the peak is
1.53, 3.05, 6.10, 12.2, 24.4 MB. Measurement is concurrent tensor bytes on
CPU (no CUDA on this machine).

Wall clock: the earlier minibatch GD win was a counter bug. Reusing the same
QuadraticBilevelLS objects for MAID and fixed made n_gd_iters sum both runs
(355611 + 123646 = 479257). With fresh problems per method, MAID loses on
both GD (2.88x) and wall (2.90x) at m=4, cond 20, chunk 2. About 99 percent
of wall is lower-level work; hypergradient and certified omega are under
1 percent. The cause is line search multiplying sample_lower calls
(1168 vs 72). Report both metrics.

PR body drafted at .scratch/pr_body.md (not opened). Mechanism first,
proven/measured/heuristic table, regimes where the method should not be used.
The minibatch comparison previously ran only vanilla MAID. The structural
cost is line search: each trial re-solves all m samples, so every avoided
backtracking failure saves m lower-level solves. Acceleration is the
measurement that has to be run on that path.

At m=4, cond 20, chunk 2, max_iter 10 (fresh counters, fixed to each
method's f):

  vanilla MAID  355611 GD  1168 sample_lower  BT 11  wall 4.24s  GD/fixed 2.88
  fixed (van)   123646 GD    72 sample_lower          wall 1.48s
  accel MAID    144294 GD   688 sample_lower  BT  7  wall 1.73s  GD/fixed 0.90
  fixed (acc)   160380 GD   100 sample_lower          wall 1.92s

Acceleration flips the GD ratio from 2.88 to 0.90 and the wall ratio from
2.87 to 0.90. The cost table over cond in {5,10,20} shows the same pattern.
Vanilla numbers stay in the example: a table where vanilla loses and
accelerated wins is a better argument for the acceleration than showing
only the winner.

PR body updated in .scratch/pr_body.md (not opened).
Self-contained gallery script examples/optimization/demo_maid_imaging.py
learns a Tikhonov weight on Set3C greyscale 128x128 for random-mask
inpainting and Gaussian deblurring. Three methods share the same initial
theta: fixed residual 1e-4, MAID, and accelerated MAID.

Figures (written to .scratch/figs/ when the script runs): reconstruction
panels with PSNR, and upper-level objective against cumulative lower-level
GD for the three methods. TV weight learning is not used because a
nonsmooth TV IFT residual path is not wired; the script states that.

Also floor residual_tol at 1e-8 in TikhonovWeightProblem so GD does not hang
when eps * mu falls below the floating-point residual floor on weak-lambda
blur.

Measured (CPU float64): deblur PSNR 9.52 -> 20.23 (fixed), 21.48 (MAID),
21.68 (accelerated) at higher GD cost for the adaptive methods. Inpainting
PSNR gains stay under 1 dB for all three (Tikhonov is a weak prior there).
Learned prior of the Goujon-Unser form

  R(x) = sum_k lambda_k sum_j rho((W_k * x)_j),  rho = mu log cosh
  lambda_k = exp(vartheta_k)

with 8 kernels of 5x5 (208 parameters) and a gamma ||x||^2 / 2 floor so
mu >= gamma is known for residual stopping. The exp weight parameterisation
matches TikhonovWeightProblem and keeps the lower level convex for every
finite parameter value.

Hessian and mixed-Jacobian products use autograd; the smooth MAID oracle
and gradient residual are unchanged. Minibatch training over reconstruction
samples is via build_crr_minibatch_oracle.

Gallery demo_maid_crr.py trains on two Set3C 64x64 greyscale images and
evaluates held-out on a third. Scalar Tikhonov is grid-tuned on the train
set per problem. Measured held-out PSNR:

  inpainting: grid Tikhonov 5.06 dB, CRR-MAID 22.25 dB
  deblur:     grid Tikhonov 22.00 dB, CRR-MAID 20.83 dB

Inpainting is where spatial coupling matters; deblur shows a well-tuned
scalar can still win on a short CRR budget. Both numbers are reported.
Without a norm constraint the pair (lambda_k, W_k) is degenerate under
the small-signal expansion of log-cosh, so the hypergradient on log
lambda is nearly flat and the outer method only moves kernels. Normalise
each free kernel to unit Frobenius norm after zero-mean centring; scale
then lives in lambda_k = exp(vartheta_k). Free coordinates are
initialised at free_kernel_scale = 1 so the Euclidean chart balances
kernel and weight blocks under a single outer step size. Document the
degeneracy, print init/final lambda tables in the demo, and re-chart
after outer steps.
The previous commit message claimed unit-norm kernels, but a concurrent
edit had replaced the staged convex_ridge module with a multiconv draft.
Restore the Goujon-Unser log-cosh CRR with unit-norm kernels, free
kernel chart scale 1, demo lambda tables, and the lambda-motion test.
Avoids Python 3.12 SyntaxWarning on LaTeX backslashes in the module
docstring without changing the rendered text.
Replace the single-layer log-cosh FoE with the reference multiconv CRR:
Lipschitz-normalised stacked convolutions, log scaling, smooth L1,
zero-mean first layer, and weak_convexity=0. Bilevel parameters stay
flat for MAID. The CRR demo uses colour Set3C, labels reconstruction
panels by the prior rather than MAID, drops Tikhonov-only figures,
reports converged residuals next to PSNR, and prints the exp(scaling)
motion table. Imaging demo figure labels follow the same framing.
State precisely that log-scale s cancels in the quadratic region of
smooth_l1 (with the response-scale table) so a flat exp(s) is an
architecture property, not a failed outer step.

Give the fixed-accuracy bilevel arm the same Armijo line search as MAID
and hold only eps and delta constant (option 1). Matched outer budget
N_OUTER=6 is printed with the summary table so the short-run numbers
are not over-read.
Replace the hand-rolled GD loop in CRRSampleProblem.solve_lower with
DeepInverse BaseOptim via base_optim_lower. Default solver is FISTA on
a single smooth fidelity (CRR energy folded into DataFidelity, null
prior so the proximal residual equals the gradient residual).

The certificate is residual-based and solver-independent:
||x - xstar|| <= ||grad h|| / mu with mu = gamma. Any algorithm that
reaches the residual earns the same bound.

The CRR demo prints a GD-vs-FISTA ablation (iterations and wall clock)
before learning so the solver choice is evidenced rather than assumed.
mu is now mu_data + gamma. Denoising with A = I contributes mu_data = 1,
so gamma may be set to zero and the ridge floor no longer competes with
the learned prior. General physics keeps mu_data = 0 and relies on the
explicit floor.

The CRR demo is reordered with denoising first (gamma = 0), then
inpainting and deblur (gamma = 1e-2). Every residual is printed with the
strong-convexity distance bound ||grad h|| / mu. A ridge ablation
(prior off, data term plus gamma floor only) reports held-out PSNR on
each problem; when the floor alone explains most of the reconstruction
PSNR the prior comparison is labelled uninformative.
Demo had sigma_init=0.5 and beta_init=1.0. At sigma=0.5 the smooth_l1
knee sits five times further out than the reference (sigma=0.1), so filter
responses stay in the quadratic region where the regulariser degenerates
to a generalised Tikhonov penalty on filter responses, which is the class
it exists to beat. Restore sigma_init=0.1 and beta_init=4.0.

Also replace the fixed twelve-iteration L_data power method (optimistic
by construction, 2.94 percent low on deblurring) with adaptive power
iteration plus a safety factor, add a Hess-R probe for L_prior, residual
histories, certificate checks against a tight reference solve, and
FISTA_RESTART / truncated Newton solvers so the step size and the
minimiser claim can be audited before any PSNR is reported.
Recomputing lip(theta) and detaching it made the forward solve and the
differentiated model disagree on the weight block. Store lip once on the
prior, reuse it in energy, gradient, HVP and mixed Jacobian, and use
2*exp(beta)*lip(W)/lip_0+gamma for the step-size chart. Refresh only
between outer iterations.

Add a central finite-difference test on scaling, beta, random weight and
full directions, and three weight coordinates. Add grid-tuned isotropic
TV as the mandatory baseline in the CRR demo.
Freezing lip_0 made the hypergradient exact but destroyed the scale
invariance the normalisation provides. With a stale lip_0 the prior
curvature grows as lip(W)/lip_0, so a line-search trial that inflated the
weights by ~1000x drove L to 1.3e8 and the step to 7.8e-09, and the demo
died with Newton unable to reach its residual.

Recompute lip per theta and compute it with the graph attached inside the
differentiable energy, so autograd carries d lip / d w rather than dropping
it. Forward and backward stay one model, and L_prior returns to
2*exp(beta)+gamma, independent of the weight scale.

Finite differences on a training sample, lower level at residual 6.6e-9:
  scaling block   -7.91491026e+00 vs -7.91491029e+00   rel 3.0e-09
  beta             4.20544024e-03 vs  4.20541793e-03   rel 5.3e-06
  weight dir       1.08079019e-02 vs  1.08079177e-02   rel 1.5e-06
  full dir         1.55231895e-03 vs  1.55232085e-03   rel 1.2e-06

Replace the frozen-lip test with one asserting lip tracks the weights and
that the step-size chart is invariant under W -> 1000 W.
x0 = A^* y lies in the range of A, so every projection operator satisfies
A(x0) = x0 exactly. An inpainting mask is idempotent and was therefore
misread as the identity, taking the single-prox denoising shortcut for a
problem it does not solve. The primal-decrease assertion caught it rather
than letting a wrong baseline be reported.

Probe with a random vector instead. Inpainting now runs proximal gradient
and decreases the objective at every lambda tested.
eps was an absolute threshold on ||grad h||, so it meant different accuracy
at different resolutions: 1.8e-07 per element at 32x32 against 2.3e-08 at
256x256. It is now a per-element rms tolerance,

    ||grad h|| / (sqrt(n) * mu) <= eps

which is resolution-independent and directly comparable to finfo(dtype).eps.

The floor is derived from the dtype rather than hard-coded at 1e-8. Measured
float32 floor is 2.7e-06 per element; the clamp at 100*finfo.eps gives
1.2e-05, above it with margin, and 2.2e-14 in float64, well under its
solver-limited floor. Previously float32 could not satisfy the defaults at
all: eps=1e-4 was unreachable, so the reporting solve raised on CUDA and MPS.

Falling short of eps is no longer an error by itself. The certificate
||x* - x|| <= ||grad h||/mu holds at any residual, so a solve stalled at the
precision floor still returns a valid, wider bound; raising discarded a
correct reconstruction along with a usable certificate. A shortfall too large
to be the dtype floor still raises, and says so.

Tests updated to the normalised contract rather than loosened, and the
distance bound is asserted per element so it compares against the image
scale. REPORT_EPS, LEARN_EPS0 and FIXED_EPS rescaled by sqrt(n).
The reachable residual depends on the problem's conditioning and, now that
lip is recomputed inside the gradient, on FFT error too. It cannot be
predicted from finfo(dtype).eps alone: a floor calibrated as 100*eps from one
sample (1.2e-05 per element) was out by 4x on another, where the achievable
rms was 4.95e-05, so float32 still raised on a tight request.

Track the best residual instead and stop when it fails to improve by 0.1%
for 25 consecutive iterations. Return the best iterate, since the residual
wobbles once precision-limited. A solve still improving when it exhausts
max_iter remains a hard error, and the message distinguishes the two cases.

Verified across device x dtype: every combination now returns a valid
certificate and none raise. float32 is inexact but not biased against a
float64 reference on identical inputs (cosine 0.99999713, relative norm
2.4e-03), and MPS float32 matches CPU float32 to 7 significant figures.
minibatch.chunk_size bounds peak memory but still solves samples one at a
time, so an outer iteration is N sequential Newton solves on tensors of a few
thousand elements. That is dispatch bound, not arithmetic bound: three chained
convolutions on (1,3,32,32) take 145 us on MPS, while (24,3,32,32) takes
446 us for all 24.

The samples are independent given theta, so the batched Hessian is block
diagonal and every operator applies to all of them at once. Only the scalars
stay per sample: CG alpha and beta, the Armijo step, and the convergence and
stall tests. Sharing any of those would couple independent problems and break
the Krylov property.

cg_solve_batched freezes converged samples by zeroing their alpha rather than
dropping them, so the tensor stays rectangular. Verified bit-exact against
sequential CG on distinct SPD operators (rel diff 0.0).

Batch size comes from a measured probe, not a formula. Estimating autograd
graph retention through three convolutions plus a CG basis is the kind of
guess that was already wrong twice here.

Verified end to end on 8 CBSD68 samples:
  cpu/float64  solve 9.9e-17, hypergrad 5e-06 (cos 1.0000000000), 1.31x
  mps/float32  solve 1.2e-05, hypergrad 5.4e-04 (cos 0.9999998), 4.45x
and splitting into batches of 3 moves the hypergradient by 1.9e-15 (float64),
so batch size is a memory decision with no numerical consequence.
The restriction to A = I was mine, not the method's. DeepInverse operators
already act per sample along the batch dimension, including with per-sample
parameters: an Inpainting built with a (B,C,H,W) mask and a BlurFFT with
batched kernels each apply their own operator to their own sample.

BatchedCRR now takes a physics and forms the data term from A and A_adjoint,
with mu_data = 1 for A = I and 0 otherwise, matching CRRSampleProblem, so a
general operator needs gamma > 0 for a positive modulus.

Batching is only valid if A is block diagonal across the batch, so that is
checked at construction rather than assumed: perturb one sample, confirm the
others' measurements do not move, and refuse the physics otherwise. An
operator that mixed samples would silently corrupt every gradient in the
batch.

Samples needing genuinely different operators are grouped into separate
batches and summed, which costs nothing because accumulation is exact.

Verified against the sequential path, float64, 4 samples at 24x24:
  denoising  A=I          dx=8.7e-17  dz=1.2e-05  cos=0.9999999999
  inpainting per-sample   dx=9.9e-16  dz=2.5e-07  cos=1.0000000000
  deblur BlurFFT          dx=2.8e-13  dz=9.4e-08  cos=1.0000000000
BatchedMinibatchOracle is a drop-in for MinibatchOracle with the same API and
the same mean reduction, but solves each group as one batch and sums the
hypergradients. Batch size comes from a measured probe against
cuda.mem_get_info or mps.recommended_max_memory, so it adapts to the device
rather than being tuned by hand, and it is numerically free: accumulation is
exact to 1.9e-15.

One deliberate approximation. mixed_jac_T returns the sum over samples, so the
||J|| fed to the error bound is the norm of the summed operator. Probing it
with a vector supported on one sample recovers that sample's J_i, so it upper
bounds every per-sample norm and omega comes out at least as large as the
sequential mean. That is conservative: extra backtracking at worst, never an
understated hypergradient error.

Consequently the two oracles are not bitwise identical. MAID is adaptive, so a
larger omega changes line-search decisions and the trajectories drift: over 25
outer iterations they agree to 6e-09 at iteration 0 and 6.4e-03 at the end,
cos(theta) = 0.9999884, both reaching the same objective (1.2375 sequential,
1.2296 batched). Same algorithm under slightly different error bounds, not a
discrepancy, but worth stating plainly.

Also fixes two device bugs in tv_baseline that made it unusable off CPU: a CPU
generator was asked to fill a device tensor ('Expected a mps device type for
generator but found cpu'). Drawing on CPU then moving also keeps a seed
meaning the same thing on every backend.
The bilevel package had no progress output at all, so a 300-iteration run was
a silent black box: there was no way to tell a job that was converging from
one that was wedged, and no way to answer how long it had left.

Follows the library's own convention rather than general Python practice.
DeepInverse does not use the logging module anywhere; FixedPoint and Trainer
use a verbose flag with print, plus tqdm gated on show_progress_bar. So:
verbose, show_progress_bar and log_every on MAIDConfig, all off by default, a
tqdm bar carrying f, ||z||, alpha, eps and omega, and a per-iteration line
when verbose without a bar.

Stopping is now explicit about which case occurred. Reaching max_iter while
still descending is a budget limit, not convergence, and the message says so:
every training run in this work so far hit that case and reported a
non-converged number as if it were a result.

omega prints as 'off' rather than 'nan' when check_descent_direction is False.
Algorithm 3.1 never forms it by contract, but a bare nan reads as a numerical
failure. It is finite on the certified path (1.1e-05 against a threshold of
0.35), which the new output made easy to confirm.
Usage appeared at line 967 of 1006, after every theoretical and diagnostic
section, so a reader met the full derivation before seeing how to run
anything. Practical material now comes first: problem form, quick start,
learning a prior, precision and batching. Oracles, certificates, goal-oriented
estimation, the accelerated heuristics and minibatch accumulation follow as
reference.

The quick start learns a convex ridge regulariser end to end, from
constructing the prior through training to reconstructing an unseen
measurement and reading its certificate, and shows the scalar-weight case as a
smaller starting point. It relies on eps0 and alpha0 defaulting to None, so no
hyper-parameter is chosen by hand.

Three passages still described how conclusions were reached rather than what
holds. They now state the behaviour and what follows from it, in particular
that a prior starting far from the data-term scale benefits from a warm start
and that MAID reports which limit a run reached.
…trace

The guide carried tables but no images, so the difference between a learned
prior and total variation was stated rather than shown, and MAID's adaptive
accuracy was described without being visible.

The denoising figure magnifies a crop under each panel, where total variation
flattens texture that the learned prior keeps; this is the mechanism behind
its larger margin in SSIM than in PSNR. The convergence figure shows the step
size and lower-level accuracy rising and falling over a run, which is the
behaviour that separates MAID from a fixed-accuracy scheme.

Both are sized for the repository: 558 KB and 165 KB against 637 KB for the
existing schematic, the denoising panel palette-quantised from 5.2 MB while
keeping the insets legible.
Deriving eps0 from the initial residual was available only on the batched
path, so every use of MinibatchOracle took the fallback and emitted a warning
it could not act on. CRRSampleProblem now exposes initial_residual_rms,
CRRSampleOracle forwards it, and MinibatchOracle averages over its samples.

The sequential path now derives eps0 = 1.33e-03 from a residual of 6.64e-02
with no warning. The fallback remains for oracles that genuinely cannot supply
the quantity, which is what the uncertified-oracle test exercises.
# Conflicts:
#	docs/source/changelog.rst
@codecov

codecov Bot commented Aug 11, 2026

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 80.55194% with 747 lines in your changes missing coverage. Please review.
✅ Project coverage is 88.26%. Comparing base (522bba9) to head (6c78e16).
⚠️ Report is 4 commits behind head on main.

Files with missing lines Patch % Lines
deepinv/optim/bilevel/batched.py 17.50% 278 Missing ⚠️
deepinv/optim/bilevel/priors.py 18.34% 178 Missing ⚠️
deepinv/optim/bilevel/maid.py 78.50% 66 Missing ⚠️
deepinv/optim/bilevel/crr_bilevel.py 84.77% 65 Missing ⚠️
deepinv/optim/bilevel/cg_utils.py 63.63% 36 Missing ⚠️
deepinv/optim/bilevel/convex_ridge.py 88.08% 23 Missing ⚠️
deepinv/optim/bilevel/minibatch.py 90.45% 23 Missing ⚠️
deepinv/optim/bilevel/tv_baseline.py 85.32% 16 Missing ⚠️
deepinv/optim/bilevel/prior_learning.py 90.54% 14 Missing ⚠️
deepinv/optim/bilevel/nonquadratic.py 87.50% 11 Missing ⚠️
... and 8 more
Additional details and impacted files
@@            Coverage Diff             @@
##             main    #1318      +/-   ##
==========================================
- Coverage   89.26%   88.26%   -1.01%     
==========================================
  Files         245      270      +25     
  Lines       28529    32370    +3841     
==========================================
+ Hits        25466    28570    +3104     
- Misses       3063     3800     +737     
Flag Coverage Δ
full-cpu 87.36% <80.55%> (-0.89%) ⬇️

Flags with carried forward coverage won't be shown. Click here to find out more.

☔ View full report in Codecov by Harness.
📢 Have feedback on the report? Share it here.

The two largest new modules had no tests: a case-insensitive search for
their class names across the MAID suite returned nothing, so the 17-18%
coverage codecov reports was import-time execution alone.

Add 29 tests. For every supplied prior, assemble the Hessian on a small
image and check its smallest eigenvalue, and compare the hypergradient
against a central finite difference; both properties were claimed but
never checked. For the batched path, check that splitting a sample set
leaves the hypergradient unchanged and that it agrees with the sequential
oracle, which is the reference implementation. Verified by mutation: a
sign flip on the hypergradient fails all three finite-difference tests,
and dropping a batch from the accumulation fails both exactness tests.

The docs build failed because two examples take Path(__file__), which
sphinx-gallery leaves undefined since it execs examples rather than
importing them. Fall back to the working directory.
WEIGHT_SCALE was 0.02, fifty times below the reference scale the trained
runs use. That leaves the prior nearly inert: the data term dominates, the
lower level is easy at any tolerance, and an adaptive accuracy schedule has
nothing to adapt to.

The example then failed to show what it exists to show. On denoising MAID
and the fixed-accuracy arm came out identical to two decimal places, and on
the two ill-conditioned problems MAID came out worse, so the example
contradicted its own premise. At the reference scale MAID is ahead in all
six comparisons, by 0.05 to 2.29 dB, and the run is also shorter because a
well-scaled prior needs fewer inner iterations.
The example trained on two Set3C images and held out one, stopped at 30
outer steps, and ran through the sequential oracle rather than the batched
one this branch adds. It also carried a comment claiming load_theta freezes
the Lipschitz constant, which stopped being true when the hypergradient was
fixed to differentiate through it.

Rebuilt on CBSD68, sixteen 32x32 training crops and eight held out. Sixteen
is the smallest set that generalises: with eight the regulariser fits the
training crops and loses to TV on held-out data by 0.46 dB, and raising the
budget from 60 to 100 outer steps only moves that to 0.26 dB while the
train-to-held gap stays at 1.93 dB.

The comparison against fixed accuracy was wrong three times over, each time
flattering this branch. Holding the control at MAID's own starting tolerance
of 1e-3 left it unable to certify descent, so it exhausted the line search
every iteration and never stepped. Giving MAID bb_init while hand-rolling
the control without it made the step rule the difference rather than the
accuracy. Hardcoding alpha0 = 1e-3, four hundred times below what
auto_initial_step derives, starved both arms so the accuracy schedule could
not matter either way.

Both arms now take the derived step and the same step machinery, and the
fixed arm gets a tolerance a practitioner would actually pick. On a level
field adaptive accuracy alone matches fixed accuracy in PSNR at this budget,
31.07 against 31.08, so the example reports all three arms rather than
claiming a gain it cannot show. The accelerated switches are what clear TV,
31.75 against 31.24 held out.

Also set pack_init_theta's default weight_scale to 1.0 to match
ConvexRidgePrior2.init_theta, documenting that the scale leaves the energy
exactly invariant and only moves the hypergradient, and fix a test that
asserted a solve reaches 1e-11, which sits at the float64 floor.
The docstring still claimed the second arm shows adaptive accuracy is what
does the work, and described two arms where there are now three. The run
shows the opposite: against a well-chosen fixed tolerance the two land
within 0.01 dB and MAID spends about twice the lower-level iterations. The
user guide already says as much, that fixed accuracy wins on a cheap lower
level and that Barzilai-Borwein initialisation carries almost all of the
benefit, so the example was contradicting the documentation it ships with.
# Conflicts:
#	docs/source/changelog.rst
#	docs/source/refs.bib
accelerated_maid_config(**cfg.__dict__) returns a config with acceleration
off. The expansion makes every field an explicit keyword, including
bb_init=False from the defaults, and setdefault leaves explicit values
alone. The call reads as enabling acceleration and does nothing.

Found while writing the CRR example, where two arms that should have
differed produced byte-identical output down to the iteration count.
Documented on the function and pinned by a test, since the shape of the
mistake is invisible at the call site.
@MohammadSadeghSalehi

Copy link
Copy Markdown
Author

A few changes since the last CI run, and one request.

Request: the workflows need approval again to run. CI last ran on 6c78e16; the commits since have no check runs, so the Codecov report above is measured against that commit and predates the tests mentioned below.

Coverage. Codecov flagged batched.py at 17.5% and priors.py at 18.3%, and it was right: those two modules had no tests at all. They now do, 30 tests across test_maid_batched.py and test_maid_priors.py, covering the batched hypergradient accumulation, the block-diagonal check, per-sample CG, and for every supplied prior a Hessian-eigenvalue convexity check and a finite-difference check on the hypergradient. Measured locally over deepinv/optim/bilevel, priors.py is at 94% and batched.py at 75%, the remainder there being the memory-probe and CUDA/MPS branches that a CPU run cannot reach. The suite is 116 tests.

Docs build. The failure on the previous run was real and is fixed: two examples used Path(__file__), which is undefined when sphinx-gallery execs a file rather than importing it.

The CRR example has been rebuilt. It trained on two Set3C images and held out one, stopped at 30 outer iterations, and ran through the sequential oracle rather than the batched one. It is now CBSD68, sixteen 32x32 training crops and eight held out, on the batched path. Sixteen is the smallest set that generalises here; with eight the regulariser fits the training crops and loses to TV on held-out data, and raising the budget to 100 outer steps does not close that.

Worth flagging explicitly, since it changes what the example claims: on a level comparison, adaptive accuracy alone does not beat a well-chosen fixed accuracy on PSNR at this budget. The two land within 0.01 dB and MAID spends about twice the lower-level iterations. What clears the grid-tuned TV baseline is the accelerated pair, Barzilai-Borwein initialisation and the nonmonotone test, at 31.75 dB against 31.24 held out. This matches what the user guide already says, that fixed accuracy wins on a cheap lower level and that BB carries almost all of the benefit; the example had been asserting otherwise. It now reports all three arms so each switch can be attributed separately.

The value of adapting the accuracy is robustness rather than PSNR: the fixed arm only matches because its tolerance was picked after seeing which values fail. Given MAID's own starting accuracy of 1e-3, the same arm cannot certify descent, exhausts its line search every iteration and never takes a step.

Also fixed: pack_init_theta defaulted to weight_scale=0.05 while ConvexRidgePrior2.init_theta used 1.0. The scale leaves the energy exactly invariant, since the features are Lipschitz-normalised, but moves the hypergradient by 30x, so the mismatch silently changed step-size behaviour depending on which entry point you came through. Both are 1.0 now, with the invariance documented.

@Andrewwango
Andrewwango marked this pull request as draft August 12, 2026 14:05
@Andrewwango Andrewwango changed the title Bilevel hyperparameter learning with MAID WIP: Bilevel hyperparameter learning with MAID Aug 12, 2026
@MohammadSadeghSalehi MohammadSadeghSalehi changed the title WIP: Bilevel hyperparameter learning with MAID Bilevel hyperparameter learning with MAID Aug 23, 2026
@MohammadSadeghSalehi
MohammadSadeghSalehi marked this pull request as ready for review August 23, 2026 18:14
@MohammadSadeghSalehi

Copy link
Copy Markdown
Author

@Andrewwango you marked this WIP on 12 August, which was fair: my comment that morning described the example being rebuilt and the test coverage being filled in, so it read as a moving target. That work is finished, so I have taken it out of draft.

What changed since you looked:

  • The CRR example was rebuilt. It had trained on two Set3C images and held out one, stopped at 30 outer iterations well short of convergence, and ran through the sequential oracle rather than the batched one this PR adds. It is now CBSD68, sixteen 32x32 training crops and eight held out, on the batched path, and it is 371 lines rather than 1415.
  • The example's conclusion changed and is now stated plainly rather than asserted. On a matched comparison, adaptive accuracy alone does not beat a well-chosen fixed accuracy on PSNR at this budget: the two land within 0.01 dB and MAID spends roughly twice the lower-level iterations. What clears the grid-tuned TV baseline is the accelerated pair, Barzilai-Borwein initialisation and the nonmonotone test, at 31.75 dB against 31.24 held out. This agrees with what the user guide already said about cheap lower levels, so the example had been contradicting our own documentation. It now reports all three arms so each switch can be attributed.
  • Codecov was right that batched.py and priors.py were untested. They now have 30 tests, including a Hessian-eigenvalue convexity check and a finite-difference hypergradient check for every supplied prior. The suite is 116 tests.
  • The docs build failure was real and is fixed: two examples used Path(__file__), which is undefined when sphinx-gallery execs rather than imports.
  • Merged main four times since; currently conflict-free.

One request: the workflows still need approval to run. CI has only ever run on 6c78e16, so the Codecov report on this PR predates the tests above.

Happy to split anything out if the diff is too large to review in one piece.

@Tmodrzyk

Tmodrzyk commented Sep 4, 2026

Copy link
Copy Markdown
Contributor

Hi @MohammadSadeghSalehi,

Thanks for the work you put into this. Unfortunately, at over 12 000 lines, this PR is too large for us to review or maintain and exceeds our approximate 3 000 lines limit (see our guidelines with regard to PR scope https://deepinv.org/contributing.html#code-scope).
We’re therefore going to close this PR.

If you would like to continue working on this feature, we suggest:

  • opening an issue to motivate the addition of this feature, and for us to discuss the possible integration and which existing abstractions can be reused 
  • dividing the work into smaller PRs
  • be more concise with respect to the documentation added, you can refer to the doc guidelines for more details: https://deepinv.org/contributing.html#docstring-guidelines 

We would be happy to discuss this plan with you before you begin. Thanks for your understanding and for your interest in contributing !

@MohammadSadeghSalehi

Copy link
Copy Markdown
Author

Thanks, that is fair. I have opened #1390 with the plan you suggested: what the feature adds, which existing abstractions it reuses instead of reimplementing, a four PR split each under 3000 lines, and two design questions I would like settled before opening the first one.

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

Labels

None yet

Projects

Status: Done

Development

Successfully merging this pull request may close these issues.

2 participants