Skip to content

refactor(variational)!: remove vestigial Natural/Expectation families [stack 4/10] - #715

Merged
thomaspinder merged 23 commits into
v1.0from
natgrads-02-remove-vestigial
Aug 7, 2026
Merged

thomaspinder merged 23 commits into
v1.0from
natgrads-02-remove-vestigial

Conversation

@thomaspinder

@thomaspinder thomaspinder commented Jul 27, 2026 •

Copy link
Copy Markdown
Collaborator

Checklist

  • I've formatted the new code by running uv run poe format before committing.
  • I've added tests for new code.
  • I've added docstrings for the new code.

Description

PR 4/10 of the v1.0 stack (natural-gradients 2 of 5), stacked on #714 — merge #743, #744 and #714 before this one.

Removes NaturalVariationalGaussian and ExpectationVariationalGaussian (−335 lines in variational_families.py, plus the now-dead _psd helper and an unused import).

Why: these were parameterization-only classes with no optimiser attached — vestiges of the pre-rewrite natural-gradients code. With #714, natural gradients are computed for the standard VariationalGaussian/WhitenedVariationalGaussian directly via the θ↔η exponential-family maps (the natgrad update is parameterization-invariant), so a dedicated stored-θ class is redundant and its reverse-index Cholesky inversion trick was the numerically weakest code in the module.

Also in this PR:

  • Test migration in tests/test_variational_families.py (24 → 12 parametrized cases, new absence-guard test, cholesky-count monkeypatch fix) and tests/test_benchmarks_smoke.py.
  • benchmarks/objectives.py pruned to the surviving families (this was a hard CI blocker otherwise).
  • CHANGELOG.md: ## [Unreleased] with ### Added (fit_natgrads, from feat(natural-gradients): fit_natgrads + exponential-family machinery [stack 3/10] #714) and ### Removed entries including a migration note.
  • CLAUDE.md family list + fit docs updated.

Full poe all-tests green at this stack point (2,721 passed; the single local failure is a stray untracked working-tree file, absent in CI). Exhaustive grep confirms the only remaining references to the deleted classes are the CHANGELOG entry and the absence-guard test.

Issue Number: N/A

🤖 Generated with Claude Code

https://claude.ai/code/session_01VHxYz2P7JCSqpax5RmDwWD

thomaspinder and others added 13 commits August 6, 2026 01:27
conjugate_mll had no value-level test anywhere in the suite, yet serves as
the oracle for the Kalman MLL and collapsed_elbo. These closed-form pins,
computed through an independent jnp.linalg path, give the reference frame
its ground truth ahead of the v1.0 conditioning refactor.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
…tter bug

