HeteroscedasticPrediction#

class gpjax.variational_families.HeteroscedasticPrediction(mean_f, variance_f, mean_g, variance_g)[source]#

Bases: NamedTuple

Mean and variance of the signal and noise latent processes.

Parameters:
  • mean_f (Float[jaxlib._jax.Array, 'N 1'] | Float[ndarray, 'N 1'])

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

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

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

mean_f: Float[jaxlib._jax.Array, 'N 1'] | Float[ndarray, 'N 1']#

Alias for field number 0

mean_g: Float[jaxlib._jax.Array, 'N 1'] | Float[ndarray, 'N 1']#

Alias for field number 2

variance_f: Float[jaxlib._jax.Array, 'N 1'] | Float[ndarray, 'N 1']#

Alias for field number 1

variance_g: Float[jaxlib._jax.Array, 'N 1'] | Float[ndarray, 'N 1']#

Alias for field number 3