Source code for gpjax.likelihoods

# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
from __future__ import annotations

import abc
from dataclasses import dataclass
from typing import TYPE_CHECKING

import beartype.typing as tp
import equinox as eqx
import jax
from jax import vmap
import jax.nn as jnn
import jax.numpy as jnp
import jax.scipy as jsp
from jaxtyping import Float
import lineax as lx
import numpy as np
import numpyro.distributions as npd

from gpjax.distributions import GaussianDistribution
from gpjax.integrators import (
    AbstractIntegrator,
    AnalyticalGaussianIntegrator,
    GHQuadratureIntegrator,
)
from gpjax.parameters import (
    NonNegativeReal,
    _val,
)
from gpjax.summary import _SummaryMixin
from gpjax.typing import (
    Array,
    ScalarFloat,
)

if TYPE_CHECKING:
    from gpjax.gps import Prior


def _diagonal_scale(op):
    """Unwrap a single TaggedLinearOperator layer and return the inner diagonal, or None.

    Returns the input itself if it is already a DiagonalLinearOperator; returns the
    inner DiagonalLinearOperator if `op` is a TaggedLinearOperator wrapping one;
    returns None otherwise (e.g. dense MatrixLinearOperator, nested tags).
    """
    if isinstance(op, lx.DiagonalLinearOperator):
        return op
    if isinstance(op, lx.TaggedLinearOperator) and isinstance(
        op.operator, lx.DiagonalLinearOperator
    ):
        return op.operator
    return None


