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:
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