HeteroscedasticVariationalFamily#

class gpjax.variational_families.HeteroscedasticVariationalFamily(posterior, 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, jitter=1e-06, 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:
  • posterior (AbstractPosterior)

  • 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)

  • jitter (float)

  • signal_init (VariationalGaussianInit | None)

  • noise_init (VariationalGaussianInit | None)

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