diff --git a/autofit/graphical/expectation_propagation/diagnostics.py b/autofit/graphical/expectation_propagation/diagnostics.py index a71e48e6e..5279db116 100644 --- a/autofit/graphical/expectation_propagation/diagnostics.py +++ b/autofit/graphical/expectation_propagation/diagnostics.py @@ -10,10 +10,16 @@ writes them as machine-readable CSVs and a matplotlib evolution plot. - ``mean_field_summary`` — a human-readable table of a mean field, suitable for printing at the end of any example or script. -- ``check_sigma_collapse`` — guards against the known pathology where - repeated undamped EP updates over-count shared information and every - sigma collapses towards zero around the starting point (rather than - the data); see PyAutoFit issue #1332 (F10). +- ``check_sigma_collapse`` — guards against two collapse pathologies. + First, the one from PyAutoFit issue #1332 (F10): repeated undamped EP + updates over-count shared information and every sigma collapses + towards zero around the starting point (rather than the data). + Second, the *hierarchical parent-scale* collapse of PyAutoFit issue + #1405: the scale hyperparameter of a ``HierarchicalFactor``'s parent + distribution settles near zero with an error bar that is + over-confident only *relative to that mean* — reporting "no scatter" + as a confident answer. The two need different tests; see + ``check_sigma_collapse``. Outputs written to the EP output folder by ``EPOptimiser`` when paths are enabled: @@ -27,12 +33,18 @@ import csv import logging from pathlib import Path -from typing import Dict, List, Optional, Tuple +from typing import Dict, List, Optional, Set, Tuple import numpy as np logger = logging.getLogger(__name__) +#: Argument names by which a parent distribution's *scale* parameter is +#: recognised on a ``HierarchicalFactor``. ``GaussianPrior`` and +#: ``LogGaussianPrior`` call it ``sigma``; the others are accepted so a +#: distribution that names it differently is still covered. +_SCALE_ARGUMENT_NAMES = frozenset({"sigma", "scale", "std", "stddev"}) + def _scalar_mean_std(message) -> Tuple[float, float]: """ @@ -65,9 +77,37 @@ def __init__(self): """ self.factor_rows: List[dict] = [] self.variable_rows: List[dict] = [] + self.scale_variables: Set[str] = set() self._previous_mean_field = None self._step = 0 + def register_hierarchical_scales(self, factor_graph) -> None: + """ + Record which variables are hierarchical parent *scale* + hyperparameters, so ``check_sigma_collapse`` can apply the + scale-specific test to them (PyAutoFit #1405). + + A ``HierarchicalFactor`` parameterises a parent distribution + with priors named after that distribution's arguments — e.g. + ``af.HierarchicalFactor(af.GaussianPrior, mean=..., sigma=...)``. + The scale argument is the one that collapses, so it is picked + out by name (``_SCALE_ARGUMENT_NAMES``). + + Silently records nothing for a graph with no hierarchical + factors, or one that does not expose them (a plain + ``FactorGraph`` rather than a ``DeclarativeFactorGraph``) — + diagnostics must never kill the fit. + """ + try: + distribution_models = factor_graph.hierarchical_factors + except AttributeError: + return + + for distribution_model in distribution_models: + for name, prior in distribution_model.prior_tuples: + if name in _SCALE_ARGUMENT_NAMES: + self.scale_variables.add(prior.name) + def snapshot(self, factor, model_approx, status) -> None: """ Record the state of the approximation after one factor update. @@ -210,23 +250,94 @@ def mean_field_summary(mean_field) -> str: return "\n".join(lines) +def _drop_consecutive_repeats(values: np.ndarray) -> np.ndarray: + """ + Collapse runs of identical consecutive values to a single entry. + + ``EPDiagnostics.snapshot`` records *every* variable on *every* + factor update, but a factor update only moves the marginals of the + variables adjacent to that factor. A variable's history is + therefore dominated by steps at which it did not move at all, and a + strict ``diff < 0`` monotonicity test can essentially never be + satisfied in a multi-factor graph. Dropping the repeats restores + the test to what it was written to mean: consecutive *updates of + this variable* that shrank it. + + The first and last values are always preserved, so magnitude + comparisons against them are unaffected. + """ + if len(values) < 2: + return values + return values[np.insert(np.diff(values) != 0, 0, True)] + + +def _flag_scale_collapse( + means: np.ndarray, + stds: np.ndarray, + mean_fraction: float, + relative_error: float, +) -> bool: + """ + Whether a hierarchical parent scale has collapsed (PyAutoFit #1405). + + True when the scale's mean has fallen below ``mean_fraction`` of + its initial value *and* its error is small relative to that mean — + i.e. the fit is confidently reporting near-zero parent scatter. + + A non-positive initial mean gives no baseline to collapse from, so + no judgement is made. A latest mean at or below zero is outside the + scale's support and is always flagged. + """ + initial, latest = means[0], means[-1] + + if not initial > 0: + return False + if latest <= 0: + return True + if latest >= mean_fraction * initial: + return False + + return stds[-1] / latest < relative_error + + def check_sigma_collapse( diagnostics: EPDiagnostics, std_floor: float = 1e-8, monotone_steps: int = 5, shrink_factor: float = 1e-3, + scale_mean_fraction: float = 0.2, + scale_relative_error: float = 0.5, ) -> List[str]: """ - Detect the EP sigma-collapse pathology (PyAutoFit #1332, F10). + Detect the two EP collapse pathologies. - Repeated undamped EP updates can over-count shared-variable - information: every std shrinks monotonically towards zero around - the *starting* means, while the KL convergence criterion never - triggers. This check flags a variable when either: + **Sigma collapse (PyAutoFit #1332, F10).** Repeated undamped EP + updates can over-count shared-variable information: every std + shrinks monotonically towards zero around the *starting* means, + while the KL convergence criterion never triggers. This check flags + a variable when either: - its latest std is below ``std_floor``, or - - its std has shrunk monotonically for the last ``monotone_steps`` - updates *and* by more than a factor ``1 / shrink_factor`` overall. + - its std has shrunk monotonically over its last ``monotone_steps`` + *updates* (steps at which it did not move are not counted — see + ``_drop_consecutive_repeats``) *and* by more than a factor + ``1 / shrink_factor`` overall. + + **Hierarchical parent-scale collapse (PyAutoFit #1405).** Both + tests above are *absolute* and variable-agnostic, because #1332 is + a pathology in which every std goes to zero. The parent scale of a + ``HierarchicalFactor`` collapses in a different shape: its *mean* + goes to ~0 while its std stays moderate in absolute terms and is + over-confident only *relative to that mean* — a confident claim of + "no scatter". Measured on the #1405 toy, the collapsed runs sit at + mean 0.80 (std 0.11) and mean 0.0030 against a parent scale + hyper-prior of mean 10, where healthy runs recover 9.1-12.8; an + absolute std test cannot separate those, and does not fire on + either. So for variables registered by + ``EPDiagnostics.register_hierarchical_scales`` a variable is + additionally flagged when its mean has fallen below + ``scale_mean_fraction`` of its initial value *and* its relative + error ``std / |mean|`` is below ``scale_relative_error``. Returns ------- @@ -235,12 +346,35 @@ def check_sigma_collapse( results text at the end of a run. """ warnings_list = [] + scale_variables = getattr(diagnostics, "scale_variables", set()) for name, rows in diagnostics.variable_history.items(): stds = np.array([std for _, _, std in rows], dtype=float) + means = np.array([mean for _, mean, _ in rows], dtype=float) if len(stds) == 0: continue + if name in scale_variables and _flag_scale_collapse( + means, stds, scale_mean_fraction, scale_relative_error + ): + relative_error = ( + stds[-1] / abs(means[-1]) if means[-1] != 0 else float("inf") + ) + warnings_list.append( + f"scale-collapse: hierarchical parent scale '{name}' has " + f"collapsed to {means[-1]:.3g} " + f"({means[-1] / means[0]:.1%} of its initial {means[0]:.3g}) " + f"with a relative error of {relative_error:.2g} — the fit is " + f"reporting near-zero parent scatter as a confident answer " + f"(see PyAutoFit #1405). This is a known EP instability for " + f"scale hyperparameters and the value should NOT be trusted; " + f"cross-check the parent scatter against a joint sampler fit " + f"of the same graph." + ) + continue + + stds = _drop_consecutive_repeats(stds) + if stds[-1] < std_floor: warnings_list.append( f"sigma-collapse: variable '{name}' has std {stds[-1]:.3g} " @@ -258,7 +392,7 @@ def check_sigma_collapse( if np.all(np.diff(tail) < 0) and stds[-1] < shrink_factor * stds[0]: warnings_list.append( f"sigma-collapse: variable '{name}' std has shrunk " - f"monotonically over the last {monotone_steps} updates to " + f"monotonically over its last {monotone_steps} updates to " f"{stds[-1]:.3g} ({stds[-1] / stds[0]:.1e} of its initial " f"value) — possible information over-counting (PyAutoFit " f"#1332 F10)." diff --git a/autofit/graphical/expectation_propagation/optimiser.py b/autofit/graphical/expectation_propagation/optimiser.py index 9b82b4dc4..deb3cf54e 100644 --- a/autofit/graphical/expectation_propagation/optimiser.py +++ b/autofit/graphical/expectation_propagation/optimiser.py @@ -238,6 +238,7 @@ def __init__( self.ep_history = ep_history or EPHistory() self.diagnostics = EPDiagnostics() + self.diagnostics.register_hierarchical_scales(factor_graph) # Per-factor count of consecutive failed updates; see # `_check_consecutive_failures`. Reset at the start of every `run`. diff --git a/test_autofit/graphical/functionality/test_diagnostics.py b/test_autofit/graphical/functionality/test_diagnostics.py index 736e40bfd..57b5f90a8 100644 --- a/test_autofit/graphical/functionality/test_diagnostics.py +++ b/test_autofit/graphical/functionality/test_diagnostics.py @@ -202,3 +202,130 @@ def test_parallel_end_of_run_guards(): opt._output_diagnostics() opt._output_diagnostics(final=True, model_approx=model_approx) opt._warn_sigma_collapse() + + +def _scale_diagnostics(means, stds, variable="parent_sigma"): + """ + An `EPDiagnostics` carrying one registered parent-scale variable + with the given mean/std trajectory. + """ + diagnostics = EPDiagnostics() + diagnostics.scale_variables = {variable} + diagnostics.variable_rows = [ + {"step": step, "factor": "hierarchical", "variable": variable, + "mean": float(mean), "std": float(std)} + for step, (mean, std) in enumerate(zip(means, stds)) + ] + return diagnostics + + +# The three states below are the measured outcomes of the PyAutoFit #1405 toy +# (parent scale hyper-prior mean 10, truth 10): two COLLAPSE runs and the +# RECOVER band. See PyAutoMind complete/2026/07/ep_scale_collapse_assets/. +@pytest.mark.parametrize( + "final_mean, final_std", + [ + (0.80, 0.11), # shallow collapse — the std alone looks unremarkable + (0.0030, 1e-5), # deep collapse + ], +) +def test_scale_collapse_flags_measured_collapses(final_mean, final_std): + diagnostics = _scale_diagnostics( + means=[10.0, 8.0, 4.0, final_mean], + stds=[5.0, 3.0, 1.0, final_std], + ) + + warnings_list = graph.check_sigma_collapse(diagnostics) + + assert len(warnings_list) == 1 + assert "scale-collapse" in warnings_list[0] + assert "parent_sigma" in warnings_list[0] + assert "#1405" in warnings_list[0] + + +@pytest.mark.parametrize("final_mean, final_std", [(9.1, 0.9), (12.8, 2.4)]) +def test_scale_collapse_silent_on_measured_recoveries(final_mean, final_std): + diagnostics = _scale_diagnostics( + means=[10.0, 8.0, 11.0, final_mean], + stds=[5.0, 3.0, 2.0, final_std], + ) + + assert graph.check_sigma_collapse(diagnostics) == [] + + +def test_scale_collapse_needs_confidence_not_just_a_small_mean(): + """ + A small parent scale that is honestly uncertain is not a collapse — + the pathology is a small scale reported *confidently*. + """ + diagnostics = _scale_diagnostics( + means=[10.0, 8.0, 4.0, 0.80], + stds=[5.0, 3.0, 1.0, 2.0], + ) + + assert graph.check_sigma_collapse(diagnostics) == [] + + +def test_scale_check_applies_only_to_registered_scale_variables(): + """ + The same trajectory on an unregistered variable must not be flagged: + the relative test is meaningful for a scale hyperparameter, not for + an arbitrary variable that happens to approach zero. + """ + diagnostics = _scale_diagnostics(means=[10.0, 4.0, 0.0030], stds=[5.0, 1.0, 1e-5]) + diagnostics.scale_variables = set() + + assert graph.check_sigma_collapse(diagnostics) == [] + + +def test_monotone_limb_survives_unchanged_steps(): + """ + A variable only moves when a factor adjacent to it is updated, but a + snapshot is recorded for every variable on every factor update. The + monotone test must not be defeated by the resulting repeated rows. + """ + shrinking = np.geomspace(1.0, 1e-5, num=8) + # Interleave each real update with two steps at which this variable + # did not move — the shape a real multi-factor graph produces. + with_repeats = [std for std in shrinking for _ in range(3)] + + diagnostics = EPDiagnostics() + diagnostics.variable_rows = [ + {"step": step, "factor": "f", "variable": "shrinking", "mean": 1.0, + "std": float(std)} + for step, std in enumerate(with_repeats) + ] + + warnings_list = graph.check_sigma_collapse(diagnostics) + + assert len(warnings_list) == 1 + assert "monotonically" in warnings_list[0] + + +def test_register_hierarchical_scales_finds_the_parent_sigma(): + import autofit as af + + hierarchical_factor = af.HierarchicalFactor( + af.GaussianPrior, + mean=af.GaussianPrior(mean=50.0, sigma=10.0), + sigma=af.GaussianPrior(mean=10.0, sigma=5.0), + ) + for _ in range(3): + hierarchical_factor.add_drawn_variable(af.GaussianPrior(mean=50.0, sigma=10.0)) + + factor_graph = af.FactorGraphModel(hierarchical_factor) + + diagnostics = EPDiagnostics() + diagnostics.register_hierarchical_scales(factor_graph.graph) + + sigma_prior = dict(hierarchical_factor.prior_tuples)["sigma"] + assert diagnostics.scale_variables == {sigma_prior.name} + + +def test_register_hierarchical_scales_tolerates_a_plain_factor_graph(): + model_approx, _ = make_model_approx() + + diagnostics = EPDiagnostics() + diagnostics.register_hierarchical_scales(model_approx.factor_graph) + + assert diagnostics.scale_variables == set()