Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
20 commits
Select commit Hold shift + click to select a range
4e8f8d3
test(oracle): pin conjugate MLL, predict, and LOOCV to closed form
thomaspinder Aug 5, 2026
0ff8f2d
test(equivalence): pin cross-derivation agreement; xfail two-owner ji…
thomaspinder Aug 5, 2026
d195be8
test(integration): fail loudly when golden values drift
thomaspinder Aug 5, 2026
5d96ee2
test(integration): re-pin golden values to current-main behaviour
thomaspinder Aug 5, 2026
880c58f
feat(linalg): stabilised_cholesky — single stabilise-and-factor seam
thomaspinder Aug 5, 2026
5814aa3
feat(dataset): static n_total metadata; get_batch stamps full size
thomaspinder Aug 5, 2026
8c7ba10
feat(likelihoods)!: pure conditional families — drop num_datapoints a…
thomaspinder Aug 5, 2026
9b9a43c
feat!: v1.0 conditioning architecture — JointModel, Posterior, condition
thomaspinder Aug 5, 2026
b5c9884
refactor!: rename sweep — tests, benchmarks, examples onto the v1.0 API
thomaspinder Aug 6, 2026
1ce4ca1
style: ruff format after the rename sweep
thomaspinder Aug 6, 2026
f113932
docs: v1.0 migration guide, CONTEXT.md glossary, ADR-0001
thomaspinder Aug 6, 2026
bb44a85
docs: self-contained migration snippets
thomaspinder Aug 6, 2026
78debf3
test(equivalence): pin predict/MLL same-matrix consistency at non-def…
thomaspinder Aug 6, 2026
72a7212
docs: fix Sphinx build for the v1.0 API
thomaspinder Aug 6, 2026
84a133e
feat(natural-gradients): add fit_natgrads and exponential-family mach…
thomaspinder Jul 26, 2026
12fa290
fix(natural-gradients): address review findings on PR#1
thomaspinder Jul 26, 2026
bc74e69
refactor(natgrads): adapt to the v1.0 conditioning API
thomaspinder Aug 6, 2026
d53c808
docs(fit): replace Markdown code fences with RST doctest block in fit…
thomaspinder Aug 6, 2026
f270f9c
ci(security): scan every PR, not only those targeting main
thomaspinder Aug 7, 2026
58e8845
ci(security): exclude TruffleHog's Lob detector, which flags pytest n…
thomaspinder Aug 7, 2026
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
24 changes: 22 additions & 2 deletions .github/workflows/security-analysis.yml
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,14 @@ name: "Security Analysis"
on:
push:
branches: [ "main" ]
# Deliberately unfiltered: the two jobs below are required contexts on `main`,
# and a `branches: [main]` filter here meant they only ever reported on PRs
# that already targeted `main`. Stacked PRs — which target their predecessor,
# not `main` — produced 16 checks instead of 18, so the whole of a stack got
# reviewed and approved without ever being secret- or dependency-scanned; the
# gate first fired at retarget time, after review. Scan the code, not the
# target branch.
pull_request:
branches: [ "main" ]
schedule:
- cron: '30 4 * * 1' # Weekly on Mondays at 4:30 AM UTC

Expand Down Expand Up @@ -75,4 +81,18 @@ jobs:
uses: trufflesecurity/trufflehog@main
with:
path: ./
extra_args: --debug --only-verified
# `--exclude-detectors=lob`: TruffleHog's Lob key pattern is
# `(live|test)_` followed by 35 more characters, which matches any
# pytest function name that happens to be exactly 40 characters long —
# and its verifier then reports those as *verified*, so `--only-verified`
# does not filter them out. Four real examples in this repo:
# test_fit_natgrads_accepts_optax_schedule,
# test_dual_natgrad_matches_salimbeni_step,
# test_dual_working_matrices_reconstruct_r,
# test_prior_kl_matches_textbook_reference.
# Reproduced standalone: a fresh repo containing only one of those names
# yields `Lob verified=true`. Since `Secrets Detection` is a required
# context on `main`, leaving this on means an arbitrary, length-dependent
# subset of PRs can never merge. GPJax has no Lob (direct-mail API)
# integration, so excluding the detector costs no real coverage.
extra_args: --debug --only-verified --exclude-detectors=lob
62 changes: 62 additions & 0 deletions CONTEXT.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
# GPJax domain glossary

The ubiquitous language of this codebase. Code, docs, tests, and reviews use
these terms exactly; when a concept is missing, add it here in the same PR
that introduces it.

## Design principle

