HeteroscedasticVariationalFamily#

class gpjax.variational_families.HeteroscedasticVariationalFamily(model, inducing_inputs=None, inducing_inputs_g=None, variational_mean_f=None, variational_root_covariance_f=None, variational_mean_g=None, variational_root_covariance_g=None, signal_init=None, noise_init=None)[source]#

Bases: AbstractVariationalFamily[HL]

Variational family for two independent latent processes f and g.

Expand for references to gpjax.variational_families.HeteroscedasticVariationalFamily

Heteroscedastic Inference / Background

Parameters:
  • model (JointModel)

  • inducing_inputs (Int[jaxlib._jax.Array, 'N D'] | Int[ndarray, 'N D'] | Float[jaxlib._jax.Array, 'N D'] | Float[ndarray, 'N D'])

  • inducing_inputs_g (Int[jaxlib._jax.Array, 'M D'] | Int[ndarray, 'M D'] | Float[jaxlib._jax.Array, 'M D'] | Float[ndarray, 'M D'] | None)

  • variational_mean_f (Float[jaxlib._jax.Array, 'N 1'] | Float[ndarray, 'N 1'] | None)

  • variational_root_covariance_f (Float[jaxlib._jax.Array, 'N N'] | Float[ndarray, 'N N'] | None)

  • variational_mean_g (Float[jaxlib._jax.Array, 'M 1'] | Float[ndarray, 'M 1'] | None)

  • variational_root_covariance_g (Float[jaxlib._jax.Array, 'M M'] | Float[ndarray, 'M M'] | None)

  • signal_init (VariationalGaussianInit | None)

  • noise_init (VariationalGaussianInit | None)

condition(train_data)[source]#

Not available: the heteroscedastic family has no single posterior.

This family approximates two latent processes – signal and noise – so there is no one conditioned process to return, matching the exclusion recorded for gpjax.gps.HeteroscedasticModel. Condition the components instead, via signal_variational and noise_variational, or call predict_latents() for the two predictive distributions and predict() for their moments.

Parameters:

train_data (Dataset | None) – Unused; present for interface uniformity.

Raises:

NotImplementedError – Always.

Return type:

Posterior

predict(test_inputs)[source]#

Predict the GP’s output given the input.

Parameters:
  • *args (Any) – Arguments of the variational family’s predict method.

  • **kwargs (Any) – Keyword arguments of the variational family’s predict method.

  • test_inputs (Int[jaxlib._jax.Array, 'N D'] | Int[ndarray, 'N D'] | Float[jaxlib._jax.Array, 'N D'] | Float[ndarray, 'N D'])

Returns:

The output of the variational family’s predict method.

Return type:

GaussianDistribution

prior_kl()[source]#

The KL divergence from the variational distribution to the prior.

Every ELBO-style objective subtracts this term, so each concrete family must provide it.

Returns:

The KL divergence.

Return type:

ScalarFloat