Repository navigation
Bilevel hyperparameter learning with MAID - #1318
MohammadSadeghSalehi wants to merge 61 commits into
Conversation
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 Report❌ Patch coverage is 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
Flags with carried forward coverage won't be shown. Click here to find out more. ☔ View full report in Codecov by Harness. |
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.
|
A few changes since the last CI run, and one request. Request: the workflows need approval again to run. CI last ran on Coverage. Codecov flagged Docs build. The failure on the previous run was real and is fixed: two examples used 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: |
# Conflicts: # docs/source/refs.bib
# Conflicts: # docs/source/changelog.rst
|
@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:
One request: the workflows still need approval to run. CI has only ever run on Happy to split anything out if the diff is too large to review in one piece. |
# Conflicts: # docs/source/changelog.rst
# Conflicts: # docs/source/changelog.rst
|
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). If you would like to continue working on this feature, we suggest:
We would be happy to discuss this plan with you before you begin. Thanks for your understanding and for your interest in contributing ! |
|
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. |
Bilevel hyperparameter learning with MAID
Adds the Method of Adaptive Inexact Descent to
deepinv.optim.bilevel, with aprior-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
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 abackward pass in
xthen intheta. No derivative is written by hand, and themodel that is solved and the model that is differentiated are the same object.
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.
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,delta0andalpha0default to
None, meaning derive attheta0: the accuracy from thestationarity residual at
A* y, the step from the hypergradient norm. Noabsolute constant serves every prior, because both references are properties of
the regulariser. Measured at the same initialisation,
eps0 = 1e-1leaves adenoising 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.7for 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.
epsis a per-element rmsgradient tolerance, so it means the same accuracy at 32x32 and 256x256, and the
floor comes from
torch.finfo(dtype).eps. Falling short ofepsis not anerror: 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 thebatched 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-levelsolvers, 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.pyandpriors.pyat 17-18% patch coverage: they had no tests, and the coverage shown was
import-time execution alone.
test_maid_priors.pyandtest_maid_batched.pyclose 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 adisagreement 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:
non-negative, establishing convexity in
xdirection
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.pygives 14failed, 221 passed, 32 skipped. Every failure is on the MPS device, and the
same suite on unmodified
mainat 522bba9 gives 14 failed, 221 passed, 32skipped with a test-for-test identical failure set, so none of them is
introduced here. Eight are
test_least_squares_implicit_backward, the resttest_tvprior_gradient,test_linear_system,test_MLEM,test_CP_datafidsplit,test_CP_Kandtest_least_squares_implicit_backward_nonleaf_buffer_grad. They look likegaps in MPS operator coverage rather than anything specific to this branch.
test_datasets.pywas not run locally: it downloads full public datasets(Flickr2K alone exceeded 3 GB before being stopped). CI covers both.
Breaking change
MAIDConfig.eps0,delta0andalpha0change fromfloattofloat | Noneand now default to
None. Code passing explicit values is unaffected. Theprevious default of
eps0 = 1e-1was harmful under the per-element toleranceintroduced here, so anyone relying on it was already getting poor behaviour.
Documentation
docs/source/user_guide/reconstruction/bilevel.rst, covering the oracleinterface, certified and non-certified paths, the accelerated heuristics with
an ablation, numerical precision and GPU backends, batching, and learning a
new prior
examples/optimization/Checks to be done before submitting your PR
python3 -m pytest deepinv/testsruns successfully.black .andruff check .run successfully.make htmlruns successfully (in thedocs/directory).LLM policy
LLM usage is ok, but not PRs generated 100% by AI. See our LLM policy Tick below as appropriate: