create_oilmm#

gpjax.models.create_oilmm(num_outputs, num_latent_gps, key, kernel=None, mean_function=None)[source]#

Create OILMM model with shared kernel across latents.

Parameters:
  • num_outputs (int) – Number of output dimensions (p)

  • num_latent_gps (int) – Number of latent GPs (m)

  • key (Array) – JAX PRNG key

  • kernel (AbstractKernel | list[AbstractKernel] | None) – Kernel for latent GPs (default: RBF)

  • mean_function (tp.Any) – Mean function for latent GPs (default: Zero)

Returns:

Initialized OILMMModel

Return type:

OILMMModel

Example

>>> import gpjax as gpx
>>> import jax.random as jr
>>> model = gpx.models.create_oilmm(
...     num_outputs=5,
...     num_latent_gps=2,
...     key=jr.key(42),
...     kernel=gpx.kernels.Matern52()
... )

Expand for references to gpjax.models.create_oilmm

create_oilmm