Source code for gpjax.kernels.stationary.gneiting

# Copyright 2026 The thomaspinder Contributors. All Rights Reserved.
#
# 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 typing import ClassVar

import beartype.typing as tp
import equinox as eqx
import jax.numpy as jnp
from jaxtyping import Float
from paramax import AbstractUnwrappable

from gpjax.kernels.base import (
    AbstractKernel,
    val,
)
from gpjax.kernels.computations import (
    AbstractKernelComputation,
    DenseKernelComputation,
)
from gpjax.parameters import (
    NonNegativeReal,
    PositiveReal,
    SigmoidBounded,
)
from gpjax.typing import (
    Array,
    ScalarFloat,
)


[docs] class Gneiting(AbstractKernel): r"""The Gneiting nonseparable space–time kernel. Computes the covariance for a pair of inputs with spatial separation $h$ over the space columns and time lag $u$ over the time column (Gneiting, 2002, eq. 14): $$ k(h, u) = \frac{\sigma^2}{\psi(u)^{d/2}} \exp\!\left(-\frac{(\lVert h\rVert/\ell_s)^{2\gamma}} {\psi(u)^{\beta\gamma}}\right), \qquad \psi(u) = \left(\frac{\lvert u\rvert}{\ell_t}\right)^{2\alpha} + 1, $$ where $d$ is the number of space columns. The kernel is stationary, but it is not separable: as the time lag grows, $\psi(u)$ grows and the spatial correlation decays more slowly. The interaction parameter $\beta \in [0, 1]$ controls this effect, and $\beta = 0$ gives the separable product of a powered exponential kernel in space and a generalised Cauchy kernel in time. $\alpha \in (0, 1]$ and $\gamma \in (0, 1]$ set the smoothness in time and in space. The trainable parameters $\alpha$, $\beta$ and $\gamma$ are bounded to the open interval $(0, 1)$. To fix one of them at a bound, pass a non-trainable value, for example `paramax.non_trainable(jnp.array(1.0))`. The kernel has two lengthscales and no closed-form spectral density, so it is not a :class:`StationaryKernel` subclass and it does not support random Fourier features. """ name: ClassVar[str] = "Gneiting" space_dims: list[int] = eqx.field(static=True) time_dim: int = eqx.field(static=True) variance: tp.Any space_lengthscale: tp.Any time_lengthscale: tp.Any alpha: tp.Any beta: tp.Any gamma: tp.Any def __init__( self, space_dims: list[int], time_dim: int, variance: tp.Union[ScalarFloat, AbstractUnwrappable] = 1.0, space_lengthscale: tp.Union[ScalarFloat, AbstractUnwrappable] = 1.0, time_lengthscale: tp.Union[ScalarFloat, AbstractUnwrappable] = 1.0, alpha: tp.Union[ScalarFloat, AbstractUnwrappable] = 0.5, beta: tp.Union[ScalarFloat, AbstractUnwrappable] = 0.5, gamma: tp.Union[ScalarFloat, AbstractUnwrappable] = 0.5, compute_engine: AbstractKernelComputation = DenseKernelComputation(), ): r"""Initialise the kernel. Args: space_dims: the indices of the space columns. time_dim: the index of the time column. variance: the variance $\sigma^2$. space_lengthscale: the spatial lengthscale $\ell_s$. time_lengthscale: the temporal lengthscale $\ell_t$. alpha: the smoothness in time, $\alpha \in (0, 1]$. beta: the space–time interaction, $\beta \in [0, 1]$. gamma: the smoothness in space, $\gamma \in (0, 1]$. compute_engine: the computation engine that the kernel uses to compute its covariance matrices. Raises: ValueError: if the columns are not valid, or if a float value of `alpha`, `beta` or `gamma` is not in the open interval $(0, 1)$. """ _check_columns(space_dims, time_dim) self.space_dims = list(space_dims) self.time_dim = time_dim self.variance = _wrap(variance, NonNegativeReal) self.space_lengthscale = _wrap(space_lengthscale, PositiveReal) self.time_lengthscale = _wrap(time_lengthscale, PositiveReal) self.alpha = _wrap_unit(alpha, "alpha") self.beta = _wrap_unit(beta, "beta") self.gamma = _wrap_unit(gamma, "gamma") super().__init__(compute_engine=compute_engine) def __call__(self, x: Float[Array, " D"], y: Float[Array, " D"]) -> ScalarFloat: h = (x[..., self.space_dims] - y[..., self.space_dims]) / val( self.space_lengthscale ) u = (x[..., self.time_dim] - y[..., self.time_dim]) / val(self.time_lengthscale) gamma = val(self.gamma) psi = _power(u**2, val(self.alpha)) + 1.0 space_term = _power(jnp.sum(h**2), gamma) / psi ** (val(self.beta) * gamma) dims = len(self.space_dims) K = val(self.variance) * psi ** (-0.5 * dims) * jnp.exp(-space_term) return K.squeeze()
def _power(t: Float[Array, ""], p: Float[Array, ""]) -> Float[Array, ""]: # t**p with t >= 0. The gradient of 0**p with respect to p, or of t**p at # t = 0 for p < 1, is not finite, so the zero case takes a separate branch. positive = t > 0 safe = jnp.where(positive, t, 1.0) return jnp.where(positive, safe**p, 0.0) def _check_columns(space_dims: tp.Any, time_dim: tp.Any) -> None: if ( not isinstance(space_dims, (list, tuple)) or not space_dims or not all(isinstance(i, int) for i in space_dims) ): raise ValueError( "Expected `space_dims` to be a non-empty list of column indices. " f"Got {space_dims!r}." ) if len(set(space_dims)) != len(space_dims): raise ValueError(f"`space_dims` has repeated columns: {space_dims!r}.") if not isinstance(time_dim, int): raise ValueError( f"Expected `time_dim` to be one column index. Got {time_dim!r}." ) if time_dim in space_dims: raise ValueError( f"Column {time_dim} is both a space column and the time column." ) def _wrap(value: tp.Any, parameter: type) -> tp.Any: if isinstance(value, AbstractUnwrappable): return value return parameter(jnp.asarray(value, dtype=float)) def _wrap_unit(value: tp.Any, label: str) -> tp.Any: if isinstance(value, AbstractUnwrappable): return value value = jnp.asarray(value, dtype=float) if not 0.0 < float(value) < 1.0: raise ValueError( f"Expected `{label}` in the open interval (0, 1), so that it can be " f"trained. Got {float(value)}. To fix it at a bound, pass " "`paramax.non_trainable(jnp.array(value))`." ) return SigmoidBounded(value, low=0.0, high=1.0) __all__ = ["Gneiting"]