# 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
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,
)
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.
"""
integrator: AbstractIntegrator = eqx.field(static=True)
def __init__(
self,
integrator: AbstractIntegrator = GHQuadratureIntegrator(),
):
"""Initializes the likelihood.
Args:
integrator (AbstractIntegrator): The integrator to be used for computing expected log
likelihoods. Must be an instance of `AbstractIntegrator`.
"""
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]
@abc.abstractmethod
def link_function(self, f: Float[Array, ...]) -> npd.Distribution:
r"""Return the link function of the likelihood function.
Args:
f (Float[Array, "..."]): the latent Gaussian process values.
Returns:
npd.Distribution: The distribution of observations, y, given values of the
Gaussian process, f.
"""
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 AbstractHeteroscedasticLikelihood(AbstractLikelihood):
r"""Base class for heteroscedastic likelihoods with latent noise processes."""
noise_transform: AbstractNoiseTransform
def __init__(
self,
noise_transform: tp.Union[
AbstractNoiseTransform,
tp.Callable[[Float[Array, ...]], Float[Array, ...]],
] = SoftplusTransform(),
integrator: AbstractIntegrator = GHQuadratureIntegrator(),
):
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__(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,
obs_stddev: tp.Union[ScalarFloat, Float[Array, "#N"], NonNegativeReal] = 1.0,
integrator: AbstractIntegrator = AnalyticalGaussianIntegrator(),
):
r"""Initializes the Gaussian likelihood.
Args:
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__(integrator)
[docs]
def link_function(self, f: Float[Array, ...]) -> npd.Normal:
r"""The link function of the Gaussian likelihood.
Args:
f (Float[Array, "..."]): Function values.
Returns:
npd.Normal: The likelihood function.
"""
return npd.Normal(loc=f, scale=_val(self.obs_stddev).astype(f.dtype))
[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_outputs: Number of output dimensions (P).
obs_stddev: Per-output noise standard deviation. Scalar broadcasts to [P].
"""
def __init__(
self,
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__(
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 link_function(
self,
f: Float[Array, ...],
g: tp.Optional[Float[Array, ...]] = None,
) -> npd.Normal:
r"""The conditional observation density $p(y \mid f, g)$.
For a heteroscedastic likelihood the observation noise is itself a
function of a second latent process $g$, so the conditional is
$\mathcal{N}(y \mid f, \sigma^2(g))$ (Lázaro-Gredilla & Titsias, 2011).
Unlike the homoscedastic likelihoods, $f$ alone does not determine the
density, so `g` is required.
Args:
f (Float[Array, "..."]): the latent signal process values.
g (Float[Array, "..."] | None): the latent noise process values. Required —
there is no conditional density without it.
Returns:
npd.Normal: The observation density given both latent processes.
Raises:
ValueError: If `g` is not supplied.
"""
if g is None:
raise ValueError(
f"{type(self).__name__}.link_function requires the noise latent `g` "
"as well as the signal latent `f`: the observation noise is "
"sigma^2(g), so p(y | f) alone is not defined. Pass "
"`link_function(f, g)`, or use `expected_log_likelihood(..., "
"mean_g=..., variance_g=...)` / `predict(dist, noise_dist)` which "
"handle the noise process for you."
)
return npd.Normal(loc=f, scale=jnp.sqrt(self.noise_transform(g)))
[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 link_function(self, f: Float[Array, ...]) -> npd.BernoulliProbs:
r"""The probit link function of the Bernoulli likelihood.
Args:
f (Float[Array, "..."]): Function values.
Returns:
npd.Bernoulli: The likelihood function.
"""
return npd.Bernoulli(probs=inv_probit(f))
[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 link_function(self, f: Float[Array, ...]) -> npd.Poisson:
r"""The link function of the Poisson likelihood.
Args:
f (Float[Array, "..."]): Function values.
Returns:
npd.Poisson: The likelihood function.
"""
return npd.Poisson(rate=jnp.exp(f))
[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",
]