CoregionalizationMatrix#

class gpjax.parameters.CoregionalizationMatrix(num_outputs, rank, key)[source]#

Bases: Module

Parameterises a PSD output-correlation matrix B = WW^T + diag(kappa).

Parameters:
  • num_outputs (int)

  • rank (int)

  • key (Array)