Skip to content

Hierarchical EP: parent scale hyperparameter collapses to ~0 with over-confident error (+ init-crash) #1405

Description

@Jammy2211

Summary

A hierarchical EP fit of a parent scale hyperparameter (the sigma of a
HierarchicalFactor parent Gaussian) is unstable: repeated fits of an identical,
known-answer
problem land in three qualitatively different outcomes. On a clean CPU toy
(no JAX, no lensing, parent scatter truth σ=10 — far from the σ→0 boundary):

outcome freq parent scatter (truth 10)
RECOVER ~70% 9–13, sane errors; matches a joint nested-sampler fit
COLLAPSE ~7% 0.003–0.80 with an over-confident ~0 error
CRASH ~23% InitializerException hard-aborts mid-EP

Measured over 30 identical-problem runs across max_steps ∈ {20,25,30,50,60}:
21 RECOVER / 2 COLLAPSE / 7 CRASH. Both COLLAPSEs were at the shortest setting
(max_steps=20); none at 25–60. CRASH appears at every max_steps.

The joint fit — Dynesty on factor_graph.global_prior_model for the same graph — is
stable and correct every run (scatter 12.8 [9.9, 16.3]). The graph and data are fine; EP
is the sole source of the instability.
This is on main including the #1383 projection
fix — it is a distinct, deeper issue.

The COLLAPSE outcome is the exact readout that a downstream hierarchical science fit hit
(parent scatter 4–12× too low, errors ~1000× too tight); this toy shows it is a framework
property, reproduced off-boundary and off-lensing.

Two distinct defects

1. Scale-hyperparameter COLLAPSE (over-shrinkage basin). When it collapses, the
per-group drawn-variable posterior means cluster far tighter than truth (spread ~0.6 vs
true ~±10); the parent factor then sees near-identical draws and infers a tiny scatter,
tightening the shrinkage further — positive feedback to a fixed point at scatter ≈ 0 with
~0 reported error. Not a delta-method/boundary artefact: truth σ=10, the base-space
scatter message is unbounded (TruncatedNormal(mean≈0, lower=-inf, upper=inf)), so
TransformedMessage.variance faithfully reads a genuinely over-confident posterior.
ep_history.csv shows recurring BAD_PROJECTION on the HierarchicalFactor and wildly
swinging log-evidence during collapse. The F10 sigma-collapse guard does not flag it.

2. InitializerException CRASH. ~1-in-4 runs hard-abort when EP drives a factor to a
degenerate all-equal-likelihood state its per-factor DynestyStatic cannot initialise
(autofit/non_linear/initializer.py:185). Should degrade to a flagged BAD_PROJECTION /
skipped update, not kill the whole fit.

Triage questions

  • Is COLLAPSE curable by a more robust / damped HierarchicalFactor scale
    moment-match
    (or a deterministic per-factor optimiser to cut the sampler noise that
    drives basin selection)? Note: naive damping (delta=0.5) made a real downstream fit
    worse — the standing "consider damping, delta < 1" diagnostics hint likely
    mis-diagnoses this mode.
  • Should the F10 sigma-collapse guard fire on a scatter≈0 / error≈0 parent? It
    currently passes silently.
  • Should the InitializerException be caught inside the EP loop and recorded as a bad
    projection rather than aborting the fit?

Minimal repro

Self-contained, numpy-only, runs from the HowToFit repo root in minutes on CPU. Builds
the HowToFit chapter-3 hierarchical toy (5 low-SNR 1-D Gaussians, centre ~ N(50, 10)),
runs EP and a joint Dynesty fit on the same graph, prints an OUTCOME: tag. Loop it ~10×
at TOY_MAX_STEPS=20 to see the recover/collapse/crash split.

ep_toy_diagnostic.py
"""
Toy hierarchical EP diagnostic
==============================

Cheap CPU reproduction of the slope_hierarchy goal-2 pathology:
EP recovers the parent MEAN but reports a parent SCATTER that is too low with
error bars that are orders-of-magnitude too tight.

Known-answer toy (HowToFit chapter 3 hierarchical dataset):
  - N Gaussians, each centre drawn from parent N(mean=50, sigma=10)
  - TRUE parent mean   = 50.0
  - TRUE parent scatter= 10.0   <-- FAR from the sigma->0 boundary (the key contrast
                                     with slope_hierarchy whose truth was 0.1)
  - per-Gaussian normalization = 0.5, sigma = 5.0 (fixed at truth to keep the joint
    nested-sampler fit cheap and to isolate the parent-scatter question)

Run from the HowToFit repo root.
"""
import os
import sys
import numpy as np
from os import path

