ConjugateModel#

class gpjax.gps.ConjugateModel(prior, likelihood)[source]#

Bases: JointModel[M, K, GL]

A joint model with Gaussian likelihood: conditioning is exact.

For a Gaussian process prior \(p(\mathbf{f})\) and a Gaussian likelihood \(p(y | \mathbf{f}) = \mathcal{N}(y\mid \mathbf{f}, \sigma^2))\), the latent function can be analytically integrated out. Conditioning returns the closed-form posterior

\[\begin{split}\begin{aligned} p(\mathbf{f}^{\star}\mid \mathbf{y}) & =\mathcal{N}(\mathbf{f}^{\star}; \boldsymbol{\mu}_{\mid \mathbf{y}}, \boldsymbol{\Sigma}_{\mid \mathbf{y}}),\\ \boldsymbol{\mu}_{\mid \mathbf{y}} & = k(\mathbf{x}^{\star}, \mathbf{x})\left(k(\mathbf{x}, \mathbf{x}')+\sigma^2\mathbf{I}_n\right)^{-1}\mathbf{y}, \\ \boldsymbol{\Sigma}_{\mid \mathbf{y}} & =k(\mathbf{x}^{\star}, \mathbf{x}^{\star\prime}) -k(\mathbf{x}^{\star}, \mathbf{x})\left( k(\mathbf{x}, \mathbf{x}') + \sigma^2\mathbf{I}_n \right)^{-1}k(\mathbf{x}, \mathbf{x}^{\star}). \end{aligned}\end{split}\]

Example

>>> import gpjax as gpx
>>> import jax.numpy as jnp
>>>
>>> xtrain = jnp.linspace(0, 1).reshape(-1, 1)
>>> D = gpx.Dataset(X=xtrain, y=jnp.sin(xtrain))
>>>
>>> prior = gpx.gps.Prior(
...     mean_function = gpx.mean_functions.Zero(),
...     kernel = gpx.kernels.RBF()
... )
>>> model = prior * gpx.likelihoods.Gaussian()
>>> posterior = model.condition(D)
>>> predictive = posterior(xtrain)
>>> evidence = posterior.log_marginal_likelihood
Parameters:
condition(train_data)[source]#

Condition on data exactly.

Returns:

The closed-form posterior process, with the

training-covariance factorisation cached. Exposes the predictive (via __call__), log_marginal_likelihood, loo and sample_approx.

Return type:

ExactPosterior

Parameters:

train_data (Dataset)

sample_approx(num_samples, train_data, key, num_features=100)[source]#

Sugar: self.condition(train_data).sample_approx(...).

Draw approximate posterior samples via pathwise conditioning (Wilson et al., 2020).

Parameters:
  • num_samples (int)

  • train_data (Dataset)

  • key (UInt32[jaxlib._jax.Array, '2'] | Key[jaxlib._jax.Array, ''])

  • num_features (int | None)

Return type:

Callable[[Float[jaxlib._jax.Array, ‘N D’] | Float[ndarray, ‘N D’]], Float[jaxlib._jax.Array, ‘N B’] | Float[ndarray, ‘N B’]]