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- 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, viasignal_variationalandnoise_variational, or callpredict_latents()for the two predictive distributions andpredict()for their moments.- Parameters:
train_data (Dataset | None) – Unused; present for interface uniformity.
- Raises:
NotImplementedError – Always.
- Return type:
- 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: