MultiOutputGaussian#

class gpjax.likelihoods.MultiOutputGaussian(num_datapoints, num_outputs, obs_stddev=1.0)[source]#

Bases: Gaussian

Gaussian likelihood with per-output noise variance.

Parameters:
  • num_datapoints (int) – Total number of observations (N, not N*P).

  • num_outputs (int) – Number of output dimensions (P).

  • obs_stddev (Any) – Per-output noise standard deviation. Scalar broadcasts to [P].

noise_vector(n)[source]#

Per-observation noise variance in output-major (Kronecker) order.

Returns sigma_p^2 with each output’s variance repeated N times, concatenated across outputs: [sigma_1^2…sigma_1^2, sigma_2^2…sigma_2^2, …].

Parameters:

n (int)

Return type:

Float[jaxlib._jax.Array, ‘NP’] | Float[ndarray, ‘NP’]

prepare_targets(y, mx)[source]#

Reshape multi-output targets to output-major long format.

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

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

Return type:

tuple[Float[jaxlib._jax.Array, ‘NP 1’] | Float[ndarray, ‘NP 1’], Float[jaxlib._jax.Array, ‘NP 1’] | Float[ndarray, ‘NP 1’]]