import autofit as af
from autofit.graphical import mean_field_summary, check_sigma_collapse

# ---- config ---------------------------------------------------------------
TOTAL_DATASETS = int(os.environ.get("TOY_N", "5"))
MAX_STEPS = int(os.environ.get("TOY_MAX_STEPS", "20"))
RUN_JOINT = os.environ.get("TOY_JOINT", "1") == "1"
OUT_TAG = os.environ.get("TOY_TAG", f"ep_toy_n{TOTAL_DATASETS}_s{MAX_STEPS}")

TRUE_MEAN = 50.0
TRUE_SCATTER = 10.0

print(f"=== TOY EP DIAGNOSTIC  N={TOTAL_DATASETS}  max_steps={MAX_STEPS} ===")
print(f"TRUE parent: mean={TRUE_MEAN}  scatter={TRUE_SCATTER}")

# ---- data -----------------------------------------------------------------
data_list, noise_map_list = [], []
for i in range(TOTAL_DATASETS):
    dp = path.join("dataset", "example_1d", "gaussian_x1__hierarchical", f"dataset_{i}")
    data_list.append(af.util.numpy_array_from_json(file_path=path.join(dp, "data.json")))
    noise_map_list.append(af.util.numpy_array_from_json(file_path=path.join(dp, "noise_map.json")))

analysis_list = [af.ex.Analysis(data=d, noise_map=n) for d, n in zip(data_list, noise_map_list)]

# ---- per-dataset models: only `centre` free (nuisances fixed at truth) -----
model_list = []
for _ in range(TOTAL_DATASETS):
    g = af.Model(af.ex.Gaussian)
    g.centre = af.TruncatedGaussianPrior(mean=50.0, sigma=20.0, lower_limit=0.0, upper_limit=100.0)
    g.normalization = 0.5
    g.sigma = 5.0
    model_list.append(g)

# ---- EP fit ---------------------------------------------------------------
dynesty = af.DynestyStatic(nlive=100, sample="rwalk")
analysis_factor_list = [
    af.AnalysisFactor(prior_model=m, analysis=a, optimiser=dynesty, name=f"dataset_{i}")
    for i, (m, a) in enumerate(zip(model_list, analysis_list))
]

hierarchical_factor = af.HierarchicalFactor(
    af.GaussianPrior,
    mean=af.TruncatedGaussianPrior(mean=50.0, sigma=10.0, lower_limit=0.0, upper_limit=100.0),
    sigma=af.TruncatedGaussianPrior(mean=10.0, sigma=5.0, lower_limit=0.0, upper_limit=100.0),
)
for m in model_list:
    hierarchical_factor.add_drawn_variable(m.centre)

factor_graph = af.FactorGraphModel(*analysis_factor_list, hierarchical_factor)

laplace = af.LaplaceOptimiser()
try:
    ep_result = factor_graph.optimise(
        laplace,
        paths=af.DirectoryPaths(name=path.join("ep_toy_diagnostic", OUT_TAG)),
        ep_history=af.EPHistory(kl_tol=0.05),
        max_steps=MAX_STEPS,
    )
except Exception as e:
    print(f"\nOUTCOME: CRASH  ({type(e).__name__}: {str(e).splitlines()[0][:80]})")
    print("DONE-SENTINEL")
    sys.exit(0)

mean_field = ep_result.updated_ep_mean_field.mean_field
ep_mean = mean_field.mean[hierarchical_factor.mean]
ep_mean_err = float(np.sqrt(mean_field.variance[hierarchical_factor.mean]))
ep_sigma = mean_field.mean[hierarchical_factor.sigma]
ep_sigma_err = float(np.sqrt(mean_field.variance[hierarchical_factor.sigma]))

print("\n================ EP RESULT ================")
print(f"parent mean    : {ep_mean:.4f} +/- {ep_mean_err:.4g}   (truth {TRUE_MEAN})")
print(f"parent scatter : {ep_sigma:.4f} +/- {ep_sigma_err:.4g}   (truth {TRUE_SCATTER})")
# outcome tag: COLLAPSE if scatter << truth, else RECOVER
_tag = "COLLAPSE" if ep_sigma < 0.4 * TRUE_SCATTER else "RECOVER"
print(f"OUTCOME: {_tag}  scatter={ep_sigma:.4f}  err={ep_sigma_err:.4f}  mean={ep_mean:.3f}")
print("\n--- mean_field_summary ---")
try:
    print(mean_field_summary(mean_field))