collapsed_elbo(z=X) vs conjugate_mll and whitened-vs-unwhitened predicts at
matched parameters now guard the five independent derivations of the
conjugate conditioning algebra. The strict xfail documents that at
non-default jitter the derivations factorise different matrices — the bug
the v1.0 conditioning module removes.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
_compare previously swallowed AssertionError with a print, so the harness
could never fail. Failures are now collected per-example and raised at the
end of test(), making the four golden-value pins a real no-behaviour-change
net for the v1.0 refactor.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
The newly-loud harness exposed pre-existing drift in all four examples: the
collapsed/uncollapsed goldens predated the real-data example swap (#696) and
the regression/heteroscedastic goldens predated subsequent behaviour fixes
(#707/#708/#713 and dependency bumps). The toothless harness never noticed.
Re-pinned so the net measures the v1.0 refactor, not history.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
The full-dataset size a minibatch ELBO needs now travels on the one object
that knows it, as static pytree aux_data, instead of being smuggled through
likelihood constructors.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
…nd noise_prior

Of ~245 occurrences of num_datapoints, only four were real reads: the ELBO
minibatch scale (now served by Dataset.n_total) and latent sizing (moves to
data-contact time in the JointModel rewrite). No likelihood used the value
internally, and nothing validated it — a wrong value silently mis-scaled
the ELBO. noise_prior moves to the model layer, where priors live.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
The API now mirrors the maths: prior * likelihood -> JointModel (the joint
p(f,y), the trainable object); model.condition(D) — sugar: model | D —
returns an immutable Posterior pytree caching the Cholesky factor and
representer weights. The predictive, log_marginal_likelihood, loo, and
pathwise sample_approx are views of that one factorisation, deleting the
eleven independent derivations and the two-owner jitter split
(prior.jitter is now the single knob, applied once inside conditioning).

- gpjax/conditioning.py: deep module (Posterior, ExactPosterior,
  LatentPosterior); MO validation moves to condition time; sample_approx
  refuses multi-output loudly instead of silently broadcasting wrong.
- gps.py: Prior (AbstractPrior folded in), ConjugateModel,
  NonConjugateModel (lazy latent, sized at data contact),
  HeteroscedasticModel (owns noise_prior — likelihoods are pure
  conditionals again, killing the likelihoods->gps circular import).
  Deleted: AbstractPrior, AbstractPosterior, LatentPosterior marker,
  ChainedPosterior marker, construct_posterior (now construct_model).
- objectives: conjugate_mll/conjugate_loocv/log_posterior_density are
  one-line views of the conditioned posterior.
- fit: _prepare_model hook sizes lazily-initialised state from data.
- predict(t, D) survives as documented one-line sugar everywhere.
- return_covariance_type kwarg renamed to covariance.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
Mechanical: ConjugatePosterior->ConjugateModel and friends,
construct_posterior->construct_model, return_covariance_type->covariance,
num_datapoints/noise_prior constructor ceremony deleted (~240 sites).
Semantic: heteroscedastic tests build HeteroscedasticModel directly;
non-conjugate tests size the latent via init_latent; docs/index.md
quickstart shows condition(); regression example narrates the
condition API; StateSpaceConjugatePosterior renamed
StateSpaceConjugateModel.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
…ault jitter

The model-side two-owner jitter bug is fixed; the strict xfail narrows to
the family-side knob, which unifies in the variational stack PR.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
@thomaspinder
thomaspinder force-pushed the natgrads-02-remove-vestigial branch from 095b0a7 to d083652 Compare August 6, 2026 09:02
thomaspinder and others added 8 commits August 6, 2026 12:15
- reference/gps.md lists the JointModel hierarchy and the conditioning
  module; state_space.md and linalg.md updated for renamed/new symbols
- stale glossary/sharp_bits/classification xrefs renamed
- poisson example initialises the lazy non-conjugate latent before MCMC
- ADR directory excluded from the docs site (in-repo records for now)
- codeautolink match_block warnings suppressed on every path: a matcher
  limitation on doctest-SKIP blocks, predating this stack — the docs
  workflow had not run cold since the Sphinx migration, so tonight's PR
  pushes surfaced it for the first time

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
…inery

Implements Salimbeni et al. 2018 (arXiv:1803.09151) natural-gradient VI for
VariationalGaussian and WhitenedVariationalGaussian.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01VHxYz2P7JCSqpax5RmDwWD
Adversarial review of the natural-gradient core turned up two correctness
defects, one performance defect, three contract mismatches and a set of
documentation and test gaps. Fixes, in order of severity:

* `_first_valid_trial` leaked float64 into the scan carry. Under x64 the
  exponent `jnp.arange(K + 1)` is int64, so `backoff ** arange` was a non-weak
  float64 that promoted a float32 model and made `lax.scan` reject the carry.
  The trial ladder is now cast to the dtype of Theta_2, so `fit_natgrads` is no
  longer strictly narrower than `fit`.

* The backoff replicated the whole theta -> xi map across all K+1 trials, at a
  measured 13% of total training wall clock at M=200 -- not the "negligible"
  cost impl-plan 1.2.5 assumed. Only the admissibility probe is replicated now;
  the inversion, the X^T X product and the second Cholesky run once, at the
  accepted step size. Measured overhead at M=200 falls to 4.6%.

* `fit_natgrads`' signature rejected everything `_check_natgrad_lr` blessed:
  `natgrad_lr=1`, `map_jitter=0`, `backoff=1` and a 0-d array all raised under
  the beartype import hook. Annotations widened, the validator now accepts 0-d
  arrays and rejects bool, and the entry point (not just the validator) is
  tested with each.

* `_reject_frozen_coordinates` matched coordinates against top-level dataclass
  fields by identity, so a future nested registration would have passed the
  guard silently. It now re-walks the tree with the selector as the `is_leaf`
  predicate and reports the full key path. Its message also pluralises and
  points at a remedy that exists -- the old one recommended freezing the whole
  family, which re-raises the same error.

* `fit_natgrads` now calls the guard under `safe=True` as impl-plan 1.2.4
  prescribes, and forwards `log_rate` to `vscan` instead of documenting a knob
  that did nothing.

Tests: `test_natgrad_backoff_recovers_from_large_step` never exercised the
backoff (the k=0 trial was already admissible at natgrad_lr=100), so it now
starts from a tight S_0, asserts the un-shrunk step genuinely leaves the cone,
and checks the accepted step size. The Cholesky-budget test could not see the
vmap width; it is now an exact count parametrised over max_backoff, plus a
lowered-IR test asserting the batched Cholesky has leading dimension K+1. Added
dtype-preservation tests, and moved the duplicated conjugate oracle into
`tests/_reference/conjugate_svgp.py` so the two transcriptions cannot drift.

Docs: all seven exported functions gained runnable `Example:` blocks (impl-plan
section 6), the map_jitter bias on `history` and the eta -> xi cancellation
regime are now documented on the public surface, and two docstrings became raw
strings.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01VHxYz2P7JCSqpax5RmDwWD
…_natgrads

The ```pycon fences render fine under MkDocs but are invalid RST under
Sphinx/napoleon, producing docutils warnings that fail the -W docs-ci
gate. Drop the fences; the indented doctest block matches fit()'s style.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
…tionVariationalGaussian

Both classes were parameterisation-only: they stored the natural or
expectation coordinates of q(u) but shipped no way to take a
natural-gradient step in them, so they bought nothing over the standard
families.

Natural-gradient geometry belongs to the optimiser, not the family. The
Fisher matrix is exactly the Jacobian dn/dt, so the natural gradient with
respect to the natural parameters equals the ordinary gradient with
respect to the expectation parameters, in any parameterisation. fit_natgrads
(PR#1, gpjax/natural_gradients.py) therefore computes the transforms
on the fly and operates directly on VariationalGaussian and
WhitenedVariationalGaussian, which store constraint-respecting
coordinates. Users of the removed classes should switch to
VariationalGaussian with gpjax.fit_natgrads.

Also drops the now-dead _psd helper and the cholesky_factor import, whose
only call sites lived inside the deleted classes, and the "natural" and
"expectation" arms of the VariationalParametrisationSuite ASV benchmark.

BREAKING CHANGE: NaturalVariationalGaussian and ExpectationVariationalGaussian
are removed from gpjax.variational_families.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01VHxYz2P7JCSqpax5RmDwWD
CHANGELOG: add the missing "### Added" entry for fit_natgrads and
gpjax.natural_gradients. PR#1 shipped both without a changelog entry, so the
Removed entry added here forward-referenced an API the changelog never
announced. Bringing it forward from PR#3 keeps any release cut mid-stack
self-consistent; PR#3 appends the dual entries to the same section.

CHANGELOG: correct the justification prose. "the natural gradient with respect
to theta equals the ordinary gradient with respect to eta, in any
parameterisation" is false as literally written -- for a reparameterisation xi
with J = dtheta/dxi, the natural gradient in xi is J^-1 grad_eta L, not
grad_eta L. The identity is specific to the natural/expectation pair of an
exponential family. Reworded to state that pairing and the point it supports:
either coordinate system is recoverable on the fly, so no dedicated class is
needed.

benchmarks: drop diff-relative wording from the
VariationalParametrisationSuite docstring. "surviving" only means something to
someone reading this commit's diff, and "Both" is a count PR#3 invalidates when
it re-adds the dual arm. The module docstring keeps its explicit
"(standard, whitened)" list, which PR#3 must extend regardless.

tests: split the _psd guard out of test_removed_families_are_gone. The _psd arm
asserted `"_psd" not in __all__`, vacuous for a helper that was never exported,
under a failure message about superseded parameterisations. It is now its own
test with a docstring saying what it actually guards.

CLAUDE.md: "Three optimisers" -> four. fit_natgrads landed in gpjax/fit.py in
PR#1; the sentence has been stale since.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01VHxYz2P7JCSqpax5RmDwWD
The Sphinx docs arrived with v1.0, after this branch was cut, so the
deletion of NaturalVariationalGaussian and ExpectationVariationalGaussian
now has to reach three doc files the original commits could not know
about: drop the two classes from the variational-families autosummary
page, and repoint the glossary's "natural parameters" entry from the
removed classes to fit_natgrads on the surviving families. fit_natgrads
is added to the fit reference page so that glossary link resolves —
an omission from PR#1, which shipped the function without a reference
entry.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_019d7TF7oQt2Du4EQ74bMBQB
@thomaspinder
thomaspinder force-pushed the natgrads-02-remove-vestigial branch from d083652 to 21c4ec0 Compare August 6, 2026 10:41
@github-actions

github-actions Bot commented Aug 6, 2026

Copy link
Copy Markdown

📖 Docs preview: https://pr-715--endearing-crepe-c2d5fe.netlify.app

Smoke render — the expensive notebooks run with reduced budgets, so
figures are not publication fidelity. /render-mode.txt says smoke.

`Dependency Vulnerability Scan` and `Secrets Detection` are required contexts
on `main`, but the workflow filtered `pull_request` to `branches: [main]`. A
stacked PR targets its predecessor rather than `main`, so neither job ever
ran on one: PRs in the current v1 stack report 16 check runs where a
main-targeting PR reports 18.

The effect is that a stack gets reviewed, approved and declared green with no
security gate at any point — the scans first fire when each PR auto-retargets
to `main` at merge time, which is the worst moment to discover a verified
secret. Of the 45 commits in the v1 stack only the 5 on the bottom branch had
ever been scanned.

Drop the filter so the trigger follows the code rather than the target branch.
Both jobs are already branch-agnostic — TruffleHog scans `./` over full
history and the dependency scan reads the checkout — so neither needs a base
ref to work off `main`.

`push` stays pinned to `main`: post-merge scanning of every branch push would
duplicate the PR run for no extra signal.

(cherry picked from commit a8d6603)
…ames

Unfiltering the `pull_request` trigger made `Secrets Detection` run on the v1
stack for the first time, and it immediately failed three PRs (#714, #716,
#748) on four "verified" secrets. All four are pytest function names:

  test_fit_natgrads_accepts_optax_schedule
  test_dual_natgrad_matches_salimbeni_step
  test_dual_working_matrices_reconstruct_r
  test_prior_kl_matches_textbook_reference

TruffleHog's Lob key pattern is `(live|test)_` plus 35 more characters, and
each of those names is exactly 40 characters long. Worse, the Lob verifier
reports them as *verified*, so `--only-verified` does not filter them: a fresh
repository containing nothing but one of those names reproduces
`Lob verified=true` under trufflehog 3.96.0, the same version the action pins.

`Secrets Detection` is a required context on `main`, so this is not cosmetic —
it means an arbitrary subset of PRs, selected by test-name length, can never
merge. GPJax has no Lob (direct-mail API) integration, so excluding the
detector costs no coverage. Verified that the flag is parsed rather than
silently ignored: an unrecognised detector name makes trufflehog exit with
"unrecognized detector type", and the four findings disappear with the flag
and return without it.

(cherry picked from commit 83ad296)
@thomaspinder thomaspinder changed the title refactor(variational)!: remove vestigial Natural/Expectation families [natgrads stack 2/5] refactor(variational)!: remove vestigial Natural/Expectation families [stack 4/10] Aug 7, 2026
@thomaspinder
thomaspinder changed the base branch from natgrads-01-core to v1.0 August 7, 2026 13:38
@thomaspinder
thomaspinder merged commit 4e48381 into v1.0 Aug 7, 2026
26 of 34 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant