create_oilmm_with_kernels#

gpjax.models.create_oilmm_with_kernels(latent_kernels, num_outputs, key, mean_function=None)[source]#

Create OILMM with custom kernel per latent GP.

Parameters:
  • latent_kernels (list[AbstractKernel]) – List of M kernels, one per latent GP

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

  • key (Array) – JAX PRNG key

  • mean_function (tp.Any) – Mean function (shared, default: Zero)

Returns:

OILMMModel with heterogeneous latent kernels

Return type:

OILMMModel

Example

>>> import gpjax as gpx
>>> import jax.random as jr
>>> model = gpx.models.create_oilmm_with_kernels(
...     latent_kernels=[gpx.kernels.RBF(), gpx.kernels.Matern52()],
...     num_outputs=6,
...     key=jr.key(42)
... )

Expand for references to gpjax.models.create_oilmm_with_kernels

create_oilmm_with_kernels