except Exception as e:
    print("mean_field_summary failed:", e)
print("\n--- check_sigma_collapse ---")
try:
    print(check_sigma_collapse(mean_field))
except Exception as e:
    print("check_sigma_collapse failed:", e)

# base-space (untransformed) sigma message: is it near the sigma->0 boundary?
try:
    base_msg = mean_field[hierarchical_factor.sigma]
    print("\n--- parent-sigma message (base space) ---")
    print("  repr:", repr(base_msg))
    for attr in ("mean", "sigma", "scale", "variance", "natural_parameters"):
        v = getattr(base_msg, attr, None)
        if callable(v):
            try:
                v = v()
            except Exception:
                continue
        if v is not None:
            print(f"  {attr}: {v}")
except Exception as e:
    print("base sigma message introspection failed:", e)

# ---- JOINT nested-sampler fit of the SAME graph ---------------------------
def weighted_percentiles(values, weights, qs=(0.16, 0.5, 0.84)):
    order = np.argsort(values)
    v = np.asarray(values)[order]
    w = np.asarray(weights)[order]
    cw = np.cumsum(w)
    cw = cw / cw[-1]
    return np.interp(qs, cw, v)

if RUN_JOINT:
    print("\n================ JOINT NESTED-SAMPLER FIT (SAME graph) ================")
    joint_model = factor_graph.global_prior_model
    print("joint dims (prior_count):", joint_model.prior_count)
    search = af.DynestyStatic(
        name=path.join("ep_toy_diagnostic", OUT_TAG + "_joint"),
        nlive=150,
        sample="rwalk",
    )
    joint_result = search.fit(model=joint_model, analysis=factor_graph)
    samples = joint_result.samples

    ordered = samples.model.priors_ordered_by_id
    params = np.array(samples.parameter_lists)
    weights = np.array(samples.weight_list)

    def marg(prior):
        col = [i for i, p in enumerate(ordered) if p.id == prior.id][0]
        lo, med, hi = weighted_percentiles(params[:, col], weights)
        return med, lo, hi

    jm_med, jm_lo, jm_hi = marg(hierarchical_factor.mean)
    js_med, js_lo, js_hi = marg(hierarchical_factor.sigma)
    print("\n---------------- JOINT RESULT (percentile 16/50/84) ----------------")
    print(f"parent mean    : {jm_med:.4f}  [{jm_lo:.4f}, {jm_hi:.4f}]   (truth {TRUE_MEAN})")
    print(f"parent scatter : {js_med:.4f}  [{js_lo:.4f}, {js_hi:.4f}]   (truth {TRUE_SCATTER})")

    # ---- side-by-side verdict table ----
    print("\n================ EP vs JOINT ================")
    print(f"{'param':10s} {'truth':>7s} {'EP':>22s} {'JOINT':>26s}")
    print(f"{'mean':10s} {TRUE_MEAN:7.2f} {ep_mean:8.4f}+/-{ep_mean_err:<10.3g} "
          f"{jm_med:8.4f} [{jm_lo:.3f},{jm_hi:.3f}]")
    print(f"{'scatter':10s} {TRUE_SCATTER:7.2f} {ep_sigma:8.4f}+/-{ep_sigma_err:<10.3g} "
          f"{js_med:8.4f} [{js_lo:.3f},{js_hi:.3f}]")
    joint_sigma_err = 0.5 * (js_hi - js_lo)
    ratio = joint_sigma_err / ep_sigma_err if ep_sigma_err > 0 else float("inf")
    print(f"\nscatter error ratio JOINT/EP = {ratio:.1f}x  "
          f"(EP err {ep_sigma_err:.4g}  vs  JOINT half-CI {joint_sigma_err:.4g})")

print("\nDONE-SENTINEL")
cd HowToFit
export NUMBA_CACHE_DIR=/tmp/numba_cache MPLCONFIGDIR=/tmp/matplotlib PYAUTO_SKIP_VISUALIZATION=1
for r in $(seq 1 10); do TOY_MAX_STEPS=20 TOY_JOINT=0 TOY_TAG=rep_$r \
  python3 ep_toy_diagnostic.py 2>/dev/null | grep OUTCOME; done

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