**Maths-first, comprehensible to the maths-familiar non-expert.** API names
follow the textbook equation unless the textbook word is jargon a
practitioner would not own. When the two conflict, choose the word a
scikit-learn/PyMC user already speaks and state the maths once in the
docstring (e.g. the class is `JointModel`, and its docstring says "the joint
distribution p(f, y)").

## Terms

**Prior** — the Gaussian process prior p(f), pairing a kernel with a mean
function. Queryable at any inputs: `prior(x)` returns the prior predictive.
Owns the model's single numerical-stabilisation knob, `jitter`.

**Likelihood** — the conditional distribution p(y | f). A *pure* conditional:
it holds no priors, no dataset facts, and no data sizes. (Both were tried and
removed at v1.0 — see ADR-0001.)

**JointModel** — the joint distribution p(f, y) = p(y | f) · p(f), created by
`prior * likelihood`. The *trainable* object: `gpx.fit` optimises its
hyperparameters. It carries state that must exist before conditioning
(hyperparameters; the non-conjugate latent; the heteroscedastic noise-process
model) and no derived quantities. Concrete kinds: `ConjugateModel`,
`NonConjugateModel`, `HeteroscedasticModel` — each holds state the others
cannot.

**condition** — the operation p(f, y) + 𝒟 → p(f | 𝒟), spelled
`model.condition(D)` or the operator form `model | D` (read: "f given D").
Signatures follow the maths: models take data; already-fit variational
families will condition without arguments (post-natgrads stack PR).

**Posterior** — the conditioned process p(f | 𝒟) returned by `condition`. An
*immutable* pytree: the training-covariance factorisation is computed once and
cached; every query is a view of it. The uniform query surface:
`posterior(xtest, covariance="dense"|"diagonal")`, plus per-mode views —
`log_marginal_likelihood` / `loo` / `sample_approx` on the exact mode,
`log_posterior_density` on the latent mode. Users never name the concrete
implementations behind the interface.

**evidence / log marginal likelihood** — p(𝒟), the normalising constant of
conditioning, exposed as `posterior.log_marginal_likelihood`. "Evidence" and
"marginal likelihood" are the same quantity; the attribute uses the
GP-community's term.

**sugar** — a documented one-line composition kept for ergonomics, never a
second implementation. `model.predict(x, D)` and `model(x, D)` are sugar for
`model.condition(D)(x)`; `model | D` is sugar for `model.condition(D)`.

**objective** — a scalar function `(model, Dataset) -> ScalarFloat` consumed
by `gpx.fit`. Objectives are thin: `conjugate_mll` is the evidence view of
the conditioned posterior, not a second derivation.

**Dataset** — the data container. `n_total` records the full-dataset size
when the object is a minibatch view (stamped by `get_batch`); the minibatch
ELBO scale is derived from it, never supplied by hand.
2 changes: 1 addition & 1 deletion benchmarks/compile.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,7 @@ def setup(self):
kernel = gpx.kernels.RBF()
mean = gpx.mean_functions.Zero()
prior = gpx.gps.Prior(kernel=kernel, mean_function=mean)
likelihood = gpx.likelihoods.Gaussian(num_datapoints=n)
likelihood = gpx.likelihoods.Gaussian()
self.posterior = prior * likelihood
self.q = VariationalGaussian(
posterior=self.posterior, inducing_inputs=X[:M_INDUCING]
Expand Down
4 changes: 2 additions & 2 deletions benchmarks/objectives.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,7 @@ def _conjugate_posterior(n: int):
kernel = gpx.kernels.RBF()
mean = gpx.mean_functions.Zero()
prior = gpx.gps.Prior(kernel=kernel, mean_function=mean)
likelihood = gpx.likelihoods.Gaussian(num_datapoints=n)
likelihood = gpx.likelihoods.Gaussian()
return prior * likelihood, data


Expand Down Expand Up @@ -157,7 +157,7 @@ def setup(self, n):
noise_prior = gpx.gps.Prior(
kernel=noise_kernel, mean_function=gpx.mean_functions.Zero()
)
likelihood = HeteroscedasticGaussian(num_datapoints=n, noise_prior=noise_prior)
likelihood = HeteroscedasticGaussian(noise_prior=noise_prior)
posterior = signal_prior * likelihood
Z = data.X[:M_INDUCING]
self.q = HeteroscedasticVariationalFamily(
Expand Down
8 changes: 3 additions & 5 deletions benchmarks/state_space.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@

import gpjax as gpx
from gpjax.state_space import StateSpacePrior, state_space_mll
from gpjax.state_space.gps import StateSpaceConjugatePosterior
from gpjax.state_space.gps import StateSpaceConjugateModel
import jax.numpy as jnp
import jax.random as jr

Expand All @@ -35,10 +35,8 @@ def setup(self, n):
mean_function=gpx.mean_functions.Zero(),
kernel=gpx.kernels.Matern32(lengthscale=1.0, variance=1.0),
)
likelihood = gpx.likelihoods.Gaussian(num_datapoints=n, obs_stddev=0.1)
self.posterior = StateSpaceConjugatePosterior(
prior=prior, likelihood=likelihood
)
likelihood = gpx.likelihoods.Gaussian(obs_stddev=0.1)
self.posterior = StateSpaceConjugateModel(prior=prior, likelihood=likelihood)
self.data = _temporal_dataset(n)
realise(state_space_mll(self.posterior, self.data))

Expand Down
88 changes: 88 additions & 0 deletions docs/adr/0001-conditioning-architecture.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
# ADR-0001: The v1.0 conditioning architecture

- **Status:** accepted
- **Date:** 2026-08-06
- **Deciders:** Thomas Pinder (design settled in an architecture-review +
grilling session; full decision tree recorded there)

## Context

An architecture review (2026-08-06) found that no module owned "condition a
GP on data": the derivation *stabilise → factor → solve → predictive
moments* was written out at eleven call sites across `gps.py`,
`objectives.py`, `variational_families.py`, and `models/oilmm.py`. Fixes
reached some copies and missed others (the diagonal-predict fix landed in two
of eleven), jitter had two owners (`Prior.jitter` vs `Posterior.jitter`), so
`predict` and `conjugate_mll` could factorise *different matrices* for the
same model, and `conjugate_mll` — the oracle for the Kalman MLL and
`collapsed_elbo` — had no value-level test of its own.

## Decision

One deep conditioning module, with the public API as its veneer, rolled out
universally at v1.0:

- `prior * likelihood` returns a **`JointModel`** — the joint p(f, y), the
trainable object (`ConjugateModel`, `NonConjugateModel` with a lazily-sized
latent, `HeteroscedasticModel` owning the noise-process prior).
- `model.condition(D)` (operator sugar `model | D`) returns a **`Posterior`**:
an immutable pytree caching the factorisation, with the predictive,
`log_marginal_likelihood`, `loo`, and pathwise `sample_approx` as views of
it. One abstract interface; per-mode internal implementations.
- `prior.jitter` is the single stabilisation knob, applied exactly once,
inside conditioning, through `linalg.stabilised_cholesky` (the seed of the
linalg deepening).
- `predict(x, D)` survives as documented one-line sugar. Objectives become
one-line views. `return_covariance_type` is renamed `covariance`.
- Likelihoods are pure conditionals: `num_datapoints` deleted (`Dataset`
carries a static `n_total`, stamped by `get_batch`, from which the
minibatch ELBO scale is derived) and `noise_prior` moved to
`HeteroscedasticModel` (removing the `likelihoods -> gps` circular import).
- Deleted as failed deletion-tests: `AbstractPrior` (one-adapter seam),
`AbstractPosterior` (splits into JointModel/Posterior), the
`LatentPosterior` and `ChainedPosterior` markers.
- Universality is a property of the **v1.0 release**, assembled from stacked
PRs: safety net → core conditioning → variational universalisation. The
variational PR is sequenced **after** the natural-gradients stack
(#714–#730) merges, since that stack rewrites `variational_families.py`;
until then the families keep their `posterior` field name and their own
predict derivations (renamed mechanically only).

The safety net landed first, in its own PR: closed-form oracles for the MLL,
predict, and LOOCV; cross-derivation equivalence pins; and an integration
harness whose failures actually raise.

## Rejected alternatives

- **Dropping the JointModel object** ("just `prior.condition(data,
likelihood)`"): the trainable/derived split is load-bearing. `fit` needs a
named pytree of everything trainable; the Posterior carries derived cache
that a gradient step must never touch. Any container invented to make
`fit` ergonomic *is* the JointModel under another name.
- **Staged or partial public rollout**: a patchy API (some objects
conditioning, others not) was ruled out; stacking PRs into one v1.0 release
achieves universality without a big-bang branch.
- **Prior inside the likelihood**: the codebase already ran this experiment —
`HeteroscedasticGaussian.noise_prior` produced circular ownership across
three modules. Priors live in models; likelihoods stay conditional
families.
- **Keeping `num_datapoints`**: only four real reads existed; nothing
validated the value, so a wrong one silently mis-scaled the ELBO; no major
GP library puts dataset size on likelihoods (research:
`plans/2026-08-06-drop-num-datapoints-research.md`).
- **Naming the joint `Model`** (rejected in favour of `JointModel`): the
prefix self-documents the maths and answers "why can't I predict with this
yet?"; "model" remains in the name for practitioners.

## Consequences

- Locality: one home for the algebra; the two-owner jitter bug is
structurally impossible; `sample_approx` reuses the same factor as
`predict` and refuses multi-output loudly instead of broadcasting wrongly.
- The evidence is cached with the factorisation, so repeated
predict-then-score workflows stop re-factorising.
- Breaking changes at v1.0 are recorded in `docs/migration.md`; the
vocabulary lives in `CONTEXT.md`.
- Follow-ups tracked for the stack: variational universalisation
(post-natgrads), the linalg structure-preserving deepening, one training
loop with stepper adapters, and the compute-engine seam.
9 changes: 8 additions & 1 deletion docs/conf.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,6 +72,7 @@

exclude_patterns = [
"_build",
"adr/*", # ADRs are in-repo records, not (yet) part of the docs site
"Thumbs.db",
".DS_Store",
"conf.py", # this config module is not a document
Expand Down Expand Up @@ -181,8 +182,14 @@
# Everything else (broken xrefs, bad anchors, malformed directives) stays fatal
# on deploy, which it was not when that build ran without `-W` at all.
# The PR gate suppresses nothing.
# codeautolink cannot match doctest blocks carrying `# doctest: +SKIP` markers
# against their rendered HTML (a matcher limitation, not a doc defect — xdoctest
# validates the examples). Suppressed on every path; predates the v1.0 stack but
# first surfaced when this workflow ran cold post-Sphinx-migration.
suppress_warnings = ["codeautolink.match_block"]

if os.environ.get("GPJAX_DOCS_RESILIENT") == "1":
suppress_warnings = [
suppress_warnings = suppress_warnings + [
"mystnb.exec", # execution failure + "traceback saved in:" follow-up
"mystnb.glue", # a glue key that never got produced by a failed notebook
# A notebook that fails to execute renders with no outputs, and MyST-NB
Expand Down
2 changes: 1 addition & 1 deletion docs/examples/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -157,7 +157,7 @@ def glue(*args, **kwargs):

prior = gpx.gps.Prior(mean_function=meanf, kernel=kernel)

likelihood = gpx.likelihoods.Gaussian(100)
likelihood = gpx.likelihoods.Gaussian()
posterior = likelihood * prior
print(posterior)

Expand Down
4 changes: 2 additions & 2 deletions docs/examples/barycentres.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,7 +132,7 @@ def glue(*args, **kwargs):
# ## Dataset
#
# We'll simulate five datasets and develop a Gaussian process
# [posterior](#gpjax.gps.ConjugatePosterior) before
# [posterior](#gpjax.gps.ConjugateModel) before
# identifying the Gaussian process barycentre at a set of test points. Each dataset
# will be a sine function with a different vertical shift, periodicity, and quantity
# of noise.
Expand Down Expand Up @@ -184,7 +184,7 @@ def fit_gp(x: jax.Array, y: jax.Array) -> gpx.distributions.GaussianDistribution
y = y.reshape(-1, 1)
D = gpx.Dataset(X=x, y=y)

likelihood = gpx.likelihoods.Gaussian(num_datapoints=n)
likelihood = gpx.likelihoods.Gaussian()
posterior = (
gpx.gps.Prior(
mean_function=gpx.mean_functions.Constant(), kernel=gpx.kernels.RBF()
Expand Down
4 changes: 2 additions & 2 deletions docs/examples/classification.py
Original file line number Diff line number Diff line change
Expand Up @@ -105,7 +105,7 @@
kernel = gpx.kernels.RBF()
meanf = gpx.mean_functions.Constant()
prior = gpx.gps.Prior(mean_function=meanf, kernel=kernel)
likelihood = gpx.likelihoods.Bernoulli(num_datapoints=D.n)
likelihood = gpx.likelihoods.Bernoulli()

# %% [markdown]
# We construct the posterior through the product of our prior and likelihood.
Expand Down Expand Up @@ -144,7 +144,7 @@
)

# %% [markdown]
# From which we can [make predictions](#gpjax.gps.NonConjugatePosterior.predict) at
# From which we can [make predictions](#gpjax.gps.JointModel.predict) at
# novel inputs, as illustrated in {numref}`fig-classification-map-predictive`.

# %% mystnb={"figure": {"caption": "The MAP predictive mean and its one-sigma band over the binary observations.", "name": "fig-classification-map-predictive"}}
Expand Down
4 changes: 2 additions & 2 deletions docs/examples/collapsed_vi.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,7 +121,7 @@
# %%
meanf = gpx.mean_functions.Constant()
kernel = gpx.kernels.RBF() # 1-dimensional inputs
likelihood = gpx.likelihoods.Gaussian(num_datapoints=D.n)
likelihood = gpx.likelihoods.Gaussian()
prior = gpx.gps.Prior(mean_function=meanf, kernel=kernel)
posterior = prior * likelihood

Expand Down Expand Up @@ -242,7 +242,7 @@
# %%
full_rank_model = gpx.gps.Prior(
mean_function=gpx.mean_functions.Zero(), kernel=gpx.kernels.RBF()
) * gpx.likelihoods.Gaussian(num_datapoints=D.n)
) * gpx.likelihoods.Gaussian()
nmll = jit(lambda: -gpx.objectives.conjugate_mll(full_rank_model, D))
# %timeit nmll().block_until_ready()

Expand Down
2 changes: 1 addition & 1 deletion docs/examples/constructing_new_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -287,7 +287,7 @@ def __call__(
# Define polar Gaussian process
PKern = Polar()
meanf = gpx.mean_functions.Zero()
likelihood = gpx.likelihoods.Gaussian(num_datapoints=n)
likelihood = gpx.likelihoods.Gaussian()
circular_posterior = gpx.gps.Prior(mean_function=meanf, kernel=PKern) * likelihood

# Optimise GP's marginal log-likelihood using BFGS
Expand Down
2 changes: 1 addition & 1 deletion docs/examples/deep_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -196,7 +196,7 @@ def __call__(self, x: jax.Array) -> jax.Array:
kernel = DeepKernelFunction(network=forward_linear, base_kernel=base_kernel)
meanf = gpx.mean_functions.Zero()
prior = gpx.gps.Prior(mean_function=meanf, kernel=kernel)
likelihood = gpx.likelihoods.Gaussian(num_datapoints=D.n)
likelihood = gpx.likelihoods.Gaussian()
posterior = prior * likelihood
# %% [markdown]
# ### Optimisation
Expand Down
2 changes: 1 addition & 1 deletion docs/examples/graph_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -174,7 +174,7 @@ def glue(*args, **kwargs):
# [`fit_scipy`](#gpjax.fit.fit_scipy).

# %%
likelihood = gpx.likelihoods.Gaussian(num_datapoints=D.n)
likelihood = gpx.likelihoods.Gaussian()
kernel = gpx.kernels.GraphKernel(laplacian=L)
prior = gpx.gps.Prior(mean_function=gpx.mean_functions.Zero(), kernel=kernel)
posterior = prior * likelihood
Expand Down
21 changes: 9 additions & 12 deletions docs/examples/heteroscedastic_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -163,8 +163,9 @@
# \mathcal{N}\!\big(y_i \mid f(x_i), \exp(g(x_i))\big),
# $$ (eq-heteroscedastic-likelihood)
#
# to form the posterior target that we shall approximate variationally. The product
# syntax `signal_prior * likelihood` used below constructs this augmented GP model.
# to form the posterior target that we shall approximate variationally. Because the
# joint model holds *two* priors — one per latent process — it is constructed
# directly as a `HeteroscedasticModel` rather than via the two-operand product.

# %%
# Signal and noise priors.
Expand All @@ -176,12 +177,10 @@
mean_function=gpx.mean_functions.Zero(),
kernel=gpx.kernels.RBF(),
)
likelihood = HeteroscedasticGaussian(
num_datapoints=train.n,
noise_prior=noise_prior,
noise_transform=LogNormalTransform(),
likelihood = HeteroscedasticGaussian(noise_transform=LogNormalTransform())
posterior = gpx.gps.HeteroscedasticModel(
prior=signal_prior, likelihood=likelihood, noise_prior=noise_prior
)
posterior = signal_prior * likelihood

# Variational family over both processes.
z = jnp.linspace(-3.2, 3.2, 25)[:, None]
Expand Down Expand Up @@ -330,12 +329,10 @@
mean_function=gpx.mean_functions.Zero(),
kernel=gpx.kernels.RBF(),
)
likelihood_adv = HeteroscedasticGaussian(
num_datapoints=data_adv.n,
noise_prior=noise_prior_adv,
noise_transform=SoftplusTransform(),
likelihood_adv = HeteroscedasticGaussian(noise_transform=SoftplusTransform())
posterior_adv = gpx.gps.HeteroscedasticModel(
prior=mean_prior, likelihood=likelihood_adv, noise_prior=noise_prior_adv
)
posterior_adv = mean_prior * likelihood_adv

# %%
# Configure variational family
Expand Down
Loading
Loading