[docs] @dataclass(slots=True) class NoiseMoments: log_variance: Array inv_variance: Array variance: Array
jax.tree_util.register_pytree_node( NoiseMoments, lambda x: ((x.log_variance, x.inv_variance, x.variance), None), lambda _, x: NoiseMoments(*x), )
[docs] class AbstractLikelihood(_SummaryMixin, eqx.Module): r"""Abstract base class for likelihoods. All likelihoods must inherit from this class and implement the `predict` and `link_function` methods. """ num_datapoints: int = eqx.field(static=True) integrator: AbstractIntegrator = eqx.field(static=True) def __init__( self, num_datapoints: int, integrator: AbstractIntegrator = GHQuadratureIntegrator(), ): """Initializes the likelihood. Args: num_datapoints (int): the number of data points. integrator (AbstractIntegrator): The integrator to be used for computing expected log likelihoods. Must be an instance of `AbstractIntegrator`. """ self.num_datapoints = num_datapoints self.integrator = integrator def __call__( self, dist: tp.Union[npd.MultivariateNormal, GaussianDistribution] ) -> npd.Distribution: r"""Evaluate the likelihood function at a given predictive distribution. Args: dist: The predictive distribution to evaluate the likelihood at. Returns: The predictive distribution. """ return self.predict(dist)
[docs] @abc.abstractmethod def predict( self, dist: tp.Union[npd.MultivariateNormal, GaussianDistribution] ) -> npd.Distribution: r"""Evaluate the likelihood function at a given predictive distribution. Args: dist: The predictive distribution to evaluate the likelihood at. Returns: npd.Distribution: The predictive distribution. """ raise NotImplementedError
[docs] def expected_log_likelihood( self, y: Float[Array, "N D"], mean: Float[Array, "N D"], variance: Float[Array, "N D"], mean_g: tp.Optional[Float[Array, "N D"]] = None, variance_g: tp.Optional[Float[Array, "N D"]] = None, **_: tp.Any, ) -> Float[Array, " N"]: r"""Compute the expected log likelihood. For a variational distribution $q(f)\sim\mathcal{N}(m, s)$ and a likelihood $p(y|f)$, compute the expected log likelihood: .. math:: \mathbb{E}_{q(f)}\left[\log p(y|f)\right] Args: y (Float[Array, 'N D']): The observed response variable. mean (Float[Array, 'N D']): The variational mean. variance (Float[Array, 'N D']): The variational variance. mean_g (Float[Array, 'N D']): Optional moments of the latent noise process for heteroscedastic likelihoods. variance_g (Float[Array, 'N D']): Optional moments of the latent noise process for heteroscedastic likelihoods. **_: Unused extra arguments for compatibility with specialised likelihoods. Returns: ScalarFloat: The expected log likelihood. """ log_prob = vmap(lambda f, y: self.link_function(f).log_prob(y)) return self.integrator( fun=log_prob, y=y, mean=mean, variance=variance, likelihood=self )
[docs] class AbstractNoiseTransform(eqx.Module): """Abstract base class for noise transformations.""" @abc.abstractmethod def __call__(self, x: Float[Array, ...]) -> Float[Array, ...]: """Transform the input noise signal.""" raise NotImplementedError
[docs] @abc.abstractmethod def moments( self, mean: Float[Array, ...], variance: Float[Array, ...] ) -> NoiseMoments: """Compute the moments of the transformed noise signal.""" raise NotImplementedError
[docs] class LogNormalTransform(AbstractNoiseTransform): """Log-normal noise transformation.""" def __call__(self, x: Float[Array, ...]) -> Float[Array, ...]: return jnp.exp(x)
[docs] def moments( self, mean: Float[Array, ...], variance: Float[Array, ...] ) -> NoiseMoments: expected_variance = jnp.exp(mean + 0.5 * variance) expected_log_variance = mean expected_inv_variance = jnp.exp(-mean + 0.5 * variance) return NoiseMoments( log_variance=expected_log_variance, inv_variance=expected_inv_variance, variance=expected_variance, )
[docs] class SoftplusTransform(AbstractNoiseTransform): """Softplus noise transformation.""" num_points: int = eqx.field(static=True, default=20) def __init__(self, num_points: int = 20): self.num_points = num_points def __call__(self, x: Float[Array, ...]) -> Float[Array, ...]: return jnn.softplus(x)
[docs] def moments( self, mean: Float[Array, ...], variance: Float[Array, ...] ) -> NoiseMoments: quad_x, quad_w = np.polynomial.hermite.hermgauss(self.num_points) quad_w = jnp.asarray(quad_w / jnp.sqrt(jnp.pi)) quad_x = jnp.asarray(quad_x) std = jnp.sqrt(variance) samples = mean[..., None] + jnp.sqrt(2.0) * std[..., None] * quad_x sigma2 = self(samples) log_sigma2 = jnp.log(sigma2) inv_sigma2 = 1.0 / sigma2 expected_variance = jnp.sum(sigma2 * quad_w, axis=-1) expected_log_variance = jnp.sum(log_sigma2 * quad_w, axis=-1) expected_inv_variance = jnp.sum(inv_sigma2 * quad_w, axis=-1) return NoiseMoments( log_variance=expected_log_variance, inv_variance=expected_inv_variance, variance=expected_variance, )
[docs] class AbstractHeteroscedasticLikelihood(AbstractLikelihood): r"""Base class for heteroscedastic likelihoods with latent noise processes.""" noise_prior: tp.Any noise_transform: AbstractNoiseTransform def __init__( self, num_datapoints: int, noise_prior: Prior, noise_transform: tp.Union[ AbstractNoiseTransform, tp.Callable[[Float[Array, ...]], Float[Array, ...]], ] = SoftplusTransform(), integrator: AbstractIntegrator = GHQuadratureIntegrator(), ): self.noise_prior = noise_prior if isinstance(noise_transform, AbstractNoiseTransform): self.noise_transform = noise_transform else: transform_name = getattr(noise_transform, "__name__", "") if noise_transform is jnp.exp or transform_name == "exp": self.noise_transform = LogNormalTransform() else: # Default to SoftplusTransform for softplus or unknown callables (legacy behavior used quadrature) # Note: If an unknown callable is passed, we technically use SoftplusTransform which applies softplus. # Users should implement AbstractNoiseTransform for custom transforms. self.noise_transform = SoftplusTransform() super().__init__(num_datapoints=num_datapoints, integrator=integrator) def __call__( self, dist: tp.Union[npd.MultivariateNormal, GaussianDistribution], noise_dist: tp.Optional[ tp.Union[npd.MultivariateNormal, GaussianDistribution] ] = None, ) -> npd.Distribution: return self.predict(dist, noise_dist)
[docs] def supports_tight_bound(self) -> bool: """Return whether the tighter bound from Lazaro-Gredilla & Titsias (2011) is applicable.""" return False
[docs] def noise_statistics( self, mean: Float[Array, "N D"], variance: Float[Array, "N D"] ) -> NoiseMoments: r"""Moment matching of the transformed noise process. Args: mean: Mean of the latent noise GP. variance: Variance of the latent noise GP. Returns: NoiseMoments: Expected log variance, inverse variance, and variance. """ return self.noise_transform.moments(mean, variance)
[docs] def expected_log_likelihood( self, y: Float[Array, "N D"], mean: Float[Array, "N D"], variance: Float[Array, "N D"], mean_g: tp.Optional[Float[Array, "N D"]] = None, variance_g: tp.Optional[Float[Array, "N D"]] = None, **kwargs: tp.Any, ) -> Float[Array, " N"]: raise NotImplementedError
[docs] class Gaussian(AbstractLikelihood): r"""Gaussian likelihood object.""" obs_stddev: tp.Any num_outputs: int = eqx.field(static=True, default=1) def __init__( self, num_datapoints: int, obs_stddev: tp.Union[ScalarFloat, Float[Array, "#N"], NonNegativeReal] = 1.0, integrator: AbstractIntegrator = AnalyticalGaussianIntegrator(), ): r"""Initializes the Gaussian likelihood. Args: num_datapoints (int): the number of data points. obs_stddev (Union[ScalarFloat, Float[Array, "#N"]]): the standard deviation of the Gaussian observation noise. integrator (AbstractIntegrator): The integrator to be used for computing expected log likelihoods. Must be an instance of `AbstractIntegrator`. For the Gaussian likelihood, this defaults to the `AnalyticalGaussianIntegrator`, as the expected log likelihood can be computed analytically. """ if not isinstance(obs_stddev, NonNegativeReal): obs_stddev = NonNegativeReal(jnp.asarray(obs_stddev)) self.obs_stddev = obs_stddev self.num_outputs = 1 super().__init__(num_datapoints, integrator)
[docs] def predict( self, dist: tp.Union[npd.MultivariateNormal, GaussianDistribution] ) -> GaussianDistribution: r"""Evaluate the Gaussian likelihood at a predictive distribution. Preserves diagonal scale when the input carries a ``lineax.DiagonalLinearOperator`` (including when wrapped in ``lx.TaggedLinearOperator`` as emitted by ``DiagonalKernelComputation`` / ``ConstantDiagonalKernelComputation``). Always returns ``GaussianDistribution``. This widens the previous return type from ``numpyro.distributions.MultivariateNormal`` — see CHANGELOG v0.15. Args: dist: The Gaussian process posterior at a finite set of test points. Returns: GaussianDistribution: The predictive distribution with observation noise added to the diagonal of the covariance. """ obs_var = _val(self.obs_stddev) ** 2 if isinstance(dist, GaussianDistribution): diag = _diagonal_scale(dist.scale) if diag is not None: noisy_scale = lx.DiagonalLinearOperator(lx.diagonal(diag) + obs_var) return GaussianDistribution(dist.mean, noisy_scale) # Dense fallback — widen return type to GaussianDistribution for API # consistency. n_data = dist.event_shape[0] cov = dist.covariance_matrix noisy_cov = cov.at[jnp.diag_indices(n_data)].add(obs_var) return GaussianDistribution(dist.mean, lx.MatrixLinearOperator(noisy_cov))
[docs] def noise_vector(self, n: int) -> Float[Array, " N"]: """Per-observation noise variance vector (scalar broadcast for single-output).""" return jnp.full(n, jnp.square(_val(self.obs_stddev)))
[docs] def prepare_targets( self, y: Float[Array, "N 1"], mx: Float[Array, "N 1"] ) -> tuple[Float[Array, "N 1"], Float[Array, "N 1"]]: """Return targets and mean in the format expected by the unified predict/MLL path.""" return y, mx
[docs] class MultiOutputGaussian(Gaussian): """Gaussian likelihood with per-output noise variance. Args: num_datapoints: Total number of observations (N, not N*P). num_outputs: Number of output dimensions (P). obs_stddev: Per-output noise standard deviation. Scalar broadcasts to [P]. """ def __init__( self, num_datapoints: int, num_outputs: int, obs_stddev: tp.Union[float, Float[Array, " P"]] = 1.0, ): if isinstance(obs_stddev, (int, float)): obs_stddev = jnp.full(num_outputs, float(obs_stddev)) super().__init__( num_datapoints=num_datapoints, obs_stddev=NonNegativeReal(jnp.asarray(obs_stddev)), ) self.num_outputs = num_outputs
[docs] def noise_vector(self, n: int) -> Float[Array, " NP"]: """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, ...]. """ per_output_var = jnp.square(_val(self.obs_stddev)) # [P] return jnp.repeat(per_output_var, n) # [NP]
[docs] def prepare_targets( self, y: Float[Array, "N P"], mx: Float[Array, "N 1"] ) -> tuple[Float[Array, "NP 1"], Float[Array, "NP 1"]]: """Reshape multi-output targets to output-major long format.""" P = self.num_outputs y_flat = y.T.reshape(-1, 1) # [N, P] -> [NP, 1] mx_flat = jnp.tile(mx, (P, 1)) # [N, 1] -> [NP, 1] return y_flat, mx_flat
[docs] class HeteroscedasticGaussian(AbstractHeteroscedasticLikelihood):
[docs] def predict( self, dist: tp.Union[npd.MultivariateNormal, GaussianDistribution], noise_dist: tp.Optional[ tp.Union[npd.MultivariateNormal, GaussianDistribution] ] = None, ) -> GaussianDistribution: if noise_dist is None: raise ValueError( "noise_dist must be provided for heteroscedastic prediction." ) n_data = dist.event_shape[0] noise_mean = noise_dist.mean noise_variance = jnp.diag(noise_dist.covariance_matrix) noise_stats = self.noise_statistics( noise_mean[..., None], noise_variance[..., None] ) cov = dist.covariance_matrix noisy_cov = cov.at[jnp.diag_indices(n_data)].add(noise_stats.variance.squeeze()) return GaussianDistribution(dist.mean, lx.MatrixLinearOperator(noisy_cov))
[docs] def expected_log_likelihood( self, y: Float[Array, "N D"], mean: Float[Array, "N D"], variance: Float[Array, "N D"], mean_g: tp.Optional[Float[Array, "N D"]] = None, variance_g: tp.Optional[Float[Array, "N D"]] = None, noise_stats: tp.Optional[NoiseMoments] = None, return_parts: bool = False, **_: tp.Any, ) -> tp.Union[Float[Array, " N"], tuple[Float[Array, " N"], NoiseMoments]]: if mean_g is None or variance_g is None: raise ValueError( "mean_g and variance_g must be provided for heteroscedastic models." ) if noise_stats is None: noise_stats = self.noise_statistics(mean_g, variance_g) sq_error = jnp.square(y - mean) log2pi = jnp.log(2.0 * jnp.pi) expected = -0.5 * ( log2pi + noise_stats.log_variance + (sq_error + variance) * noise_stats.inv_variance ) expected_sum = jnp.sum(expected, axis=1) if return_parts: return expected_sum, noise_stats return expected_sum
[docs] def supports_tight_bound(self) -> bool: return True
[docs] class Bernoulli(AbstractLikelihood):
[docs] def predict( self, dist: tp.Union[npd.MultivariateNormal, GaussianDistribution] ) -> npd.BernoulliProbs: r"""Evaluate the pointwise predictive distribution. Evaluate the pointwise predictive distribution, given a Gaussian process posterior and likelihood parameters. Args: dist ([npd.MultivariateNormal, GaussianDistribution].): The Gaussian process posterior, evaluated at a finite set of test points. Returns: npd.Bernoulli: The pointwise predictive distribution. """ variance = jnp.diag(dist.covariance_matrix) mean = dist.mean.ravel() return self.link_function(mean / jnp.sqrt(1.0 + variance))
[docs] class Poisson(AbstractLikelihood):
[docs] def predict( self, dist: tp.Union[npd.MultivariateNormal, GaussianDistribution] ) -> npd.Poisson: r"""Evaluate the pointwise predictive distribution. Evaluate the pointwise predictive distribution, given a Gaussian process posterior and likelihood parameters. Args: dist (tp.Union[npd.MultivariateNormal, GaussianDistribution]): The Gaussian process posterior, evaluated at a finite set of test points. Returns: npd.Poisson: The pointwise predictive distribution. """ return self.link_function(dist.mean)
[docs] def inv_probit(x: Float[Array, " *N"]) -> Float[Array, " *N"]: r"""Compute the inverse probit function. Args: x (``Float[Array, "*N"]``): A vector of values. Returns: ``Float[Array, "*N"]``: The inverse probit of the input vector. """ jitter = 1e-3 # To ensure output is in interval (0, 1). return 0.5 * (1.0 + jsp.special.erf(x / jnp.sqrt(2.0))) * (1 - 2 * jitter) + jitter
NonGaussian = tp.Union[Poisson, Bernoulli] __all__ = [ "AbstractHeteroscedasticLikelihood", "AbstractLikelihood", "AbstractNoiseTransform", "Bernoulli", "Gaussian", "HeteroscedasticGaussian", "LogNormalTransform", "MultiOutputGaussian", "NoiseMoments", "NonGaussian", "Poisson", "SoftplusTransform", "inv_probit", ]