Skip to content

feat: LaplaceOptimiser(projection="moments") for hierarchical scatter #1654

Description

@Jammy2211

Overview

Design departure from the prompt (needs the human's eye): the prompt suggests Gauss–Hermite centred at the Laplace mode with the Hessian scale. That cannot work on the target case, because there is no interior mode and no Hessian there, and a plain tensor grid over (μ, xᵢ, σ) fails because for fixed σ the joint is a spike of width ~σ. The plan ports the referee's structure as nested quadrature (INLA-style): outer Gauss–Legendre over the scale variable(s) on their support, inner conditional Laplace over the remaining variables at each outer node (exact when the cavities are Gaussian, which is the hierarchical-Gaussian case).

Under the Laplace projection the hierarchical factor's tilted density in σ sits on the σ=0 boundary, so the Hessian is not negative-definite (BAD_PROJECTION) or the line search fails (FAILURE): 0/50 SUCCESS on slope_hierarchy_scale 343299, and the leg-B sigma misses on analytic_gaussian. The referee (autofit_workspace_test/scripts/graphical/analytic_ep_minimal.py, projection="moments", lines 227-346) recovers the exact posterior by integrating (μ, xᵢ) out analytically given θ and doing 1-D quadrature over θ.

The default projection stays "mode" (zero ripple); flipping the default is a separate human call.

Plan

  • Add LaplaceOptimiser(projection="mode"|"moments") (default "mode"), with n_quadrature, quadrature_half_width, moment_max_size, moment_max_outer.
  • New pure-numpy autofit/graphical/laplace/moments.py: nested quadrature — outer Gauss–Legendre over the scale variable(s) on their support, inner conditional Laplace over the rest (ported from the analytic_ep_minimal referee).
  • MeanField.from_weighted_nodes turns the weighted nodes into messages; projection.log_norm = Ẑ.
  • SUCCESS / BAD_PROJECTION / FAILURE semantics; bit-for-bit fallback to the mode path outside the hierarchical-scale case.
  • New test file test_autofit/graphical/functionality/test_moment_projection.py; README §3.2/§3.3/§3.5 updates.
  • Phase 2 (later, /start_workspace): switch the autofit_workspace_test analytic_gaussian* scripts to projection="moments" and un-park them.
Detailed implementation plan

Affected Repositories

  • PyAutoFit (primary, phase 1) — the only repo claimed now.
  • autofit_workspace_test (phase 2) — a later /start_workspace after the library PR merges; not claimed yet.

Branch Survey

Repository Current Branch Dirty?
./fit/PyAutoFit main clean (in sync with origin/main)

Suggested branch: feature/ep-moment-projection
Worktree root: ~/Code/PyAutoLabs-wt/ep-moment-projection/

Parallel claim

PyAutoFit is claimed in parallel by ep-projection-exception and ep-moment-projection (both registered 2026-09-30). File sets are disjoint — A: messages/abstract.py, mapper/prior/abstract.py, non_linear/result.py, graphical/expectation_propagation/optimiser.py:150-162, exc.py, test_autofit/messages/test_project_nonfinite.py, test_autofit/graphical/functionality/test_factor_failure_recovery.py; B: graphical/laplace/*, graphical/mean_field.py, graphical/declarative/factor/hierarchical.py, graphical/expectation_propagation/diagnostics.py:47, graphical/README.md, test_autofit/graphical/functionality/test_moment_projection.py. A ships first; B rebases. Parallel claim human-approved 2026-09-30 with the plan.

Implementation Steps

  1. LaplaceOptimiser(projection="mode"|"moments"), default "mode" (zero ripple; flipping the default is a separate human call). New kwargs n_quadrature=64, quadrature_half_width=8.0, moment_max_size=4, moment_max_outer=2. No plumbing changes: factor_graph.optimise(af.LaplaceOptimiser(projection="moments")) and per-factor af.HierarchicalFactor(..., optimiser=...) already carry it (declarative/abstract.py:172-214, hierarchical.py:186).
  2. New pure-numpy module autofit/graphical/laplace/moments.py; LaplaceOptimiser._moment_projection called at the top of optimise_approx (laplace/optimiser.py:272).
  3. Variable split: outer = free scalar variables that are factor.scale_variables (new _HierarchicalFactor property mapping _SCALE_ARGUMENT_NAMES — move from ep/diagnostics.py:47 to a shared spot — to prior Variables) or whose message has a finite lower/upper limit (catches a TruncatedNormal σ PriorFactor); inner = the rest. Fall back to the mode path bit-for-bit (with one logger.info) when there is no outer variable, > moment_max_outer outer or > moment_max_size flattened params, a non-scalar outer, or non-empty deterministic_variables. Leg A and dataset factors therefore take the existing code unchanged.
  4. Integrate in the message's base coordinate u (identity for Normal/TruncatedNormal; log σ for LogGaussian's TransformedMessage). Bounds = intersection of the parameter support (normal.py:79), _support_kwargs limits (truncated_normal.py:120, mean_field.py:259-268) and cavity mean ± half_width·cavity std in base space; lift a 0 edge to 1e-12·max(1,|c|) as theta_grid:262-264. Map nodes to physical space with _inverse_transform, add log|dx/du| to the log weight (the factor approximation is a physical density, mean_field.py:724-728). Two passes: cavity window, then re-window to tilted mean ± 8 sd if the tilted sd < 0.25× window scale or edge mass > 1e-6 at a non-support edge.
  5. Inner conditional Laplace per outer node: conditioned newton.OptimisationState (factor and gradient with the outer entries fixed, limits as prepare_state:144-147), warm-started from the previous node, newton.optimise_quasi_newton + finite_difference_hessian; no RNG (refine_state never called). Log weight = log w_GL + log|J| + ℓ(m_j,s_j) + (d/2)log 2π − ½ log det(−H_j).
  6. Moments → messages: new MeanField.from_weighted_nodes(nodes, log_weights, log_norm) next to from_mode_covariance (mean_field.py:405-430), calling each variable's project(nodes, log_w, id_=..., **_support_kwargs) (TransformedMessage maps to base space itself). Outer nodes = s_j; inner nodes = 3^d order-3 Gauss–Hermite expansion of each conditional N(m_j, Σ_j) in whitened coordinates (exact for E[x], E[x²]; reproduces the referee's law-of-total-variance, :337-345). Shift log weights so each message's log_norm is 0; set projection.log_norm = logsumexp(log_w) (Ẑₐ, README §5, sampler convention abstract_search.py:463-479). Truncated σ: keep the (E, Var)-as-parent convention of TruncatedNormalMessage.invert_sufficient_statistics (truncated_normal.py:263-281), same as the referee's nat(E, Var); state in README that exact truncated-moment inversion is out of scope.
  7. Status: SUCCESS when Ẑ finite > 0 and every E, Var finite with Var > 0 (message line "moments: n_nodes=…, passes=…"); BAD_PROJECTION (mean field returned unchanged, as optimiser.py:293-307) for empty window, zero mass, non-finite moment, residual edge mass, non-concave inner Hessian at a node with mass > 1e-10; FAILURE for an inner search failure at a weighted node. Downstream per-variable revert of wider-than-cavity projections in update_factor_mean_field (mean_field.py:457-604) unchanged. The joint mode search does not run on the moments path (removes the 23 FAILURE rows).
  8. README: §3.2 Laplace bullet (101-117), §3.3 Eq. 9 with quadrature weights, §3.5 "wait for a moment-matching projection" → "use projection="moments"".

Tests (test_autofit/graphical/functionality/test_moment_projection.py, style of functionality/test_laplace_hessian.py)

(test_autofit/graphical/functionality/test_moment_projection.py, style of functionality/test_laplace_hessian.py)

  1. Two-variable known moments: factor N(x|0,s), cavity xN(m,A), sTruncatedNormal(c,C,0,100); reference by scipy.integrate.quad of N(s|c,C)·N(m|0,s²+A) with E[x|s]=m s²/(s²+A), Var[x|s]=A s²/(s²+A); E and Var of both variables to 1e-6.
  2. Gaussian check with make_approx (mode/cov to 1e-6).
  3. Determinism across seeds (0, 1, 12345) and throwaway Variable ids (copy the two tests atop test_laplace_hessian.py).
  4. Fallback: no outer variable → bit-equal to projection="mode"; over moment_max_size falls back and logs; bad projection → ValueError.
  5. Status: cavity outside the support → BAD_PROJECTION, same mean_field object returned.
  6. End-to-end in hierarchical/ on the test_truncated_support.py toy (3 groups, TG σ) through the manual factor_approximation → optimise → project_mean_field loop; exact σ posterior by 1-D quadrature in the test; assert ≥1 SUCCESS, scatter within 0.5 std of exact, limits kept, and that the mode path records no SUCCESS on the same toy (documents the gap; bears on draft/bug/autofit/ep_factors_end_the_run_with_zero.md).

Phase 2 (workspace PR, later, via /start_workspace)

scripts/graphical/analytic_autofit.py:147-158 run_autofit_ep(..., projection="moments"); rename "(Laplace)" labels to "(moments)" in analytic_gaussian.py:11,32, analytic_gaussian_priors.py:7,42, collapse TOY_OPT :12-16; add the "scatter within 0.5 std" check to analytic_gaussian_collapse.py seeds 0–4; rewrite __Status__ and "first runs" numbers from real runs; delete config/build/no_run.yaml:27-28; add graphical/analytic_gaussian.py to smoke_tests.txt after line 14 (real DynestyStatic search — time it under the smoke profile; route through /ci_speedup if too slow, never drop it). Tolerances stay a 0.15 / b 0.25.

Key Files

  • autofit/graphical/laplace/optimiser.py — LaplaceOptimiser(projection=...), _moment_projection
  • autofit/graphical/laplace/moments.py (new) — nested quadrature
  • autofit/graphical/mean_field.py — MeanField.from_weighted_nodes
  • autofit/graphical/declarative/factor/hierarchical.py — scale_variables
  • autofit/graphical/expectation_propagation/diagnostics.py:47 — _SCALE_ARGUMENT_NAMES moved to a shared spot
  • autofit/graphical/README.md — §3.2, §3.3, §3.5
  • test_autofit/graphical/functionality/test_moment_projection.py (new)

Risks

  • JAX: the Laplace path is eager numpy (newton.py, line_search.py); no tracing concern, but eager JAX dispatch per factor call.
  • Cost: ~64 nodes × 2 passes × 30–60 factor calls ≈ 4–8k calls per hierarchical update (1–2 s numpy, 10–20 s eager JAX); minutes for N=25 × 2 steps vs a 5h20m wall. n_quadrature is the lever.
  • Ripple: no library use of LaplaceOptimiser in autolens/autogalaxy; workspace guide uses the default; default "mode" → zero ripple.
  • Accuracy: Gaussian matching on skewed σ has ~a 0.08 / b 0.15 bias (README §3.5); witness 0.15 / 0.25 has margin, only just on b. If the priors script's truncated row misses, suspect the parent-parameter convention first.

Verification

New test file green; full pytest test_autofit -x; then run analytic_gaussian*.py locally as the witness before /ship_library. Library PR first (/ship_library, RED override recorded), workspace PR second.

Heart RED override (development only)

Authorised by the live human in the Claude Code session 2026-09-30 ~10:55 BST, answer "Override for all three (Recommended)" to the question offering the RED override for PyAutoCortex#50 and "for opening the two PyAutoFit tasks (issue + worktree + plan; no merge, no release)".

Exact Heart RED reasons at the 10:51 BST tick:

  • "PyAutoArray: 2 commit(s) behind origin"
  • "PyAutoLens: 2 commit(s) behind origin"
  • "release validation FAILED (stage integrate)"

Scope: issue + worktree + plan; PR-open permitted; no merge, no release. Plans approved in-session 2026-09-30 ~11:20 BST via Plan Mode.

Original Prompt

Click to expand starting prompt

EP: moment-matching projection for the hierarchical scatter (the cure the Laplace path cannot give)

Type: feature
Target: PyAutoFit
Repos:

  • PyAutoFit
    Themes:
  • graphical-ep
    Difficulty: medium
    Autonomy: supervised
    Priority: normal
    Status: formalised — filed as the "cure" follow-on of phase 2; gated by the campaign's "JAX/gradient/Hessian EP internals" check-in (draft/research/graphical_ep/ep_campaign.md, Deferred) — adopt only if the human judges the scatter worth it
    Consequence: glance
    Witness: analytic_gaussian.py leg B sigma row and both x_i rows PASS at the autofit-EP tolerance (a 0.15 / b 0.25) with leg A unchanged (18/18); analytic_gaussian_priors.py truncated and gaussian scatter rows PASS; analytic_gaussian_collapse.py seeds 0-4 put the scatter within 0.5 std of the closed form; a two-variable factor with known tilted moments matches to 1e-6; both scripts are un-parked from no_run.yaml.
    Review-minutes: 3
    Unattended: ready
    Epic: graphical-ep
    Filed: 2026-09-02

Why

Phase 2 of the EP campaign (record
complete/2026/09/ep-scale-collapse-basin-cure-or-caveat.md) fixed the mechanism of the
parent-scale collapse (PyAutoFit#1558 gradients, #1560 message support,
#1561 Hessian at the mode + skip-not-write) and shipped the caveat
(autofit/graphical/README.md §3.5): the collapse configuration now
RECOVERS on 5/5 referee seeds. But the scatter itself is still not
estimated by EP — a hierarchical factor's tilted density in σ is
∝ 1/σ at xᵢ = μ, its mode sits on the boundary, and the fixed Laplace
path correctly refuses to write a Gaussian there. The scatter therefore
stays near its prior with an honest width (referee seed 0: 9.37 ± 3.57
vs exact 6.57 ± 2.88; |Δmean|/std 0.97), and the two per-dataset
means most affected by it miss at a ≈ 0.19.

The closed-form benchmark's minimal EP shows the cure: moment
matching
of the same Gaussian site (E and Var of the tilted
distribution by quadrature) recovers the exact posterior on every seed
with a ≤ 0.08 / b ≤ 0.15 on the scatter row
(autofit_workspace_test/scripts/graphical/analytic_ep_minimal.py,
projection="moments").

What

A moment-matching projection option for _HierarchicalFactor updates
(and, if cheap, for any factor with ≤ ~4 free variables):

Acceptance

  • analytic_gaussian.py leg B sigma row PASS at the autofit-EP
    tolerance (a 0.15 / b 0.25) and both x_i rows that miss today PASS;
    analytic_gaussian_priors.py truncated and gaussian families PASS on
    their scatter rows; leg A unchanged (18/18).
  • analytic_gaussian_collapse.py seeds 0–4: scatter within 0.5 std of
    the closed form (was 0.7–1.0 std under the mode projection).
  • Unit test: a two-variable factor with known tilted moments matches to
    1e-6; determinism across seeds and variable ids as in
    test_laplace_hessian.py.
  • Un-park analytic_gaussian.py / analytic_gaussian_priors.py in
    autofit_workspace_test/config/build/no_run.yaml and curate
    analytic_gaussian.py into smoke_tests.txt (their __Status__
    paragraphs name this prompt).

Links

2026-09-24 — slope_hierarchy_scale 343299: hierarchical factor 0/50 SUCCESS at N=25

RAL job 343299 (PyAutoCortex projects/slope_hierarchy_scale.md) completed
2026-09-16 02:03 BST: wall 5h20m, MAX_STEPS=2, 25 lenses, JAX-on-CPU, no
pool. It logged 50 "jit compiling vectorized" lines — one compile per factor
search (2 steps x 25), confirming draft/refactor/autofit/ep_analysis_level_compile_cache.md.
Every dataset factor reached SUCCESS (2/2 each). The hierarchical factor
never did:
HierarchicalFactor0 has 50 rows = 27 BAD_PROJECTION + 23
FAILURE, zero SUCCESS (PriorFactor239 and PriorFactor263 also stale),
and the library's own STALE FACTORS warning fired. The parent recovery printed
exactly the prior: mean 2.0000 ± 0.7071, sigma 0.5000 ± 0.3536 (priors
TG(2,1) / TG(0.5,0.5); truth sigma 0.1).

Under the Laplace optimiser's semantics BAD_PROJECTION = the Hessian at the
mode is not finite or not negative-definite (the scale parameter is driven to
a limit) and FAILURE = the line search failed and the mean field was handed
back unchanged. This is the Laplace-on-scatter caveat
(PyAutoFit/autofit/graphical/README.md §3.5, analytic_gaussian leg B) at
100 % on the lensing model: per-lens slope widths of 0.002-0.05 are far
tighter than the scatter prior, so the tilted density in sigma sits at the
boundary. The EP arm of campaign phase 3 therefore cannot estimate the slope
scatter with the Laplace projection; this prompt is now the blocking decision
for it (human call, Autonomy: supervised). EP output (local):
output/sample_n25_seed42/ep/expectation_propagation/a7ec60f07f4261afcb98b2adef484199/
(ep_history.csv, ep_diagnostics.results).

Cross-reference — draft/bug/autofit/ep_factors_end_the_run_with_zero.md:
that bug records a HierarchicalFactor1 ending with only BAD_PROJECTION and
never a SUCCESS on 6/200 leg-B seeds of the analytic toy (seeds 29, 71, 106,
152, 178, 193). The same shape now reproduces at 100 % on a real lensing model
(every hierarchical update over 2 steps x 25 factors), so it is no longer a
rare-seed edge case: whether a zero-SUCCESS hierarchical factor is legitimate
(the bug's open question) and whether moment matching removes it should be
judged together.

🤖 Generated with Claude Code

Activity

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

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions