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- 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
predictmethod.**kwargs (Any) – Keyword arguments of the variational family’s
predictmethod.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
predictmethod.- Return type: