JointModel#
- class gpjax.gps.JointModel(prior, likelihood)[source]#
Bases:
_SummaryMixin,Module,Generic[M,K,L]The joint distribution \(p(f, y) = p(y \mid f)\,p(f)\).
Pairs a
Priorwith a likelihood. This is the trainable object:gpx.fitoptimises its hyperparameters. Conditioning it on data —model.condition(D)ormodel | D— produces the posterior process.The base class carries no inference of its own; concrete subclasses (
ConjugateModel,NonConjugateModel,HeteroscedasticModel) define what conditioning means for their likelihood. A bareJointModelis a lightweight pairing used where inference is delegated elsewhere (e.g. variational families over a latent noise process).- predict(test_inputs, train_data, *, covariance='dense')[source]#
Sugar: condition on
train_dataand query attest_inputs.Defined as exactly
self.condition(train_data)(test_inputs). When making repeated predictions, condition once and reuse the returned posterior — the factorisation is cached there.- Parameters:
test_inputs (Num[jaxlib._jax.Array, 'N D'] | Num[ndarray, 'N D']) – A Jax array of test inputs.
train_data (Dataset) – A
gpx.Datasetto condition on.covariance (Literal['dense', 'diagonal']) – Whether to return the dense joint covariance at the test inputs or only the marginal (diagonal) variances.
- Returns:
The predictive distribution.
- Return type: