Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
60 changes: 58 additions & 2 deletions autofit/graphical/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]`:
Expand Down Expand Up @@ -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
Expand Down
21 changes: 21 additions & 0 deletions autofit/graphical/declarative/factor/hierarchical.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down
3 changes: 3 additions & 0 deletions autofit/graphical/expectation_propagation/diagnostics.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"})


Expand Down
Loading
Loading