Source code for gpjax.kernels.nonstationary.gibbs

# 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 jax.numpy as jnp
from jaxtyping import Float

from gpjax.kernels.base import AbstractKernel
from gpjax.kernels.computations import (
    AbstractKernelComputation,
    DenseKernelComputation,
)
from gpjax.kernels.location_functions import AbstractLocationFunction
from gpjax.typing import (
    Array,
    ScalarFloat,
)


[docs] class Gibbs(AbstractKernel): r"""A base kernel whose lengthscale changes with location. The Gibbs kernel (Gibbs, 1997), also known as the Paciorek–Schervish kernel (Paciorek & Schervish, 2006). A location function $g$ gives $\log\ell(x)$, and $\ell(x)$ multiplies the lengthscale of an isotropic base kernel $k_0$ with correlation $\rho$ and variance $\sigma^2$: $$ k(x, y) = \sigma^2 \left(\frac{2\,\ell(x)\,\ell(y)}{\ell(x)^2 + \ell(y)^2}\right)^{d/2} \rho\!\left(\sqrt{\frac{2}{\ell(x)^2 + \ell(y)^2}}\, \lVert x - y\rVert\right), \qquad \ell(x) = \exp g(x). $$ Here $d$ is the number of columns over which the base kernel measures distance, and $\lVert\cdot\rVert$ uses the base kernel's lengthscales, so an ARD base kernel keeps its shape and $\ell(x)$ scales it. Correlation decays faster where $\ell(x)$ is small, for example over mountains, and more slowly where it is large. The marginal variance is $\sigma^2$ at every location. The kernel is positive definite when $\rho$ is positive definite in every dimension. Only base kernels with `isotropic_radial = True` meet this condition: RBF, the Matérn kernels, RationalQuadratic and PoweredExponential. The base kernel selects the columns over which it measures distance with its own `active_dims`, and the location function selects its covariate columns. The wrapper itself always receives every column. """ name: ClassVar[str] = "Gibbs" base_kernel: AbstractKernel lengthscale: AbstractLocationFunction def __init__( self, base_kernel: AbstractKernel, lengthscale: AbstractLocationFunction, compute_engine: AbstractKernelComputation = DenseKernelComputation(), ): r"""Initialise the kernel. Args: base_kernel: the isotropic kernel $k_0$ whose lengthscale changes. lengthscale: the location function that gives $\log\ell(x)$. compute_engine: the computation engine that the kernel uses to compute its covariance matrices. Raises: TypeError: if `base_kernel` is not an isotropic radial kernel. """ if not getattr(base_kernel, "isotropic_radial", False): raise TypeError( "Gibbs needs an isotropic radial base kernel: RBF, Matern12, " "Matern32, Matern52, RationalQuadratic or PoweredExponential. " f"Got {type(base_kernel).__name__}, for which the Gibbs " "construction is not guaranteed to be positive definite." ) self.base_kernel = base_kernel self.lengthscale = lengthscale super().__init__(compute_engine=compute_engine) def __call__(self, x: Float[Array, " D"], y: Float[Array, " D"]) -> ScalarFloat: log_lx = self.lengthscale(x) log_ly = self.lengthscale(y) # log(ℓ(x)² + ℓ(y)²), computed stably for large or small lengthscales. log_sum = jnp.logaddexp(2.0 * log_lx, 2.0 * log_ly) log_ratio = jnp.log(2.0) + log_lx + log_ly - log_sum dims = self.base_kernel.slice_input(x).shape[-1] prefactor = jnp.exp(0.5 * dims * log_ratio) # Scaling both inputs by the same factor scales their distance, so the # base kernel evaluates ρ at the Gibbs distance with its own variance. scale = jnp.exp(0.5 * (jnp.log(2.0) - log_sum)) return (prefactor * self.base_kernel(scale * x, scale * y)).squeeze()
__all__ = ["Gibbs"]