diff --git a/autofit/graphical/README.md b/autofit/graphical/README.md index 7c7f729dd..e60c034af 100644 --- a/autofit/graphical/README.md +++ b/autofit/graphical/README.md @@ -115,6 +115,44 @@ tilted distribution is then *fitted* by the factor's optimiser: back, with a log line, to `hessian="quasi"`, which keeps the pre-2026-09 path (the quasi-Newton diagonal secant, refined at `n_refine` random draws from the mean field; not deterministic). +- **Moments path** (`LaplaceOptimiser(projection="moments")`, + `graphical/laplace/moments.py`): the tilted moments of Eq. (8) by + nested quadrature instead of the mode. **Outer** variables — a + hierarchical factor's scale (`_HierarchicalFactor.scale_variables`, + the `sigma`-named argument of its parent distribution) and any + variable whose message has a bounded support (a + `TruncatedGaussianPrior`) — are integrated by a tensor + Gauss–Legendre rule of `n_quadrature` (64) nodes per variable in the + message's base coordinate `u` (σ itself for Normal/TruncatedNormal + messages, log σ for a `LogGaussianPrior`), on the support clipped to + the cavity `mean ± quadrature_half_width (8) · std`; further passes + (at most 4) re-window onto the tilted `mean ± 8 · max(std, node + spacing)` while a pass finds the tilted std under a quarter of the + scale its window was built from, or mass at a non-support window + edge. **Inner** variables — the rest — are integrated at each + outer node `s_j` by a conditional Laplace approximation (mode + `m_j` by quasi-Newton warm-started from the previous node, polished + by Newton steps on central differences of the tilted log-density + *values*, covariance `Σ_j = (−H_j)⁻¹`), which is exact when the + conditional is Gaussian — the hierarchical-Gaussian case. Each node + carries + + log w_j = log w_GL,j + log|dx/du|_j + log p̂ₐ(m_j, s_j) + + (d/2) log 2π − ½ log det(−H_j) (6a) + + and `logsumexp(log w)` is the tilted normalisation `Ẑₐ` (§5). Moments + go through each message's own `project` (Eq. 9, via + `MeanField.from_weighted_nodes`). The path is deterministic (no + random draws) and costs ~0.6–2 s / ~1200–2000 factor calls per + hierarchical-factor update in numpy. A factor with no outer variable, + more than `moment_max_outer` (2) of them, a non-scalar one, more than + `moment_max_size` (4) flattened free parameters, or deterministic + variables takes the Laplace path unchanged; the default is + `projection="mode"`. An empty cavity window, zero tilted mass, a + non-finite moment, residual edge mass or a non-concave inner Hessian at + a node carrying more than 1e-10 of the mass is a `BAD_PROJECTION`, an + inner search that does not converge at such a node a `FAILURE`; both + return the mean field unchanged. - **Exact path** (`ExactFactorFit`, `expectation_propagation/factor_optimiser.py`): if the factor is itself a message of the same family as the cavity @@ -154,6 +192,21 @@ numerical stabilisation, and the mean weight supplies the projection's natural parameters per family. On the Laplace path the "projection" is the Gaussian mode/covariance construction instead. +On the moments path (§3.2) the same Eq. (9) runs over quadrature nodes +rather than samples: `log w_s` are the node weights of Eq. (6a), the +outer variables sit at their Gauss–Legendre nodes `s_j`, and each inner +conditional `N(m_j, Σ_j)` is expanded into the 3^d points +`m_j + L_j z`, `z ∈ {−√3, 0, √3}^d`, `L_j L_jᵀ = Σ_j`, with the order-3 +Gauss–Hermite weights `{1/6, 2/3, 1/6}` — exact for `E[x]` and +`E[x xᵀ]`, so the matched variance is the law of total variance +`E[Σ_j] + Var[m_j]`. The weights are shifted so every projected message +has `log_norm` 0 and the projection carries `log_norm = log Ẑₐ` +(`MeanField.from_weighted_nodes`). A `TruncatedNormalMessage` takes the +matched `(E, Var)` as its *parent* mean and variance +(`invert_sufficient_statistics`), so its own truncated moments differ +from the matched ones when the mass is near a limit; exact inversion of +truncated moments is out of scope. + ### 3.4 Factor update with damping — `MeanField.update_factor_mean_field` Divide out the cavity and damp with `δ ∈ (0, 1]`: @@ -275,8 +328,11 @@ What to do with it: - Trust EP for the parent **mean** and the per-dataset variables; read the hierarchical **scatter** from a joint sampler over `factor_graph.global_prior_model` (the `graphical/` pattern) before - quoting it, or wait for a moment-matching projection of the - hierarchical factor. + quoting it, or project the hierarchical factor by its moments: + `factor_graph.optimise(af.LaplaceOptimiser(projection="moments"))` + (§3.2), which integrates σ over its support instead of seeking a mode + and so updates the scatter where the Laplace path skips it + (PyAutoFit#1654). - Prefer a log-scale parameterisation of the scatter (`LogGaussianPrior`), whose tilted density is bounded, over a `GaussianPrior`/`TruncatedGaussianPrior` truncated at zero: it diff --git a/autofit/graphical/declarative/factor/hierarchical.py b/autofit/graphical/declarative/factor/hierarchical.py index e3b95c1fd..276325dc8 100644 --- a/autofit/graphical/declarative/factor/hierarchical.py +++ b/autofit/graphical/declarative/factor/hierarchical.py @@ -3,6 +3,9 @@ import numpy as np from autofit import exc +from autofit.graphical.expectation_propagation.diagnostics import ( + _SCALE_ARGUMENT_NAMES, +) from autofit.mapper.model import ModelInstance from autofit.mapper.prior.abstract import Prior from autofit.mapper.prior_model.collection import Collection @@ -203,6 +206,24 @@ def message_dict(self) -> Dict[Prior, NormalMessage]: def variable(self): return self.drawn_prior + @property + def scale_variables(self) -> frozenset: + """ + The priors parameterising the parent distribution's *scale* — the + argument named ``sigma`` (or ``scale``/``std``/``stddev``, see + ``_SCALE_ARGUMENT_NAMES``) of e.g. ``af.GaussianPrior``. + + These are the variables whose tilted density can sit on the ``0`` + boundary, so ``LaplaceOptimiser(projection="moments")`` integrates + them by quadrature on their positive support instead of seeking a + mode (PyAutoFit#1654). + """ + return frozenset( + prior + for name, prior in self.distribution_model.prior_tuples + if name in _SCALE_ARGUMENT_NAMES + ) + def log_likelihood_function(self, instance, shared=None): return instance.distribution_model.message(instance.drawn_prior, xp=self._xp) diff --git a/autofit/graphical/expectation_propagation/diagnostics.py b/autofit/graphical/expectation_propagation/diagnostics.py index a97e1cc08..838d4db52 100644 --- a/autofit/graphical/expectation_propagation/diagnostics.py +++ b/autofit/graphical/expectation_propagation/diagnostics.py @@ -44,6 +44,9 @@ #: recognised on a ``HierarchicalFactor``. ``GaussianPrior`` and #: ``LogGaussianPrior`` call it ``sigma``; the others are accepted so a #: distribution that names it differently is still covered. +#: Shared with ``_HierarchicalFactor.scale_variables`` (the variables the +#: moment projection integrates by quadrature); this module imports nothing +#: from ``autofit``, so both can import it without a cycle. _SCALE_ARGUMENT_NAMES = frozenset({"sigma", "scale", "std", "stddev"}) diff --git a/autofit/graphical/laplace/moments.py b/autofit/graphical/laplace/moments.py new file mode 100644 index 000000000..d4d4c39be --- /dev/null +++ b/autofit/graphical/laplace/moments.py @@ -0,0 +1,710 @@ +""" +Moment-matching projection of a factor's tilted distribution by nested +quadrature (PyAutoFit#1654). + +The Laplace ("mode") projection of `LaplaceOptimiser` fails on the factor it is +most needed for: a ``HierarchicalFactor`` whose parent scale σ is poorly +constrained. Its tilted density in σ piles up against σ = 0, so there is no +interior mode and no negative-definite Hessian (``BAD_PROJECTION``), or the +line search walks into the boundary (``FAILURE``). EP's projection is, by +definition, a *moment* match — the Gaussian closest in KL(p̃ ‖ q) to the tilted +density p̃ — and the mode/curvature pair is only its large-data approximation. + +This module computes those moments directly, INLA-style: + +- **Outer** variables — the factor's scale variables (``scale_variables`` of a + ``_HierarchicalFactor``) and any variable whose message has a bounded support + (e.g. a ``TruncatedGaussianPrior`` σ) — are integrated with a tensor + Gauss–Legendre rule of ``n_quadrature`` nodes per variable, in the message's + base coordinate u (identity for Normal/TruncatedNormal messages, log σ for a + ``LogGaussianPrior``), on the intersection of the variable's support with the + cavity window ``mean ± half_width · std``. Further passes (at most + ``MAX_PASSES``) re-window the rule onto the tilted + ``mean ± half_width · max(std, node spacing)`` while a pass finds the tilted + density narrower than a quarter of its window's scale (under-resolved) or + finds mass at a window edge that is not a support edge. +- **Inner** variables — everything else — are integrated at each outer node + s_j by a conditional Laplace approximation: the conditional mode m_j + (quasi-Newton, warm-started from the previous node, polished by Newton steps + on the finite-difference Hessian H_j) and covariance Σ_j = (−H_j)⁻¹. For a + hierarchical Gaussian model the conditional density is exactly Gaussian, so + this step is exact. + +Each outer node carries the log weight + + log w_j = log w_GL,j + log |dx/du|_j + ℓ(m_j, s_j) + (d/2) log 2π − ½ log det(−H_j) + +where ℓ is the tilted log-density (factor plus cavity) and d the number of inner +parameters, so that ``logsumexp(log w)`` is the tilted normalisation +Ẑ = ∫ f q_cavity. Moments are then handed to each message's own ``project`` +(`MeanField.from_weighted_nodes`) with the outer nodes at s_j and the inner +nodes an order-3 Gauss–Hermite expansion of each conditional N(m_j, Σ_j) in +whitened coordinates (3^d points, exact for E[x] and E[x xᵀ]), which reproduces +the law of total variance. + +No random numbers are drawn: the projection is bit-for-bit deterministic. +""" +import itertools +import logging +import math +from operator import attrgetter +from typing import Dict, List, Optional, Tuple + +import numpy as np + +from autofit import exc +from autofit.graphical.laplace import newton +from autofit.graphical.mean_field import MeanField +from autofit.graphical.utils import FlattenArrays, Status, StatusFlag +from autofit.mapper.variable import Variable +from autofit.mapper.variable_operator import VariableData +from autofit.messages.composed_transform import TransformedMessage + +logger = logging.getLogger(__name__) + +#: The parameter support of a distribution's scale argument (``sigma`` of +#: ``NormalMessage._parameter_support``): a scale is positive. +SCALE_SUPPORT = (0.0, math.inf) + +#: A node whose share of the tilted mass is above this is "weighted": a failed +#: inner search or a non-concave inner Hessian there invalidates the projection. +WEIGHTED_NODE_MASS = 1e-10 + +#: Mass on the outermost node of a window edge that is not a support edge, +#: above which the window is judged to have cut off tilted mass. +EDGE_MASS = 1e-6 + +#: A pass is re-windowed when the tilted std in base space is below this +#: fraction of the scale its window was built from (the cavity std on the +#: first pass): the tilted density is under-resolved by the rule. +NARROW_FRACTION = 0.25 + +#: At most this many quadrature passes (each re-window shrinks the node +#: spacing ~2.5x); still unresolved after the last is a BAD_PROJECTION. +MAX_PASSES = 4 + +#: Newton polish of the conditional mode stops once the predicted increase of +#: the log-density, ½ gᵀ(−H)⁻¹g, is below this (nats). +NEWTON_DECREMENT_TOL = 1e-12 +MAX_NEWTON_POLISH = 5 + +# Order-3 (probabilists') Gauss–Hermite rule: exact to the fifth moment of N(0, 1) +_GH_NODES = np.array([-math.sqrt(3.0), 0.0, math.sqrt(3.0)]) +_GH_LOG_WEIGHTS = np.log(np.array([1.0 / 6.0, 2.0 / 3.0, 1.0 / 6.0])) + + +def _logsumexp(a) -> float: + # numpy only: `autofit` must not import scipy.special at import time + a = np.asarray(a, dtype=float) + m = np.max(a) + if not np.isfinite(m): + return float(m) + return float(m + np.log(np.sum(np.exp(a - m)))) + + +def _cho_solve(L, b): + """Solve (L Lᵀ) x = b for a lower-triangular Cholesky factor L.""" + return np.linalg.solve(L.T, np.linalg.solve(L, b)) + + +def _scale_variables(factor_approx) -> frozenset: + factor = getattr(factor_approx, "factor", factor_approx) + try: + return frozenset(getattr(factor, "scale_variables", ()) or ()) + except Exception: # a factor that cannot resolve its priors is not a scale factor + return frozenset() + + +def _has_bounded_support(message) -> bool: + kw = message._support_kwargs + return bool(kw) and ( + np.isfinite(kw.get("lower_limit", -math.inf)) + or np.isfinite(kw.get("upper_limit", math.inf)) + ) + + +def split_variables( + factor_approx, mean_field +) -> Tuple[List[Variable], List[Variable], frozenset]: + """ + Split a factor's free variables into outer (scale / bounded-support) and + inner variables, each sorted by ``Variable.id`` so that the node order does + not depend on dict order. + """ + scales = _scale_variables(factor_approx) + free = sorted(factor_approx.free_variables, key=attrgetter("id")) + outer = [v for v in free if v in scales or _has_bounded_support(mean_field[v])] + inner = [v for v in free if v not in outer] + return outer, inner, scales + + +def fallback_reason( + factor_approx, mean_field, max_size: int, max_outer: int +) -> Optional[str]: + """ + Why the moments path does not apply to this factor (so the mode path runs + unchanged), or ``None`` when it does. + """ + if not ( + hasattr(factor_approx, "cavity_dist") + and hasattr(factor_approx, "func_gradient") + ): + return "not a FactorApproximation" + if factor_approx.deterministic_variables: + return "factor has deterministic variables" + + outer, inner, _ = split_variables(factor_approx, mean_field) + if not outer: + return "no scale or bounded-support variable" + if len(outer) > max_outer: + return f"{len(outer)} outer variables exceed moment_max_outer={max_outer}" + for v in outer: + if np.shape(mean_field[v].mean) != (): + return f"outer variable {v.name} is not a scalar" + n_params = sum(np.size(mean_field[v].mean) for v in outer + inner) + if n_params > max_size: + return f"{n_params} free parameters exceed moment_max_size={max_size}" + return None + + +def gauss_legendre(n: int, lo: float, hi: float) -> Tuple[np.ndarray, np.ndarray]: + """Gauss–Legendre nodes and log weights on [lo, hi].""" + x, w = np.polynomial.legendre.leggauss(n) + half = 0.5 * (hi - lo) + return lo + half * (x + 1.0), np.log(half * w) + + +class _OuterAxis: + def __init__(self, variable: Variable, message, cavity, is_scale: bool): + """ + One outer quadrature axis, in the message's base coordinate u. + + ``lo``/``hi`` are the support in u: the message's own support (the + truncation limits of a TruncatedNormal, in base space for a + TransformedMessage) intersected, for a scale variable, with the + positive half-line mapped through the transform. + """ + self.variable = variable + self.message = message + self.transformed = isinstance(message, TransformedMessage) + + lo, hi = -math.inf, math.inf + kw = message._support_kwargs + lo = max(lo, float(kw.get("lower_limit", -math.inf))) + hi = min(hi, float(kw.get("upper_limit", math.inf))) + if is_scale: + ends = np.array(SCALE_SUPPORT, dtype=float) + if self.transformed: + with np.errstate(all="ignore"): + ends = np.asarray(message._transform(ends), dtype=float) + ends = np.nan_to_num(ends, nan=-math.inf) + lo = max(lo, float(np.min(ends))) + hi = min(hi, float(np.max(ends))) + self.lo, self.hi = lo, hi + + self.mean = self.std = math.nan + if cavity is not None and hasattr(cavity, "variance"): + base = cavity.base_message if isinstance(cavity, TransformedMessage) else cavity + try: + mean = float(base.mean) + variance = float(base.variance) + except (TypeError, ValueError, AttributeError): + mean = variance = math.nan + if np.isfinite(mean) and np.isfinite(variance) and variance > 0: + self.mean, self.std = mean, math.sqrt(variance) + + def window(self, centre: float, scale: float, half_width: float): + """ + ``centre ± half_width·scale`` clipped to the support; an edge that lies + on a (finite) support limit is lifted inside it by + ``1e-12·max(1, |centre|)`` — a scale of exactly 0 has zero measure and + N(x | μ, 0) is undefined. Returns ``(lo, hi, lo_is_support, + hi_is_support)`` or ``None`` when the window is empty. + """ + if not (np.isfinite(centre) and np.isfinite(scale) and scale > 0): + if np.isfinite(self.lo) and np.isfinite(self.hi): + g_lo, g_hi = self.lo, self.hi + else: + return None + else: + g_lo = max(centre - half_width * scale, self.lo) + g_hi = min(centre + half_width * scale, self.hi) + lift = 1e-12 * max(1.0, abs(centre) if np.isfinite(centre) else 1.0) + lo_is_support = g_lo == self.lo + hi_is_support = g_hi == self.hi + if lo_is_support: + g_lo = g_lo + lift + if hi_is_support: + g_hi = g_hi - lift + if not g_hi > g_lo: + return None + return g_lo, g_hi, lo_is_support, hi_is_support + + def to_physical(self, u: float) -> Tuple[float, float]: + """The physical value at base coordinate u and log |dx/du|.""" + if not self.transformed: + return u, 0.0 + x = float(self.message._inverse_transform(np.asarray(u, dtype=float))) + _, log_du_dx = self.message._transform_det(np.asarray(x, dtype=float)) + return x, -float(log_du_dx) + + +class _Conditional: + def __init__(self, factor_approx, outer_values: Dict[Variable, float]): + """The tilted log-density with the outer variables held fixed.""" + self.factor_approx = factor_approx + self.outer_values = outer_values + self.f_count = 0 + self.g_count = 0 + + def __call__(self, parameters): + self.f_count += 1 + return self.factor_approx({**parameters, **self.outer_values}) + + def gradient(self, parameters): + self.g_count += 1 + value, gradient = self.factor_approx.func_gradient( + {**parameters, **self.outer_values} + ) + return value, VariableData({v: gradient[v] for v in parameters}) + + +class _NodeResult: + __slots__ = ("ok", "kind", "log_weight", "mode", "chol_cov", "reason") + + def __init__(self, ok, log_weight, mode=None, chol_cov=None, kind="", reason=""): + self.ok = ok + self.kind = kind # "search" or "hessian" on failure + self.log_weight = log_weight + self.mode = mode + self.chol_cov = chol_cov + self.reason = reason + + +class MomentProjection: + def __init__(self, optimiser, factor_approx, mean_field, params=None, **kwargs): + """ + One moment-matching projection of ``factor_approx``'s tilted density + with the settings of ``optimiser`` (a ``LaplaceOptimiser``). + """ + self.optimiser = optimiser + self.factor_approx = factor_approx + self.mean_field = mean_field + self.kwargs = kwargs + + self.outer, self.inner, scales = split_variables(factor_approx, mean_field) + cavity = factor_approx.cavity_dist + self.axes = [ + _OuterAxis(v, mean_field[v], cavity.get(v), v in scales) + for v in self.outer + ] + + parameters = MeanField.mean.fget(mean_field) + if params: + for v, p in params.items(): + parameters[v] = p + self.start = VariableData({v: parameters[v] for v in self.inner}) + self.shapes = FlattenArrays({v: np.shape(self.start[v]) for v in self.inner}) + self.n_inner = int(sum(np.size(self.start[v]) for v in self.inner)) + + if self.inner: + self.hessian0 = optimiser.make_hessian( + mean_field, self.inner, **optimiser.hessian_kws + ) + if optimiser.check_limits: + lower = MeanField.lower_limit.fget(mean_field) + upper = MeanField.upper_limit.fget(mean_field) + self.limits = dict( + lower_limit=VariableData({v: lower[v] for v in self.inner}), + upper_limit=VariableData({v: upper[v] for v in self.inner}), + ) + else: + self.limits = {} + + self.n_passes = 0 + self.n_nodes = 0 + self.n_polished = 0 + self.f_count = 0 + self.g_count = 0 + + @property + def factor_name(self): + factor = getattr(self.factor_approx, "factor", self.factor_approx) + return getattr(factor, "name", factor) + + # ----------------------------------------------------------------- nodes + + def _solve_node(self, outer_values, log_w_outer, start) -> _NodeResult: + conditional = _Conditional(self.factor_approx, outer_values) + if not self.inner: + value = float(conditional({})) + self.f_count += conditional.f_count + return _NodeResult(True, log_w_outer + value) + + state = newton.OptimisationState( + conditional, + conditional.gradient, + start, + self.hessian0, + None, + **self.limits, + ) + with np.errstate(all="ignore"): + next_state, status = self.optimiser.optimise_state(state, **self.kwargs) + best = max(state, next_state, key=lambda s: s.value) + + # Newton polish on central differences of the conditional tilted + # log-density *values*. The factor gradients are forward differences + # (eps=1e-8) of a log-density whose exponential-family terms cancel + # catastrophically at a small scale node (η₁x, η₂x², A(η) ~ 1e6-1e7 at + # σ ~ 0.02): their ~1e-1 noise swamps the soft direction of the + # conditional (curvature ~1e-2) and made `finite_difference_hessian` + # indefinite there. Value differences on the `fd_steps` scale carry + # ~1e-9 noise and are exact for the (Gaussian) conditional of a + # hierarchical Gaussian model. + h = self.shapes.flatten(self.optimiser.fd_steps(best, self.mean_field)) + x = self.shapes.flatten(best.parameters) + f0 = float(best.value) + converged = False + chol = None + reason = "" + for _ in range(MAX_NEWTON_POLISH + 1): + with np.errstate(all="ignore"): + f0, g, H = self._value_derivatives(conditional, x, f0, h) + if not (np.isfinite(f0) and np.all(np.isfinite(H)) and np.all(np.isfinite(g))): + chol, reason = None, "non-finite inner Hessian" + break + try: + chol = np.linalg.cholesky(-H) + except np.linalg.LinAlgError: + min_eig = np.linalg.eigvalsh(-H).min() + chol = None + reason = f"inner tilted log-density not concave (min eig {min_eig:.3g})" + break + step = _cho_solve(chol, g) + decrement = 0.5 * float(g @ step) + if decrement < NEWTON_DECREMENT_TOL: + converged = True + break + for _halving in range(20): + f1 = float(conditional(self._unflatten(x + step))) + if f1 > f0: + break + step = 0.5 * step + else: + # No further increase is available at this resolution + converged = decrement < 1e-6 + if not converged: + reason = "Newton polish of the conditional mode stalled" + break + x, f0 = x + step, f1 + self.n_polished += 1 + else: + reason = "Newton polish did not converge" + + self.f_count += conditional.f_count + self.g_count += conditional.g_count + proxy = log_w_outer + f0 + if chol is None: + return _NodeResult(False, proxy, kind="hessian", reason=reason) + if not converged: + return _NodeResult( + False, + proxy, + kind="search", + reason=f"{reason}; {status.messages[-1] if status.messages else ''}", + ) + + log_det = 2.0 * float(np.sum(np.log(np.diag(chol)))) + covariance = _cho_solve(chol, np.eye(len(chol))) + chol_cov = np.linalg.cholesky(0.5 * (covariance + covariance.T)) + log_weight = ( + proxy + 0.5 * self.n_inner * math.log(2 * math.pi) - 0.5 * log_det + ) + return _NodeResult( + True, + log_weight, + mode=x, + chol_cov=chol_cov, + ) + + def _unflatten(self, x): + return VariableData(self.shapes.unflatten(x)) + + def _value_derivatives(self, conditional, x, f0, h): + """ + The value, gradient and Hessian of the conditional tilted log-density + at flat parameters ``x`` by central differences of its values with + per-parameter steps ``h``: 1 + 2d + 2d(d-1) evaluations. + """ + d = x.size + f = lambda y: float(conditional(self._unflatten(y))) + E = np.diag(h) + fp = np.array([f(x + E[i]) for i in range(d)]) + fm = np.array([f(x - E[i]) for i in range(d)]) + g = (fp - fm) / (2 * h) + H = np.empty((d, d)) + H[np.diag_indices(d)] = (fp - 2 * f0 + fm) / h**2 + for i in range(d): + for j in range(i + 1, d): + fpp = f(x + E[i] + E[j]) + fpm = f(x + E[i] - E[j]) + fmp = f(x - E[i] + E[j]) + fmm = f(x - E[i] - E[j]) + H[i, j] = H[j, i] = (fpp - fpm - fmp + fmm) / (4 * h[i] * h[j]) + return f0, g, H + + def _integrate(self, windows): + """One tensor-product pass over the outer windows.""" + n = self.optimiser.n_quadrature + rules = [gauss_legendre(n, lo, hi) for lo, hi, _, _ in windows] + grid = [] + for index in itertools.product(range(n), repeat=len(rules)): + u = tuple(rules[k][0][i] for k, i in enumerate(index)) + log_w = sum(rules[k][1][i] for k, i in enumerate(index)) + grid.append((index, u, log_w)) + + results = [] + start = self.start + for index, u, log_w in grid: + outer_values = {} + log_j = 0.0 + for axis, u_k in zip(self.axes, u): + x, lj = axis.to_physical(u_k) + outer_values[axis.variable] = x + log_j += lj + result = self._solve_node(outer_values, log_w + log_j, start) + if result.ok and result.mode is not None: + start = VariableData(self.shapes.unflatten(result.mode)) + results.append((index, u, outer_values, result)) + + self.n_passes += 1 + return results + + # ------------------------------------------------------------- statistics + + def _check_nodes(self, results) -> Tuple[Optional[Status], Optional[float]]: + good = [r.log_weight for *_, r in results if r.ok] + failed = [r for *_, r in results if not r.ok] + log_z = _logsumexp(good) if good else -math.inf + if not np.isfinite(log_z): + if any(r.kind == "search" for r in failed if np.isfinite(r.log_weight)): + return self._status(StatusFlag.FAILURE, "inner search failed at every weighted node"), None + return self._status(StatusFlag.BAD_PROJECTION, "zero tilted mass in the window"), None + + threshold = log_z + math.log(WEIGHTED_NODE_MASS) + for r in failed: + if np.isfinite(r.log_weight) and r.log_weight > threshold: + flag = StatusFlag.FAILURE if r.kind == "search" else StatusFlag.BAD_PROJECTION + return self._status(flag, f"at a weighted node: {r.reason}"), None + return None, log_z + + def _axis_statistics(self, results, log_z, windows): + """ + Tilted mean/std in base space, edge-mass flag and node spacing at the + peak, per outer axis. + """ + n = self.optimiser.n_quadrature + stats = [] + for k, axis in enumerate(self.axes): + marginal = np.zeros(n) + u_k = np.zeros(n) + for index, u, _, r in results: + u_k[index[k]] = u[k] + if r.ok: + marginal[index[k]] += math.exp(r.log_weight - log_z) + mean = float(np.sum(marginal * u_k)) + std = math.sqrt(max(float(np.sum(marginal * (u_k - mean) ** 2)), 0.0)) + _, _, lo_support, hi_support = windows[k] + edge = (not lo_support and marginal[0] > EDGE_MASS) or ( + not hi_support and marginal[-1] > EDGE_MASS + ) + # The rule's node spacing at the peak: a floor for the re-window + # scale when the first pass under-resolves a narrow tilted density + peak = int(np.argmax(marginal)) + gaps = np.diff(u_k) + spacing = float(np.max(gaps[max(peak - 1, 0) : peak + 1])) + stats.append((mean, std, edge, spacing)) + return stats + + def _status(self, flag, reason, success=False): + return Status( + success=success, + messages=(f"moment projection for {self.factor_name}: {reason}",), + updated=False, + flag=flag, + ) + + # ------------------------------------------------------------------- run + + def __call__(self): + half_width = self.optimiser.quadrature_half_width + windows, scales = [], [] + for axis in self.axes: + window = axis.window(axis.mean, axis.std, half_width) + if window is None: + return self.mean_field, self._status( + StatusFlag.BAD_PROJECTION, + f"cavity window of {axis.variable.name} lies outside its support", + ) + windows.append(window) + scales.append( + axis.std + if np.isfinite(axis.std) + else (window[1] - window[0]) / (2 * half_width) + ) + + # Pass 1 on the cavity window; re-window onto the tilted + # mean ± half_width·max(std, node spacing) while the tilted density is + # under-resolved (std < NARROW_FRACTION of the window's scale) or has + # mass on a window edge that is not a support edge. + for n_pass in range(MAX_PASSES): + results = self._integrate(windows) + status, log_z = self._check_nodes(results) + if status is not None: + return self.mean_field, status + + stats = self._axis_statistics(results, log_z, windows) + unresolved = [ + edge or not std >= NARROW_FRACTION * scale + for (mean, std, edge, _), scale in zip(stats, scales) + ] + if not any(unresolved): + break + if n_pass == MAX_PASSES - 1: + names = ", ".join( + axis.variable.name + for axis, bad in zip(self.axes, unresolved) + if bad + ) + return self.mean_field, self._status( + StatusFlag.BAD_PROJECTION, + f"residual edge mass / unresolved tilted density for {names} " + f"after {MAX_PASSES} passes", + ) + + windows, scales = [], [] + for axis, (mean, std, _, spacing) in zip(self.axes, stats): + scale = max(std, spacing) + window = axis.window(mean, scale, half_width) if scale > 0 else None + if window is None: + return self.mean_field, self._status( + StatusFlag.BAD_PROJECTION, + f"empty tilted window for {axis.variable.name}", + ) + windows.append(window) + scales.append(scale) + + return self._project(results, log_z) + + def _project(self, results, log_z): + nodes, log_weights = self._expand_nodes(results) + self.n_nodes = len(log_weights) + + for v, values in nodes.items(): + message = self.mean_field[v] + if isinstance(message, TransformedMessage): + with np.errstate(all="ignore"): + base = message._transform(values) + if not np.all(np.isfinite(base)): + return self.mean_field, self._status( + StatusFlag.BAD_PROJECTION, + f"quadrature nodes of {v.name} fall outside its support", + ) + + try: + with np.errstate(all="ignore"): + projection = MeanField.from_weighted_nodes( + self.mean_field, nodes, log_weights, log_norm=log_z + ) + except (AssertionError, ValueError, ArithmeticError, exc.MessageException) as e: + return self.mean_field, self._status( + StatusFlag.BAD_PROJECTION, f"moment inversion failed: {e}" + ) + + for v in nodes: + message = projection[v] + mean = np.asarray(message.mean, dtype=float) + variance = np.asarray(message.variance, dtype=float) + if not ( + np.all(np.isfinite(mean)) + and np.all(np.isfinite(variance)) + and np.all(variance > 0) + ): + return self.mean_field, self._status( + StatusFlag.BAD_PROJECTION, f"non-finite moment for {v.name}" + ) + if not np.isfinite(projection.log_norm): + return self.mean_field, self._status( + StatusFlag.BAD_PROJECTION, "non-finite tilted normalisation" + ) + + return projection, Status( + success=True, + messages=( + f"moments: n_nodes={self.n_nodes}, passes={self.n_passes}, " + f"f_count={self.f_count}, g_count={self.g_count}", + ), + updated=True, + flag=StatusFlag.SUCCESS, + ) + + def _expand_nodes(self, results): + """ + Outer nodes at s_j; each inner conditional N(m_j, Σ_j) expanded into the + 3^d order-3 Gauss–Hermite points m_j + L_j z in whitened coordinates. + """ + good = [(outer_values, r) for _, _, outer_values, r in results if r.ok] + d = self.n_inner + if d: + z = np.array(list(itertools.product(_GH_NODES, repeat=d))) + log_gh = np.array( + [sum(t) for t in itertools.product(_GH_LOG_WEIGHTS, repeat=d)] + ) + else: + z = np.zeros((1, 0)) + log_gh = np.zeros(1) + k = len(log_gh) + + log_weights = np.concatenate([r.log_weight + log_gh for _, r in good]) + nodes = { + axis.variable: np.repeat( + np.array([ov[axis.variable] for ov, _ in good], dtype=float), k + ) + for axis in self.axes + } + if d: + flat = np.concatenate([r.mode[None, :] + z @ r.chol_cov.T for _, r in good]) + n_nodes = flat.shape[0] + for v, (_, s) in zip(self.inner, self._slices()): + nodes[v] = flat[:, s].reshape((n_nodes,) + self.shapes[v]) + return VariableData(nodes), log_weights + + def _slices(self): + offset = 0 + for v in self.inner: + size = int(np.prod(self.shapes[v], dtype=int)) + yield v, slice(offset, offset + size) + offset += size + + +def moment_projection(optimiser, factor_approx, mean_field, params=None, **kwargs): + """ + The moments projection of ``factor_approx`` or ``None`` when it does not + apply, in which case the caller runs the mode path unchanged. + """ + reason = fallback_reason( + factor_approx, + mean_field, + optimiser.moment_max_size, + optimiser.moment_max_outer, + ) + if reason is not None: + factor = getattr(factor_approx, "factor", factor_approx) + logger.info( + "moment projection for %s: %s; using the mode (Laplace) projection", + getattr(factor, "name", factor), + reason, + ) + return None + return MomentProjection(optimiser, factor_approx, mean_field, params, **kwargs)() diff --git a/autofit/graphical/laplace/optimiser.py b/autofit/graphical/laplace/optimiser.py index bcb36b479..f6cbd04cf 100644 --- a/autofit/graphical/laplace/optimiser.py +++ b/autofit/graphical/laplace/optimiser.py @@ -7,7 +7,7 @@ from autofit.graphical.expectation_propagation.ep_mean_field import EPMeanField from autofit.graphical.expectation_propagation.optimiser import AbstractFactorOptimiser from autofit.graphical.factor_graphs.factor import Factor -from autofit.graphical.laplace import newton +from autofit.graphical.laplace import moments, newton from autofit.graphical.mean_field import MeanField, FactorApproximation from autofit.graphical.utils import FlattenArrays, Status, StatusFlag from autofit.mapper.variable_operator import VariableData, VariableFullOperator @@ -41,6 +41,22 @@ class LaplaceOptimiser(AbstractFactorOptimiser): A failed optimisation (line-search failure) returns the mean field it was handed, unchanged, so that the caller's projection reproduces the previous message exactly. + + ``projection="moments"`` replaces the mode/curvature pair by a moment match + of the tilted distribution, computed by nested quadrature + (`autofit.graphical.laplace.moments`): an outer Gauss–Legendre rule of + ``n_quadrature`` nodes over each scale variable of a hierarchical factor + (``_HierarchicalFactor.scale_variables``) or bounded-support variable + (e.g. a ``TruncatedGaussianPrior`` σ), windowed to the support and to the + cavity ``mean ± quadrature_half_width · std``, and an inner conditional + Laplace approximation over the remaining variables at each node. This is + the projection that recovers a hierarchical scatter whose tilted density + sits on σ = 0, where the mode path has no interior mode (PyAutoFit#1654). + It is deterministic (no random draws). A factor with no such variable, more + than ``moment_max_outer`` of them, a non-scalar one, more than + ``moment_max_size`` flattened free parameters, or deterministic variables + takes the mode path unchanged. The default ``"mode"`` is the Laplace + projection described above. """ def __init__( @@ -64,6 +80,11 @@ def __init__( quasi_newton_kws: Optional[Dict[str, Any]] = None, stop_kws: Optional[Dict[str, Any]] = None, check_limits=True, + projection: str = "mode", + n_quadrature: int = 64, + quadrature_half_width: float = 8.0, + moment_max_size: int = 4, + moment_max_outer: int = 2, **kwargs ): super().__init__(**kwargs) @@ -72,6 +93,10 @@ def __init__( raise ValueError( f"hessian must be 'fd' or 'quasi', got {hessian!r}" ) + if projection not in ("mode", "moments"): + raise ValueError( + f"projection must be 'mode' or 'moments', got {projection!r}" + ) self.make_hessian = make_hessian self.make_det_hessian = make_det_hessian or make_hessian @@ -104,6 +129,14 @@ def __init__( self.stop_kws = stop_kws or {} self.check_limits = check_limits + # Tilted-distribution projection: "mode" (Laplace, default) or + # "moments" (nested-quadrature moment match, `laplace.moments`). + self.projection = projection + self.n_quadrature = n_quadrature + self.quadrature_half_width = quadrature_half_width + self.moment_max_size = moment_max_size + self.moment_max_outer = moment_max_outer + @property def default_kws(self): return dict( @@ -278,6 +311,13 @@ def optimise_approx( ) -> Tuple[MeanField, Status]: mean_field = mean_field or factor_approx.model_dist + if self.projection == "moments": + result = self._moment_projection( + factor_approx, mean_field, params, **kwargs + ) + if result is not None: + return result + state = self.prepare_state(factor_approx, mean_field, params) next_state, status = self.optimise_state(state, **kwargs) if not status.success: @@ -330,6 +370,23 @@ def optimise_approx( projection = mean_field.from_opt_state(next_state) return projection, status + def _moment_projection( + self, + factor_approx: FactorApprox, + mean_field: MeanField, + params: VariableData = None, + **kwargs + ) -> Optional[Tuple[MeanField, Status]]: + """ + The moment-matching projection of `factor_approx`'s tilted distribution + (`laplace.moments.moment_projection`), or ``None`` when the factor has + no scale / bounded-support variable to integrate over (or too many, or + deterministic variables), in which case the mode path runs unchanged. + """ + return moments.moment_projection( + self, factor_approx, mean_field, params, **kwargs + ) + def refine_state(self, state, new_param, n_refine=None): """ Refine the quasi-Newton Hessian estimate of `state` with `n_refine` diff --git a/autofit/graphical/mean_field.py b/autofit/graphical/mean_field.py index d6479ad9d..488ade2e6 100755 --- a/autofit/graphical/mean_field.py +++ b/autofit/graphical/mean_field.py @@ -429,6 +429,60 @@ def from_mode_covariance( return projection + def from_weighted_nodes( + self, + nodes: Dict[Variable, np.ndarray], + log_weights: np.ndarray, + log_norm: float = 0.0, + ) -> "MeanField": + """ + Moment-matching projection of a weighted node set onto this mean + field's message families. + + Each variable in ``nodes`` is projected with its own message's + ``project`` — the exponential-family member whose expected sufficient + statistics match the weighted nodes (a ``TransformedMessage`` maps the + physical nodes to its base space itself; a truncated message keeps its + limits). The log weights are shifted so every projected message has + ``log_norm`` 0, and the returned mean field carries ``log_norm``: the + tilted normalisation estimate Ẑ (``logsumexp`` of the weights for a + quadrature rule). Fixed-value variables are carried over unchanged. + + Parameters + ---------- + nodes + Per variable, the node values, leading axis the node index. + log_weights + One log weight per node (shape ``(n_nodes,)``), e.g. quadrature + log weights plus the tilted log-density. + log_norm + The projection's log normalisation. + """ + log_weights = np.asarray(log_weights, dtype=float) + n_nodes = log_weights.shape[0] + # `project` takes the *mean* of the unshifted weights as its log_norm: + # shift so that mean is 1. + log_w_max = np.max(log_weights) + shifted = ( + log_weights + - (log_w_max + np.log(np.sum(np.exp(log_weights - log_w_max)))) + + np.log(n_nodes) + ) + dists = {} + for v in self.keys() & nodes.keys(): + message = self[v] + values = np.asarray(nodes[v], dtype=float) + weights = np.broadcast_to( + shifted.reshape((n_nodes,) + (1,) * (values.ndim - 1)), values.shape + ) + dists[v] = message.project( + values, weights, id_=message.id, **message._support_kwargs + ) + for v, value in self.fixed_values.items(): + if v in self and v not in dists: + dists[v] = self[v] + return MeanField(dists, log_norm=log_norm) + def sample(self, n_samples=None): return VariableData({v: dist.sample(n_samples) for v, dist in self.items()}) diff --git a/test_autofit/graphical/functionality/test_moment_projection.py b/test_autofit/graphical/functionality/test_moment_projection.py new file mode 100644 index 000000000..9af2d690c --- /dev/null +++ b/test_autofit/graphical/functionality/test_moment_projection.py @@ -0,0 +1,425 @@ +""" +``LaplaceOptimiser(projection="moments")`` projects a factor's tilted +distribution by matching moments computed with nested quadrature: an outer +Gauss-Legendre rule over the factor's scale (bounded-support) variables and an +inner conditional Laplace approximation over the remaining variables at each +outer node (PyAutoFit#1654). +""" +import logging + +import numpy as np +import pytest +from scipy import integrate, stats + +from autofit import graphical as graph +from autofit.graphical.laplace.optimiser import LaplaceOptimiser +from autofit.graphical.utils import StatusFlag +from autofit.mapper.variable import Variable +from autofit.messages import NormalMessage +from autofit.messages.truncated_normal import TruncatedNormalMessage + +# Cavity of `make_scale_approx`: x ~ N(M, A), s ~ TruncatedNormal(C, sqrt(CV), 0, 100) +M, A = 1.0, 0.5 +C, CV = 0.8, 0.36 + + +def make_scale_approx(c=C, cv=CV): + """ + A two-variable factor N(x | 0, s) with an analytic Jacobian: `s` is a + scale with a truncated (bounded-support) cavity, so it is the outer + quadrature variable and `x` the inner Laplace variable, whose conditional + tilted density is exactly Gaussian. + """ + x_, s_ = Variable("x"), Variable("s") + + def f(x, s): + return -0.5 * (x / s) ** 2 - np.log(s) - 0.5 * np.log(2 * np.pi) + + def f_jac(x, s): + return f(x, s), (-x / s**2, x**2 / s**3 - 1 / s) + + factor = graph.Factor(f, x_, s_, factor_jacobian=f_jac) + cavity = graph.MeanField( + { + x_: NormalMessage(M, np.sqrt(A)), + s_: TruncatedNormalMessage( + c, np.sqrt(cv), lower_limit=0.0, upper_limit=100.0 + ), + } + ) + return graph.FactorApproximation(factor, cavity, cavity, cavity), x_, s_ + + +def exact_scale_moments(): + """ + The tilted moments of `make_scale_approx` by 1-D quadrature over s: + p(s) ∝ TN(s | C, CV) N(M | 0, s² + A), E[x | s] = M s²/(s²+A), + Var[x | s] = A s²/(s²+A). + """ + a, b = (0.0 - C) / np.sqrt(CV), (100.0 - C) / np.sqrt(CV) + prior = stats.truncnorm(a, b, loc=C, scale=np.sqrt(CV)) + + def p(s): + return prior.pdf(s) * stats.norm.pdf(M, 0.0, np.sqrt(s**2 + A)) + + hi = C + 12 * np.sqrt(CV) + kw = dict(epsabs=1e-14, epsrel=1e-12, limit=200) + + def integral(g): + return integrate.quad(lambda s: g(s) * p(s), 0.0, hi, **kw)[0] + + z = integral(lambda s: 1.0) + e_s = integral(lambda s: s) / z + v_s = integral(lambda s: s**2) / z - e_s**2 + ex = lambda s: M * s**2 / (s**2 + A) + vx = lambda s: A * s**2 / (s**2 + A) + e_x = integral(ex) / z + v_x = integral(lambda s: vx(s) + ex(s) ** 2) / z - e_x**2 + return dict(log_z=np.log(z), e_s=e_s, v_s=v_s, e_x=e_x, v_x=v_x) + + +def _bits(proj, *variables): + return tuple( + float(getattr(proj[v], attr)) + for v in variables + for attr in ("mean", "variance") + ) + (float(proj.log_norm),) + + +def test__known_moments_match_one_dimensional_quadrature(): + fa, x_, s_ = make_scale_approx() + proj, status = LaplaceOptimiser(projection="moments").optimise(fa) + assert status.success, status.messages + assert status.flag is StatusFlag.SUCCESS + assert any(m.startswith("moments: n_nodes=") for m in status.messages) + + exact = exact_scale_moments() + assert proj[x_].mean == pytest.approx(exact["e_x"], abs=1e-6) + assert proj[x_].variance == pytest.approx(exact["v_x"], abs=1e-6) + # Truncated scale: the (E, Var)-as-parent convention of + # TruncatedNormalMessage.invert_sufficient_statistics + assert isinstance(proj[s_], TruncatedNormalMessage) + assert proj[s_].mean == pytest.approx(exact["e_s"], abs=1e-6) + assert proj[s_].variance == pytest.approx(exact["v_s"], abs=1e-6) + assert (proj[s_].lower_limit, proj[s_].upper_limit) == (0.0, 100.0) + # the tilted normalisation Z = ∫ f q_cavity + assert proj.log_norm == pytest.approx(exact["log_z"], abs=1e-6) + + +def make_gaussian_approx(): + """ + `test_laplace_hessian.make_approx`, with `x`'s cavity a truncated normal + whose limits sit hundreds of sigma out: the tilted density is Gaussian to + roundoff, but `x` has a bounded support and so takes the moments path. + """ + mu_, x_ = Variable("mu"), Variable("x") + sigma_f = 10.0 + + def f(mu, x): + return -0.5 * ((x - mu) / sigma_f) ** 2 + + def f_jac(mu, x): + d = (x - mu) / sigma_f**2 + return f(mu, x), (d, -d) + + factor = graph.Factor(f, mu_, x_, factor_jacobian=f_jac) + cavity = graph.MeanField( + { + mu_: NormalMessage(50.0, 10.0), + x_: TruncatedNormalMessage( + 55.0, 20.0, lower_limit=-1e4, upper_limit=1e4 + ), + } + ) + return graph.FactorApproximation(factor, cavity, cavity, cavity), mu_, x_ + + +def test__gaussian_tilted_moments_match_mode_and_covariance(): + P = np.array([[0.02, -0.01], [-0.01, 0.0125]]) + cov = np.linalg.inv(P) + mode = np.linalg.solve(P, np.diag([0.01, 0.0025]) @ np.array([50.0, 55.0])) + + fa, mu_, x_ = make_gaussian_approx() + proj, status = LaplaceOptimiser(projection="moments").optimise(fa) + assert status.flag is StatusFlag.SUCCESS, status.messages + assert proj[mu_].mean == pytest.approx(mode[0], abs=1e-6) + assert proj[x_].mean == pytest.approx(mode[1], abs=1e-6) + assert proj[mu_].variance == pytest.approx(cov[0, 0], abs=1e-6) + assert proj[x_].variance == pytest.approx(cov[1, 1], abs=1e-6) + + +def test__narrow_tilted_density_is_rewindowed(): + """ + A tilted density ~14x narrower than the cavity window (x: cavity std 20, + tilted std 1.4) is under-resolved by the first Gauss-Legendre pass; the + rule is re-windowed onto the tilted density until it is resolved. + """ + mu_, x_ = Variable("mu"), Variable("x") + + def f(mu, x): + return -0.5 * (x - mu) ** 2 + + def f_jac(mu, x): + return f(mu, x), (x - mu, mu - x) + + factor = graph.Factor(f, mu_, x_, factor_jacobian=f_jac) + cavity = graph.MeanField( + { + mu_: NormalMessage(50.0, 1.0), + x_: TruncatedNormalMessage(55.0, 20.0, lower_limit=-1e4, upper_limit=1e4), + } + ) + fa = graph.FactorApproximation(factor, cavity, cavity, cavity) + P = np.array([[2.0, -1.0], [-1.0, 1.0 + 1 / 400]]) + cov = np.linalg.inv(P) + mode = np.linalg.solve(P, np.array([50.0, 55.0 / 400])) + + proj, status = LaplaceOptimiser(projection="moments").optimise(fa) + assert status.flag is StatusFlag.SUCCESS, status.messages + assert "passes=1," not in status.messages[0] + assert proj[mu_].mean == pytest.approx(mode[0], abs=1e-6) + assert proj[x_].mean == pytest.approx(mode[1], abs=1e-6) + assert proj[mu_].variance == pytest.approx(cov[0, 0], abs=1e-6) + assert proj[x_].variance == pytest.approx(cov[1, 1], abs=1e-6) + + +def test__moments_rng_independent(): + results = [] + for seed in (0, 1, 12345): + fa, x_, s_ = make_scale_approx() + np.random.seed(seed) + proj, status = LaplaceOptimiser(projection="moments").optimise(fa) + assert status.flag is StatusFlag.SUCCESS + results.append(_bits(proj, x_, s_)) + assert results[0] == results[1] == results[2] + + +def test__moments_variable_id_independent(): + reference = None + for n_throwaway in (1, 3, 6): + _ = [Variable("k") for _ in range(n_throwaway)] + fa, x_, s_ = make_scale_approx() + np.random.seed(0) + proj, status = LaplaceOptimiser(projection="moments").optimise(fa) + assert status.flag is StatusFlag.SUCCESS + bits = _bits(proj, x_, s_) + if reference is None: + reference = bits + assert bits == reference + + +def _laplace_hessian_make_approx(): + from test_autofit.graphical.functionality.test_laplace_hessian import make_approx + + return make_approx() + + +def test__no_outer_variable_falls_back_to_mode_bit_for_bit(): + fa, mu_, x_ = _laplace_hessian_make_approx() + np.random.seed(0) + mode_proj, mode_status = LaplaceOptimiser().optimise(fa) + mode_bits = _bits(mode_proj, mu_, x_) + + fa, mu_, x_ = _laplace_hessian_make_approx() + np.random.seed(0) + proj, status = LaplaceOptimiser(projection="moments").optimise(fa) + + assert status.flag is mode_status.flag is StatusFlag.SUCCESS + assert status.messages == mode_status.messages + assert _bits(proj, mu_, x_) == mode_bits + + +def test__over_moment_max_size_falls_back_and_logs(caplog): + fa, x_, s_ = make_scale_approx() + np.random.seed(0) + mode_proj, mode_status = LaplaceOptimiser().optimise(fa) + mode_bits = _bits(mode_proj, x_, s_) if mode_status.success else None + + fa, x_, s_ = make_scale_approx() + np.random.seed(0) + with caplog.at_level(logging.INFO, logger="autofit.graphical.laplace"): + proj, status = LaplaceOptimiser( + projection="moments", moment_max_size=1 + ).optimise(fa) + + assert any("moment_max_size" in r.getMessage() for r in caplog.records) + assert status.flag is mode_status.flag + assert status.messages == mode_status.messages + if status.success: + assert _bits(proj, x_, s_) == mode_bits + else: + assert proj is fa.model_dist + + +def test__too_many_outer_variables_falls_back(): + fa, x_, s_ = make_scale_approx() + np.random.seed(0) + _, mode_status = LaplaceOptimiser().optimise(fa) + fa, x_, s_ = make_scale_approx() + np.random.seed(0) + _, status = LaplaceOptimiser(projection="moments", moment_max_outer=0).optimise( + fa + ) + assert status.messages == mode_status.messages + + +def test__bad_projection_value_raises(): + with pytest.raises(ValueError, match="projection"): + LaplaceOptimiser(projection="median") + + +def test__cavity_outside_support_is_bad_projection(): + # The cavity window C ± 8 sd lies wholly below the scale's lower limit 0 + fa, x_, s_ = make_scale_approx(c=-100.0, cv=1.0) + proj, status = LaplaceOptimiser(projection="moments").optimise(fa) + assert status.flag is StatusFlag.BAD_PROJECTION + assert not status.success + assert status.updated is False + assert proj is fa.model_dist + assert any("window" in m for m in status.messages) + + +# -------------------------------------------------------------------------- +# End to end: the hierarchical toy of hierarchical/test_truncated_support.py +# -------------------------------------------------------------------------- + + +def make_hierarchical_toy(truths, noise=2.0): + """ + Three groups with a `TruncatedGaussianPrior` scatter, as + `test_truncated_support.py`. Returns the scatter/mean priors, the graph and + each group's data. + """ + import autofit as af + from test_autofit.graphical.hierarchical.test_truncated_support import ( + GroupAnalysis, + Level, + ) + + rng = np.random.default_rng(0) + hf = af.HierarchicalFactor( + af.GaussianPrior, + mean=af.GaussianPrior(mean=50.0, sigma=10.0), + sigma=af.TruncatedGaussianPrior( + mean=10.0, sigma=5.0, lower_limit=0.0, upper_limit=100.0 + ), + ) + analysis_factors, ys = [], [] + for truth in truths: + model = af.Model(Level) + model.x = af.GaussianPrior(50.0, 20.0) + y = truth + noise * rng.standard_normal(20) + ys.append(y) + analysis_factors.append(af.AnalysisFactor(model, GroupAnalysis(y, noise))) + hf.add_drawn_variable(model.x) + return hf, af.FactorGraphModel(*analysis_factors, hf), ys + + +def exact_scatter_posterior(ys, noise=2.0): + """ + Mean and std of the toy's exact σ posterior: given σ the model is Gaussian + in (μ, x₁..x₃), so p(σ | y) ∝ TN(σ | 10, 5; 0, 100) · p(y | σ) with the + marginal likelihood in closed form, then 1-D quadrature over σ. + """ + n = len(ys) + + def log_marginal(s): + L = np.zeros((n + 1, n + 1)) + h = np.zeros(n + 1) + L[0, 0] += 1 / 10.0**2 + h[0] += 50.0 / 10.0**2 + for i, y in enumerate(ys): + L[i + 1, i + 1] += 1 / 20.0**2 + len(y) / noise**2 + 1 / s**2 + h[i + 1] += 50.0 / 20.0**2 + y.sum() / noise**2 + L[0, 0] += 1 / s**2 + L[0, i + 1] = L[i + 1, 0] = L[0, i + 1] - 1 / s**2 + return ( + 0.5 * h @ np.linalg.solve(L, h) + - 0.5 * np.linalg.slogdet(L)[1] + - n * np.log(s) + ) + + s = np.linspace(1e-4, 60.0, 30001) + lp = np.array([log_marginal(v) for v in s]) + stats.norm.logpdf(s, 10.0, 5.0) + w = np.exp(lp - lp.max()) + w /= w.sum() + mean = float(np.sum(w * s)) + return mean, float(np.sqrt(np.sum(w * (s - mean) ** 2))) + + +def run_manual_ep(graph_model, optimiser, n_sweeps): + """ + The manual factor_approximation -> optimise -> project_mean_field loop of + `test_truncated_support.py`. Returns the final mean field and, per + hierarchical-factor update, the flag after projection. + """ + approx = graph_model.mean_field_approximation() + flags = [] + for _ in range(n_sweeps): + for factor in graph_model.graph.factors: + factor_approx = approx.factor_approximation(factor) + new_dist, status = optimiser.optimise(factor_approx) + approx, status = approx.project_mean_field( + new_dist, factor_approx, status=status + ) + if hasattr(factor, "scale_variables"): + flags.append(status.flag) + return approx.mean_field, flags + + +def test__hierarchical_scale_variables(): + hf, graph_model, _ = make_hierarchical_toy((45.0, 52.0, 58.0)) + hierarchical = [f for f in graph_model.graph.factors if hasattr(f, "scale_variables")] + assert len(hierarchical) == 3 + for factor in hierarchical: + assert factor.scale_variables == frozenset({hf.sigma}) + + +def test__hierarchical_scatter_moments_match_exact_posterior(): + np.random.seed(0) + hf, graph_model, ys = make_hierarchical_toy((45.0, 52.0, 58.0)) + mean_field, flags = run_manual_ep( + graph_model, LaplaceOptimiser(projection="moments"), n_sweeps=2 + ) + assert flags.count(StatusFlag.SUCCESS) >= 1, flags + assert StatusFlag.EXCEPTION not in flags + + scatter = mean_field[hf.sigma] + assert isinstance(scatter, TruncatedNormalMessage) + assert (scatter.lower_limit, scatter.upper_limit) == (0.0, 100.0) + + exact_mean, exact_std = exact_scatter_posterior(ys) + assert abs(scatter.mean - exact_mean) < 0.5 * exact_std, ( + scatter.mean, + exact_mean, + exact_std, + ) + + +def test__near_zero_scatter_updates_by_moments_not_by_mode(): + """ + Groups consistent with no scatter: the tilted density in σ piles up on + σ = 0, so the mode path never lands an update of a hierarchical factor + (each projection is rejected) while the moments path does. + """ + truths = (50.0, 50.5, 49.5) + + np.random.seed(0) + hf, graph_model, _ = make_hierarchical_toy(truths) + mode_mean_field, mode_flags = run_manual_ep( + graph_model, LaplaceOptimiser(), n_sweeps=2 + ) + assert StatusFlag.SUCCESS not in mode_flags, mode_flags + + np.random.seed(0) + hf, graph_model, _ = make_hierarchical_toy(truths) + mean_field, flags = run_manual_ep( + graph_model, LaplaceOptimiser(projection="moments"), n_sweeps=1 + ) + assert flags == [StatusFlag.SUCCESS] * 3, flags + # the scatter has moved off its prior (10 ± 5) towards 0 + assert mean_field[hf.sigma].mean < 9.0 + assert (mean_field[hf.sigma].lower_limit, mean_field[hf.sigma].upper_limit) == ( + 0.0, + 100.0, + )