Skip to content

GP API

The full GP stack: kernel math functions, concrete Parameterized kernel classes, model-facing entry points (GPPrior, ConditionedGP, gp_factor, gp_sample), sparse variational GPs with inter-domain inducing features, variational guides and likelihoods, non-Gaussian inference strategies (Laplace, Gauss-Newton, EP, posterior linearization, quasi-Newton — dense and Markov), pathwise (Matheron) posterior samplers, state-space (Kalman) GPs, and multi-output kernels. Scalable matrix construction and solver strategies (numerically stable assembly, implicit operators, batched matvec, Cholesky / CG / BBMM / LSMR / SLQ) live in gaussx.

Split with gaussx

pyrox owns the kernel function side — closed-form math primitives readable in a dozen lines — plus the NumPyro-aware model shell (GPPrior, gp_factor, gp_sample). gaussx owns every piece of linear algebra: stable matrix construction, solver strategies, and the underlying MultivariateNormal distribution. The model entry points accept any gaussx.AbstractSolverStrategy (default gaussx.DenseSolver()).

Model entry points

import jax.numpy as jnp
import numpyro
from pyrox_gp import GPPrior, RBF, gp_factor, gp_sample

def regression_model(X, y):
    """Collapsed Gaussian-likelihood GP regression."""
    kernel = RBF()
    prior = GPPrior(kernel=kernel, X=X)
    gp_factor("obs", prior, y, noise_var=jnp.array(0.05))


def latent_model(X):
    """Latent-function GP for non-conjugate likelihoods."""
    kernel = RBF()
    prior = GPPrior(kernel=kernel, X=X)
    f = gp_sample("f", prior)
    # ... attach any likelihood to f here, e.g. Bernoulli or Poisson.

Swap the solver strategy at construction time:

from gaussx import CGSolver, ComposedSolver, DenseLogdet, DenseSolver
prior = GPPrior(kernel=RBF(), X=X, solver=CGSolver())
# Or compose — CG for solve, dense Cholesky for logdet:
prior = GPPrior(
    kernel=RBF(), X=X,
    solver=ComposedSolver(solve_strategy=CGSolver(), logdet_strategy=DenseLogdet()),
)

GPPrior

Bases: Module

Finite-dimensional GP prior over a fixed training input set.

Holds a kernel, training inputs X, an optional mean function, an optional solver strategy, and a small diagonal jitter for numerical stability on otherwise-singular prior covariances.

Attributes:

Name Type Description
kernel Kernel

Any pyrox_gp.Kernel — evaluated on X.

X Float[Array, 'N D']

Training inputs of shape (N, D).

mean_fn Callable[[Float[Array, 'N D']], Float[Array, ' N']] | None

Callable X -> (N,) or None for the zero mean.

solver AbstractSolverStrategy | None

Any gaussx.AbstractSolverStrategy. Defaults to gaussx.DenseSolver() — swap for CGSolver, BBMMSolver, ComposedSolver(solve=..., logdet=...), etc.

jitter float

Diagonal regularization added to the prior covariance for numerical stability. Not a noise model — use noise_var on condition for that.

Source code in packages/pyrox-gp/src/pyrox_gp/_models.py
class GPPrior(eqx.Module):
    """Finite-dimensional GP prior over a fixed training input set.

    Holds a kernel, training inputs ``X``, an optional mean function, an
    optional solver strategy, and a small diagonal jitter for numerical
    stability on otherwise-singular prior covariances.

    Attributes:
        kernel: Any `pyrox_gp.Kernel` — evaluated on ``X``.
        X: Training inputs of shape ``(N, D)``.
        mean_fn: Callable ``X -> (N,)`` or ``None`` for the zero mean.
        solver: Any ``gaussx.AbstractSolverStrategy``. Defaults to
            ``gaussx.DenseSolver()`` — swap for ``CGSolver``,
            ``BBMMSolver``, ``ComposedSolver(solve=..., logdet=...)``, etc.
        jitter: Diagonal regularization added to the prior covariance
            for numerical stability. Not a noise model — use
            ``noise_var`` on `condition` for that.
    """

    kernel: Kernel
    X: Float[Array, "N D"]
    mean_fn: Callable[[Float[Array, "N D"]], Float[Array, " N"]] | None = None
    solver: AbstractSolverStrategy | None = None
    jitter: float = 1e-6

    def mean(self, X: Float[Array, "N D"]) -> Float[Array, " N"]:
        """Evaluate the mean function at ``X``; zero by default."""
        if self.mean_fn is None:
            return jnp.zeros(X.shape[0], dtype=X.dtype)
        return self.mean_fn(X)

    def _prior_operator(self) -> lx.AbstractLinearOperator:
        K = self.kernel(self.X, self.X)
        K = K.at[jnp.diag_indices_from(K)].add(self.jitter)
        return _psd_operator(K)

    def _noisy_operator(self, noise_var: Float[Array, ""]) -> lx.AbstractLinearOperator:
        K = self.kernel(self.X, self.X)
        K = K.at[jnp.diag_indices_from(K)].add(self.jitter + noise_var)
        return _psd_operator(K)

    def _resolved_solver(self) -> AbstractSolverStrategy:
        return DenseSolver() if self.solver is None else self.solver

    def log_prob(self, f: Float[Array, " N"]) -> Float[Array, ""]:
        r"""Marginal log-density of ``f`` under the GP prior.

        Computes $\log \mathcal{N}(f \mid \mu(X), K(X, X) + \text{jitter}\,I)$
        using `gaussx.log_marginal_likelihood`, so any solver strategy
        on this prior applies.
        """
        return log_marginal_likelihood(
            self.mean(self.X),
            self._prior_operator(),
            f,
            solver=self._resolved_solver(),
        )

    def sample(self, key: Array) -> Float[Array, " N"]:
        r"""Draw ``f \sim p(f) = \mathcal{N}(\mu(X), K + \text{jitter}\,I)``.

        Wraps the prior in a `gaussx.MultivariateNormal` with
        the configured `solver`. This is the non-NumPyro analogue
        of `gp_sample` — useful for tests, diagnostics, and
        prior-sample initialization without registering a sample site.
        """
        op = self._prior_operator()
        loc = self.mean(self.X)
        mvn = MultivariateNormal(loc, op, solver=self._resolved_solver())
        return mvn.sample(key)

    def condition(
        self,
        y: Float[Array, " N"],
        noise_var: Float[Array, ""],
    ) -> ConditionedGP:
        """Condition on Gaussian-likelihood observations ``y``.

        Precomputes
        ``alpha = (K + (jitter + noise_var) * I)^{-1} (y - mu(X))`` and
        caches it in the returned `ConditionedGP`. The same
        ``jitter`` regularization configured on this prior is included
        alongside ``noise_var`` in the conditioned operator and solve, so
        every downstream predict / sample call sees the regularized
        covariance.

        The operator construction and any subsequent hyperparameter
        capture share one `_kernel_context`, so for Pattern B/C
        kernels with priors the cached operator and the resolved
        hyperparameters on the returned `ConditionedGP` come from
        the same draw. Downstream consumers (notably
        `pyrox_gp.PathwiseSampler`) reuse those values to stay
        consistent with the cached operator.
        """
        with _kernel_context(self.kernel):
            operator = self._noisy_operator(noise_var)
            resolved_hyperparams = _resolve_kernel_hyperparams(self.kernel)
        residual = y - self.mean(self.X)
        cache = build_prediction_cache(
            operator, residual, solver=self._resolved_solver()
        )
        return ConditionedGP(
            prior=self,
            y=y,
            noise_var=noise_var,
            cache=cache,
            operator=operator,
            resolved_hyperparams=resolved_hyperparams,
        )

    def condition_nongauss(
        self,
        likelihood: Likelihood,
        y: Float[Array, " N"],
        *,
        strategy: _NonGaussStrategy,
    ) -> NonGaussConditionedGP:
        """Condition on a non-Gaussian likelihood via a site-based strategy.

        Convenience that forwards to ``strategy.fit(self, likelihood, y)``.
        Pick any of the site-based strategies in
        `pyrox_gp._inference_nongauss`:
        `pyrox_gp.LaplaceInference`,
        `pyrox_gp.GaussNewtonInference`,
        `pyrox_gp.PosteriorLinearization`,
        `pyrox_gp.ExpectationPropagation`, or
        `pyrox_gp.QuasiNewtonInference`. Returns a
        `pyrox_gp.NonGaussConditionedGP` with the same
        ``predict`` / ``predict_mean`` / ``predict_var`` API as the
        Gaussian-likelihood `ConditionedGP`.

        Examples:

            from pyrox_gp import (
                BernoulliLikelihood,
                ExpectationPropagation,
                GPPrior,
                RBF,
            )

            prior = GPPrior(kernel=RBF(), X=X)
            cond = prior.condition_nongauss(
                BernoulliLikelihood(), y,
                strategy=ExpectationPropagation(),
            )
            mean, var = cond.predict(X_star)
        """
        return strategy.fit(self, likelihood, y)

mean(X: Float[Array, 'N D']) -> Float[Array, ' N']

Evaluate the mean function at X; zero by default.

Source code in packages/pyrox-gp/src/pyrox_gp/_models.py
def mean(self, X: Float[Array, "N D"]) -> Float[Array, " N"]:
    """Evaluate the mean function at ``X``; zero by default."""
    if self.mean_fn is None:
        return jnp.zeros(X.shape[0], dtype=X.dtype)
    return self.mean_fn(X)

log_prob(f: Float[Array, ' N']) -> Float[Array, '']

Marginal log-density of f under the GP prior.

Computes \(\log \mathcal{N}(f \mid \mu(X), K(X, X) + \text{jitter}\,I)\) using gaussx.log_marginal_likelihood, so any solver strategy on this prior applies.

Source code in packages/pyrox-gp/src/pyrox_gp/_models.py
def log_prob(self, f: Float[Array, " N"]) -> Float[Array, ""]:
    r"""Marginal log-density of ``f`` under the GP prior.

    Computes $\log \mathcal{N}(f \mid \mu(X), K(X, X) + \text{jitter}\,I)$
    using `gaussx.log_marginal_likelihood`, so any solver strategy
    on this prior applies.
    """
    return log_marginal_likelihood(
        self.mean(self.X),
        self._prior_operator(),
        f,
        solver=self._resolved_solver(),
    )

sample(key: Array) -> Float[Array, ' N']

Draw f \sim p(f) = \mathcal{N}(\mu(X), K + \text{jitter}\,I).

Wraps the prior in a gaussx.MultivariateNormal with the configured solver. This is the non-NumPyro analogue of gp_sample — useful for tests, diagnostics, and prior-sample initialization without registering a sample site.

Source code in packages/pyrox-gp/src/pyrox_gp/_models.py
def sample(self, key: Array) -> Float[Array, " N"]:
    r"""Draw ``f \sim p(f) = \mathcal{N}(\mu(X), K + \text{jitter}\,I)``.

    Wraps the prior in a `gaussx.MultivariateNormal` with
    the configured `solver`. This is the non-NumPyro analogue
    of `gp_sample` — useful for tests, diagnostics, and
    prior-sample initialization without registering a sample site.
    """
    op = self._prior_operator()
    loc = self.mean(self.X)
    mvn = MultivariateNormal(loc, op, solver=self._resolved_solver())
    return mvn.sample(key)

condition(y: Float[Array, ' N'], noise_var: Float[Array, '']) -> ConditionedGP

Condition on Gaussian-likelihood observations y.

Precomputes alpha = (K + (jitter + noise_var) * I)^{-1} (y - mu(X)) and caches it in the returned ConditionedGP. The same jitter regularization configured on this prior is included alongside noise_var in the conditioned operator and solve, so every downstream predict / sample call sees the regularized covariance.

The operator construction and any subsequent hyperparameter capture share one _kernel_context, so for Pattern B/C kernels with priors the cached operator and the resolved hyperparameters on the returned ConditionedGP come from the same draw. Downstream consumers (notably pyrox_gp.PathwiseSampler) reuse those values to stay consistent with the cached operator.

Source code in packages/pyrox-gp/src/pyrox_gp/_models.py
def condition(
    self,
    y: Float[Array, " N"],
    noise_var: Float[Array, ""],
) -> ConditionedGP:
    """Condition on Gaussian-likelihood observations ``y``.

    Precomputes
    ``alpha = (K + (jitter + noise_var) * I)^{-1} (y - mu(X))`` and
    caches it in the returned `ConditionedGP`. The same
    ``jitter`` regularization configured on this prior is included
    alongside ``noise_var`` in the conditioned operator and solve, so
    every downstream predict / sample call sees the regularized
    covariance.

    The operator construction and any subsequent hyperparameter
    capture share one `_kernel_context`, so for Pattern B/C
    kernels with priors the cached operator and the resolved
    hyperparameters on the returned `ConditionedGP` come from
    the same draw. Downstream consumers (notably
    `pyrox_gp.PathwiseSampler`) reuse those values to stay
    consistent with the cached operator.
    """
    with _kernel_context(self.kernel):
        operator = self._noisy_operator(noise_var)
        resolved_hyperparams = _resolve_kernel_hyperparams(self.kernel)
    residual = y - self.mean(self.X)
    cache = build_prediction_cache(
        operator, residual, solver=self._resolved_solver()
    )
    return ConditionedGP(
        prior=self,
        y=y,
        noise_var=noise_var,
        cache=cache,
        operator=operator,
        resolved_hyperparams=resolved_hyperparams,
    )

condition_nongauss(likelihood: Likelihood, y: Float[Array, ' N'], *, strategy: _NonGaussStrategy) -> NonGaussConditionedGP

Condition on a non-Gaussian likelihood via a site-based strategy.

Convenience that forwards to strategy.fit(self, likelihood, y). Pick any of the site-based strategies in pyrox_gp._inference_nongauss: pyrox_gp.LaplaceInference, pyrox_gp.GaussNewtonInference, pyrox_gp.PosteriorLinearization, pyrox_gp.ExpectationPropagation, or pyrox_gp.QuasiNewtonInference. Returns a pyrox_gp.NonGaussConditionedGP with the same predict / predict_mean / predict_var API as the Gaussian-likelihood ConditionedGP.

Examples:

from pyrox_gp import (
    BernoulliLikelihood,
    ExpectationPropagation,
    GPPrior,
    RBF,
)

prior = GPPrior(kernel=RBF(), X=X)
cond = prior.condition_nongauss(
    BernoulliLikelihood(), y,
    strategy=ExpectationPropagation(),
)
mean, var = cond.predict(X_star)
Source code in packages/pyrox-gp/src/pyrox_gp/_models.py
def condition_nongauss(
    self,
    likelihood: Likelihood,
    y: Float[Array, " N"],
    *,
    strategy: _NonGaussStrategy,
) -> NonGaussConditionedGP:
    """Condition on a non-Gaussian likelihood via a site-based strategy.

    Convenience that forwards to ``strategy.fit(self, likelihood, y)``.
    Pick any of the site-based strategies in
    `pyrox_gp._inference_nongauss`:
    `pyrox_gp.LaplaceInference`,
    `pyrox_gp.GaussNewtonInference`,
    `pyrox_gp.PosteriorLinearization`,
    `pyrox_gp.ExpectationPropagation`, or
    `pyrox_gp.QuasiNewtonInference`. Returns a
    `pyrox_gp.NonGaussConditionedGP` with the same
    ``predict`` / ``predict_mean`` / ``predict_var`` API as the
    Gaussian-likelihood `ConditionedGP`.

    Examples:

        from pyrox_gp import (
            BernoulliLikelihood,
            ExpectationPropagation,
            GPPrior,
            RBF,
        )

        prior = GPPrior(kernel=RBF(), X=X)
        cond = prior.condition_nongauss(
            BernoulliLikelihood(), y,
            strategy=ExpectationPropagation(),
        )
        mean, var = cond.predict(X_star)
    """
    return strategy.fit(self, likelihood, y)

ConditionedGP

Bases: Module

GP conditioned on Gaussian-likelihood training observations.

Holds the precomputed training solve alpha (via gaussx.PredictionCache) and the noisy covariance operator so predictions at multiple test sets reuse the training solve.

Source code in packages/pyrox-gp/src/pyrox_gp/_models.py
class ConditionedGP(eqx.Module):
    """GP conditioned on Gaussian-likelihood training observations.

    Holds the precomputed training solve ``alpha`` (via
    `gaussx.PredictionCache`) and the noisy covariance operator so
    predictions at multiple test sets reuse the training solve.
    """

    prior: GPPrior
    y: Float[Array, " N"]
    noise_var: Float[Array, ""]
    cache: PredictionCache
    operator: lx.AbstractLinearOperator
    resolved_hyperparams: tuple[Float[Array, ""], Float[Array, ""]] | None = None

    def predict_mean(self, X_star: Float[Array, "M D"]) -> Float[Array, " M"]:
        r"""$\mu_* = \mu(X_*) + K_{*f}\,\alpha$."""
        with _kernel_context(self.prior.kernel):
            K_cross = self.prior.kernel(X_star, self.prior.X)
        return self.prior.mean(X_star) + predict_mean(self.cache, K_cross)

    def predict_var(self, X_star: Float[Array, "M D"]) -> Float[Array, " M"]:
        r"""Diagonal predictive variance at ``X_*``.

        $$
        \sigma^2_{*,i} = k(x_{*,i}, x_{*,i})
            - K_{*f}[i,:] \cdot (K + \sigma^2 I)^{-1} K_{f*}[:,i]
        $$

        ``K_cross`` and ``K_diag`` are computed under one shared kernel
        context so Pattern B / C kernels with prior'd hyperparameters
        register their NumPyro sites once and reuse them across both
        kernel calls (and the cached training solve).
        """
        with _kernel_context(self.prior.kernel):
            K_cross = self.prior.kernel(X_star, self.prior.X)
            K_diag = self.prior.kernel.diag(X_star)
        return predict_variance(
            K_cross,
            K_diag,
            self.operator,
            solver=self.prior._resolved_solver(),
        )

    def predict(
        self, X_star: Float[Array, "M D"]
    ) -> tuple[Float[Array, " M"], Float[Array, " M"]]:
        """Return ``(mean, variance)`` at ``X_*`` as a tuple.

        Both kernel evaluations share a single kernel context; see
        `predict_var`.
        """
        with _kernel_context(self.prior.kernel):
            return self.predict_mean(X_star), self.predict_var(X_star)

    def sample(
        self,
        key: Array,
        X_star: Float[Array, "M D"],
        n_samples: int = 1,
    ) -> Float[Array, "S M"]:
        """Sample from the diagonal predictive ``N(mean, diag(var))``.

        Returns samples independently per test point; correlated joint
        samples from the full predictive covariance are not covered by
        the Wave 2 dense surface. For correlated samples, build the full
        predictive covariance explicitly and draw from
        `gaussx.MultivariateNormal`.
        """
        with _kernel_context(self.prior.kernel):
            mean = self.predict_mean(X_star)
            var = self.predict_var(X_star)
        std = jnp.sqrt(jnp.clip(var, min=0.0))
        eps = jax.random.normal(key, (n_samples, X_star.shape[0]), dtype=mean.dtype)
        # Scale per-point std across the S sample rows: (M,) ⊙ (S, M) → (S, M).
        return einx.multiply("m, s m -> s m", std, eps) + mean

predict_mean(X_star: Float[Array, 'M D']) -> Float[Array, ' M']

\(\mu_* = \mu(X_*) + K_{*f}\,\alpha\).

Source code in packages/pyrox-gp/src/pyrox_gp/_models.py
def predict_mean(self, X_star: Float[Array, "M D"]) -> Float[Array, " M"]:
    r"""$\mu_* = \mu(X_*) + K_{*f}\,\alpha$."""
    with _kernel_context(self.prior.kernel):
        K_cross = self.prior.kernel(X_star, self.prior.X)
    return self.prior.mean(X_star) + predict_mean(self.cache, K_cross)

predict_var(X_star: Float[Array, 'M D']) -> Float[Array, ' M']

Diagonal predictive variance at X_*.

\[ \sigma^2_{*,i} = k(x_{*,i}, x_{*,i}) - K_{*f}[i,:] \cdot (K + \sigma^2 I)^{-1} K_{f*}[:,i] \]

K_cross and K_diag are computed under one shared kernel context so Pattern B / C kernels with prior'd hyperparameters register their NumPyro sites once and reuse them across both kernel calls (and the cached training solve).

Source code in packages/pyrox-gp/src/pyrox_gp/_models.py
def predict_var(self, X_star: Float[Array, "M D"]) -> Float[Array, " M"]:
    r"""Diagonal predictive variance at ``X_*``.

    $$
    \sigma^2_{*,i} = k(x_{*,i}, x_{*,i})
        - K_{*f}[i,:] \cdot (K + \sigma^2 I)^{-1} K_{f*}[:,i]
    $$

    ``K_cross`` and ``K_diag`` are computed under one shared kernel
    context so Pattern B / C kernels with prior'd hyperparameters
    register their NumPyro sites once and reuse them across both
    kernel calls (and the cached training solve).
    """
    with _kernel_context(self.prior.kernel):
        K_cross = self.prior.kernel(X_star, self.prior.X)
        K_diag = self.prior.kernel.diag(X_star)
    return predict_variance(
        K_cross,
        K_diag,
        self.operator,
        solver=self.prior._resolved_solver(),
    )

predict(X_star: Float[Array, 'M D']) -> tuple[Float[Array, ' M'], Float[Array, ' M']]

Return (mean, variance) at X_* as a tuple.

Both kernel evaluations share a single kernel context; see predict_var.

Source code in packages/pyrox-gp/src/pyrox_gp/_models.py
def predict(
    self, X_star: Float[Array, "M D"]
) -> tuple[Float[Array, " M"], Float[Array, " M"]]:
    """Return ``(mean, variance)`` at ``X_*`` as a tuple.

    Both kernel evaluations share a single kernel context; see
    `predict_var`.
    """
    with _kernel_context(self.prior.kernel):
        return self.predict_mean(X_star), self.predict_var(X_star)

sample(key: Array, X_star: Float[Array, 'M D'], n_samples: int = 1) -> Float[Array, 'S M']

Sample from the diagonal predictive N(mean, diag(var)).

Returns samples independently per test point; correlated joint samples from the full predictive covariance are not covered by the Wave 2 dense surface. For correlated samples, build the full predictive covariance explicitly and draw from gaussx.MultivariateNormal.

Source code in packages/pyrox-gp/src/pyrox_gp/_models.py
def sample(
    self,
    key: Array,
    X_star: Float[Array, "M D"],
    n_samples: int = 1,
) -> Float[Array, "S M"]:
    """Sample from the diagonal predictive ``N(mean, diag(var))``.

    Returns samples independently per test point; correlated joint
    samples from the full predictive covariance are not covered by
    the Wave 2 dense surface. For correlated samples, build the full
    predictive covariance explicitly and draw from
    `gaussx.MultivariateNormal`.
    """
    with _kernel_context(self.prior.kernel):
        mean = self.predict_mean(X_star)
        var = self.predict_var(X_star)
    std = jnp.sqrt(jnp.clip(var, min=0.0))
    eps = jax.random.normal(key, (n_samples, X_star.shape[0]), dtype=mean.dtype)
    # Scale per-point std across the S sample rows: (M,) ⊙ (S, M) → (S, M).
    return einx.multiply("m, s m -> s m", std, eps) + mean

gp_factor(name: str, prior: GPPrior, y: Float[Array, ' N'], noise_var: Float[Array, '']) -> None

Register the collapsed GP log marginal likelihood with NumPyro.

Adds log p(y | X, theta) = log N(y | mu, K + (jitter + sigma^2) I) to the NumPyro trace as numpyro.factor(name, ...). The prior's jitter is included in addition to the observation noise variance so the covariance matches what GPPrior.condition builds. Use this inside a NumPyro model when the likelihood is Gaussian and you want the latent function marginalized analytically.

Source code in packages/pyrox-gp/src/pyrox_gp/_models.py
def gp_factor(
    name: str,
    prior: GPPrior,
    y: Float[Array, " N"],
    noise_var: Float[Array, ""],
) -> None:
    """Register the collapsed GP log marginal likelihood with NumPyro.

    Adds
    ``log p(y | X, theta) = log N(y | mu, K + (jitter + sigma^2) I)``
    to the NumPyro trace as ``numpyro.factor(name, ...)``. The prior's
    ``jitter`` is included in addition to the observation noise variance
    so the covariance matches what `GPPrior.condition` builds. Use
    this inside a NumPyro model when the likelihood is Gaussian and you
    want the latent function marginalized analytically.
    """
    logp = log_marginal_likelihood(
        prior.mean(prior.X),
        prior._noisy_operator(noise_var),
        y,
        solver=prior._resolved_solver(),
    )
    numpyro.factor(name, logp)

gp_sample(name: str, prior: GPPrior, *, whitened: bool = False, guide: object | None = None) -> Float[Array, ' N']

Sample a latent function f at the prior's training inputs.

Three mutually exclusive modes:

  • whitened=False, guide=None (default) — register a single numpyro.sample(name, MVN(mu, K + jitter I)) site. The latent function is sampled directly from the prior.
  • whitened=True, guide=None — register a unit-normal latent site f"{name}_u" with shape (N,) and return the deterministic value f = mu(X) + L u where L is the Cholesky factor of K + jitter I. This reparameterization is the standard fix for mean-field SVI on GP-correlated latents (Murray & Adams, 2010): a NumPyro auto-guide such as numpyro.infer.autoguide.AutoNormal then approximates the well-conditioned isotropic posterior over u instead of the ill-conditioned correlated posterior over f.
  • guide provided — delegate to guide.register(name, prior). Concrete variational guides (Wave 3) own their own parameterization, so combining whitened=True with guide is rejected.

Use this inside a NumPyro model for non-conjugate likelihoods, where the latent function cannot be marginalized analytically.

Source code in packages/pyrox-gp/src/pyrox_gp/_models.py
def gp_sample(
    name: str,
    prior: GPPrior,
    *,
    whitened: bool = False,
    guide: object | None = None,
) -> Float[Array, " N"]:
    r"""Sample a latent function ``f`` at the prior's training inputs.

    Three mutually exclusive modes:

    * ``whitened=False``, ``guide=None`` (default) — register a single
      ``numpyro.sample(name, MVN(mu, K + jitter I))`` site. The latent
      function is sampled directly from the prior.
    * ``whitened=True``, ``guide=None`` — register a unit-normal latent
      site ``f"{name}_u"`` with shape ``(N,)`` and return the
      deterministic value ``f = mu(X) + L u`` where ``L`` is the
      Cholesky factor of ``K + jitter I``. This reparameterization is the
      standard fix for mean-field SVI on GP-correlated latents
      (Murray & Adams, 2010): a NumPyro auto-guide such as
      `numpyro.infer.autoguide.AutoNormal` then approximates the
      well-conditioned isotropic posterior over ``u`` instead of the
      ill-conditioned correlated posterior over ``f``.
    * ``guide`` provided — delegate to ``guide.register(name, prior)``.
      Concrete variational guides (Wave 3) own their own
      parameterization, so combining ``whitened=True`` with ``guide`` is
      rejected.

    Use this inside a NumPyro model for non-conjugate likelihoods, where
    the latent function cannot be marginalized analytically.
    """
    if guide is not None:
        if whitened:
            raise ValueError(
                "gp_sample: cannot combine `whitened=True` with `guide=...`. "
                "Provide one or the other; concrete guides own their own "
                "parameterization."
            )
        return guide.register(name, prior)  # type: ignore[attr-defined]  # ty: ignore[unresolved-attribute]

    if whitened:
        L = cholesky(prior._prior_operator())
        n = prior.X.shape[0]
        dtype = prior.X.dtype
        u = numpyro.sample(
            f"{name}_u",
            dist.Normal(jnp.zeros(n, dtype=dtype), jnp.ones((), dtype=dtype)).to_event(
                1
            ),
        )
        f = prior.mean(prior.X) + unwhiten(jnp.asarray(u), L)
        return numpyro.deterministic(name, f)  # ty: ignore[invalid-return-type]

    return numpyro.sample(  # ty: ignore[invalid-return-type]
        name,
        MultivariateNormal(
            prior.mean(prior.X),
            prior._prior_operator(),
            solver=prior._resolved_solver(),
        ),
    )

Concrete kernels

Each Parameterized kernel registers its hyperparameters with positivity constraints where appropriate. Attach priors with set_prior, autoguides with autoguide, and flip set_mode("model" | "guide").

RBF

Bases: _ParameterizedKernel

Radial basis function (squared exponential) kernel.

input_dim: set to the input dimension D to fit a separate lengthscale per input dimension (ARD). Leave as None for a single isotropic lengthscale.

Source code in packages/pyrox-gp/src/pyrox_gp/_kernels.py
class RBF(_ParameterizedKernel):
    """Radial basis function (squared exponential) kernel.

    ``input_dim``: set to the input dimension ``D`` to fit a separate
    lengthscale per input dimension (ARD). Leave as ``None`` for a single
    isotropic lengthscale.
    """

    _frozen_cls = kl.RBF
    _frozen_params = ("variance", "lengthscale")

    pyrox_name: str = "RBF"
    init_variance: float = 1.0
    init_lengthscale: float = 1.0
    input_dim: int | None = None

    def setup(self) -> None:
        self.register_param(
            "variance",
            jnp.asarray(self.init_variance),
            constraint=dist.constraints.positive,
        )
        lengthscale = (
            jnp.asarray(self.init_lengthscale)
            if self.input_dim is None
            else jnp.full((self.input_dim,), self.init_lengthscale)
        )
        self.register_param(
            "lengthscale",
            lengthscale,
            constraint=dist.constraints.positive,
        )

    @pyrox_method
    def __call__(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "N1 N2"]:
        return _k.rbf_kernel(
            X1, X2, self.get_param("variance"), self.get_param("lengthscale")
        )

Matern

Bases: _ParameterizedKernel

Matern kernel with nu in {0.5, 1.5, 2.5}.

nu is a static class attribute — it selects a code path in the underlying math primitive and is not a trainable parameter.

input_dim: set to the input dimension D to fit a separate lengthscale per input dimension (ARD). Leave as None for a single isotropic lengthscale.

Source code in packages/pyrox-gp/src/pyrox_gp/_kernels.py
class Matern(_ParameterizedKernel):
    """Matern kernel with ``nu in {0.5, 1.5, 2.5}``.

    ``nu`` is a static class attribute — it selects a code path in the
    underlying math primitive and is not a trainable parameter.

    ``input_dim``: set to the input dimension ``D`` to fit a separate
    lengthscale per input dimension (ARD). Leave as ``None`` for a single
    isotropic lengthscale.
    """

    _frozen_cls = kl.Matern
    _frozen_params = ("variance", "lengthscale")
    _frozen_static = ("nu",)

    pyrox_name: str = "Matern"
    init_variance: float = 1.0
    init_lengthscale: float = 1.0
    nu: float = 2.5
    input_dim: int | None = None

    def setup(self) -> None:
        self.register_param(
            "variance",
            jnp.asarray(self.init_variance),
            constraint=dist.constraints.positive,
        )
        lengthscale = (
            jnp.asarray(self.init_lengthscale)
            if self.input_dim is None
            else jnp.full((self.input_dim,), self.init_lengthscale)
        )
        self.register_param(
            "lengthscale",
            lengthscale,
            constraint=dist.constraints.positive,
        )

    @pyrox_method
    def __call__(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "N1 N2"]:
        return _k.matern_kernel(
            X1,
            X2,
            self.get_param("variance"),
            self.get_param("lengthscale"),
            self.nu,
        )

Periodic

Bases: _ParameterizedKernel

Periodic (MacKay) kernel.

Source code in packages/pyrox-gp/src/pyrox_gp/_kernels.py
class Periodic(_ParameterizedKernel):
    """Periodic (MacKay) kernel."""

    _frozen_cls = kl.Periodic
    _frozen_params = ("variance", "lengthscale", "period")

    pyrox_name: str = "Periodic"
    init_variance: float = 1.0
    init_lengthscale: float = 1.0
    init_period: float = 1.0

    def setup(self) -> None:
        self.register_param(
            "variance",
            jnp.asarray(self.init_variance),
            constraint=dist.constraints.positive,
        )
        self.register_param(
            "lengthscale",
            jnp.asarray(self.init_lengthscale),
            constraint=dist.constraints.positive,
        )
        self.register_param(
            "period",
            jnp.asarray(self.init_period),
            constraint=dist.constraints.positive,
        )

    @pyrox_method
    def __call__(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "N1 N2"]:
        return _k.periodic_kernel(
            X1,
            X2,
            self.get_param("variance"),
            self.get_param("lengthscale"),
            self.get_param("period"),
        )

Linear

Bases: _ParameterizedKernel

Linear kernel sigma^2 x^T x' + bias.

bias is constrained nonnegative because k = sigma^2 X X^T + b 1 1^T is only PSD for b >= 0 (e.g. X = 0 gives eigenvalue N*b).

Source code in packages/pyrox-gp/src/pyrox_gp/_kernels.py
class Linear(_ParameterizedKernel):
    """Linear kernel ``sigma^2 x^T x' + bias``.

    ``bias`` is constrained nonnegative because ``k = sigma^2 X X^T + b 1 1^T``
    is only PSD for ``b >= 0`` (e.g. ``X = 0`` gives eigenvalue ``N*b``).
    """

    _frozen_cls = kl.Linear
    _frozen_params = ("variance", "bias")

    pyrox_name: str = "Linear"
    init_variance: float = 1.0
    init_bias: float = 0.0

    def setup(self) -> None:
        self.register_param(
            "variance",
            jnp.asarray(self.init_variance),
            constraint=dist.constraints.positive,
        )
        self.register_param(
            "bias",
            jnp.asarray(self.init_bias),
            constraint=dist.constraints.nonnegative,
        )

    @pyrox_method
    def __call__(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "N1 N2"]:
        return _k.linear_kernel(
            X1, X2, self.get_param("variance"), self.get_param("bias")
        )

    def diag(self, X: Float[Array, "N D"]) -> Float[Array, " N"]:
        # Non-stationary: diagonal depends on |X[i]|^2.
        v = self.get_param("variance")
        b = self.get_param("bias")
        return v * jnp.sum(X * X, axis=-1) + b

RationalQuadratic

Bases: _ParameterizedKernel

Rational quadratic kernel.

input_dim: set to the input dimension D to fit a separate lengthscale per input dimension (ARD). Leave as None for a single isotropic lengthscale.

Source code in packages/pyrox-gp/src/pyrox_gp/_kernels.py
class RationalQuadratic(_ParameterizedKernel):
    """Rational quadratic kernel.

    ``input_dim``: set to the input dimension ``D`` to fit a separate
    lengthscale per input dimension (ARD). Leave as ``None`` for a single
    isotropic lengthscale.
    """

    _frozen_cls = kl.RationalQuadratic
    _frozen_params = ("variance", "lengthscale", "alpha")

    pyrox_name: str = "RationalQuadratic"
    init_variance: float = 1.0
    init_lengthscale: float = 1.0
    init_alpha: float = 1.0
    input_dim: int | None = None

    def setup(self) -> None:
        self.register_param(
            "variance",
            jnp.asarray(self.init_variance),
            constraint=dist.constraints.positive,
        )
        lengthscale = (
            jnp.asarray(self.init_lengthscale)
            if self.input_dim is None
            else jnp.full((self.input_dim,), self.init_lengthscale)
        )
        self.register_param(
            "lengthscale",
            lengthscale,
            constraint=dist.constraints.positive,
        )
        self.register_param(
            "alpha",
            jnp.asarray(self.init_alpha),
            constraint=dist.constraints.positive,
        )

    @pyrox_method
    def __call__(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "N1 N2"]:
        return _k.rational_quadratic_kernel(
            X1,
            X2,
            self.get_param("variance"),
            self.get_param("lengthscale"),
            self.get_param("alpha"),
        )

Polynomial

Bases: _ParameterizedKernel

Polynomial kernel sigma^2 (x^T x' + bias)^degree.

degree is a static class field (it selects an integer power, not an optimization target). bias is constrained nonnegative — the degree=1 case reduces to Linear and has the same PSD-requires-b>=0 failure mode.

Source code in packages/pyrox-gp/src/pyrox_gp/_kernels.py
class Polynomial(_ParameterizedKernel):
    """Polynomial kernel ``sigma^2 (x^T x' + bias)^degree``.

    ``degree`` is a static class field (it selects an integer power, not
    an optimization target). ``bias`` is constrained nonnegative — the
    ``degree=1`` case reduces to `Linear` and has the same
    PSD-requires-``b>=0`` failure mode.
    """

    _frozen_cls = kl.Polynomial
    _frozen_params = ("variance", "bias")
    _frozen_static = ("degree",)

    pyrox_name: str = "Polynomial"
    init_variance: float = 1.0
    init_bias: float = 0.0
    degree: int = 2

    def setup(self) -> None:
        self.register_param(
            "variance",
            jnp.asarray(self.init_variance),
            constraint=dist.constraints.positive,
        )
        self.register_param(
            "bias",
            jnp.asarray(self.init_bias),
            constraint=dist.constraints.nonnegative,
        )

    @pyrox_method
    def __call__(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "N1 N2"]:
        return _k.polynomial_kernel(
            X1,
            X2,
            self.get_param("variance"),
            self.get_param("bias"),
            self.degree,
        )

    def diag(self, X: Float[Array, "N D"]) -> Float[Array, " N"]:
        v = self.get_param("variance")
        b = self.get_param("bias")
        return v * (jnp.sum(X * X, axis=-1) + b) ** self.degree

Cosine

Bases: _ParameterizedKernel

Cosine kernel sigma^2 cos(2 pi ||x - x'|| / period).

Source code in packages/pyrox-gp/src/pyrox_gp/_kernels.py
class Cosine(_ParameterizedKernel):
    """Cosine kernel ``sigma^2 cos(2 pi ||x - x'|| / period)``."""

    _frozen_cls = kl.Cosine
    _frozen_params = ("variance", "period")

    pyrox_name: str = "Cosine"
    init_variance: float = 1.0
    init_period: float = 1.0

    def setup(self) -> None:
        self.register_param(
            "variance",
            jnp.asarray(self.init_variance),
            constraint=dist.constraints.positive,
        )
        self.register_param(
            "period",
            jnp.asarray(self.init_period),
            constraint=dist.constraints.positive,
        )

    @pyrox_method
    def __call__(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "N1 N2"]:
        return _k.cosine_kernel(
            X1, X2, self.get_param("variance"), self.get_param("period")
        )

White

Bases: _ParameterizedKernel

White-noise kernel sigma^2 delta(x, x').

Source code in packages/pyrox-gp/src/pyrox_gp/_kernels.py
class White(_ParameterizedKernel):
    """White-noise kernel ``sigma^2 delta(x, x')``."""

    _frozen_cls = kl.White
    _frozen_params = ("variance",)

    pyrox_name: str = "White"
    init_variance: float = 1.0

    def setup(self) -> None:
        self.register_param(
            "variance",
            jnp.asarray(self.init_variance),
            constraint=dist.constraints.positive,
        )

    @pyrox_method
    def __call__(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "N1 N2"]:
        return _k.white_kernel(X1, X2, self.get_param("variance"))

Constant

Bases: _ParameterizedKernel

Constant kernel k(x, x') = sigma^2.

Source code in packages/pyrox-gp/src/pyrox_gp/_kernels.py
class Constant(_ParameterizedKernel):
    """Constant kernel ``k(x, x') = sigma^2``."""

    _frozen_cls = kl.Constant
    _frozen_params = ("variance",)

    pyrox_name: str = "Constant"
    init_variance: float = 1.0

    def setup(self) -> None:
        self.register_param(
            "variance",
            jnp.asarray(self.init_variance),
            constraint=dist.constraints.positive,
        )

    @pyrox_method
    def __call__(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "N1 N2"]:
        return _k.constant_kernel(X1, X2, self.get_param("variance"))

Sparse-GP inducing features (#49)

Inter-domain inducing-feature families used to build scalable sparse GPs where the inducing-prior covariance K_uu becomes diagonal. Pass any of these to SparseGPPrior via the inducing= keyword in place of a raw point matrix Z.

from pyrox_gp import RBF, FourierInducingFeatures, SparseGPPrior

kernel   = RBF(init_lengthscale=0.3, init_variance=1.0)
features = FourierInducingFeatures.init(in_features=1, num_basis_per_dim=64, L=5.0)
prior    = SparseGPPrior(kernel=kernel, inducing=features)   # K_uu is diagonal!

InducingFeatures

Bases: Protocol

Protocol for inter-domain inducing features.

Implementations expose the inducing-prior covariance K_uu and the cross-covariance k_ux(X) between data points and inducing features. Diagonal-friendly concretions return lineax.DiagonalLinearOperator so the downstream solve dispatches to elementwise division.

Input shape is family-dependent. k_ux takes a batch of data points X in whatever representation the family consumes:

  • FourierInducingFeatures: coordinates (N, D).
  • SphericalHarmonicInducingFeatures: unit vectors (N, 3).
  • LaplacianInducingFeatures: integer node indices (N,).

Each implementation validates its own expected shape and dtype.

Source code in packages/pyrox-gp/src/pyrox_gp/_inducing.py
@runtime_checkable
class InducingFeatures(Protocol):
    """Protocol for inter-domain inducing features.

    Implementations expose the inducing-prior covariance ``K_uu`` and the
    cross-covariance ``k_ux(X)`` between data points and inducing
    features. Diagonal-friendly concretions return
    `lineax.DiagonalLinearOperator` so the downstream solve dispatches
    to elementwise division.

    **Input shape is family-dependent.** ``k_ux`` takes a batch of data
    points ``X`` in whatever representation the family consumes:

    - `FourierInducingFeatures`: coordinates ``(N, D)``.
    - `SphericalHarmonicInducingFeatures`: unit vectors ``(N, 3)``.
    - `LaplacianInducingFeatures`: integer node indices ``(N,)``.

    Each implementation validates its own expected shape and dtype.
    """

    @property
    def num_features(self) -> int: ...

    def K_uu(
        self, kernel: Kernel, *, jitter: float = 1e-6
    ) -> lx.AbstractLinearOperator: ...

    def k_ux(self, x: Array, kernel: Kernel) -> Float[Array, "N M"]: ...

FourierInducingFeatures

Bases: Module

VFF — Variational Fourier inducing features on \([-L, L]^D\).

For a stationary kernel with spectral density \(S(\cdot)\), the basis \(\{\phi_j\}\) of Laplacian eigenfunctions on the box gives

\[ K_{uu} = \mathrm{diag}\!\big(S(\sqrt{\lambda_j})\big), \qquad K_{ux}(x)_j = S(\sqrt{\lambda_j})\,\phi_j(x). \]

With this convention \(K_{ux} K_{uu}^{-1} = \phi_j(x)\), so the SVGP predictive mean reduces to a basis evaluation. K_{uu} is returned as a lineax.DiagonalLinearOperator to preserve the O(M) solve dispatch end-to-end.

Attributes:

Name Type Description
in_features int

Input dimension \(D\).

num_basis_per_dim tuple[int, ...]

Per-axis number of 1D eigenfunctions; total count is prod(num_basis_per_dim).

L tuple[float, ...]

Per-axis box half-width.

Source code in packages/pyrox-gp/src/pyrox_gp/_inducing.py
class FourierInducingFeatures(eqx.Module):
    r"""VFF — Variational Fourier inducing features on $[-L, L]^D$.

    For a stationary kernel with spectral density $S(\cdot)$, the
    basis $\{\phi_j\}$ of Laplacian eigenfunctions on the box gives

    $$
    K_{uu} = \mathrm{diag}\!\big(S(\sqrt{\lambda_j})\big),
    \qquad
    K_{ux}(x)_j = S(\sqrt{\lambda_j})\,\phi_j(x).
    $$

    With this convention $K_{ux} K_{uu}^{-1} = \phi_j(x)$, so the
    SVGP predictive mean reduces to a basis evaluation. ``K_{uu}`` is
    returned as a `lineax.DiagonalLinearOperator` to preserve the
    O(M) solve dispatch end-to-end.

    Attributes:
        in_features: Input dimension $D$.
        num_basis_per_dim: Per-axis number of 1D eigenfunctions; total
            count is ``prod(num_basis_per_dim)``.
        L: Per-axis box half-width.
    """

    in_features: int = eqx.field(static=True)
    num_basis_per_dim: tuple[int, ...] = eqx.field(static=True)
    L: tuple[float, ...] = eqx.field(static=True)

    @classmethod
    def init(
        cls,
        in_features: int,
        num_basis_per_dim: int | tuple[int, ...],
        L: float | tuple[float, ...],
    ) -> FourierInducingFeatures:
        M_per = _to_tuple(num_basis_per_dim, in_features, "num_basis_per_dim")
        L_per = _to_tuple(L, in_features, "L")
        if any(L_d <= 0 for L_d in L_per):
            raise ValueError(f"L must be all positive; got {L_per}.")
        if any(M_d < 1 for M_d in M_per):
            raise ValueError(f"num_basis_per_dim must be all >= 1; got {M_per}.")
        return cls(
            in_features=in_features,
            num_basis_per_dim=M_per,
            L=tuple(float(L_d) for L_d in L_per),
        )

    @property
    def num_features(self) -> int:
        n = 1
        for m in self.num_basis_per_dim:
            n *= m
        return n

    def _check_stationary(self, kernel: Kernel) -> None:
        if not _is_stationary(kernel):
            raise ValueError(
                f"FourierInducingFeatures requires a stationary kernel with a "
                f"registered spectral density (RBF or Matern); got "
                f"{type(kernel).__name__}."
            )

    def K_uu(
        self, kernel: Kernel, *, jitter: float = 1e-6
    ) -> lx.DiagonalLinearOperator:
        """Diagonal $K_{uu}$ — entries ``S(sqrt(lambda_j))`` plus jitter."""
        self._check_stationary(kernel)
        with _kernel_context(kernel):
            lam = fourier_eigenvalues(self.num_basis_per_dim, self.L, self.in_features)
            S = spectral_density(kernel, lam, D=self.in_features)
        return _diagonal_with_jitter(S, jitter)

    def k_ux(self, x: Float[Array, "N D"], kernel: Kernel) -> Float[Array, "N M"]:
        """Cross-covariance entries $S(\\sqrt{\\lambda_j})\\,\\phi_j(x)$."""
        self._check_stationary(kernel)
        if x.ndim != 2 or x.shape[-1] != self.in_features:
            raise ValueError(f"x must be (N, {self.in_features}); got shape {x.shape}.")
        with _kernel_context(kernel):
            Phi, lam = fourier_basis(x, self.num_basis_per_dim, self.L)
            S = spectral_density(kernel, lam, D=self.in_features)
        return Phi * S[None, :]

K_uu(kernel: Kernel, *, jitter: float = 1e-06) -> lx.DiagonalLinearOperator

Diagonal \(K_{uu}\) — entries S(sqrt(lambda_j)) plus jitter.

Source code in packages/pyrox-gp/src/pyrox_gp/_inducing.py
def K_uu(
    self, kernel: Kernel, *, jitter: float = 1e-6
) -> lx.DiagonalLinearOperator:
    """Diagonal $K_{uu}$ — entries ``S(sqrt(lambda_j))`` plus jitter."""
    self._check_stationary(kernel)
    with _kernel_context(kernel):
        lam = fourier_eigenvalues(self.num_basis_per_dim, self.L, self.in_features)
        S = spectral_density(kernel, lam, D=self.in_features)
    return _diagonal_with_jitter(S, jitter)

k_ux(x: Float[Array, 'N D'], kernel: Kernel) -> Float[Array, 'N M']

Cross-covariance entries \(S(\sqrt{\lambda_j})\,\phi_j(x)\).

Source code in packages/pyrox-gp/src/pyrox_gp/_inducing.py
def k_ux(self, x: Float[Array, "N D"], kernel: Kernel) -> Float[Array, "N M"]:
    """Cross-covariance entries $S(\\sqrt{\\lambda_j})\\,\\phi_j(x)$."""
    self._check_stationary(kernel)
    if x.ndim != 2 or x.shape[-1] != self.in_features:
        raise ValueError(f"x must be (N, {self.in_features}); got shape {x.shape}.")
    with _kernel_context(kernel):
        Phi, lam = fourier_basis(x, self.num_basis_per_dim, self.L)
        S = spectral_density(kernel, lam, D=self.in_features)
    return Phi * S[None, :]

SphericalHarmonicInducingFeatures

Bases: Module

VISH — inducing harmonics on \(S^2\) (Dutordoir et al. 2020).

For any zonal kernel \(k(x, x') = \kappa(x \cdot x')\) on the unit 2-sphere, the Funk-Hecke theorem gives a diagonal \(K_{uu}\) whose eigenvalues are the kernel's Funk-Hecke coefficients \(a_l\). The cross-covariance is \(a_l\,Y_{lm}(x)\).

Funk-Hecke coefficients are computed by Gauss-Legendre quadrature (arbitrary kernels supported, no closed form required). For kernels that have a closed-form Funk-Hecke series (RBF on S² via Bessel functions etc.), the numerical and analytic answers should agree to the quadrature tolerance.

Attributes:

Name Type Description
l_max int

Maximum harmonic degree, inclusive.

num_quadrature int

Gauss-Legendre nodes for the Funk-Hecke integral.

Source code in packages/pyrox-gp/src/pyrox_gp/_inducing.py
class SphericalHarmonicInducingFeatures(eqx.Module):
    r"""VISH — inducing harmonics on $S^2$ (Dutordoir et al. 2020).

    For any zonal kernel $k(x, x') = \kappa(x \cdot x')$ on the
    unit 2-sphere, the Funk-Hecke theorem gives a diagonal $K_{uu}$
    whose eigenvalues are the kernel's Funk-Hecke coefficients
    $a_l$. The cross-covariance is $a_l\,Y_{lm}(x)$.

    Funk-Hecke coefficients are computed by Gauss-Legendre quadrature
    (arbitrary kernels supported, no closed form required). For
    kernels that have a closed-form Funk-Hecke series (RBF on S² via
    Bessel functions etc.), the numerical and analytic answers should
    agree to the quadrature tolerance.

    Attributes:
        l_max: Maximum harmonic degree, inclusive.
        num_quadrature: Gauss-Legendre nodes for the Funk-Hecke integral.
    """

    l_max: int = eqx.field(static=True)
    num_quadrature: int = eqx.field(static=True, default=256)

    @classmethod
    def init(
        cls, l_max: int, *, num_quadrature: int = 256
    ) -> SphericalHarmonicInducingFeatures:
        if l_max < 0:
            raise ValueError(f"l_max must be >= 0; got {l_max}.")
        if num_quadrature < 1:
            raise ValueError(f"num_quadrature must be >= 1; got {num_quadrature}.")
        return cls(l_max=l_max, num_quadrature=num_quadrature)

    @property
    def num_features(self) -> int:
        return (self.l_max + 1) ** 2

    def _per_feature_coeffs(self, kernel: Kernel) -> Float[Array, " M"]:
        a = funk_hecke_coefficients(
            kernel, self.l_max, num_quadrature=self.num_quadrature
        )
        # Each l contributes 2l+1 features with the same coefficient.
        return jnp.concatenate(
            [jnp.full((2 * ell + 1,), a[ell]) for ell in range(self.l_max + 1)]
        )

    def K_uu(
        self, kernel: Kernel, *, jitter: float = 1e-6
    ) -> lx.DiagonalLinearOperator:
        """Diagonal $K_{uu}$ — Funk-Hecke coefficients per harmonic."""
        diag = self._per_feature_coeffs(kernel)
        return _diagonal_with_jitter(diag, jitter)

    def k_ux(
        self,
        unit_xyz: Float[Array, "N 3"],
        kernel: Kernel,
    ) -> Float[Array, "N M"]:
        r"""Cross-covariance: $a_l\,Y_{lm}(x)$."""
        if unit_xyz.ndim != 2 or unit_xyz.shape[-1] != 3:
            raise ValueError(f"unit_xyz must be (N, 3); got {unit_xyz.shape}.")
        Y = real_spherical_harmonics(unit_xyz, self.l_max)
        a_per_feature = self._per_feature_coeffs(kernel)
        return Y * a_per_feature[None, :]

K_uu(kernel: Kernel, *, jitter: float = 1e-06) -> lx.DiagonalLinearOperator

Diagonal \(K_{uu}\) — Funk-Hecke coefficients per harmonic.

Source code in packages/pyrox-gp/src/pyrox_gp/_inducing.py
def K_uu(
    self, kernel: Kernel, *, jitter: float = 1e-6
) -> lx.DiagonalLinearOperator:
    """Diagonal $K_{uu}$ — Funk-Hecke coefficients per harmonic."""
    diag = self._per_feature_coeffs(kernel)
    return _diagonal_with_jitter(diag, jitter)

k_ux(unit_xyz: Float[Array, 'N 3'], kernel: Kernel) -> Float[Array, 'N M']

Cross-covariance: \(a_l\,Y_{lm}(x)\).

Source code in packages/pyrox-gp/src/pyrox_gp/_inducing.py
def k_ux(
    self,
    unit_xyz: Float[Array, "N 3"],
    kernel: Kernel,
) -> Float[Array, "N M"]:
    r"""Cross-covariance: $a_l\,Y_{lm}(x)$."""
    if unit_xyz.ndim != 2 or unit_xyz.shape[-1] != 3:
        raise ValueError(f"unit_xyz must be (N, 3); got {unit_xyz.shape}.")
    Y = real_spherical_harmonics(unit_xyz, self.l_max)
    a_per_feature = self._per_feature_coeffs(kernel)
    return Y * a_per_feature[None, :]

SlepianInducingFeatures

Bases: Module

Region-localized Slepian inducing features on \(S^2\).

The retained Slepian functions are linear combinations G = Y C of the real spherical-harmonic basis evaluated in the cap-centred frame. For a zonal kernel with Funk-Hecke coefficients a_l this gives dense inducing covariance K_uu = C.T diag(a_l) C and cross-covariance K_ux = Y(R x) diag(a_l) C, where R is the rotation aligning the cap centre with the north pole. The basis (a SlepianCapBasis) is built once at init time and stored on the module so that K_uu, k_ux and num_features are cheap matrix multiplies.

Source code in packages/pyrox-gp/src/pyrox_gp/_inducing.py
class SlepianInducingFeatures(eqx.Module):
    r"""Region-localized Slepian inducing features on $S^2$.

    The retained Slepian functions are linear combinations ``G = Y C`` of the
    real spherical-harmonic basis evaluated in the cap-centred frame. For a
    zonal kernel with Funk-Hecke coefficients ``a_l`` this gives dense
    inducing covariance ``K_uu = C.T diag(a_l) C`` and cross-covariance
    ``K_ux = Y(R x) diag(a_l) C``, where ``R`` is the rotation aligning the
    cap centre with the north pole. The basis (a `SlepianCapBasis`)
    is built once at `init` time and stored on the module so that
    ``K_uu``, ``k_ux`` and ``num_features`` are cheap matrix multiplies.
    """

    l_max: int = eqx.field(static=True)
    cap_radius_deg: float = eqx.field(static=True)
    cap_centre_lonlat_deg: tuple[float, float] = eqx.field(static=True)
    eig_threshold: float = eqx.field(static=True)
    n_modes: int | None = eqx.field(static=True)
    num_quadrature: int = eqx.field(static=True)
    basis_num_quadrature: int | None = eqx.field(static=True)
    basis: SlepianCapBasis

    @classmethod
    def init(
        cls,
        *,
        l_max: int,
        cap_radius_deg: float,
        cap_centre_lonlat_deg: tuple[float, float],
        eig_threshold: float = 0.05,
        n_modes: int | None = None,
        num_quadrature: int = 256,
        basis_num_quadrature: int | None = None,
    ) -> SlepianInducingFeatures:
        if l_max < 0:
            raise ValueError(f"l_max must be >= 0; got {l_max}.")
        if cap_radius_deg <= 0.0 or cap_radius_deg > 180.0:
            raise ValueError(
                f"cap_radius_deg must lie in (0, 180]; got {cap_radius_deg}."
            )
        if num_quadrature < 1:
            raise ValueError(f"num_quadrature must be >= 1; got {num_quadrature}.")
        if n_modes is not None and n_modes < 1:
            raise ValueError(f"n_modes must be >= 1; got {n_modes}.")
        if len(cap_centre_lonlat_deg) != 2:
            raise ValueError(
                "cap_centre_lonlat_deg must contain (lon, lat); "
                f"got {cap_centre_lonlat_deg}."
            )
        lon, lat = cap_centre_lonlat_deg
        # Build the basis once at construction so K_uu / k_ux are matrix
        # multiplies; cap geometry is static, so the eigensolve does not
        # rerun per call.
        basis = slepian_cap_basis(
            l_max,
            jnp.deg2rad(cap_radius_deg),
            n_modes=n_modes,
            eig_threshold=eig_threshold,
            lonlat_centre=jnp.deg2rad(jnp.asarray((float(lon), float(lat)))),
            num_quadrature=basis_num_quadrature,
        )
        return cls(
            l_max=l_max,
            cap_radius_deg=float(cap_radius_deg),
            cap_centre_lonlat_deg=(float(lon), float(lat)),
            eig_threshold=float(eig_threshold),
            n_modes=n_modes,
            num_quadrature=num_quadrature,
            basis_num_quadrature=basis_num_quadrature,
            basis=basis,
        )

    @property
    def num_features(self) -> int:
        return self.basis.num_modes

    def _per_feature_coeffs(self, kernel: Kernel) -> Float[Array, " M"]:
        a = funk_hecke_coefficients(
            kernel, self.l_max, num_quadrature=self.num_quadrature
        )
        return a[jnp.asarray(harmonic_degrees(self.l_max))]

    def K_uu(self, kernel: Kernel, *, jitter: float = 1e-6) -> lx.MatrixLinearOperator:
        """Dense Slepian inducing covariance with diagonal jitter."""
        a_per_feature = self._per_feature_coeffs(kernel)
        # Scale each feature row by its Funk-Hecke coefficient, then form the
        # Gram K = Φᵀ diag(a) Φ by contracting the feature axis f.
        weighted_coeffs = einx.multiply(
            "f m, f -> f m", self.basis.coeffs, a_per_feature
        )
        K = einx.dot("f i, f j -> i j", self.basis.coeffs, weighted_coeffs)
        K = K.at[jnp.diag_indices_from(K)].add(jitter)
        return lx.MatrixLinearOperator(K, lx.positive_semidefinite_tag)

    def k_ux(
        self,
        unit_xyz: Float[Array, "N 3"],
        kernel: Kernel,
    ) -> Float[Array, "N K"]:
        """Cross-covariance between unit-sphere inputs and Slepian features.

        Spherical harmonics are evaluated in the cap-centred frame to match
        `SlepianCapBasis.evaluate`; without this rotation, two
        ``SlepianInducingFeatures`` differing only in cap centre would
        produce the same cross-covariance.
        """
        if unit_xyz.ndim != 2 or unit_xyz.shape[-1] != 3:
            raise ValueError(f"unit_xyz must be (N, 3); got {unit_xyz.shape}.")
        centred = self.basis.centred_coordinates(unit_xyz)
        Y = real_spherical_harmonics(centred, self.l_max)
        a_per_feature = self._per_feature_coeffs(kernel)
        # Weight each harmonic by its Funk-Hecke coefficient, then project
        # onto the Slepian modes by contracting the feature axis f.
        Y_weighted = einx.multiply("n f, f -> n f", Y, a_per_feature)
        return einx.dot("n f, f m -> n m", Y_weighted, self.basis.coeffs)

K_uu(kernel: Kernel, *, jitter: float = 1e-06) -> lx.MatrixLinearOperator

Dense Slepian inducing covariance with diagonal jitter.

Source code in packages/pyrox-gp/src/pyrox_gp/_inducing.py
def K_uu(self, kernel: Kernel, *, jitter: float = 1e-6) -> lx.MatrixLinearOperator:
    """Dense Slepian inducing covariance with diagonal jitter."""
    a_per_feature = self._per_feature_coeffs(kernel)
    # Scale each feature row by its Funk-Hecke coefficient, then form the
    # Gram K = Φᵀ diag(a) Φ by contracting the feature axis f.
    weighted_coeffs = einx.multiply(
        "f m, f -> f m", self.basis.coeffs, a_per_feature
    )
    K = einx.dot("f i, f j -> i j", self.basis.coeffs, weighted_coeffs)
    K = K.at[jnp.diag_indices_from(K)].add(jitter)
    return lx.MatrixLinearOperator(K, lx.positive_semidefinite_tag)

k_ux(unit_xyz: Float[Array, 'N 3'], kernel: Kernel) -> Float[Array, 'N K']

Cross-covariance between unit-sphere inputs and Slepian features.

Spherical harmonics are evaluated in the cap-centred frame to match SlepianCapBasis.evaluate; without this rotation, two SlepianInducingFeatures differing only in cap centre would produce the same cross-covariance.

Source code in packages/pyrox-gp/src/pyrox_gp/_inducing.py
def k_ux(
    self,
    unit_xyz: Float[Array, "N 3"],
    kernel: Kernel,
) -> Float[Array, "N K"]:
    """Cross-covariance between unit-sphere inputs and Slepian features.

    Spherical harmonics are evaluated in the cap-centred frame to match
    `SlepianCapBasis.evaluate`; without this rotation, two
    ``SlepianInducingFeatures`` differing only in cap centre would
    produce the same cross-covariance.
    """
    if unit_xyz.ndim != 2 or unit_xyz.shape[-1] != 3:
        raise ValueError(f"unit_xyz must be (N, 3); got {unit_xyz.shape}.")
    centred = self.basis.centred_coordinates(unit_xyz)
    Y = real_spherical_harmonics(centred, self.l_max)
    a_per_feature = self._per_feature_coeffs(kernel)
    # Weight each harmonic by its Funk-Hecke coefficient, then project
    # onto the Slepian modes by contracting the feature axis f.
    Y_weighted = einx.multiply("n f, f -> n f", Y, a_per_feature)
    return einx.dot("n f, f m -> n m", Y_weighted, self.basis.coeffs)

LaplacianInducingFeatures

Bases: Module

Inducing features from low-frequency graph Laplacian eigenvectors.

For a graph with normalized Laplacian \(L\), take the smallest num_basis eigenpairs \((\mu_j, v_j)\). Treating the kernel as a function of the graph distance — specifically, applying the kernel spectrum \(g(\mu)\) to the Laplacian eigenvalues — gives a diagonal \(K_{uu}\).

This implementation supports the heat-kernel family \(g(\mu) = \exp(-\mu / (2 \ell^2))\) (matching pyrox_gp.RBF in spectrum) by reusing pyrox_gp._basis.spectral_density with the eigenvalues as input.

Attributes:

Name Type Description
eigvals Float[Array, ' M']

(M,) Laplacian eigenvalues.

eigvecs Float[Array, 'V M']

(V, M) Laplacian eigenvectors.

num_quadrature Float[Array, 'V M']

Unused (kept for protocol uniformity).

Note

X is a vector of node indices (integer-valued), not coordinates. The returned cross-covariance gathers the relevant rows of eigvecs.

Source code in packages/pyrox-gp/src/pyrox_gp/_inducing.py
class LaplacianInducingFeatures(eqx.Module):
    r"""Inducing features from low-frequency graph Laplacian eigenvectors.

    For a graph with normalized Laplacian $L$, take the smallest
    ``num_basis`` eigenpairs $(\mu_j, v_j)$. Treating the kernel as
    a function of the graph distance — specifically, applying the kernel
    *spectrum* $g(\mu)$ to the Laplacian eigenvalues — gives a
    diagonal $K_{uu}$.

    This implementation supports the *heat-kernel* family
    $g(\mu) = \exp(-\mu / (2 \ell^2))$ (matching `pyrox_gp.RBF`
    in spectrum) by reusing `pyrox_gp._basis.spectral_density` with the
    eigenvalues as input.

    Attributes:
        eigvals: ``(M,)`` Laplacian eigenvalues.
        eigvecs: ``(V, M)`` Laplacian eigenvectors.
        num_quadrature: Unused (kept for protocol uniformity).

    Note:
        ``X`` is a vector of *node indices* (integer-valued), not
        coordinates. The returned cross-covariance gathers the relevant
        rows of ``eigvecs``.
    """

    eigvals: Float[Array, " M"]
    eigvecs: Float[Array, "V M"]

    @classmethod
    def fit(
        cls,
        adjacency: Float[Array, "V V"],
        num_basis: int,
        *,
        normalized: bool = True,
    ) -> LaplacianInducingFeatures:
        eigvals, eigvecs = graph_laplacian_eigpairs(
            adjacency, num_basis, normalized=normalized
        )
        return cls(eigvals=eigvals, eigvecs=eigvecs)

    @property
    def num_features(self) -> int:
        return int(self.eigvals.shape[0])

    def _check_stationary(self, kernel: Kernel) -> None:
        if not _is_stationary(kernel):
            raise ValueError(
                "LaplacianInducingFeatures requires a stationary kernel with a "
                f"registered spectral density; got {type(kernel).__name__}."
            )

    def K_uu(
        self, kernel: Kernel, *, jitter: float = 1e-6
    ) -> lx.DiagonalLinearOperator:
        self._check_stationary(kernel)
        with _kernel_context(kernel):
            S = spectral_density(kernel, self.eigvals, D=1)
        return _diagonal_with_jitter(S, jitter)

    def k_ux(
        self, node_indices: Int[Array, " N"], kernel: Kernel
    ) -> Float[Array, "N M"]:
        self._check_stationary(kernel)
        if node_indices.ndim != 1:
            raise ValueError(
                "node_indices must be a 1D integer array; got shape "
                f"{node_indices.shape}."
            )
        with _kernel_context(kernel):
            S = spectral_density(kernel, self.eigvals, D=1)
        rows = self.eigvecs[node_indices]
        return rows * S[None, :]

DecoupledInducingFeatures

Bases: Module

Decoupled mean / covariance inducing-feature bases (Cheng & Boots 2017).

Two independent inducing-feature sets:

  • mean_features: a large alpha-basis used by the SVGP posterior mean (cheap — predictive mean cost is linear in the mean-basis size).
  • cov_features: a small beta-basis used for the posterior covariance (the true bottleneck; keep this small).

The two bases need not share the same family — a common pattern is a large Fourier basis for the mean and a small spherical-harmonic basis for the covariance, or vice versa. The downstream guide consumes both via the standard SVGP machinery.

Attributes:

Name Type Description
mean_features InducingFeatures

Inducing-feature object backing the predictive mean.

cov_features InducingFeatures

Inducing-feature object backing the predictive covariance.

Note

DecoupledInducingFeatures itself does not implement InducingFeatures (no single K_uu makes sense for two bases). Consumers should access .mean_features and .cov_features directly.

Source code in packages/pyrox-gp/src/pyrox_gp/_inducing.py
class DecoupledInducingFeatures(eqx.Module):
    r"""Decoupled mean / covariance inducing-feature bases (Cheng & Boots 2017).

    Two independent inducing-feature sets:

    - ``mean_features``: a large ``alpha``-basis used by the SVGP
      posterior *mean* (cheap — predictive mean cost is linear in the
      mean-basis size).
    - ``cov_features``: a small ``beta``-basis used for the posterior
      *covariance* (the true bottleneck; keep this small).

    The two bases need not share the same family — a common pattern is a
    large Fourier basis for the mean and a small spherical-harmonic
    basis for the covariance, or vice versa. The downstream guide
    consumes both via the standard SVGP machinery.

    Attributes:
        mean_features: Inducing-feature object backing the predictive mean.
        cov_features: Inducing-feature object backing the predictive covariance.

    Note:
        ``DecoupledInducingFeatures`` itself does *not* implement
        `InducingFeatures` (no single ``K_uu`` makes sense for two
        bases). Consumers should access ``.mean_features`` and
        ``.cov_features`` directly.
    """

    mean_features: InducingFeatures
    cov_features: InducingFeatures

    @property
    def num_mean_features(self) -> int:
        return self.mean_features.num_features

    @property
    def num_cov_features(self) -> int:
        return self.cov_features.num_features

funk_hecke_coefficients(kernel: Kernel, l_max: int, *, num_quadrature: int = 256) -> Float[Array, ' l_max_plus_1']

Funk-Hecke coefficients of a zonal kernel on \(S^2\).

For a kernel of the form \(k(x, x') = \kappa(x \cdot x')\) on the unit 2-sphere, the Funk-Hecke theorem gives:

\[ a_l = 2\pi \int_{-1}^{1} \kappa(t)\,P_l(t)\,dt. \]

Returns (l_max + 1,) coefficients indexed by l. We treat any Euclidean kernel as zonal-on-the-sphere via \(\kappa(t) = k_{\mathrm{euc}}(\hat{n}_0, \hat{n}_t)\) for unit vectors at angular separation arccos(t).

Source code in packages/pyrox-gp/src/pyrox_gp/_inducing.py
def funk_hecke_coefficients(
    kernel: Kernel,
    l_max: int,
    *,
    num_quadrature: int = 256,
) -> Float[Array, " l_max_plus_1"]:
    r"""Funk-Hecke coefficients of a zonal kernel on $S^2$.

    For a kernel of the form $k(x, x') = \kappa(x \cdot x')$ on the
    unit 2-sphere, the Funk-Hecke theorem gives:

    $$
    a_l = 2\pi \int_{-1}^{1} \kappa(t)\,P_l(t)\,dt.
    $$

    Returns ``(l_max + 1,)`` coefficients indexed by ``l``. We treat any
    Euclidean kernel as zonal-on-the-sphere via
    $\kappa(t) = k_{\mathrm{euc}}(\hat{n}_0, \hat{n}_t)$ for unit
    vectors at angular separation ``arccos(t)``.
    """
    # Gauss-Legendre quadrature nodes on [-1, 1] (host-side setup constants).
    t, w = _gauss_legendre_nodes(num_quadrature)
    # Build pairs of unit vectors: x0 = (0, 0, 1), x_t = (sin(arccos t), 0, t).
    sin_t = jnp.sqrt(jnp.maximum(1.0 - t**2, 0.0))
    n0 = jnp.array([0.0, 0.0, 1.0])
    nT = jnp.stack([sin_t, jnp.zeros_like(t), t], axis=-1)  # (Q, 3)
    # Single batched kernel call — stays on-device and keeps autodiff edges
    # to any hyperparameters sampled inside ``kernel``. Taking row 0 of the
    # ``(1, Q)`` Gram is O(Q), not O(Q^2).
    with _kernel_context(kernel):
        _reject_anisotropic(kernel)
        kt = kernel(n0[None, :], nT)[0]  # (Q,)
    # Evaluate P_l(t) for l = 0, ..., l_max via three-term recurrence.
    P_lm1 = jnp.ones_like(t)  # P_0
    P_l = t  # P_1
    coeffs = [2.0 * jnp.pi * jnp.sum(w * kt)]  # a_0 = 2pi * int kt * 1 dt
    if l_max >= 1:
        coeffs.append(2.0 * jnp.pi * jnp.sum(w * kt * P_l))  # a_1
    for ell in range(2, l_max + 1):
        P_lp1 = ((2 * ell - 1) * t * P_l - (ell - 1) * P_lm1) / ell
        coeffs.append(2.0 * jnp.pi * jnp.sum(w * kt * P_lp1))
        P_lm1, P_l = P_l, P_lp1
    return jnp.stack(coeffs, axis=0)

Sparse GP prior

SparseGPPrior

Bases: Module

GP prior parameterized over inducing inputs Z.

Represents the zero-mean prior over inducing values u = f(Z) used by sparse variational guides:

\[ p(u) = \mathcal{N}(0,\, K_{ZZ} + \mathrm{jitter}\,I). \]

The standard SVGP convention is to subtract any global mean function before forming the prior over u and to add it back at predict time, so the inducing-prior mean is fixed to zero (this is what the guides' KL terms assume — see FullRankGuide.kl_divergence, MeanFieldGuide.kl_divergence, WhitenedGuide.kl_divergence). The mean_fn attribute on this class is exposed as a convenience for callers that want to add mu(X_*) back onto the predictive mean returned by Guide.predict; it is not incorporated in inducing_operator or in the guides' KL.

Pair with a sparse variational guide that owns q(u) = N(m, S) to obtain the standard SVGP predictive

\[ \mu_*(x) = K_{xZ} K_{ZZ}^{-1} m, \qquad \sigma_*^2(x) = k(x, x) - K_{xZ} K_{ZZ}^{-1} K_{Zx} + K_{xZ} K_{ZZ}^{-1} S K_{ZZ}^{-1} K_{Zx}. \]

Attributes:

Name Type Description
kernel Kernel

Any pyrox_gp.Kernel — evaluated on Z.

Z Float[Array, 'M D'] | None

Inducing inputs of shape (M, D).

mean_fn Callable[[Float[Array, 'N D']], Float[Array, ' N']] | None

Callable X -> (N,) or None for the zero mean. Convenience accessor; not folded into the inducing prior.

solver AbstractSolverStrategy | None

Any gaussx.AbstractSolverStrategy. Defaults to gaussx.DenseSolver(). Used by guides that need to solve against K_zz (e.g.\ for KL or unwhitening).

jitter float

Diagonal regularization added to K_zz for numerical stability. Not a noise model — sparse SVGP does not put observation noise on the inducing-value prior.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse.py
class SparseGPPrior(eqx.Module):
    r"""GP prior parameterized over inducing inputs ``Z``.

    Represents the *zero-mean* prior over inducing values ``u = f(Z)``
    used by sparse variational guides:

    $$
    p(u) = \mathcal{N}(0,\, K_{ZZ} + \mathrm{jitter}\,I).
    $$

    The standard SVGP convention is to subtract any global mean function
    before forming the prior over ``u`` and to add it back at predict
    time, so the inducing-prior mean is fixed to zero (this is what the
    guides' KL terms assume — see `FullRankGuide.kl_divergence`,
    `MeanFieldGuide.kl_divergence`, `WhitenedGuide.kl_divergence`).
    The `mean_fn` attribute on this class is exposed as a
    convenience for callers that want to add ``mu(X_*)`` back onto the
    predictive mean returned by `Guide.predict`; it is **not**
    incorporated in `inducing_operator` or in the guides' KL.

    Pair with a sparse variational guide that owns ``q(u) = N(m, S)`` to
    obtain the standard SVGP predictive

    $$
    \mu_*(x) = K_{xZ} K_{ZZ}^{-1} m, \qquad
    \sigma_*^2(x) = k(x, x) - K_{xZ} K_{ZZ}^{-1} K_{Zx}
                   + K_{xZ} K_{ZZ}^{-1} S K_{ZZ}^{-1} K_{Zx}.
    $$

    Attributes:
        kernel: Any `pyrox_gp.Kernel` — evaluated on ``Z``.
        Z: Inducing inputs of shape ``(M, D)``.
        mean_fn: Callable ``X -> (N,)`` or ``None`` for the zero mean.
            Convenience accessor; not folded into the inducing prior.
        solver: Any ``gaussx.AbstractSolverStrategy``. Defaults to
            ``gaussx.DenseSolver()``. Used by guides that need to solve
            against ``K_zz`` (e.g.\ for KL or unwhitening).
        jitter: Diagonal regularization added to ``K_zz`` for numerical
            stability. Not a noise model — sparse SVGP does not put
            observation noise on the inducing-value prior.
    """

    kernel: Kernel
    Z: Float[Array, "M D"] | None = None
    inducing: InducingFeatures | None = None
    mean_fn: Callable[[Float[Array, "N D"]], Float[Array, " N"]] | None = None
    solver: AbstractSolverStrategy | None = None
    jitter: float = 1e-6

    def __check_init__(self) -> None:
        if (self.Z is None) == (self.inducing is None):
            raise ValueError(
                "SparseGPPrior must be constructed with exactly one of `Z` "
                "(point inducing) or `inducing` (inducing features)."
            )

    @property
    def num_inducing(self) -> int:
        """Number of inducing inputs / features ``M``."""
        if self.inducing is not None:
            return self.inducing.num_features
        assert self.Z is not None
        return self.Z.shape[0]

    def mean(self, X: Float[Array, "N D"]) -> Float[Array, " N"]:
        """Evaluate the mean function at ``X``; zero by default."""
        if self.mean_fn is None:
            return jnp.zeros(X.shape[0], dtype=X.dtype)
        return self.mean_fn(X)

    def inducing_operator(self) -> lx.AbstractLinearOperator:
        r"""Return ``K_{ZZ} + \text{jitter}\,I`` as a ``lineax`` operator.

        For point-inducing priors, returns a dense
        `lineax.MatrixLinearOperator` with ``positive_semidefinite_tag``.
        For inducing-feature priors, delegates to
        `InducingFeatures.K_uu` — typically a
        `lineax.DiagonalLinearOperator` so the downstream
        `gaussx.solve` dispatches in O(M) instead of O(M^3).

        Single kernel call; safe standalone for kernels with priors. For
        building several SVGP blocks together, prefer
        `predictive_blocks`, which scopes one shared kernel
        context across ``K_zz``, ``K_xz``, and ``K_xx_diag`` so
        Pattern B / C kernels register their NumPyro hyperparameter
        sites once instead of resampling per call.
        """
        if self.inducing is not None:
            with _kernel_context(self.kernel):
                return self.inducing.K_uu(self.kernel, jitter=self.jitter)
        assert self.Z is not None
        with _kernel_context(self.kernel):
            K = self.kernel(self.Z, self.Z)
        K = K.at[jnp.diag_indices_from(K)].add(self.jitter)
        return _psd_operator(K)

    def cross_covariance(self, X: Array) -> Float[Array, "N M"]:
        r"""$K_{XZ}$ — covariance between ``X`` and the inducing inputs/features.

        The expected shape of ``X`` is inducing-family-dependent:

        - Point-inducing (``Z``) or `FourierInducingFeatures`:
          coordinates ``(N, D)``.
        - `SphericalHarmonicInducingFeatures`: unit vectors ``(N, 3)``.
        - `LaplacianInducingFeatures`: integer node indices ``(N,)``.

        See `predictive_blocks` for the shared-context batch
        helper to use when assembling several SVGP blocks together.
        """
        if self.inducing is not None:
            with _kernel_context(self.kernel):
                return self.inducing.k_ux(X, self.kernel)
        assert self.Z is not None
        with _kernel_context(self.kernel):
            return self.kernel(X, self.Z)

    def kernel_diag(self, X: Float[Array, "N D"]) -> Float[Array, " N"]:
        r"""Prior diagonal ``\mathrm{diag}\,K(X, X)`` — variance at each ``x``.

        See `predictive_blocks` for the shared-context batch
        helper to use when assembling several SVGP blocks together.
        """
        with _kernel_context(self.kernel):
            return self.kernel.diag(X)

    def predictive_blocks(
        self, X: Array
    ) -> tuple[
        lx.AbstractLinearOperator,
        Float[Array, "N M"],
        Float[Array, " N"],
    ]:
        r"""Return ``(K_zz_op, K_xz, K_xx_diag)`` under one shared kernel context.

        For Pattern B / C kernels with prior'd hyperparameters, the three
        kernel evaluations needed for an SVGP predictive must share a
        single `pyrox.PyroxModule` context so the underlying
        ``pyrox_sample`` sites register once and yield consistent
        hyperparameter draws across ``K_{ZZ}``, ``K_{XZ}``, and the
        diagonal ``\mathrm{diag}\,K(X, X)``. Without this scoping, three
        separate calls would draw three independent hyperparameter
        samples (under seed) or raise NumPyro duplicate-site errors
        (under tracing) — either way invalidating the SVGP math.

        For pure `equinox.Module` kernels (no ``_get_context``),
        this is equivalent to calling `inducing_operator`,
        `cross_covariance`, and `kernel_diag` independently.

        For inducing-feature priors, ``K_zz_op`` is a
        `lineax.DiagonalLinearOperator` (jitter folded into the
        diagonal vector — never ``+ jnp.eye``) so the downstream solve
        stays O(M).
        """
        with _kernel_context(self.kernel):
            if self.inducing is not None:
                K_zz_op = self.inducing.K_uu(self.kernel, jitter=self.jitter)
                K_xz = self.inducing.k_ux(X, self.kernel)
            else:
                assert self.Z is not None
                K_zz_raw = self.kernel(self.Z, self.Z)
                K_xz = self.kernel(X, self.Z)
                K_zz = K_zz_raw.at[jnp.diag_indices_from(K_zz_raw)].add(self.jitter)
                K_zz_op = _psd_operator(K_zz)
            K_xx_diag = self.kernel.diag(X)
        return K_zz_op, K_xz, K_xx_diag

    def _resolved_solver(self) -> AbstractSolverStrategy:
        return DenseSolver() if self.solver is None else self.solver

    def log_prob(self, u: Float[Array, " M"]) -> Float[Array, ""]:
        r"""Log-density under $p(u) = \mathcal{N}(0, K_{ZZ} + \text{jitter}\,I)$.

        Delegates to `gaussx.gaussian_log_prob` with the
        configured `solver` so the user-supplied solver controls
        the ``solve`` / ``logdet`` work on ``K_zz_op``. Useful for
        scoring inducing values against the SVGP prior in non-NumPyro
        contexts (e.g.\\ tests, diagnostics).
        """
        m = jnp.zeros(self.num_inducing, dtype=u.dtype)
        return gaussian_log_prob(
            m, self.inducing_operator(), u, solver=self._resolved_solver()
        )

    def sample(self, key: Array) -> Float[Array, " M"]:
        r"""Draw ``u \sim p(u)`` from the inducing prior.

        Wraps the inducing operator in a
        `gaussx.MultivariateNormal` with the configured
        `solver`. ``MultivariateNormal.sample`` factors the
        covariance via `gaussx.cholesky` and reparameterizes;
        the returned draw has shape ``(M,)``.

        Note: the SVGP variational workflow samples ``u`` from the
        *guide* $q(u)$, not the prior. This method exists so the
        prior surface is symmetric with the guide surface and so users
        can score / draw inducing values against the prior directly
        (e.g.\\ for tests or for prior-sample initialization).
        """
        n = self.num_inducing
        op = self.inducing_operator()
        loc = jnp.zeros(n, dtype=op.out_structure().dtype)
        mvn = MultivariateNormal(loc, op, solver=self._resolved_solver())
        return mvn.sample(key)

num_inducing: int property

Number of inducing inputs / features M.

mean(X: Float[Array, 'N D']) -> Float[Array, ' N']

Evaluate the mean function at X; zero by default.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse.py
def mean(self, X: Float[Array, "N D"]) -> Float[Array, " N"]:
    """Evaluate the mean function at ``X``; zero by default."""
    if self.mean_fn is None:
        return jnp.zeros(X.shape[0], dtype=X.dtype)
    return self.mean_fn(X)

inducing_operator() -> lx.AbstractLinearOperator

Return K_{ZZ} + \text{jitter}\,I as a lineax operator.

For point-inducing priors, returns a dense lineax.MatrixLinearOperator with positive_semidefinite_tag. For inducing-feature priors, delegates to InducingFeatures.K_uu — typically a lineax.DiagonalLinearOperator so the downstream gaussx.solve dispatches in O(M) instead of O(M^3).

Single kernel call; safe standalone for kernels with priors. For building several SVGP blocks together, prefer predictive_blocks, which scopes one shared kernel context across K_zz, K_xz, and K_xx_diag so Pattern B / C kernels register their NumPyro hyperparameter sites once instead of resampling per call.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse.py
def inducing_operator(self) -> lx.AbstractLinearOperator:
    r"""Return ``K_{ZZ} + \text{jitter}\,I`` as a ``lineax`` operator.

    For point-inducing priors, returns a dense
    `lineax.MatrixLinearOperator` with ``positive_semidefinite_tag``.
    For inducing-feature priors, delegates to
    `InducingFeatures.K_uu` — typically a
    `lineax.DiagonalLinearOperator` so the downstream
    `gaussx.solve` dispatches in O(M) instead of O(M^3).

    Single kernel call; safe standalone for kernels with priors. For
    building several SVGP blocks together, prefer
    `predictive_blocks`, which scopes one shared kernel
    context across ``K_zz``, ``K_xz``, and ``K_xx_diag`` so
    Pattern B / C kernels register their NumPyro hyperparameter
    sites once instead of resampling per call.
    """
    if self.inducing is not None:
        with _kernel_context(self.kernel):
            return self.inducing.K_uu(self.kernel, jitter=self.jitter)
    assert self.Z is not None
    with _kernel_context(self.kernel):
        K = self.kernel(self.Z, self.Z)
    K = K.at[jnp.diag_indices_from(K)].add(self.jitter)
    return _psd_operator(K)

cross_covariance(X: Array) -> Float[Array, 'N M']

\(K_{XZ}\) — covariance between X and the inducing inputs/features.

The expected shape of X is inducing-family-dependent:

  • Point-inducing (Z) or FourierInducingFeatures: coordinates (N, D).
  • SphericalHarmonicInducingFeatures: unit vectors (N, 3).
  • LaplacianInducingFeatures: integer node indices (N,).

See predictive_blocks for the shared-context batch helper to use when assembling several SVGP blocks together.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse.py
def cross_covariance(self, X: Array) -> Float[Array, "N M"]:
    r"""$K_{XZ}$ — covariance between ``X`` and the inducing inputs/features.

    The expected shape of ``X`` is inducing-family-dependent:

    - Point-inducing (``Z``) or `FourierInducingFeatures`:
      coordinates ``(N, D)``.
    - `SphericalHarmonicInducingFeatures`: unit vectors ``(N, 3)``.
    - `LaplacianInducingFeatures`: integer node indices ``(N,)``.

    See `predictive_blocks` for the shared-context batch
    helper to use when assembling several SVGP blocks together.
    """
    if self.inducing is not None:
        with _kernel_context(self.kernel):
            return self.inducing.k_ux(X, self.kernel)
    assert self.Z is not None
    with _kernel_context(self.kernel):
        return self.kernel(X, self.Z)

kernel_diag(X: Float[Array, 'N D']) -> Float[Array, ' N']

Prior diagonal \mathrm{diag}\,K(X, X) — variance at each x.

See predictive_blocks for the shared-context batch helper to use when assembling several SVGP blocks together.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse.py
def kernel_diag(self, X: Float[Array, "N D"]) -> Float[Array, " N"]:
    r"""Prior diagonal ``\mathrm{diag}\,K(X, X)`` — variance at each ``x``.

    See `predictive_blocks` for the shared-context batch
    helper to use when assembling several SVGP blocks together.
    """
    with _kernel_context(self.kernel):
        return self.kernel.diag(X)

predictive_blocks(X: Array) -> tuple[lx.AbstractLinearOperator, Float[Array, 'N M'], Float[Array, ' N']]

Return (K_zz_op, K_xz, K_xx_diag) under one shared kernel context.

For Pattern B / C kernels with prior'd hyperparameters, the three kernel evaluations needed for an SVGP predictive must share a single pyrox.PyroxModule context so the underlying pyrox_sample sites register once and yield consistent hyperparameter draws across K_{ZZ}, K_{XZ}, and the diagonal \mathrm{diag}\,K(X, X). Without this scoping, three separate calls would draw three independent hyperparameter samples (under seed) or raise NumPyro duplicate-site errors (under tracing) — either way invalidating the SVGP math.

For pure equinox.Module kernels (no _get_context), this is equivalent to calling inducing_operator, cross_covariance, and kernel_diag independently.

For inducing-feature priors, K_zz_op is a lineax.DiagonalLinearOperator (jitter folded into the diagonal vector — never + jnp.eye) so the downstream solve stays O(M).

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse.py
def predictive_blocks(
    self, X: Array
) -> tuple[
    lx.AbstractLinearOperator,
    Float[Array, "N M"],
    Float[Array, " N"],
]:
    r"""Return ``(K_zz_op, K_xz, K_xx_diag)`` under one shared kernel context.

    For Pattern B / C kernels with prior'd hyperparameters, the three
    kernel evaluations needed for an SVGP predictive must share a
    single `pyrox.PyroxModule` context so the underlying
    ``pyrox_sample`` sites register once and yield consistent
    hyperparameter draws across ``K_{ZZ}``, ``K_{XZ}``, and the
    diagonal ``\mathrm{diag}\,K(X, X)``. Without this scoping, three
    separate calls would draw three independent hyperparameter
    samples (under seed) or raise NumPyro duplicate-site errors
    (under tracing) — either way invalidating the SVGP math.

    For pure `equinox.Module` kernels (no ``_get_context``),
    this is equivalent to calling `inducing_operator`,
    `cross_covariance`, and `kernel_diag` independently.

    For inducing-feature priors, ``K_zz_op`` is a
    `lineax.DiagonalLinearOperator` (jitter folded into the
    diagonal vector — never ``+ jnp.eye``) so the downstream solve
    stays O(M).
    """
    with _kernel_context(self.kernel):
        if self.inducing is not None:
            K_zz_op = self.inducing.K_uu(self.kernel, jitter=self.jitter)
            K_xz = self.inducing.k_ux(X, self.kernel)
        else:
            assert self.Z is not None
            K_zz_raw = self.kernel(self.Z, self.Z)
            K_xz = self.kernel(X, self.Z)
            K_zz = K_zz_raw.at[jnp.diag_indices_from(K_zz_raw)].add(self.jitter)
            K_zz_op = _psd_operator(K_zz)
        K_xx_diag = self.kernel.diag(X)
    return K_zz_op, K_xz, K_xx_diag

log_prob(u: Float[Array, ' M']) -> Float[Array, '']

Log-density under \(p(u) = \mathcal{N}(0, K_{ZZ} + \text{jitter}\,I)\).

Delegates to gaussx.gaussian_log_prob with the configured solver so the user-supplied solver controls the solve / logdet work on K_zz_op. Useful for scoring inducing values against the SVGP prior in non-NumPyro contexts (e.g.\ tests, diagnostics).

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse.py
def log_prob(self, u: Float[Array, " M"]) -> Float[Array, ""]:
    r"""Log-density under $p(u) = \mathcal{N}(0, K_{ZZ} + \text{jitter}\,I)$.

    Delegates to `gaussx.gaussian_log_prob` with the
    configured `solver` so the user-supplied solver controls
    the ``solve`` / ``logdet`` work on ``K_zz_op``. Useful for
    scoring inducing values against the SVGP prior in non-NumPyro
    contexts (e.g.\\ tests, diagnostics).
    """
    m = jnp.zeros(self.num_inducing, dtype=u.dtype)
    return gaussian_log_prob(
        m, self.inducing_operator(), u, solver=self._resolved_solver()
    )

sample(key: Array) -> Float[Array, ' M']

Draw u \sim p(u) from the inducing prior.

Wraps the inducing operator in a gaussx.MultivariateNormal with the configured solver. MultivariateNormal.sample factors the covariance via gaussx.cholesky and reparameterizes; the returned draw has shape (M,).

Note: the SVGP variational workflow samples u from the guide \(q(u)\), not the prior. This method exists so the prior surface is symmetric with the guide surface and so users can score / draw inducing values against the prior directly (e.g.\ for tests or for prior-sample initialization).

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse.py
def sample(self, key: Array) -> Float[Array, " M"]:
    r"""Draw ``u \sim p(u)`` from the inducing prior.

    Wraps the inducing operator in a
    `gaussx.MultivariateNormal` with the configured
    `solver`. ``MultivariateNormal.sample`` factors the
    covariance via `gaussx.cholesky` and reparameterizes;
    the returned draw has shape ``(M,)``.

    Note: the SVGP variational workflow samples ``u`` from the
    *guide* $q(u)$, not the prior. This method exists so the
    prior surface is symmetric with the guide surface and so users
    can score / draw inducing values against the prior directly
    (e.g.\\ for tests or for prior-sample initialization).
    """
    n = self.num_inducing
    op = self.inducing_operator()
    loc = jnp.zeros(n, dtype=op.out_structure().dtype)
    mvn = MultivariateNormal(loc, op, solver=self._resolved_solver())
    return mvn.sample(key)

Variational guides

Variational families q(u) over the inducing values of a SparseGPPrior. All five expose the same building-block interface — sample(key), log_prob(u), kl_divergence(prior_cov), and predict(K_xz, K_zz_op, K_xx_diag) — so they swap freely inside the SVGP ELBO. WhitenedGuide parameterizes in whitened coordinates u = L_zz v (the standard choice for stable optimization); NaturalGuide parameterizes in natural form for natural-gradient / CVI workflows; DeltaGuide is a point mass for MAP-style training.

FullRankGuide

Bases: Guide

Full-rank Gaussian variational posterior over inducing values u.

\[ q(u) = \mathcal{N}(m,\, L_S L_S^\top), \]

parameterized by the mean mean of shape (M,) and the lower-triangular Cholesky factor scale_tril of shape (M, M).

Attributes:

Name Type Description
mean Float[Array, ' M']

Variational mean m of shape (M,).

scale_tril Float[Array, 'M M']

Lower-triangular Cholesky factor of the variational covariance, shape (M, M). The covariance is S = scale_tril @ scale_tril.T.

solver AbstractSolverStrategy | None

Optional gaussx.AbstractSolverStrategy exposed so downstream consumers (e.g.\ inference loops that solve against this guide's covariance) can pick the solver. The guide's own log_prob / kl_divergence use the hand-rolled Cholesky path; the field is None by default and is read by callers who need it.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
class FullRankGuide(Guide):
    r"""Full-rank Gaussian variational posterior over inducing values ``u``.

    $$
    q(u) = \mathcal{N}(m,\, L_S L_S^\top),
    $$

    parameterized by the mean ``mean`` of shape ``(M,)`` and the
    *lower-triangular* Cholesky factor ``scale_tril`` of shape ``(M, M)``.

    Attributes:
        mean: Variational mean ``m`` of shape ``(M,)``.
        scale_tril: Lower-triangular Cholesky factor of the variational
            covariance, shape ``(M, M)``. The covariance is
            ``S = scale_tril @ scale_tril.T``.
        solver: Optional `gaussx.AbstractSolverStrategy` exposed
            so downstream consumers (e.g.\\ inference loops that solve
            against this guide's covariance) can pick the solver. The
            guide's own ``log_prob`` / ``kl_divergence`` use the
            hand-rolled Cholesky path; the field is ``None`` by default
            and is read by callers who need it.
    """

    mean: Float[Array, " M"]
    scale_tril: Float[Array, "M M"]
    solver: AbstractSolverStrategy | None = None

    @classmethod
    def init(
        cls,
        num_inducing: int,
        *,
        scale: float = 1.0,
        solver: AbstractSolverStrategy | None = None,
    ) -> FullRankGuide:
        """Construct a guide initialized to ``N(0, scale^2 I)``.

        Uses JAX's default float dtype — set
        ``jax.config.update("jax_enable_x64", True)`` once at the top of
        your script if you want float64 inducing parameters.
        """
        m = jnp.zeros(num_inducing)
        L = scale * jnp.eye(num_inducing)
        return cls(mean=m, scale_tril=L, solver=solver)

    def sample(self, key: Array) -> Float[Array, " M"]:
        r"""Draw ``u = m + L_S \epsilon`` with ``\epsilon ~ N(0, I)``."""
        eps = jax.random.normal(key, self.mean.shape, dtype=self.mean.dtype)
        return self.mean + einx.dot("i j, j -> i", self.scale_tril, eps)

    def log_prob(self, u: Float[Array, " ..."]) -> Float[Array, ""]:  # ty: ignore[invalid-method-override]
        r"""Variational log density ``\log q(u)`` via `gaussx.gaussian_log_prob`.

        Delegates the ``solve`` and ``logdet`` work to the configured
        solver so the user-supplied `solver` actually controls the
        numerical path.
        """
        return gaussian_log_prob(
            self.mean,
            _full_cov_operator(self.scale_tril),
            u,
            solver=_resolve_solver(self.solver),
        )

    def kl_divergence(self, prior_cov: lx.AbstractLinearOperator) -> Float[Array, ""]:
        r"""``KL(q(u) || p(u))`` against an inducing prior with zero mean.

        Falls back to `gaussx.dist_kl_divergence`, which dispatches
        on operator structure for the trace and logdet terms but does
        not yet take an explicit solver.
        """
        q_loc = self.mean
        q_cov = _full_cov_operator(self.scale_tril)
        p_loc = jnp.zeros_like(self.mean)
        return dist_kl_divergence(q_loc, q_cov, p_loc, prior_cov)

    def predict(
        self,
        K_xz: Float[Array, "N M"],
        K_zz_op: lx.AbstractLinearOperator,
        K_xx_diag: Float[Array, " N"],
    ) -> tuple[Float[Array, " N"], Float[Array, " N"]]:
        r"""Predictive ``(mean, variance)`` at points with cross-cov ``K_xz``.

        Routes through `_svgp_predict_unwhitened`, which dispatches
        on the structure of ``K_zz_op`` via `gaussx.cholesky` /
        `gaussx.whitened_svgp_predict`. These primitives do not
        currently accept an explicit solver — the `solver` field
        is exposed for downstream consumers.
        """
        u_cov = einx.dot("i j, k j -> i k", self.scale_tril, self.scale_tril)
        return _svgp_predict_unwhitened(K_xz, K_zz_op, K_xx_diag, self.mean, u_cov)

init(num_inducing: int, *, scale: float = 1.0, solver: AbstractSolverStrategy | None = None) -> FullRankGuide classmethod

Construct a guide initialized to N(0, scale^2 I).

Uses JAX's default float dtype — set jax.config.update("jax_enable_x64", True) once at the top of your script if you want float64 inducing parameters.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
@classmethod
def init(
    cls,
    num_inducing: int,
    *,
    scale: float = 1.0,
    solver: AbstractSolverStrategy | None = None,
) -> FullRankGuide:
    """Construct a guide initialized to ``N(0, scale^2 I)``.

    Uses JAX's default float dtype — set
    ``jax.config.update("jax_enable_x64", True)`` once at the top of
    your script if you want float64 inducing parameters.
    """
    m = jnp.zeros(num_inducing)
    L = scale * jnp.eye(num_inducing)
    return cls(mean=m, scale_tril=L, solver=solver)

sample(key: Array) -> Float[Array, ' M']

Draw u = m + L_S \epsilon with \epsilon ~ N(0, I).

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def sample(self, key: Array) -> Float[Array, " M"]:
    r"""Draw ``u = m + L_S \epsilon`` with ``\epsilon ~ N(0, I)``."""
    eps = jax.random.normal(key, self.mean.shape, dtype=self.mean.dtype)
    return self.mean + einx.dot("i j, j -> i", self.scale_tril, eps)

log_prob(u: Float[Array, ' ...']) -> Float[Array, '']

Variational log density \log q(u) via gaussx.gaussian_log_prob.

Delegates the solve and logdet work to the configured solver so the user-supplied solver actually controls the numerical path.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def log_prob(self, u: Float[Array, " ..."]) -> Float[Array, ""]:  # ty: ignore[invalid-method-override]
    r"""Variational log density ``\log q(u)`` via `gaussx.gaussian_log_prob`.

    Delegates the ``solve`` and ``logdet`` work to the configured
    solver so the user-supplied `solver` actually controls the
    numerical path.
    """
    return gaussian_log_prob(
        self.mean,
        _full_cov_operator(self.scale_tril),
        u,
        solver=_resolve_solver(self.solver),
    )

kl_divergence(prior_cov: lx.AbstractLinearOperator) -> Float[Array, '']

KL(q(u) || p(u)) against an inducing prior with zero mean.

Falls back to gaussx.dist_kl_divergence, which dispatches on operator structure for the trace and logdet terms but does not yet take an explicit solver.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def kl_divergence(self, prior_cov: lx.AbstractLinearOperator) -> Float[Array, ""]:
    r"""``KL(q(u) || p(u))`` against an inducing prior with zero mean.

    Falls back to `gaussx.dist_kl_divergence`, which dispatches
    on operator structure for the trace and logdet terms but does
    not yet take an explicit solver.
    """
    q_loc = self.mean
    q_cov = _full_cov_operator(self.scale_tril)
    p_loc = jnp.zeros_like(self.mean)
    return dist_kl_divergence(q_loc, q_cov, p_loc, prior_cov)

predict(K_xz: Float[Array, 'N M'], K_zz_op: lx.AbstractLinearOperator, K_xx_diag: Float[Array, ' N']) -> tuple[Float[Array, ' N'], Float[Array, ' N']]

Predictive (mean, variance) at points with cross-cov K_xz.

Routes through _svgp_predict_unwhitened, which dispatches on the structure of K_zz_op via gaussx.cholesky / gaussx.whitened_svgp_predict. These primitives do not currently accept an explicit solver — the solver field is exposed for downstream consumers.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def predict(
    self,
    K_xz: Float[Array, "N M"],
    K_zz_op: lx.AbstractLinearOperator,
    K_xx_diag: Float[Array, " N"],
) -> tuple[Float[Array, " N"], Float[Array, " N"]]:
    r"""Predictive ``(mean, variance)`` at points with cross-cov ``K_xz``.

    Routes through `_svgp_predict_unwhitened`, which dispatches
    on the structure of ``K_zz_op`` via `gaussx.cholesky` /
    `gaussx.whitened_svgp_predict`. These primitives do not
    currently accept an explicit solver — the `solver` field
    is exposed for downstream consumers.
    """
    u_cov = einx.dot("i j, k j -> i k", self.scale_tril, self.scale_tril)
    return _svgp_predict_unwhitened(K_xz, K_zz_op, K_xx_diag, self.mean, u_cov)

MeanFieldGuide

Bases: Guide

Diagonal Gaussian variational posterior over inducing values u.

\[ q(u) = \mathcal{N}(m,\, \mathrm{diag}(s^2)), \]

parameterized by the mean mean and the per-coordinate standard deviations scale.

Attributes:

Name Type Description
mean Float[Array, ' M']

Variational mean m of shape (M,).

scale Float[Array, ' M']

Per-coordinate standard deviations s of shape (M,). Must be strictly positive.

solver AbstractSolverStrategy | None

Optional gaussx.AbstractSolverStrategy — see FullRankGuide for usage. None by default.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
class MeanFieldGuide(Guide):
    r"""Diagonal Gaussian variational posterior over inducing values ``u``.

    $$
    q(u) = \mathcal{N}(m,\, \mathrm{diag}(s^2)),
    $$

    parameterized by the mean ``mean`` and the per-coordinate standard
    deviations ``scale``.

    Attributes:
        mean: Variational mean ``m`` of shape ``(M,)``.
        scale: Per-coordinate standard deviations ``s`` of shape ``(M,)``.
            Must be strictly positive.
        solver: Optional `gaussx.AbstractSolverStrategy` — see
            `FullRankGuide` for usage. ``None`` by default.
    """

    mean: Float[Array, " M"]
    scale: Float[Array, " M"]
    solver: AbstractSolverStrategy | None = None

    @classmethod
    def init(
        cls,
        num_inducing: int,
        *,
        scale: float = 1.0,
        solver: AbstractSolverStrategy | None = None,
    ) -> MeanFieldGuide:
        """Construct a guide initialized to ``N(0, scale^2 I)`` — see
        `FullRankGuide.init` for the dtype convention."""
        m = jnp.zeros(num_inducing)
        s = jnp.full(num_inducing, scale)
        return cls(mean=m, scale=s, solver=solver)

    def sample(self, key: Array) -> Float[Array, " M"]:
        r"""Draw ``u = m + s \odot \epsilon`` with ``\epsilon ~ N(0, I)``."""
        eps = jax.random.normal(key, self.mean.shape, dtype=self.mean.dtype)
        return self.mean + self.scale * eps

    def log_prob(self, u: Float[Array, " ..."]) -> Float[Array, ""]:  # ty: ignore[invalid-method-override]
        r"""Variational log density ``\log q(u)`` via `gaussx.gaussian_log_prob`.

        The covariance operator is built from the per-coordinate scales
        as a `lineax.MatrixLinearOperator` tagged
        `positive_semidefinite_tag`; structural dispatch routes
        through the configured solver.
        """
        return gaussian_log_prob(
            self.mean,
            _diag_cov_operator(self.scale),
            u,
            solver=_resolve_solver(self.solver),
        )

    def kl_divergence(self, prior_cov: lx.AbstractLinearOperator) -> Float[Array, ""]:
        r"""``KL(q(u) || p(u))`` against an inducing prior with zero mean."""
        q_loc = self.mean
        q_cov = _diag_cov_operator(self.scale)
        p_loc = jnp.zeros_like(self.mean)
        return dist_kl_divergence(q_loc, q_cov, p_loc, prior_cov)

    def predict(
        self,
        K_xz: Float[Array, "N M"],
        K_zz_op: lx.AbstractLinearOperator,
        K_xx_diag: Float[Array, " N"],
    ) -> tuple[Float[Array, " N"], Float[Array, " N"]]:
        """Predictive ``(mean, variance)`` at the points whose cross-cov is ``K_xz``."""
        u_cov = jnp.diag(self.scale**2)
        return _svgp_predict_unwhitened(K_xz, K_zz_op, K_xx_diag, self.mean, u_cov)

init(num_inducing: int, *, scale: float = 1.0, solver: AbstractSolverStrategy | None = None) -> MeanFieldGuide classmethod

Construct a guide initialized to N(0, scale^2 I) — see FullRankGuide.init for the dtype convention.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
@classmethod
def init(
    cls,
    num_inducing: int,
    *,
    scale: float = 1.0,
    solver: AbstractSolverStrategy | None = None,
) -> MeanFieldGuide:
    """Construct a guide initialized to ``N(0, scale^2 I)`` — see
    `FullRankGuide.init` for the dtype convention."""
    m = jnp.zeros(num_inducing)
    s = jnp.full(num_inducing, scale)
    return cls(mean=m, scale=s, solver=solver)

sample(key: Array) -> Float[Array, ' M']

Draw u = m + s \odot \epsilon with \epsilon ~ N(0, I).

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def sample(self, key: Array) -> Float[Array, " M"]:
    r"""Draw ``u = m + s \odot \epsilon`` with ``\epsilon ~ N(0, I)``."""
    eps = jax.random.normal(key, self.mean.shape, dtype=self.mean.dtype)
    return self.mean + self.scale * eps

log_prob(u: Float[Array, ' ...']) -> Float[Array, '']

Variational log density \log q(u) via gaussx.gaussian_log_prob.

The covariance operator is built from the per-coordinate scales as a lineax.MatrixLinearOperator tagged positive_semidefinite_tag; structural dispatch routes through the configured solver.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def log_prob(self, u: Float[Array, " ..."]) -> Float[Array, ""]:  # ty: ignore[invalid-method-override]
    r"""Variational log density ``\log q(u)`` via `gaussx.gaussian_log_prob`.

    The covariance operator is built from the per-coordinate scales
    as a `lineax.MatrixLinearOperator` tagged
    `positive_semidefinite_tag`; structural dispatch routes
    through the configured solver.
    """
    return gaussian_log_prob(
        self.mean,
        _diag_cov_operator(self.scale),
        u,
        solver=_resolve_solver(self.solver),
    )

kl_divergence(prior_cov: lx.AbstractLinearOperator) -> Float[Array, '']

KL(q(u) || p(u)) against an inducing prior with zero mean.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def kl_divergence(self, prior_cov: lx.AbstractLinearOperator) -> Float[Array, ""]:
    r"""``KL(q(u) || p(u))`` against an inducing prior with zero mean."""
    q_loc = self.mean
    q_cov = _diag_cov_operator(self.scale)
    p_loc = jnp.zeros_like(self.mean)
    return dist_kl_divergence(q_loc, q_cov, p_loc, prior_cov)

predict(K_xz: Float[Array, 'N M'], K_zz_op: lx.AbstractLinearOperator, K_xx_diag: Float[Array, ' N']) -> tuple[Float[Array, ' N'], Float[Array, ' N']]

Predictive (mean, variance) at the points whose cross-cov is K_xz.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def predict(
    self,
    K_xz: Float[Array, "N M"],
    K_zz_op: lx.AbstractLinearOperator,
    K_xx_diag: Float[Array, " N"],
) -> tuple[Float[Array, " N"], Float[Array, " N"]]:
    """Predictive ``(mean, variance)`` at the points whose cross-cov is ``K_xz``."""
    u_cov = jnp.diag(self.scale**2)
    return _svgp_predict_unwhitened(K_xz, K_zz_op, K_xx_diag, self.mean, u_cov)

WhitenedGuide

Bases: Guide

Whitened-coordinate Gaussian variational posterior over inducing values.

Parameterizes q(v) = N(m_v, L_v L_v^T) in whitened coordinates, so the inducing values are u = L_{ZZ} v with L_{ZZ} = chol(K_{ZZ} + jitter I). In whitened coordinates the prior is p(v) = N(0, I) and the KL term has a simple closed form that is independent of the kernel — the standard reparameterization for numerically stable SVGP optimization (Hensman et al., 2015).

Attributes:

Name Type Description
mean Float[Array, ' M']

Variational mean m_v in whitened coordinates, shape (M,).

scale_tril Float[Array, 'M M']

Lower-triangular Cholesky factor of the whitened variational covariance, shape (M, M). The covariance in whitened space is L_v @ L_v.T.

solver AbstractSolverStrategy | None

Optional gaussx.AbstractSolverStrategy — see FullRankGuide for usage. None by default.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
class WhitenedGuide(Guide):
    r"""Whitened-coordinate Gaussian variational posterior over inducing values.

    Parameterizes ``q(v) = N(m_v, L_v L_v^T)`` in *whitened* coordinates,
    so the inducing values are ``u = L_{ZZ} v`` with
    ``L_{ZZ} = chol(K_{ZZ} + jitter I)``. In whitened coordinates the
    prior is ``p(v) = N(0, I)`` and the KL term has a simple closed form
    that is independent of the kernel — the standard reparameterization
    for numerically stable SVGP optimization (Hensman et al., 2015).

    Attributes:
        mean: Variational mean ``m_v`` in whitened coordinates,
            shape ``(M,)``.
        scale_tril: Lower-triangular Cholesky factor of the whitened
            variational covariance, shape ``(M, M)``. The covariance in
            whitened space is ``L_v @ L_v.T``.
        solver: Optional `gaussx.AbstractSolverStrategy` — see
            `FullRankGuide` for usage. ``None`` by default.
    """

    mean: Float[Array, " M"]
    scale_tril: Float[Array, "M M"]
    solver: AbstractSolverStrategy | None = None

    @classmethod
    def init(
        cls,
        num_inducing: int,
        *,
        scale: float = 1.0,
        solver: AbstractSolverStrategy | None = None,
    ) -> WhitenedGuide:
        """Construct a guide initialized to ``N(0, scale^2 I)`` in whitened
        space — see `FullRankGuide.init` for the dtype convention."""
        m = jnp.zeros(num_inducing)
        L = scale * jnp.eye(num_inducing)
        return cls(mean=m, scale_tril=L, solver=solver)

    def sample(self, key: Array) -> Float[Array, " M"]:
        """Draw a single whitened sample ``v = m_v + L_v \\epsilon``."""
        eps = jax.random.normal(key, self.mean.shape, dtype=self.mean.dtype)
        return self.mean + einx.dot("i j, j -> i", self.scale_tril, eps)

    def log_prob(self, v: Float[Array, " ..."]) -> Float[Array, ""]:  # ty: ignore[invalid-method-override]
        r"""Whitened ``\log q(v)`` via `gaussx.gaussian_log_prob`.

        Delegates ``solve`` and ``logdet`` to the configured solver so
        the user-supplied `solver` controls the numerical path
        (the ``KL(q(v) \| N(0, I))`` term remains the kernel-free
        closed form).
        """
        return gaussian_log_prob(
            self.mean,
            _full_cov_operator(self.scale_tril),
            v,
            solver=_resolve_solver(self.solver),
        )

    def kl_divergence(
        self,
        prior_cov: lx.AbstractLinearOperator | None = None,
    ) -> Float[Array, ""]:
        r"""``KL(q(v) || N(0, I))`` — kernel-free closed form.

        Computes

        $$
        \mathrm{KL}(\mathcal{N}(m_v, L_v L_v^\top) \,\|\,
                   \mathcal{N}(0, I))
        = \tfrac{1}{2} \bigl(
            \|m_v\|^2 + \|L_v\|_F^2 - M
            - 2 \sum_i \log |[L_v]_{ii}|
        \bigr).
        $$

        The ``prior_cov`` argument is accepted for signature parity with
        `FullRankGuide.kl_divergence` and
        `MeanFieldGuide.kl_divergence`, but is ignored — the
        whitened prior is the standard normal regardless of ``K_{ZZ}``.
        """
        del prior_cov
        m = self.mean
        L = self.scale_tril
        n = m.shape[0]
        quad = jnp.sum(m**2)
        trace = jnp.sum(L**2)  # ||L||_F^2
        log_det = 2.0 * jnp.sum(jnp.log(jnp.abs(jnp.diag(L))))
        return 0.5 * (quad + trace - n - log_det)

    def predict(
        self,
        K_xz: Float[Array, "N M"],
        K_zz_op: lx.AbstractLinearOperator,
        K_xx_diag: Float[Array, " N"],
    ) -> tuple[Float[Array, " N"], Float[Array, " N"]]:
        """Predictive ``(mean, variance)`` via `gaussx.whitened_svgp_predict`."""
        return whitened_svgp_predict(
            K_zz_op, K_xz, self.mean, self.scale_tril, K_xx_diag
        )

init(num_inducing: int, *, scale: float = 1.0, solver: AbstractSolverStrategy | None = None) -> WhitenedGuide classmethod

Construct a guide initialized to N(0, scale^2 I) in whitened space — see FullRankGuide.init for the dtype convention.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
@classmethod
def init(
    cls,
    num_inducing: int,
    *,
    scale: float = 1.0,
    solver: AbstractSolverStrategy | None = None,
) -> WhitenedGuide:
    """Construct a guide initialized to ``N(0, scale^2 I)`` in whitened
    space — see `FullRankGuide.init` for the dtype convention."""
    m = jnp.zeros(num_inducing)
    L = scale * jnp.eye(num_inducing)
    return cls(mean=m, scale_tril=L, solver=solver)

sample(key: Array) -> Float[Array, ' M']

Draw a single whitened sample v = m_v + L_v \epsilon.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def sample(self, key: Array) -> Float[Array, " M"]:
    """Draw a single whitened sample ``v = m_v + L_v \\epsilon``."""
    eps = jax.random.normal(key, self.mean.shape, dtype=self.mean.dtype)
    return self.mean + einx.dot("i j, j -> i", self.scale_tril, eps)

log_prob(v: Float[Array, ' ...']) -> Float[Array, '']

Whitened \log q(v) via gaussx.gaussian_log_prob.

Delegates solve and logdet to the configured solver so the user-supplied solver controls the numerical path (the KL(q(v) \| N(0, I)) term remains the kernel-free closed form).

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def log_prob(self, v: Float[Array, " ..."]) -> Float[Array, ""]:  # ty: ignore[invalid-method-override]
    r"""Whitened ``\log q(v)`` via `gaussx.gaussian_log_prob`.

    Delegates ``solve`` and ``logdet`` to the configured solver so
    the user-supplied `solver` controls the numerical path
    (the ``KL(q(v) \| N(0, I))`` term remains the kernel-free
    closed form).
    """
    return gaussian_log_prob(
        self.mean,
        _full_cov_operator(self.scale_tril),
        v,
        solver=_resolve_solver(self.solver),
    )

kl_divergence(prior_cov: lx.AbstractLinearOperator | None = None) -> Float[Array, '']

KL(q(v) || N(0, I)) — kernel-free closed form.

Computes

\[ \mathrm{KL}(\mathcal{N}(m_v, L_v L_v^\top) \,\|\, \mathcal{N}(0, I)) = \tfrac{1}{2} \bigl( \|m_v\|^2 + \|L_v\|_F^2 - M - 2 \sum_i \log |[L_v]_{ii}| \bigr). \]

The prior_cov argument is accepted for signature parity with FullRankGuide.kl_divergence and MeanFieldGuide.kl_divergence, but is ignored — the whitened prior is the standard normal regardless of K_{ZZ}.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def kl_divergence(
    self,
    prior_cov: lx.AbstractLinearOperator | None = None,
) -> Float[Array, ""]:
    r"""``KL(q(v) || N(0, I))`` — kernel-free closed form.

    Computes

    $$
    \mathrm{KL}(\mathcal{N}(m_v, L_v L_v^\top) \,\|\,
               \mathcal{N}(0, I))
    = \tfrac{1}{2} \bigl(
        \|m_v\|^2 + \|L_v\|_F^2 - M
        - 2 \sum_i \log |[L_v]_{ii}|
    \bigr).
    $$

    The ``prior_cov`` argument is accepted for signature parity with
    `FullRankGuide.kl_divergence` and
    `MeanFieldGuide.kl_divergence`, but is ignored — the
    whitened prior is the standard normal regardless of ``K_{ZZ}``.
    """
    del prior_cov
    m = self.mean
    L = self.scale_tril
    n = m.shape[0]
    quad = jnp.sum(m**2)
    trace = jnp.sum(L**2)  # ||L||_F^2
    log_det = 2.0 * jnp.sum(jnp.log(jnp.abs(jnp.diag(L))))
    return 0.5 * (quad + trace - n - log_det)

predict(K_xz: Float[Array, 'N M'], K_zz_op: lx.AbstractLinearOperator, K_xx_diag: Float[Array, ' N']) -> tuple[Float[Array, ' N'], Float[Array, ' N']]

Predictive (mean, variance) via gaussx.whitened_svgp_predict.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def predict(
    self,
    K_xz: Float[Array, "N M"],
    K_zz_op: lx.AbstractLinearOperator,
    K_xx_diag: Float[Array, " N"],
) -> tuple[Float[Array, " N"], Float[Array, " N"]]:
    """Predictive ``(mean, variance)`` via `gaussx.whitened_svgp_predict`."""
    return whitened_svgp_predict(
        K_zz_op, K_xz, self.mean, self.scale_tril, K_xx_diag
    )

NaturalGuide

Bases: Guide

Natural-parameter Gaussian variational posterior over inducing values.

Parameterizes \(q(u) = \mathcal{N}(m, S)\) in natural form with parameters

\[ \eta_1 = S^{-1} m, \qquad \eta_2 = -\tfrac{1}{2} S^{-1}. \]

The moments are recovered on demand via mean and covariance, which delegate to gaussx.natural_to_mean_cov.

The natural form is the parameterization of choice for natural- gradient and conjugate-computation variational inference (CVI) workflows: when the true posterior is in the same exponential family as the prior, the natural-gradient direction equals the difference in natural parameters. natural_update exposes that step with a damping factor rho and delegates to gaussx.damped_natural_update so the same primitive is shared with future natural-gradient EP / VI / Newton workflows.

Attributes:

Name Type Description
nat1 Float[Array, ' M']

First natural parameter eta_1 of shape (M,).

nat2 Float[Array, 'M M']

Second natural parameter eta_2 of shape (M, M), symmetric negative-definite.

solver AbstractSolverStrategy | None

Optional gaussx.AbstractSolverStrategy used for solve and logdet against the precision operator Lambda = -2 nat2. sample and log_prob route through gaussx.MultivariateNormalPrecision with this solver — efficient because the precision form avoids ever materializing \(\Sigma\). None defaults to gaussx.DenseSolver.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
class NaturalGuide(Guide):
    r"""Natural-parameter Gaussian variational posterior over inducing values.

    Parameterizes $q(u) = \mathcal{N}(m, S)$ in *natural form*
    with parameters

    $$
    \eta_1 = S^{-1} m, \qquad
    \eta_2 = -\tfrac{1}{2} S^{-1}.
    $$

    The moments are recovered on demand via `mean` and
    `covariance`, which delegate to `gaussx.natural_to_mean_cov`.

    The natural form is the parameterization of choice for *natural-
    gradient* and *conjugate-computation variational inference* (CVI)
    workflows: when the true posterior is in the same exponential family
    as the prior, the natural-gradient direction equals the difference
    in natural parameters. `natural_update` exposes that step
    with a damping factor ``rho`` and delegates to
    `gaussx.damped_natural_update` so the same primitive is shared
    with future natural-gradient EP / VI / Newton workflows.

    Attributes:
        nat1: First natural parameter ``eta_1`` of shape ``(M,)``.
        nat2: Second natural parameter ``eta_2`` of shape ``(M, M)``,
            symmetric negative-definite.
        solver: Optional `gaussx.AbstractSolverStrategy` used for
            ``solve`` and ``logdet`` against the precision operator
            ``Lambda = -2 nat2``. ``sample`` and ``log_prob`` route
            through `gaussx.MultivariateNormalPrecision` with
            this solver — efficient because the precision form avoids
            ever materializing $\Sigma$. ``None`` defaults to
            `gaussx.DenseSolver`.
    """

    nat1: Float[Array, " M"]
    nat2: Float[Array, "M M"]
    solver: AbstractSolverStrategy | None = None

    @classmethod
    def init(
        cls,
        num_inducing: int,
        *,
        scale: float = 1.0,
        solver: AbstractSolverStrategy | None = None,
    ) -> NaturalGuide:
        """Construct a guide initialized to ``N(0, scale^2 I)`` in moment space.

        Mapped to natural form this is ``eta_1 = 0`` and
        ``eta_2 = -1 / (2 scale^2) * I`` — see `FullRankGuide.init`
        for the dtype convention.
        """
        n = num_inducing
        nat1 = jnp.zeros(n)
        nat2 = (-0.5 / (scale**2)) * jnp.eye(n)
        return cls(nat1=nat1, nat2=nat2, solver=solver)

    def _precision_operator(self) -> lx.AbstractLinearOperator:
        """Build the precision operator ``Lambda = -2 nat2`` (PSD)."""
        return _possemi_cov_operator(-2.0 * self.nat2)

    def _moment_mean(self) -> Float[Array, " M"]:
        """Moment-form mean via one structured solve — no covariance build.

        `gaussx.natural_to_mean_cov` returns the covariance as a
        *lazy* inverse operator, so discarding it here is free.
        """
        nat2_op = _negsemi_cov_operator(self.nat2)
        m, _ = natural_to_mean_cov(
            self.nat1, nat2_op, solver=_resolve_solver(self.solver)
        )
        return m

    def _moments(self) -> tuple[Float[Array, " M"], Float[Array, "M M"]]:
        """Return ``(mean, cov_array)`` via `gaussx.natural_to_mean_cov`."""
        nat2_op = _negsemi_cov_operator(self.nat2)
        m, cov_op = natural_to_mean_cov(
            self.nat1, nat2_op, solver=_resolve_solver(self.solver)
        )
        return m, symmetrize(cov_op.as_matrix())

    def _mvn(self) -> MultivariateNormalPrecision:
        """Wrap the natural form as a `gaussx.MultivariateNormalPrecision`.

        Carries the precision operator directly, so ``sample`` /
        ``log_prob`` never materialize the covariance.
        """
        return MultivariateNormalPrecision(
            loc=self._moment_mean(),
            prec_operator=self._precision_operator(),
            solver=_resolve_solver(self.solver),
        )

    @property
    def covariance(self) -> Float[Array, "M M"]:
        """Recover the moment-form covariance ``S = (-2 nat2)^{-1}``."""
        return self._moments()[1]

    @property
    def mean(self) -> Float[Array, " M"]:
        """Recover the moment-form mean ``m = S nat1``."""
        return self._moment_mean()

    def sample(self, key: Array) -> Float[Array, " M"]:
        r"""Draw ``u \sim q`` via `gaussx.MultivariateNormalPrecision`.

        Uses the precision Cholesky directly: ``L = chol(\Lambda)``,
        ``y = L^{-T} \epsilon``, ``u = m + y``. No moment-space
        Cholesky and no ``\Sigma`` materialization.
        """
        return self._mvn().sample(key)

    def log_prob(self, u: Float[Array, " ..."]) -> Float[Array, ""]:  # ty: ignore[invalid-method-override]
        r"""Variational log density ``\log q(u)`` via the precision MVN.

        Avoids the moment-space ``\Sigma^{-1}`` solve: the quadratic
        form is one matvec ``\Lambda (u - m)``, and the log-determinant
        comes from the configured solver on ``\Lambda``.
        """
        return self._mvn().log_prob(u)

    def kl_divergence(self, prior_cov: lx.AbstractLinearOperator) -> Float[Array, ""]:
        r"""``KL(q(u) || p(u))`` against an inducing prior with zero mean.

        Falls back to `gaussx.dist_kl_divergence` in moment form —
        ``dist_kl_divergence`` does not yet take an explicit solver, but
        it dispatches on operator structure for the trace and logdet
        terms.
        """
        m, cov = self._moments()
        p_loc = jnp.zeros_like(m)
        return dist_kl_divergence(m, _possemi_cov_operator(cov), p_loc, prior_cov)

    def predict(
        self,
        K_xz: Float[Array, "N M"],
        K_zz_op: lx.AbstractLinearOperator,
        K_xx_diag: Float[Array, " N"],
    ) -> tuple[Float[Array, " N"], Float[Array, " N"]]:
        """Predictive ``(mean, variance)`` at points with cross-cov ``K_xz``."""
        m, cov = self._moments()
        return _svgp_predict_unwhitened(K_xz, K_zz_op, K_xx_diag, m, cov)

    def natural_update(
        self,
        nat1_hat: Float[Array, " M"],
        nat2_hat: Float[Array, "M M"],
        rho: float | Float[Array, ""] = 1.0,
    ) -> NaturalGuide:
        r"""Damped natural-parameter update via `gaussx.damped_natural_update`.

        Returns a new `NaturalGuide` whose natural parameters are
        the convex combination

        $$
        \eta_i \leftarrow (1 - \rho)\,\eta_i + \rho\,\hat{\eta}_i,
        \quad i \in \{1, 2\}.
        $$

        The damping factor ``rho`` interpolates between the current
        guide (``rho=0``) and the candidate update
        ``(nat1_hat, nat2_hat)`` (``rho=1``). Convex combinations in
        natural-parameter space preserve membership in the Gaussian
        exponential family — in particular ``rho * nat2_hat + (1 - rho)
        * self.nat2`` stays symmetric negative-definite when both
        endpoints are. CVI-style site updates rely on exactly this
        property.
        """
        new_nat1, new_nat2 = damped_natural_update(
            self.nat1,
            self.nat2,
            nat1_hat,
            nat2_hat,
            lr=rho,  # ty: ignore[invalid-argument-type]
        )
        # gaussx returns Array | AbstractLinearOperator; we always pass Arrays
        # in, so the runtime type is Array.
        assert not isinstance(new_nat2, lx.AbstractLinearOperator)
        return NaturalGuide(nat1=new_nat1, nat2=new_nat2)

covariance: Float[Array, 'M M'] property

Recover the moment-form covariance S = (-2 nat2)^{-1}.

mean: Float[Array, ' M'] property

Recover the moment-form mean m = S nat1.

init(num_inducing: int, *, scale: float = 1.0, solver: AbstractSolverStrategy | None = None) -> NaturalGuide classmethod

Construct a guide initialized to N(0, scale^2 I) in moment space.

Mapped to natural form this is eta_1 = 0 and eta_2 = -1 / (2 scale^2) * I — see FullRankGuide.init for the dtype convention.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
@classmethod
def init(
    cls,
    num_inducing: int,
    *,
    scale: float = 1.0,
    solver: AbstractSolverStrategy | None = None,
) -> NaturalGuide:
    """Construct a guide initialized to ``N(0, scale^2 I)`` in moment space.

    Mapped to natural form this is ``eta_1 = 0`` and
    ``eta_2 = -1 / (2 scale^2) * I`` — see `FullRankGuide.init`
    for the dtype convention.
    """
    n = num_inducing
    nat1 = jnp.zeros(n)
    nat2 = (-0.5 / (scale**2)) * jnp.eye(n)
    return cls(nat1=nat1, nat2=nat2, solver=solver)

sample(key: Array) -> Float[Array, ' M']

Draw u \sim q via gaussx.MultivariateNormalPrecision.

Uses the precision Cholesky directly: L = chol(\Lambda), y = L^{-T} \epsilon, u = m + y. No moment-space Cholesky and no \Sigma materialization.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def sample(self, key: Array) -> Float[Array, " M"]:
    r"""Draw ``u \sim q`` via `gaussx.MultivariateNormalPrecision`.

    Uses the precision Cholesky directly: ``L = chol(\Lambda)``,
    ``y = L^{-T} \epsilon``, ``u = m + y``. No moment-space
    Cholesky and no ``\Sigma`` materialization.
    """
    return self._mvn().sample(key)

log_prob(u: Float[Array, ' ...']) -> Float[Array, '']

Variational log density \log q(u) via the precision MVN.

Avoids the moment-space \Sigma^{-1} solve: the quadratic form is one matvec \Lambda (u - m), and the log-determinant comes from the configured solver on \Lambda.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def log_prob(self, u: Float[Array, " ..."]) -> Float[Array, ""]:  # ty: ignore[invalid-method-override]
    r"""Variational log density ``\log q(u)`` via the precision MVN.

    Avoids the moment-space ``\Sigma^{-1}`` solve: the quadratic
    form is one matvec ``\Lambda (u - m)``, and the log-determinant
    comes from the configured solver on ``\Lambda``.
    """
    return self._mvn().log_prob(u)

kl_divergence(prior_cov: lx.AbstractLinearOperator) -> Float[Array, '']

KL(q(u) || p(u)) against an inducing prior with zero mean.

Falls back to gaussx.dist_kl_divergence in moment form — dist_kl_divergence does not yet take an explicit solver, but it dispatches on operator structure for the trace and logdet terms.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def kl_divergence(self, prior_cov: lx.AbstractLinearOperator) -> Float[Array, ""]:
    r"""``KL(q(u) || p(u))`` against an inducing prior with zero mean.

    Falls back to `gaussx.dist_kl_divergence` in moment form —
    ``dist_kl_divergence`` does not yet take an explicit solver, but
    it dispatches on operator structure for the trace and logdet
    terms.
    """
    m, cov = self._moments()
    p_loc = jnp.zeros_like(m)
    return dist_kl_divergence(m, _possemi_cov_operator(cov), p_loc, prior_cov)

predict(K_xz: Float[Array, 'N M'], K_zz_op: lx.AbstractLinearOperator, K_xx_diag: Float[Array, ' N']) -> tuple[Float[Array, ' N'], Float[Array, ' N']]

Predictive (mean, variance) at points with cross-cov K_xz.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def predict(
    self,
    K_xz: Float[Array, "N M"],
    K_zz_op: lx.AbstractLinearOperator,
    K_xx_diag: Float[Array, " N"],
) -> tuple[Float[Array, " N"], Float[Array, " N"]]:
    """Predictive ``(mean, variance)`` at points with cross-cov ``K_xz``."""
    m, cov = self._moments()
    return _svgp_predict_unwhitened(K_xz, K_zz_op, K_xx_diag, m, cov)

natural_update(nat1_hat: Float[Array, ' M'], nat2_hat: Float[Array, 'M M'], rho: float | Float[Array, ''] = 1.0) -> NaturalGuide

Damped natural-parameter update via gaussx.damped_natural_update.

Returns a new NaturalGuide whose natural parameters are the convex combination

\[ \eta_i \leftarrow (1 - \rho)\,\eta_i + \rho\,\hat{\eta}_i, \quad i \in \{1, 2\}. \]

The damping factor rho interpolates between the current guide (rho=0) and the candidate update (nat1_hat, nat2_hat) (rho=1). Convex combinations in natural-parameter space preserve membership in the Gaussian exponential family — in particular rho * nat2_hat + (1 - rho) * self.nat2 stays symmetric negative-definite when both endpoints are. CVI-style site updates rely on exactly this property.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def natural_update(
    self,
    nat1_hat: Float[Array, " M"],
    nat2_hat: Float[Array, "M M"],
    rho: float | Float[Array, ""] = 1.0,
) -> NaturalGuide:
    r"""Damped natural-parameter update via `gaussx.damped_natural_update`.

    Returns a new `NaturalGuide` whose natural parameters are
    the convex combination

    $$
    \eta_i \leftarrow (1 - \rho)\,\eta_i + \rho\,\hat{\eta}_i,
    \quad i \in \{1, 2\}.
    $$

    The damping factor ``rho`` interpolates between the current
    guide (``rho=0``) and the candidate update
    ``(nat1_hat, nat2_hat)`` (``rho=1``). Convex combinations in
    natural-parameter space preserve membership in the Gaussian
    exponential family — in particular ``rho * nat2_hat + (1 - rho)
    * self.nat2`` stays symmetric negative-definite when both
    endpoints are. CVI-style site updates rely on exactly this
    property.
    """
    new_nat1, new_nat2 = damped_natural_update(
        self.nat1,
        self.nat2,
        nat1_hat,
        nat2_hat,
        lr=rho,  # ty: ignore[invalid-argument-type]
    )
    # gaussx returns Array | AbstractLinearOperator; we always pass Arrays
    # in, so the runtime type is Array.
    assert not isinstance(new_nat2, lx.AbstractLinearOperator)
    return NaturalGuide(nat1=new_nat1, nat2=new_nat2)

DeltaGuide

Bases: Guide

Point-mass (MAP-style) variational posterior over inducing values.

\[ q(u) = \delta(u - \text{loc}), \]

so all the variational mass concentrates on a single point. Pairs with the same SVGP ELBO infrastructure as the Gaussian guides but reduces it to MAP estimation: when ELL - kl_divergence is the objective, kl_divergence returns the loc-dependent \(-\log p(\text{loc})\) so the objective becomes the joint log-density log p(y, loc) = ELL(loc) + log p(loc), recovering standard MAP estimation of the inducing values.

Attributes:

Name Type Description
loc Float[Array, ' M']

The point at which the variational mass concentrates, shape (M,). This is also the MAP estimate of the inducing values when used inside a maximization loop.

solver AbstractSolverStrategy | None

Optional gaussx.AbstractSolverStrategy passed to gaussx.gaussian_log_prob when computing -log p(loc) against the prior covariance in kl_divergence. None defaults to gaussx.DenseSolver.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
class DeltaGuide(Guide):
    r"""Point-mass (MAP-style) variational posterior over inducing values.

    $$
    q(u) = \delta(u - \text{loc}),
    $$

    so all the variational mass concentrates on a single point. Pairs
    with the same SVGP ELBO infrastructure as the Gaussian guides but
    reduces it to MAP estimation: when ``ELL - kl_divergence`` is the
    objective, `kl_divergence` returns the loc-dependent
    $-\log p(\text{loc})$ so the objective becomes the joint
    log-density ``log p(y, loc) = ELL(loc) + log p(loc)``, recovering
    standard MAP estimation of the inducing values.

    Attributes:
        loc: The point at which the variational mass concentrates,
            shape ``(M,)``. This is also the MAP estimate of the
            inducing values when used inside a maximization loop.
        solver: Optional `gaussx.AbstractSolverStrategy` passed
            to `gaussx.gaussian_log_prob` when computing
            ``-log p(loc)`` against the prior covariance in
            `kl_divergence`. ``None`` defaults to
            `gaussx.DenseSolver`.
    """

    loc: Float[Array, " M"]
    solver: AbstractSolverStrategy | None = None

    @classmethod
    def init(
        cls,
        num_inducing: int,
        *,
        solver: AbstractSolverStrategy | None = None,
    ) -> DeltaGuide:
        """Construct a guide initialized to ``loc = 0`` — see
        `FullRankGuide.init` for the dtype convention."""
        return cls(loc=jnp.zeros(num_inducing), solver=solver)

    def sample(self, key: Array) -> Float[Array, " M"]:
        """Return ``self.loc`` — the variational draw is deterministic."""
        del key
        return self.loc

    def log_prob(self, u: Float[Array, " ..."]) -> Float[Array, ""]:  # ty: ignore[invalid-method-override]
        r"""Variational log density — constant ``0`` by convention.

        The strict density of a Dirac delta is ``+inf`` at ``loc`` and
        ``-inf`` everywhere else, which carries no useful gradient
        information. Returning ``0`` matches the Pyro / NumPyro
        ``AutoDelta`` convention: the differentiable MAP signal lives
        entirely in `kl_divergence`.
        """
        del u
        return jnp.zeros((), dtype=self.loc.dtype)

    def kl_divergence(self, prior_cov: lx.AbstractLinearOperator) -> Float[Array, ""]:
        r"""Return the loc-dependent ``-log p(loc)`` for the MAP objective.

        The strict KL of a Dirac delta against a continuous prior is
        ``+inf``, but the only loc-dependent piece is the negative log
        prior density at ``loc``,

        $$
        -\log p(\text{loc}) =
            \tfrac{1}{2} \text{loc}^\top K^{-1} \text{loc}
            + \tfrac{1}{2} \log |K|
            + \tfrac{M}{2} \log(2\pi),
        $$

        so we return that. With this convention the standard ELBO
        ``ELL - kl_divergence`` reduces to the joint log-density
        ``log p(y, loc)``, which is exactly the MAP objective. Computed
        via `gaussx.gaussian_log_prob` so the same solver / logdet
        primitives back this path as the rest of the GP surface.
        """
        loc = self.loc
        return -gaussian_log_prob(
            jnp.zeros_like(loc), prior_cov, loc, solver=_resolve_solver(self.solver)
        )

    def predict(
        self,
        K_xz: Float[Array, "N M"],
        K_zz_op: lx.AbstractLinearOperator,
        K_xx_diag: Float[Array, " N"],
    ) -> tuple[Float[Array, " N"], Float[Array, " N"]]:
        r"""Predictive ``(mean, variance)`` conditioning on ``u = loc``.

        With ``u = loc`` deterministically, the predictive is the prior
        conditional ``p(f_* | u = loc)``: the mean is
        ``K_{xZ} K_{ZZ}^{-1} loc`` and the variance is the prior
        reduction ``k(x, x) - K_{xZ} K_{ZZ}^{-1} K_{Zx}`` *exactly* —
        no posterior-uncertainty contribution and no Cholesky jitter.

        Implemented by calling `gaussx.whitened_svgp_predict`
        directly with the whitened mean ``L_{ZZ}^{-1} loc`` and a zero
        whitened Cholesky factor. The shared
        `_svgp_predict_unwhitened` helper would route the zero
        variational covariance through `gaussx.safe_cholesky`,
        which injects jitter to make the input PD and would add a tiny
        spurious variance term. Bypassing that path keeps the
        `DeltaGuide` predictive numerically equal to the prior
        conditional.
        """
        m_size = self.loc.shape[0]
        L_zz = cholesky(K_zz_op)
        u_mean_white = lx.linear_solve(L_zz, self.loc).value
        zero_chol = jnp.zeros((m_size, m_size), dtype=self.loc.dtype)
        return whitened_svgp_predict(K_zz_op, K_xz, u_mean_white, zero_chol, K_xx_diag)

init(num_inducing: int, *, solver: AbstractSolverStrategy | None = None) -> DeltaGuide classmethod

Construct a guide initialized to loc = 0 — see FullRankGuide.init for the dtype convention.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
@classmethod
def init(
    cls,
    num_inducing: int,
    *,
    solver: AbstractSolverStrategy | None = None,
) -> DeltaGuide:
    """Construct a guide initialized to ``loc = 0`` — see
    `FullRankGuide.init` for the dtype convention."""
    return cls(loc=jnp.zeros(num_inducing), solver=solver)

sample(key: Array) -> Float[Array, ' M']

Return self.loc — the variational draw is deterministic.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def sample(self, key: Array) -> Float[Array, " M"]:
    """Return ``self.loc`` — the variational draw is deterministic."""
    del key
    return self.loc

log_prob(u: Float[Array, ' ...']) -> Float[Array, '']

Variational log density — constant 0 by convention.

The strict density of a Dirac delta is +inf at loc and -inf everywhere else, which carries no useful gradient information. Returning 0 matches the Pyro / NumPyro AutoDelta convention: the differentiable MAP signal lives entirely in kl_divergence.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def log_prob(self, u: Float[Array, " ..."]) -> Float[Array, ""]:  # ty: ignore[invalid-method-override]
    r"""Variational log density — constant ``0`` by convention.

    The strict density of a Dirac delta is ``+inf`` at ``loc`` and
    ``-inf`` everywhere else, which carries no useful gradient
    information. Returning ``0`` matches the Pyro / NumPyro
    ``AutoDelta`` convention: the differentiable MAP signal lives
    entirely in `kl_divergence`.
    """
    del u
    return jnp.zeros((), dtype=self.loc.dtype)

kl_divergence(prior_cov: lx.AbstractLinearOperator) -> Float[Array, '']

Return the loc-dependent -log p(loc) for the MAP objective.

The strict KL of a Dirac delta against a continuous prior is +inf, but the only loc-dependent piece is the negative log prior density at loc,

\[ -\log p(\text{loc}) = \tfrac{1}{2} \text{loc}^\top K^{-1} \text{loc} + \tfrac{1}{2} \log |K| + \tfrac{M}{2} \log(2\pi), \]

so we return that. With this convention the standard ELBO ELL - kl_divergence reduces to the joint log-density log p(y, loc), which is exactly the MAP objective. Computed via gaussx.gaussian_log_prob so the same solver / logdet primitives back this path as the rest of the GP surface.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def kl_divergence(self, prior_cov: lx.AbstractLinearOperator) -> Float[Array, ""]:
    r"""Return the loc-dependent ``-log p(loc)`` for the MAP objective.

    The strict KL of a Dirac delta against a continuous prior is
    ``+inf``, but the only loc-dependent piece is the negative log
    prior density at ``loc``,

    $$
    -\log p(\text{loc}) =
        \tfrac{1}{2} \text{loc}^\top K^{-1} \text{loc}
        + \tfrac{1}{2} \log |K|
        + \tfrac{M}{2} \log(2\pi),
    $$

    so we return that. With this convention the standard ELBO
    ``ELL - kl_divergence`` reduces to the joint log-density
    ``log p(y, loc)``, which is exactly the MAP objective. Computed
    via `gaussx.gaussian_log_prob` so the same solver / logdet
    primitives back this path as the rest of the GP surface.
    """
    loc = self.loc
    return -gaussian_log_prob(
        jnp.zeros_like(loc), prior_cov, loc, solver=_resolve_solver(self.solver)
    )

predict(K_xz: Float[Array, 'N M'], K_zz_op: lx.AbstractLinearOperator, K_xx_diag: Float[Array, ' N']) -> tuple[Float[Array, ' N'], Float[Array, ' N']]

Predictive (mean, variance) conditioning on u = loc.

With u = loc deterministically, the predictive is the prior conditional p(f_* | u = loc): the mean is K_{xZ} K_{ZZ}^{-1} loc and the variance is the prior reduction k(x, x) - K_{xZ} K_{ZZ}^{-1} K_{Zx} exactly — no posterior-uncertainty contribution and no Cholesky jitter.

Implemented by calling gaussx.whitened_svgp_predict directly with the whitened mean L_{ZZ}^{-1} loc and a zero whitened Cholesky factor. The shared _svgp_predict_unwhitened helper would route the zero variational covariance through gaussx.safe_cholesky, which injects jitter to make the input PD and would add a tiny spurious variance term. Bypassing that path keeps the DeltaGuide predictive numerically equal to the prior conditional.

Source code in packages/pyrox-gp/src/pyrox_gp/_guides.py
def predict(
    self,
    K_xz: Float[Array, "N M"],
    K_zz_op: lx.AbstractLinearOperator,
    K_xx_diag: Float[Array, " N"],
) -> tuple[Float[Array, " N"], Float[Array, " N"]]:
    r"""Predictive ``(mean, variance)`` conditioning on ``u = loc``.

    With ``u = loc`` deterministically, the predictive is the prior
    conditional ``p(f_* | u = loc)``: the mean is
    ``K_{xZ} K_{ZZ}^{-1} loc`` and the variance is the prior
    reduction ``k(x, x) - K_{xZ} K_{ZZ}^{-1} K_{Zx}`` *exactly* —
    no posterior-uncertainty contribution and no Cholesky jitter.

    Implemented by calling `gaussx.whitened_svgp_predict`
    directly with the whitened mean ``L_{ZZ}^{-1} loc`` and a zero
    whitened Cholesky factor. The shared
    `_svgp_predict_unwhitened` helper would route the zero
    variational covariance through `gaussx.safe_cholesky`,
    which injects jitter to make the input PD and would add a tiny
    spurious variance term. Bypassing that path keeps the
    `DeltaGuide` predictive numerically equal to the prior
    conditional.
    """
    m_size = self.loc.shape[0]
    L_zz = cholesky(K_zz_op)
    u_mean_white = lx.linear_solve(L_zz, self.loc).value
    zero_chol = jnp.zeros((m_size, m_size), dtype=self.loc.dtype)
    return whitened_svgp_predict(K_zz_op, K_xz, u_mean_white, zero_chol, K_xx_diag)

Likelihoods

Observation models for latent-GP workflows. Each maps latent function values to a summed log-density log_prob(f, y); DistLikelihood wraps any numpyro.distributions.Distribution factory for one-off models.

GaussianLikelihood

Bases: Likelihood

Gaussian observation model \(p(y \mid f) = N(y \mid f, \sigma^2)\).

The only likelihood with a closed-form expected log-likelihood, enabling the analytical Titsias ELBO via gaussx.variational_elbo_gaussian.

Attributes:

Name Type Description
noise_var float | Float[Array, '']

Observation noise variance \(\sigma^2\).

Source code in packages/pyrox-gp/src/pyrox_gp/_likelihoods.py
class GaussianLikelihood(Likelihood):
    r"""Gaussian observation model $p(y \mid f) = N(y \mid f, \sigma^2)$.

    The only likelihood with a closed-form expected log-likelihood,
    enabling the analytical Titsias ELBO via
    `gaussx.variational_elbo_gaussian`.

    Attributes:
        noise_var: Observation noise variance $\sigma^2$.
    """

    noise_var: float | Float[Array, ""]

    def log_prob(
        self,
        f: Float[Array, " ..."],
        y: Float[Array, " ..."],
        X: Float[Array, " ..."] | None = None,
    ) -> Float[Array, ""]:
        r"""Sum of per-point Gaussian log-densities."""
        return nd.Normal(f, jnp.sqrt(self.noise_var)).log_prob(y).sum()

log_prob(f: Float[Array, ' ...'], y: Float[Array, ' ...'], X: Float[Array, ' ...'] | None = None) -> Float[Array, '']

Sum of per-point Gaussian log-densities.

Source code in packages/pyrox-gp/src/pyrox_gp/_likelihoods.py
def log_prob(
    self,
    f: Float[Array, " ..."],
    y: Float[Array, " ..."],
    X: Float[Array, " ..."] | None = None,
) -> Float[Array, ""]:
    r"""Sum of per-point Gaussian log-densities."""
    return nd.Normal(f, jnp.sqrt(self.noise_var)).log_prob(y).sum()

HeteroscedasticGaussianLikelihood

Bases: Likelihood

Gaussian regression with input-dependent noise.

Each observation consumes two latents — the mean and the log-noise-standard-deviation: \(p(y_n | f_n^{(0)}, f_n^{(1)}) = N(y_n | f_n^{(0)}, e^{2 f_n^{(1)}})\).

Multi-latent (latent_dim = 2). The scalar-latent advanced inference strategies reject this likelihood; use SVGP for now.

Source code in packages/pyrox-gp/src/pyrox_gp/_likelihoods.py
class HeteroscedasticGaussianLikelihood(Likelihood):
    r"""Gaussian regression with input-dependent noise.

    Each observation consumes two latents — the mean and the
    log-noise-standard-deviation:
    $p(y_n | f_n^{(0)}, f_n^{(1)}) = N(y_n | f_n^{(0)}, e^{2 f_n^{(1)}})$.

    Multi-latent (``latent_dim = 2``). The scalar-latent advanced
    inference strategies reject this likelihood; use SVGP for now.
    """

    latent_dim: int = eqx.field(static=True, default=2)

    def log_prob(
        self,
        f: Float[Array, "N 2"],
        y: Float[Array, " N"],
        X: Float[Array, " ..."] | None = None,
    ) -> Float[Array, ""]:
        loc = f[..., 0]
        log_scale = f[..., 1]
        return nd.Normal(loc=loc, scale=jnp.exp(log_scale)).log_prob(y).sum()

BernoulliLikelihood

Bases: Likelihood

Binary classification likelihood with logit link.

p(y \mid f) = \mathrm{Bernoulli}(\sigma(f)) where \(\sigma\) is the logistic function. Targets y are {0, 1} valued. Scalar latent (latent_dim = 1).

Source code in packages/pyrox-gp/src/pyrox_gp/_likelihoods.py
class BernoulliLikelihood(Likelihood):
    r"""Binary classification likelihood with logit link.

    ``p(y \mid f) = \mathrm{Bernoulli}(\sigma(f))`` where
    $\sigma$ is the logistic function. Targets ``y`` are
    ``{0, 1}`` valued. Scalar latent (``latent_dim = 1``).
    """

    def log_prob(
        self,
        f: Float[Array, " ..."],
        y: Float[Array, " ..."],
        X: Float[Array, " ..."] | None = None,
    ) -> Float[Array, ""]:
        return nd.Bernoulli(logits=f).log_prob(y).sum()

PoissonLikelihood

Bases: Likelihood

Count likelihood with log-link.

p(y \mid f) = \mathrm{Poisson}(\exp(f)). Targets y are non-negative integers. Scalar latent (latent_dim = 1).

Source code in packages/pyrox-gp/src/pyrox_gp/_likelihoods.py
class PoissonLikelihood(Likelihood):
    r"""Count likelihood with log-link.

    ``p(y \mid f) = \mathrm{Poisson}(\exp(f))``. Targets ``y`` are
    non-negative integers. Scalar latent (``latent_dim = 1``).
    """

    def log_prob(
        self,
        f: Float[Array, " ..."],
        y: Float[Array, " ..."],
        X: Float[Array, " ..."] | None = None,
    ) -> Float[Array, ""]:
        return nd.Poisson(rate=jnp.exp(f)).log_prob(y).sum()

SoftmaxLikelihood

Bases: Likelihood

Multi-class classification likelihood with softmax link.

Each observation has one latent function value per class: f has shape (N, num_classes) and p(y_n \mid f_n) = \mathrm{Categorical}(\mathrm{softmax}(f_n)). Targets y are integer class indices in [0, num_classes).

Multi-latent (latent_dim = num_classes). The scalar-latent advanced inference strategies in pyrox_gp._inference_nongauss reject this likelihood with a clear error; use SVGP / MAP for now, or wait for the multi-latent inference follow-up.

Attributes:

Name Type Description
num_classes int

Number of output classes \(C \geq 2\).

Source code in packages/pyrox-gp/src/pyrox_gp/_likelihoods.py
class SoftmaxLikelihood(Likelihood):
    r"""Multi-class classification likelihood with softmax link.

    Each observation has one latent function value per class:
    ``f`` has shape ``(N, num_classes)`` and
    ``p(y_n \mid f_n) = \mathrm{Categorical}(\mathrm{softmax}(f_n))``.
    Targets ``y`` are integer class indices in ``[0, num_classes)``.

    Multi-latent (``latent_dim = num_classes``). The scalar-latent
    advanced inference strategies in `pyrox_gp._inference_nongauss`
    reject this likelihood with a clear error; use SVGP / MAP for now,
    or wait for the multi-latent inference follow-up.

    Attributes:
        num_classes: Number of output classes $C \geq 2$.
    """

    num_classes: int = eqx.field(static=True)
    latent_dim: int = eqx.field(static=True)

    def __init__(self, num_classes: int) -> None:
        if num_classes < 2:
            msg = f"num_classes must be >= 2, got {num_classes}"
            raise ValueError(msg)
        self.num_classes = num_classes
        self.latent_dim = num_classes

    def log_prob(
        self,
        f: Float[Array, "N C"],
        y: Int[Array, " N"],
        X: Float[Array, " ..."] | None = None,
    ) -> Float[Array, ""]:
        return nd.Categorical(logits=f).log_prob(y).sum()

StudentTLikelihood

Bases: Likelihood

Heavy-tailed regression: p(y | f) = StudentT(nu, f, sigma).

Robust to outliers — the heavier-than-Gaussian tails downweight observations that are far from the latent. Scalar latent (latent_dim = 1). Both df and scale are positive.

Attributes:

Name Type Description
df float | Float[Array, '']

Degrees of freedom \(\nu > 0\). Smaller values give heavier tails. df -> infinity recovers the Gaussian.

scale float | Float[Array, '']

Scale parameter \(\sigma > 0\).

Source code in packages/pyrox-gp/src/pyrox_gp/_likelihoods.py
class StudentTLikelihood(Likelihood):
    r"""Heavy-tailed regression: ``p(y | f) = StudentT(nu, f, sigma)``.

    Robust to outliers — the heavier-than-Gaussian tails downweight
    observations that are far from the latent. Scalar latent
    (``latent_dim = 1``). Both ``df`` and ``scale`` are positive.

    Attributes:
        df: Degrees of freedom $\nu > 0$. Smaller values give
            heavier tails. ``df -> infinity`` recovers the Gaussian.
        scale: Scale parameter $\sigma > 0$.
    """

    df: float | Float[Array, ""]
    scale: float | Float[Array, ""]

    def log_prob(
        self,
        f: Float[Array, " ..."],
        y: Float[Array, " ..."],
        X: Float[Array, " ..."] | None = None,
    ) -> Float[Array, ""]:
        return nd.StudentT(df=self.df, loc=f, scale=self.scale).log_prob(y).sum()

DistLikelihood

Bases: Likelihood

Generic likelihood wrapping any numpyro.distributions.Distribution.

The user supplies a link function that maps the latent function value f to a numpyro distribution over observations:

# Bernoulli with logit link
lik = DistLikelihood(lambda f: dist.Bernoulli(logits=f))

# Poisson with log link
lik = DistLikelihood(lambda f: dist.Poisson(rate=jnp.exp(f)))

# Student-t noise
lik = DistLikelihood(lambda f: dist.StudentT(df=3, loc=f, scale=0.5))

The link function cannot carry trainable parameters

dist_fn is a static field, so anything the callable closes over is invisible to eqx.filter_grad. A link function holding learnable parameters will silently never train — there is no error, and the loss still decreases because the kernel and guide keep fitting.

# WRONG -- `warp` is frozen at its initial value forever
lik = DistLikelihood(lambda f: dist.Normal(warp(f), sigma))

For a parameterized link, write a Likelihood subclass and hold the parameters as child modules, so they are ordinary trainable leaves.

The resulting object satisfies the Likelihood protocol and can be passed to svgp_elbo. Because no closed-form expected log-likelihood is available, the ELBO uses numerical integration (GaussHermiteIntegrator or MonteCarloIntegrator from gaussx).

Attributes:

Name Type Description
dist_fn Callable[..., Any]

Callable mapping f to a numpyro.distributions.Distribution.

Source code in packages/pyrox-gp/src/pyrox_gp/_likelihoods.py
class DistLikelihood(Likelihood):
    r"""Generic likelihood wrapping any ``numpyro.distributions.Distribution``.

    The user supplies a *link function* that maps the latent function
    value ``f`` to a numpyro distribution over observations:

    ```python
    # Bernoulli with logit link
    lik = DistLikelihood(lambda f: dist.Bernoulli(logits=f))

    # Poisson with log link
    lik = DistLikelihood(lambda f: dist.Poisson(rate=jnp.exp(f)))

    # Student-t noise
    lik = DistLikelihood(lambda f: dist.StudentT(df=3, loc=f, scale=0.5))
    ```

    !!! warning "The link function cannot carry trainable parameters"
        ``dist_fn`` is a **static** field, so anything the callable closes
        over is invisible to ``eqx.filter_grad``. A link function holding
        learnable parameters will silently never train — there is no error,
        and the loss still decreases because the kernel and guide keep
        fitting.

        ```python
        # WRONG -- `warp` is frozen at its initial value forever
        lik = DistLikelihood(lambda f: dist.Normal(warp(f), sigma))
        ```

        For a parameterized link, write a `Likelihood` subclass and
        hold the parameters as child modules, so they are ordinary
        trainable leaves.

    The resulting object satisfies the `Likelihood` protocol and
    can be passed to `svgp_elbo`. Because no closed-form expected
    log-likelihood is available, the ELBO uses numerical integration
    (``GaussHermiteIntegrator`` or ``MonteCarloIntegrator`` from gaussx).

    Attributes:
        dist_fn: Callable mapping ``f`` to a
            `numpyro.distributions.Distribution`.
    """

    dist_fn: Callable[..., Any] = eqx.field(static=True)

    def log_prob(
        self,
        f: Float[Array, " ..."],
        y: Float[Array, " ..."],
        X: Float[Array, " ..."] | None = None,
    ) -> Float[Array, ""]:
        r"""Sum of per-point log-densities under the wrapped distribution."""
        return self.dist_fn(f).log_prob(y).sum()

log_prob(f: Float[Array, ' ...'], y: Float[Array, ' ...'], X: Float[Array, ' ...'] | None = None) -> Float[Array, '']

Sum of per-point log-densities under the wrapped distribution.

Source code in packages/pyrox-gp/src/pyrox_gp/_likelihoods.py
def log_prob(
    self,
    f: Float[Array, " ..."],
    y: Float[Array, " ..."],
    X: Float[Array, " ..."] | None = None,
) -> Float[Array, ""]:
    r"""Sum of per-point log-densities under the wrapped distribution."""
    return self.dist_fn(f).log_prob(y).sum()

Warped (transformed-GP) likelihood

A transformed Gaussian process (Maronas et al., AISTATS 2021) is an SVGP whose likelihood composes an elementwise monotone warp G; the warp never appears in the KL term, so every existing inference path composes with it unchanged. Needs an integrator (Gauss-Hermite recommended) and, for the smooth recommended warp MixtureGaussianCDF, the pyrox-gp[flows] extra.

The warp may also be input-dependent (G_{phi(x)}) — a non-stationary process from a stationary kernel. Pass a conditional bijection (e.g. gauss_flows.Conditioner wrapping an unconditional warp); svgp_elbo threads the inputs through Likelihood.log_prob(f, y, X) automatically. Cost note: the expected log-likelihood evaluates the warp at every quadrature node, so a conditional warp pays order conditioner forward passes per point.

Only svgp_elbo threads the inputs through. The CVI / natural-gradient strategies, sparse_markov_svgp_elbo and the multi-output ELBO do not, so a conditional warp raises there rather than silently dropping the condition — use an unconditional warp on those paths.

WarpedGaussianLikelihood

Bases: Likelihood

\(p(y \mid f) = N(y \mid G(f), \sigma^2)\) -- the transformed-GP model.

The warp is a child module, so its parameters are ordinary trainable leaves. Passing a warp through a lambda to DistLikelihood does not work: its dist_fn is a static field, so the warp is silently frozen and never trains.

The expected log-likelihood has no closed form, so svgp_elbo requires an integrator.

Prefer Gauss-Hermite, and prefer a smooth warp

The warped integrand has heavy tails, so Monte Carlo integration is high variance -- roughly 4M samples to match Gauss-Hermite at order 20. Gauss-Hermite in turn converges spectrally only for analytic warps: piecewise ones such as RationalQuadraticSpline stall around a 3e-3 error floor and their error is non-monotone in quadrature order, so raising the order is not a valid convergence check. gauss_flows.MixtureGaussianCDF is smooth and reaches machine precision.

Note eqx.filter_jit is required over plain jax.jit when jitting around this likelihood -- flowjax spline pytrees carry string leaves.

Attributes:

Name Type Description
warp AbstractBijection

Any flowjax.bijections.AbstractBijection with event shape () or (1,).

noise_var Float[Array, '']

Observation noise variance \(\sigma^2\).

Examples:

>>> import jax.numpy as jnp
>>> from flowjax.bijections import RationalQuadraticSpline
>>> lik = WarpedGaussianLikelihood(
...     warp=RationalQuadraticSpline(knots=8, interval=4.0),
...     noise_var=jnp.asarray(0.1),
... )
>>> float(lik.log_prob(jnp.zeros(3), jnp.zeros(3))) < 0.0
True
Source code in packages/pyrox-gp/src/pyrox_gp/_warped.py
class WarpedGaussianLikelihood(Likelihood):
    r"""$p(y \mid f) = N(y \mid G(f), \sigma^2)$ -- the transformed-GP model.

    The warp is a **child module**, so its parameters are ordinary trainable
    leaves. Passing a warp through a lambda to
    [`DistLikelihood`][pyrox_gp.DistLikelihood] does *not* work: its
    ``dist_fn`` is a static field, so the warp is silently frozen and never
    trains.

    The expected log-likelihood has no closed form, so
    [`svgp_elbo`][pyrox_gp.svgp_elbo] requires an integrator.

    !!! warning "Prefer Gauss-Hermite, and prefer a smooth warp"
        The warped integrand has heavy tails, so Monte Carlo integration is
        high variance -- roughly 4M samples to match Gauss-Hermite at order
        20. Gauss-Hermite in turn converges spectrally only for *analytic*
        warps: piecewise ones such as ``RationalQuadraticSpline`` stall
        around a ``3e-3`` error floor and their error is **non-monotone** in
        quadrature order, so raising the order is not a valid convergence
        check. ``gauss_flows.MixtureGaussianCDF`` is smooth and reaches
        machine precision.

    Note ``eqx.filter_jit`` is required over plain ``jax.jit`` when jitting
    around this likelihood -- flowjax spline pytrees carry string leaves.

    Attributes:
        warp: Any ``flowjax.bijections.AbstractBijection`` with event shape
            ``()`` or ``(1,)``.
        noise_var: Observation noise variance $\sigma^2$.

    Examples:
        >>> import jax.numpy as jnp
        >>> from flowjax.bijections import RationalQuadraticSpline
        >>> lik = WarpedGaussianLikelihood(
        ...     warp=RationalQuadraticSpline(knots=8, interval=4.0),
        ...     noise_var=jnp.asarray(0.1),
        ... )
        >>> float(lik.log_prob(jnp.zeros(3), jnp.zeros(3))) < 0.0
        True
    """

    warp: AbstractBijection
    noise_var: Float[Array, ""]

    def log_prob(
        self,
        f: Float[Array, " ..."],
        y: Float[Array, " ..."],
        X: Float[Array, " ..."] | None = None,
    ) -> Float[Array, ""]:
        r"""Sum of per-point Gaussian log-densities about $G(f)$.

        ``X`` is required when the warp is conditional (``cond_shape``
        set) and ignored otherwise. Note the cost: the expected
        log-likelihood evaluates the warp at every quadrature node, so a
        conditional warp pays ``order`` conditioner forward passes per
        data point.
        """
        g = _apply_warp(self.warp, f, X)
        return nd.Normal(g, jnp.sqrt(self.noise_var)).log_prob(y).sum()

log_prob(f: Float[Array, ' ...'], y: Float[Array, ' ...'], X: Float[Array, ' ...'] | None = None) -> Float[Array, '']

Sum of per-point Gaussian log-densities about \(G(f)\).

X is required when the warp is conditional (cond_shape set) and ignored otherwise. Note the cost: the expected log-likelihood evaluates the warp at every quadrature node, so a conditional warp pays order conditioner forward passes per data point.

Source code in packages/pyrox-gp/src/pyrox_gp/_warped.py
def log_prob(
    self,
    f: Float[Array, " ..."],
    y: Float[Array, " ..."],
    X: Float[Array, " ..."] | None = None,
) -> Float[Array, ""]:
    r"""Sum of per-point Gaussian log-densities about $G(f)$.

    ``X`` is required when the warp is conditional (``cond_shape``
    set) and ignored otherwise. Note the cost: the expected
    log-likelihood evaluates the warp at every quadrature node, so a
    conditional warp pays ``order`` conditioner forward passes per
    data point.
    """
    g = _apply_warp(self.warp, f, X)
    return nd.Normal(g, jnp.sqrt(self.noise_var)).log_prob(y).sum()

warped_predictive_moments(lik: WarpedGaussianLikelihood, f_loc: Float[Array, ' N'], f_var: Float[Array, ' N'], X: Float[Array, 'N D'] | None = None, *, order: int = 32) -> tuple[Float[Array, ' N'], Float[Array, ' N']]

Moments of \(y = G(f) + \epsilon\) under \(q(f)\), by quadrature.

\[ m_1 = \mathbb{E}_{q(f)}[G(f)], \qquad m_2 = \sigma^2 + \mathbb{E}_{q(f)}[G(f)^2] - m_1^2 \]

Note \(m_1 \neq G(\mathbb{E}[f])\) whenever \(G\) is nonlinear -- warping the posterior mean is the natural-looking mistake and is badly biased for a skewed warp.

Parameters:

Name Type Description Default
lik WarpedGaussianLikelihood

The warped likelihood.

required
f_loc Float[Array, ' N']

Posterior means of the base GP, shape (N,).

required
f_var Float[Array, ' N']

Posterior marginal variances of the base GP, shape (N,).

required
X Float[Array, 'N D'] | None

Inputs of shape (N, D); required when the warp is conditional, ignored otherwise. Broadcast over the quadrature nodes.

None
order int

Gauss-Hermite nodes; at least 2, since one node carries no spread. Values above ~256 are numerically unreliable (hermegauss overflows to NaN by order 512).

32

Returns:

Type Description
tuple[Float[Array, ' N'], Float[Array, ' N']]

Tuple of (mean, variance) in observation space, both (N,).

Source code in packages/pyrox-gp/src/pyrox_gp/_warped.py
def warped_predictive_moments(
    lik: WarpedGaussianLikelihood,
    f_loc: Float[Array, " N"],
    f_var: Float[Array, " N"],
    X: Float[Array, "N D"] | None = None,
    *,
    order: int = 32,
) -> tuple[Float[Array, " N"], Float[Array, " N"]]:
    r"""Moments of $y = G(f) + \epsilon$ under $q(f)$, by quadrature.

    $$
    m_1 = \mathbb{E}_{q(f)}[G(f)], \qquad
    m_2 = \sigma^2 + \mathbb{E}_{q(f)}[G(f)^2] - m_1^2
    $$

    Note $m_1 \neq G(\mathbb{E}[f])$ whenever $G$ is nonlinear -- warping the
    posterior mean is the natural-looking mistake and is badly biased for a
    skewed warp.

    Args:
        lik: The warped likelihood.
        f_loc: Posterior means of the base GP, shape ``(N,)``.
        f_var: Posterior marginal variances of the base GP, shape ``(N,)``.
        X: Inputs of shape ``(N, D)``; required when the warp is
            conditional, ignored otherwise. Broadcast over the quadrature
            nodes.
        order: Gauss-Hermite nodes; at least 2, since one node carries no
            spread. Values above ~256 are numerically unreliable
            (``hermegauss`` overflows to ``NaN`` by order 512).

    Returns:
        Tuple of ``(mean, variance)`` in observation space, both ``(N,)``.
    """
    if not 2 <= order <= 256:
        # order == 1 places a single node at the mean, so the centered
        # second moment is identically zero and every bit of latent
        # uncertainty would be silently discarded.
        raise ValueError(
            f"order must be in [2, 256]; got {order}. A single quadrature "
            "node cannot represent any spread."
        )
    x, w = np.polynomial.hermite_e.hermegauss(order)
    x = jnp.asarray(x)
    w = jnp.asarray(w) / np.sqrt(2.0 * np.pi)
    fs = f_loc[None, :] + jnp.sqrt(f_var)[None, :] * x[:, None]
    # _apply_warp broadcasts the (N, D) inputs across the node axis.
    g = _apply_warp(lik.warp, fs, X)
    m1 = jnp.sum(w[:, None] * g, axis=0)
    # Centered second moment. Subtracting m1**2 from the raw second moment
    # cancels catastrophically once G's output carries a large offset
    # relative to its spread (an Affine warp centred near 1e4 loses the
    # whole variance in float32, and can go negative).
    m2 = lik.noise_var + jnp.sum(w[:, None] * (g - m1[None, :]) ** 2, axis=0)
    return m1, m2

SVGP inference

The structured SVGP ELBO as a differentiable scalar (svgp_elbo), its NumPyro registration (svgp_factor), and the natural-gradient / CVI update loop (ConjugateVI) that exploits the NaturalGuide parameterization for conjugate-style coordinate ascent.

svgp_elbo(prior: SparseGPPrior, guide: Guide, likelihood: Likelihood, X: Float[Array, 'N D'], y: Float[Array, ' N'], *, integrator: GaussxIntegrator | None = None) -> Float[Array, '']

Structured SVGP ELBO as a differentiable scalar.

\[ \mathcal{L} = \sum_n \mathbb{E}_{q(f_n)} [\log p(y_n \mid f_n)] - \mathrm{KL}[q(u) \| p(u)] \]

For GaussianLikelihood the expected log-likelihood has a closed form and no integrator is needed:

loss = svgp_elbo(prior, guide, GaussianLikelihood(0.1), X, y)
grad = eqx.filter_grad(lambda g: -svgp_elbo(prior, g, ...))(guide)

For non-conjugate likelihoods supply a gaussx integrator:

from gaussx import GaussHermiteIntegrator
loss = svgp_elbo(
    prior, guide,
    DistLikelihood(lambda f: dist.Bernoulli(logits=f)),
    X, y, integrator=GaussHermiteIntegrator(order=20),
)

Parameters:

Name Type Description Default
prior SparseGPPrior

Sparse GP prior (kernel + inducing inputs).

required
guide Guide

Variational guide over inducing values.

required
likelihood Likelihood

Observation model.

required
X Float[Array, 'N D']

Training inputs, shape (N, D).

required
y Float[Array, ' N']

Training targets, shape (N,).

required
integrator AbstractIntegrator | None

gaussx integrator for the per-point ELL. Required for non-conjugate likelihoods; ignored for GaussianLikelihood.

None

Returns:

Type Description
Float[Array, '']

Scalar ELBO value (higher is better).

Raises:

Type Description
ValueError

If a non-conjugate likelihood is used without an integrator.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference.py
def svgp_elbo(
    prior: SparseGPPrior,
    guide: Guide,
    likelihood: Likelihood,
    X: Float[Array, "N D"],
    y: Float[Array, " N"],
    *,
    integrator: GaussxIntegrator | None = None,
) -> Float[Array, ""]:
    r"""Structured SVGP ELBO as a differentiable scalar.

    $$
    \mathcal{L} = \sum_n \mathbb{E}_{q(f_n)}
                  [\log p(y_n \mid f_n)]
                - \mathrm{KL}[q(u) \| p(u)]
    $$

    For `GaussianLikelihood` the expected log-likelihood has a
    closed form and no integrator is needed:

    ```python
    loss = svgp_elbo(prior, guide, GaussianLikelihood(0.1), X, y)
    grad = eqx.filter_grad(lambda g: -svgp_elbo(prior, g, ...))(guide)
    ```

    For non-conjugate likelihoods supply a gaussx integrator:

    ```python
    from gaussx import GaussHermiteIntegrator
    loss = svgp_elbo(
        prior, guide,
        DistLikelihood(lambda f: dist.Bernoulli(logits=f)),
        X, y, integrator=GaussHermiteIntegrator(order=20),
    )
    ```

    Args:
        prior: Sparse GP prior (kernel + inducing inputs).
        guide: Variational guide over inducing values.
        likelihood: Observation model.
        X: Training inputs, shape ``(N, D)``.
        y: Training targets, shape ``(N,)``.
        integrator: gaussx integrator for the per-point ELL. Required
            for non-conjugate likelihoods; ignored for
            `GaussianLikelihood`.

    Returns:
        Scalar ELBO value (higher is better).

    Raises:
        ValueError: If a non-conjugate likelihood is used without an
            integrator.
    """
    with _kernel_context(prior.kernel):
        K_zz_op, K_xz, K_xx_diag = prior.predictive_blocks(X)

    f_loc, f_var = guide.predict(K_xz, K_zz_op, K_xx_diag)  # ty: ignore[unresolved-attribute]
    f_loc = f_loc + prior.mean(X)

    kl = guide.kl_divergence(K_zz_op)  # ty: ignore[unresolved-attribute]

    if isinstance(likelihood, GaussianLikelihood):
        return variational_elbo_gaussian(
            y,
            f_loc,
            f_var,
            likelihood.noise_var,  # ty: ignore[invalid-argument-type]
            kl,
        )

    if integrator is None:
        raise ValueError(
            "Non-conjugate likelihoods require an integrator "
            "(e.g. gaussx.GaussHermiteIntegrator). "
            "Pass integrator=GaussHermiteIntegrator(order=20)."
        )
    ell = _ell_numerical(likelihood, y, f_loc, f_var, integrator, X)
    return ell - kl

svgp_factor(name: str, prior: SparseGPPrior, guide: Guide, likelihood: Likelihood, X: Float[Array, 'N D'], y: Float[Array, ' N'], *, integrator: GaussxIntegrator | None = None) -> None

Register the structured SVGP ELBO as a NumPyro factor site.

Wraps svgp_elbo in numpyro.factor so it plugs into numpyro.infer.SVI + Trace_ELBO. NumPyro sees one deterministic factor site — the actual ELBO uses the efficient closed-form KL + structured ELL computation.

def model(X, y):
    svgp_factor("elbo", prior, guide, lik, X, y)

svi = SVI(model, lambda X, y: None, Adam(1e-3), Trace_ELBO())
Source code in packages/pyrox-gp/src/pyrox_gp/_inference.py
def svgp_factor(
    name: str,
    prior: SparseGPPrior,
    guide: Guide,
    likelihood: Likelihood,
    X: Float[Array, "N D"],
    y: Float[Array, " N"],
    *,
    integrator: GaussxIntegrator | None = None,
) -> None:
    """Register the structured SVGP ELBO as a NumPyro factor site.

    Wraps `svgp_elbo` in ``numpyro.factor`` so it plugs into
    ``numpyro.infer.SVI`` + ``Trace_ELBO``. NumPyro sees one
    deterministic factor site — the *actual* ELBO uses the efficient
    closed-form KL + structured ELL computation.

    ```python
    def model(X, y):
        svgp_factor("elbo", prior, guide, lik, X, y)

    svi = SVI(model, lambda X, y: None, Adam(1e-3), Trace_ELBO())
    ```
    """
    numpyro.factor(
        name,
        svgp_elbo(prior, guide, likelihood, X, y, integrator=integrator),
    )

ConjugateVI

Natural-gradient / CVI update for sparse variational GPs.

Operates in natural-parameter space: each step computes the per-point expected gradients and Hessians of the log-likelihood under \(q(f_n)\), projects them into the inducing-value natural parameters, and applies a damped update via NaturalGuide.natural_update.

For GaussianLikelihood the gradients and Hessians are analytical. For non-conjugate likelihoods an integrator is required to evaluate the expectations numerically.

guide = NaturalGuide.init(num_inducing=M)
cvi = ConjugateVI(damping=0.5)

for epoch in range(100):
    guide = cvi.step(prior, guide, GaussianLikelihood(0.1), X, y)

Attributes:

Name Type Description
damping float

Learning rate / damping factor \(\rho \in (0, 1]\). 1.0 replaces the natural parameters with the target; < 1 interpolates for stability.

integrator AbstractIntegrator | None

gaussx integrator for non-conjugate expected gradients. None is fine for GaussianLikelihood.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference.py
class ConjugateVI:
    r"""Natural-gradient / CVI update for sparse variational GPs.

    Operates in natural-parameter space: each `step` computes the
    per-point expected gradients and Hessians of the log-likelihood
    under $q(f_n)$, projects them into the inducing-value natural
    parameters, and applies a damped update via
    `NaturalGuide.natural_update`.

    For `GaussianLikelihood` the gradients and Hessians are
    analytical. For non-conjugate likelihoods an integrator is required
    to evaluate the expectations numerically.

    ```python
    guide = NaturalGuide.init(num_inducing=M)
    cvi = ConjugateVI(damping=0.5)

    for epoch in range(100):
        guide = cvi.step(prior, guide, GaussianLikelihood(0.1), X, y)
    ```

    Attributes:
        damping: Learning rate / damping factor $\rho \in (0, 1]$.
            ``1.0`` replaces the natural parameters with the target;
            ``< 1`` interpolates for stability.
        integrator: gaussx integrator for non-conjugate expected
            gradients. ``None`` is fine for `GaussianLikelihood`.
    """

    damping: float
    integrator: GaussxIntegrator | None

    def __init__(
        self,
        damping: float = 1.0,
        integrator: GaussxIntegrator | None = None,
    ) -> None:
        if not 0.0 <= damping <= 1.0:
            raise ValueError(f"damping must be in [0, 1], got {damping}.")
        self.damping = damping
        self.integrator = integrator

    def step(
        self,
        prior: SparseGPPrior,
        guide: NaturalGuide,
        likelihood: Likelihood,
        X: Float[Array, "N D"],
        y: Float[Array, " N"],
    ) -> NaturalGuide:
        r"""One CVI step: compute sites, project, update.

        1. Predict ``(f_loc, f_var) = guide.predict(...)``
        2. Compute per-point natural-gradient targets
           $(\lambda_n^{(1)}, \Lambda_n^{(2)})$.
        3. Project into inducing space:

        $$
        \begin{aligned}
        \hat{\eta}_1 &= \eta_1^{\text{prior}}
            + K_{ZZ}^{-1} K_{ZX}\, \lambda^{(1)}, \\
        \hat{\eta}_2 &= \eta_2^{\text{prior}}
            + K_{ZZ}^{-1} K_{ZX}\,
              \mathrm{diag}(\Lambda^{(2)})\, K_{XZ} K_{ZZ}^{-1}.
        \end{aligned}
        $$

        4. Damped update via `NaturalGuide.natural_update`.

        Args:
            prior: Sparse GP prior.
            guide: Current natural-parameter guide.
            likelihood: Observation model.
            X: Training inputs, shape ``(N, D)``.
            y: Training targets, shape ``(N,)``.

        Returns:
            Updated `NaturalGuide`.
        """
        with _kernel_context(prior.kernel):
            K_zz_op, K_xz, K_xx_diag = prior.predictive_blocks(X)

        f_loc, f_var = guide.predict(K_xz, K_zz_op, K_xx_diag)
        f_loc = f_loc + prior.mean(X)

        grad1, grad2 = self._site_gradients(likelihood, y, f_loc, f_var)

        # B = K_zz^{-1} K_zx via gaussx.solve_matrix: one Cholesky
        # factorization shared across the matrix RHS for dense PSD
        # K_zz, structured dispatch (O(M) diagonal) otherwise.
        K_zx = einx.id("n m -> m n", K_xz)  # (N, M) → (M, N)
        B = solve_matrix(K_zz_op, K_zx)

        # Prior natural parameters: p(u) = N(0, K_zz)
        # => eta1_prior = 0, eta2_prior = -0.5 K_zz^{-1}
        # K_zz^{-1} must be dense here (it seeds the dense nat2), so
        # solve against I and re-symmetrize.
        M = K_zz_op.in_size()
        nat1_prior = jnp.zeros_like(guide.nat1)
        eye = jnp.eye(M, dtype=guide.nat2.dtype)
        nat2_prior = -0.5 * symmetrize(solve_matrix(K_zz_op, eye))

        # Project per-point site natural parameters into inducing
        # space. The site λ₁ = grad1 - f_loc * grad2 (not just grad1);
        # omitting the f_loc correction makes the update depend on the
        # current guide mean instead of being a fixed-point update.
        site_lambda1 = grad1 - f_loc * grad2
        # η̂₁ = η₁ᵖʳⁱᵒʳ + B λ⁽¹⁾,  with B: (M, N), λ⁽¹⁾: (N,) → (M,).
        nat1_hat = nat1_prior + einx.dot("m n, n -> m", B, site_lambda1)
        # η̂₂ = η₂ᵖʳⁱᵒʳ + B diag(½ grad2) Bᵀ. Fold the per-point scale into
        # B then contract the shared n axis: (M, N)·(K, N) → (M, K).
        scaled_B = einx.multiply("n, m n -> m n", 0.5 * grad2, B)
        nat2_hat = nat2_prior + einx.dot("m n, k n -> m k", scaled_B, B)

        return guide.natural_update(nat1_hat, nat2_hat, rho=self.damping)

    def _site_gradients(
        self,
        likelihood: Likelihood,
        y: Float[Array, " N"],
        f_loc: Float[Array, " N"],
        f_var: Float[Array, " N"],
    ) -> tuple[Float[Array, " N"], Float[Array, " N"]]:
        r"""Per-point ELL gradients w.r.t. the marginal mean.

        Returns ``(grad1, grad2)`` where ``grad1[n]`` is
        $\partial / \partial \mu_n \,
        \mathbb{E}_{q(f_n)}[\log p(y_n \mid f_n)]$ and ``grad2[n]``
        is the second derivative.

        For `GaussianLikelihood` these are analytical. For
        non-conjugate likelihoods they are computed via JAX autodiff
        through the per-point ELL.
        """
        if isinstance(likelihood, GaussianLikelihood):
            noise_var = likelihood.noise_var
            grad1 = (y - f_loc) / noise_var
            grad2 = jnp.full_like(f_loc, -1.0 / noise_var)
            return grad1, grad2

        if self.integrator is None:
            raise ValueError(
                "Non-conjugate likelihoods require an integrator for "
                "CVI. Pass integrator=GaussHermiteIntegrator(order=20)."
            )

        def _per_point_ell(
            mu_n: Float[Array, ""],
            var_n: Float[Array, ""],
            y_n: Float[Array, ""],
        ) -> Float[Array, ""]:
            state = GaussianState(
                mean=mu_n[None],
                cov=lx.MatrixLinearOperator(
                    var_n[None, None], lx.positive_semidefinite_tag
                ),
            )
            return log_likelihood_expectation(
                lambda f: lik.log_prob(f, y_n[None]),
                state,
                self.integrator,  # type: ignore[arg-type]  # ty: ignore[invalid-argument-type]
            )

        lik = likelihood
        grad1 = jax.vmap(jax.grad(lambda mu, v, yn: _per_point_ell(mu, v, yn)))(
            f_loc, f_var, y
        )
        grad2 = jax.vmap(
            jax.grad(jax.grad(lambda mu, v, yn: _per_point_ell(mu, v, yn)))
        )(f_loc, f_var, y)
        return grad1, grad2

step(prior: SparseGPPrior, guide: NaturalGuide, likelihood: Likelihood, X: Float[Array, 'N D'], y: Float[Array, ' N']) -> NaturalGuide

One CVI step: compute sites, project, update.

  1. Predict (f_loc, f_var) = guide.predict(...)
  2. Compute per-point natural-gradient targets \((\lambda_n^{(1)}, \Lambda_n^{(2)})\).
  3. Project into inducing space:
\[ \begin{aligned} \hat{\eta}_1 &= \eta_1^{\text{prior}} + K_{ZZ}^{-1} K_{ZX}\, \lambda^{(1)}, \\ \hat{\eta}_2 &= \eta_2^{\text{prior}} + K_{ZZ}^{-1} K_{ZX}\, \mathrm{diag}(\Lambda^{(2)})\, K_{XZ} K_{ZZ}^{-1}. \end{aligned} \]
  1. Damped update via NaturalGuide.natural_update.

Parameters:

Name Type Description Default
prior SparseGPPrior

Sparse GP prior.

required
guide NaturalGuide

Current natural-parameter guide.

required
likelihood Likelihood

Observation model.

required
X Float[Array, 'N D']

Training inputs, shape (N, D).

required
y Float[Array, ' N']

Training targets, shape (N,).

required

Returns:

Type Description
NaturalGuide

Updated NaturalGuide.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference.py
def step(
    self,
    prior: SparseGPPrior,
    guide: NaturalGuide,
    likelihood: Likelihood,
    X: Float[Array, "N D"],
    y: Float[Array, " N"],
) -> NaturalGuide:
    r"""One CVI step: compute sites, project, update.

    1. Predict ``(f_loc, f_var) = guide.predict(...)``
    2. Compute per-point natural-gradient targets
       $(\lambda_n^{(1)}, \Lambda_n^{(2)})$.
    3. Project into inducing space:

    $$
    \begin{aligned}
    \hat{\eta}_1 &= \eta_1^{\text{prior}}
        + K_{ZZ}^{-1} K_{ZX}\, \lambda^{(1)}, \\
    \hat{\eta}_2 &= \eta_2^{\text{prior}}
        + K_{ZZ}^{-1} K_{ZX}\,
          \mathrm{diag}(\Lambda^{(2)})\, K_{XZ} K_{ZZ}^{-1}.
    \end{aligned}
    $$

    4. Damped update via `NaturalGuide.natural_update`.

    Args:
        prior: Sparse GP prior.
        guide: Current natural-parameter guide.
        likelihood: Observation model.
        X: Training inputs, shape ``(N, D)``.
        y: Training targets, shape ``(N,)``.

    Returns:
        Updated `NaturalGuide`.
    """
    with _kernel_context(prior.kernel):
        K_zz_op, K_xz, K_xx_diag = prior.predictive_blocks(X)

    f_loc, f_var = guide.predict(K_xz, K_zz_op, K_xx_diag)
    f_loc = f_loc + prior.mean(X)

    grad1, grad2 = self._site_gradients(likelihood, y, f_loc, f_var)

    # B = K_zz^{-1} K_zx via gaussx.solve_matrix: one Cholesky
    # factorization shared across the matrix RHS for dense PSD
    # K_zz, structured dispatch (O(M) diagonal) otherwise.
    K_zx = einx.id("n m -> m n", K_xz)  # (N, M) → (M, N)
    B = solve_matrix(K_zz_op, K_zx)

    # Prior natural parameters: p(u) = N(0, K_zz)
    # => eta1_prior = 0, eta2_prior = -0.5 K_zz^{-1}
    # K_zz^{-1} must be dense here (it seeds the dense nat2), so
    # solve against I and re-symmetrize.
    M = K_zz_op.in_size()
    nat1_prior = jnp.zeros_like(guide.nat1)
    eye = jnp.eye(M, dtype=guide.nat2.dtype)
    nat2_prior = -0.5 * symmetrize(solve_matrix(K_zz_op, eye))

    # Project per-point site natural parameters into inducing
    # space. The site λ₁ = grad1 - f_loc * grad2 (not just grad1);
    # omitting the f_loc correction makes the update depend on the
    # current guide mean instead of being a fixed-point update.
    site_lambda1 = grad1 - f_loc * grad2
    # η̂₁ = η₁ᵖʳⁱᵒʳ + B λ⁽¹⁾,  with B: (M, N), λ⁽¹⁾: (N,) → (M,).
    nat1_hat = nat1_prior + einx.dot("m n, n -> m", B, site_lambda1)
    # η̂₂ = η₂ᵖʳⁱᵒʳ + B diag(½ grad2) Bᵀ. Fold the per-point scale into
    # B then contract the shared n axis: (M, N)·(K, N) → (M, K).
    scaled_B = einx.multiply("n, m n -> m n", 0.5 * grad2, B)
    nat2_hat = nat2_prior + einx.dot("m n, k n -> m k", scaled_B, B)

    return guide.natural_update(nat1_hat, nat2_hat, rho=self.damping)

Non-Gaussian inference strategies

Site-based Gaussian approximations q(f) = N(m, V) of the posterior under a non-conjugate likelihood. All five strategies share the same diagonal-site view and differ only in where the per-site curvature comes from: exact Hessian at the mode (Laplace), PSD-projected Gauss-Newton curvature, statistical linearization under the cavity (posterior linearization / CVI), moment matching against the tilted distribution (EP), or L-BFGS to the MAP with a Laplace covariance at convergence (quasi-Newton). Each fit(prior, likelihood, y) returns a NonGaussConditionedGP that quacks like ConditionedGP.

LaplaceInference

Bases: Module

Laplace approximation via Newton iteration on the log posterior.

Iterates the standard GP-Laplace fixed-point loop (Rasmussen & Williams Algorithm 3.1): at each iteration evaluate per-point gradient g and Hessian-diag h of log p(y | f), form the Newton update with site precision \(\Lambda = -h\) (clipped to a small positive floor for numerical safety), and recompute f as the mean of the implied Gaussian posterior.

The reported log_marginal_approx is the standard Laplace log-marginal-likelihood approximation \(\log p(y) \approx \log p(y | \hat f) - \tfrac12 \hat f^\top K^{-1} \hat f - \tfrac12 \log |I + K \Lambda|\).

Attributes:

Name Type Description
max_iter int

Newton iterations. Default 20.

tol float

inf-norm convergence tolerance on f. Default 1e-6.

damping float

Step-size in (0, 1] applied to each Newton update: f_{k+1} = (1 - alpha) f_k + alpha f_k^{Newton}. Default 1.0 (full Newton step). Drop below 1 for non-log-concave likelihoods where pure Newton oscillates.

precision_floor float

Lower bound on the diagonal precision to keep K + 1/Λ well-conditioned even for log-concave-but-flat likelihoods (e.g. Bernoulli at extreme logits). Default 1e-6.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss.py
class LaplaceInference(eqx.Module):
    r"""Laplace approximation via Newton iteration on the log posterior.

    Iterates the standard GP-Laplace fixed-point loop (Rasmussen &
    Williams Algorithm 3.1): at each iteration evaluate per-point
    gradient ``g`` and Hessian-diag ``h`` of ``log p(y | f)``, form the
    Newton update with site precision $\Lambda = -h$ (clipped to
    a small positive floor for numerical safety), and recompute ``f``
    as the mean of the implied Gaussian posterior.

    The reported ``log_marginal_approx`` is the standard Laplace
    log-marginal-likelihood approximation
    $\log p(y) \approx \log p(y | \hat f) - \tfrac12 \hat f^\top
    K^{-1} \hat f - \tfrac12 \log |I + K \Lambda|$.

    Attributes:
        max_iter: Newton iterations. Default ``20``.
        tol: ``inf``-norm convergence tolerance on ``f``. Default ``1e-6``.
        damping: Step-size in (0, 1] applied to each Newton update:
            ``f_{k+1} = (1 - alpha) f_k + alpha f_k^{Newton}``. Default
            ``1.0`` (full Newton step). Drop below 1 for non-log-concave
            likelihoods where pure Newton oscillates.
        precision_floor: Lower bound on the diagonal precision to keep
            ``K + 1/Λ`` well-conditioned even for log-concave-but-flat
            likelihoods (e.g. Bernoulli at extreme logits). Default
            ``1e-6``.
    """

    max_iter: int = eqx.field(static=True, default=20)
    tol: float = eqx.field(static=True, default=1e-6)
    damping: float = eqx.field(static=True, default=1.0)
    precision_floor: float = eqx.field(static=True, default=1e-6)

    def fit(
        self,
        prior: GPPrior,
        likelihood: Likelihood,
        y: Float[Array, " N"],
    ) -> NonGaussConditionedGP:
        _check_scalar_latent(likelihood)
        K = _prior_K(prior)
        prior_mean = prior.mean(prior.X)
        log_prob_per_point = _log_prob_per_point_factory(likelihood)

        f = jnp.asarray(prior_mean)
        converged = False
        n_iter = 0
        for it in range(self.max_iter):
            g, h = _per_point_grad_hess(log_prob_per_point, f, y)
            # Site precision Λ = -h (positive for log-concave likelihoods).
            nat1, Lam = newton_update(f, g, h, precision_floor=self.precision_floor)
            f_newton, _ = _posterior_from_diag_sites(K, nat1, Lam, prior_mean)
            f_new = (1.0 - self.damping) * f + self.damping * f_newton
            delta = jnp.max(jnp.abs(f_new - f))
            f = f_new
            n_iter = it + 1
            if delta < self.tol:
                converged = True
                break

        # Final site naturals at convergence.
        g, h = _per_point_grad_hess(log_prob_per_point, f, y)
        nat1, Lam = newton_update(f, g, h, precision_floor=self.precision_floor)
        q_mean, q_var = _posterior_from_diag_sites(K, nat1, Lam, prior_mean)

        log_marg = _laplace_log_marginal(log_prob_per_point, f, y, prior_mean, K, Lam)

        return NonGaussConditionedGP(
            prior=prior,
            y=y,
            site_nat1=nat1,
            site_nat2=Lam,
            q_mean=q_mean,
            q_var=q_var,
            log_marginal_approx=log_marg,
            n_iter=n_iter,
            converged=converged,
        )

GaussNewtonInference

Bases: Module

Gauss-Newton inference: Newton loop with PSD-projected curvature.

Identical to LaplaceInference for log-concave likelihoods, where -d^2 log p / df^2 is already positive. For non-log-concave likelihoods (e.g. StudentTLikelihood, where the Hessian becomes positive in the tails) GN aggressively floors the curvature to a strictly-positive value via precision_floor, guaranteeing a PSD site precision and stable Newton steps. Laplace uses the same floor but typically with a smaller default.

For Bernoulli / Poisson the Fisher information equals the negative Hessian, so GGN ≡ Laplace; for StudentT the floor matters.

Attributes:

Name Type Description
max_iter int

Iterations. Default 20.

tol float

inf-norm convergence tolerance on f. Default 1e-6.

damping float

Step-size in (0, 1] applied to each Newton update. Default 1.0 (full Newton step). Drop below 1 for non-log-concave likelihoods where pure Newton oscillates.

precision_floor float

Lower bound on the diagonal precision. Default 1e-3 — larger than LaplaceInference to ensure stable updates on non-log-concave likelihoods.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss.py
class GaussNewtonInference(eqx.Module):
    r"""Gauss-Newton inference: Newton loop with PSD-projected curvature.

    Identical to `LaplaceInference` for log-concave likelihoods,
    where ``-d^2 log p / df^2`` is already positive. For non-log-concave
    likelihoods (e.g. `StudentTLikelihood`, where the Hessian
    becomes positive in the tails) GN aggressively floors the curvature
    to a strictly-positive value via ``precision_floor``, guaranteeing a
    PSD site precision and stable Newton steps. Laplace uses the same
    floor but typically with a smaller default.

    For Bernoulli / Poisson the Fisher information equals the negative
    Hessian, so GGN ≡ Laplace; for StudentT the floor matters.

    Attributes:
        max_iter: Iterations. Default ``20``.
        tol: ``inf``-norm convergence tolerance on ``f``. Default ``1e-6``.
        damping: Step-size in (0, 1] applied to each Newton update.
            Default ``1.0`` (full Newton step). Drop below 1 for
            non-log-concave likelihoods where pure Newton oscillates.
        precision_floor: Lower bound on the diagonal precision. Default
            ``1e-3`` — larger than `LaplaceInference` to ensure
            stable updates on non-log-concave likelihoods.
    """

    max_iter: int = eqx.field(static=True, default=20)
    tol: float = eqx.field(static=True, default=1e-6)
    damping: float = eqx.field(static=True, default=1.0)
    precision_floor: float = eqx.field(static=True, default=1e-3)

    def fit(
        self,
        prior: GPPrior,
        likelihood: Likelihood,
        y: Float[Array, " N"],
    ) -> NonGaussConditionedGP:
        _check_scalar_latent(likelihood)
        K = _prior_K(prior)
        prior_mean = prior.mean(prior.X)
        log_prob_per_point = _log_prob_per_point_factory(likelihood)

        f = jnp.asarray(prior_mean)
        converged = False
        n_iter = 0
        for it in range(self.max_iter):
            g, h = _per_point_grad_hess(log_prob_per_point, f, y)
            # PSD curvature: clip the negative Hessian to a strictly
            # positive floor so the Newton step is always well-defined
            # even when ``-h`` goes negative (StudentT tails).
            nat1, Lam = newton_update(f, g, h, precision_floor=self.precision_floor)
            f_newton, _ = _posterior_from_diag_sites(K, nat1, Lam, prior_mean)
            f_new = (1.0 - self.damping) * f + self.damping * f_newton
            delta = jnp.max(jnp.abs(f_new - f))
            f = f_new
            n_iter = it + 1
            if delta < self.tol:
                converged = True
                break

        g, h = _per_point_grad_hess(log_prob_per_point, f, y)
        nat1, Lam = newton_update(f, g, h, precision_floor=self.precision_floor)
        q_mean, q_var = _posterior_from_diag_sites(K, nat1, Lam, prior_mean)

        log_marg = _laplace_log_marginal(log_prob_per_point, f, y, prior_mean, K, Lam)

        return NonGaussConditionedGP(
            prior=prior,
            y=y,
            site_nat1=nat1,
            site_nat2=Lam,
            q_mean=q_mean,
            q_var=q_var,
            log_marginal_approx=log_marg,
            n_iter=n_iter,
            converged=converged,
        )

PosteriorLinearization

Bases: Module

Iterated statistical-linearization site updates.

At each iteration: form the cavity q_n^\\(¬n) = N(m^c_n, v^c_n) by removing the current site, then under the cavity compute statistical-linearization moments \(\bar g = E_{\rm cav}[\partial_f \log p]\) and \(\bar h = E_{\rm cav}[\partial^2_f \log p]\) via the integrator, update the site naturals via gaussx.blr_diag_update with a damping factor, recompute the global posterior, repeat. This is PL/CVI in the sense of Adam/Garcia-Fernandez/Sarkka — equivalent in the Gaussian-cavity limit to taking one EP-style step but using derivative expectations instead of moment matching.

Attributes:

Name Type Description
integrator AbstractIntegrator

Cavity integrator. Default gaussx.GaussHermiteIntegrator(order=20).

max_iter int

Iterations. Default 20.

damping float

Step size in (0, 1]. Default 0.5.

tol float

inf-norm convergence on the posterior mean.

precision_floor float

Lower bound on the diagonal precision.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss.py
class PosteriorLinearization(eqx.Module):
    r"""Iterated statistical-linearization site updates.

    At each iteration: form the cavity ``q_n^\\(¬n) = N(m^c_n, v^c_n)``
    by removing the current site, then under the cavity compute
    statistical-linearization moments
    $\bar g = E_{\rm cav}[\partial_f \log p]$ and
    $\bar h = E_{\rm cav}[\partial^2_f \log p]$ via the integrator,
    update the site naturals via `gaussx.blr_diag_update` with
    a damping factor, recompute the global posterior, repeat. This is
    PL/CVI in the sense of Adam/Garcia-Fernandez/Sarkka — equivalent
    in the Gaussian-cavity limit to taking one EP-style step but using
    derivative expectations instead of moment matching.

    Attributes:
        integrator: Cavity integrator. Default
            ``gaussx.GaussHermiteIntegrator(order=20)``.
        max_iter: Iterations. Default ``20``.
        damping: Step size in (0, 1]. Default ``0.5``.
        tol: ``inf``-norm convergence on the posterior mean.
        precision_floor: Lower bound on the diagonal precision.
    """

    integrator: AbstractIntegrator = eqx.field(
        default_factory=lambda: GaussHermiteIntegrator(order=20)
    )
    max_iter: int = eqx.field(static=True, default=20)
    damping: float = eqx.field(static=True, default=0.5)
    tol: float = eqx.field(static=True, default=1e-6)
    precision_floor: float = eqx.field(static=True, default=1e-6)

    def fit(
        self,
        prior: GPPrior,
        likelihood: Likelihood,
        y: Float[Array, " N"],
    ) -> NonGaussConditionedGP:
        _check_scalar_latent(likelihood)
        K = _prior_K(prior)
        prior_mean = prior.mean(prior.X)
        log_prob_per_point = _log_prob_per_point_factory(likelihood)

        N = K.shape[0]
        # Initialize sites at zero — posterior == prior.
        nat1 = jnp.zeros(N, dtype=K.dtype)
        nat2 = jnp.full((N,), self.precision_floor, dtype=K.dtype)
        q_mean = jnp.asarray(prior_mean)
        q_var = jnp.diag(K)

        def grad_at(f_n: Float[Array, ""], y_n: Float[Array, ""]) -> Float[Array, ""]:
            return jax.grad(lambda f: log_prob_per_point(f[None], y_n[None])[0])(f_n)

        def hess_at(f_n: Float[Array, ""], y_n: Float[Array, ""]) -> Float[Array, ""]:
            return jax.grad(
                jax.grad(lambda f: log_prob_per_point(f[None], y_n[None])[0])
            )(f_n)

        converged = False
        n_iter = 0
        for it in range(self.max_iter):
            # Cavity for diagonal sites: q_n / site_n.
            cav_mean, cav_var = cavity_distribution(
                q_mean, q_var, nat1, nat2, precision_floor=self.precision_floor
            )
            assert isinstance(cav_var, jax.Array)  # diagonal path in, diagonal out

            # Statistical-linearization moments under the cavity. Each
            # site gets its own 1-D ``GaussianState`` and the integrator
            # returns the propagated mean per site.
            E_grad = _per_site_expectation(
                self.integrator, grad_at, cav_mean, cav_var, y
            )
            E_hess = _per_site_expectation(
                self.integrator, hess_at, cav_mean, cav_var, y
            )

            # Site update via BLR (diag): nat1_new = grad - H mu, nat2_new = -H,
            # damped.
            nat1_target, nat2_target = newton_update(
                cav_mean, E_grad, E_hess, precision_floor=self.precision_floor
            )
            nat1, nat2 = damped_natural_update(
                nat1, nat2, nat1_target, nat2_target, lr=self.damping
            )
            assert isinstance(nat2, jax.Array)  # pyrox sites are diagonal arrays
            nat2 = jnp.maximum(nat2, self.precision_floor)

            q_mean_new, q_var = _posterior_from_diag_sites(K, nat1, nat2, prior_mean)
            delta = jnp.max(jnp.abs(q_mean_new - q_mean))
            q_mean = q_mean_new
            n_iter = it + 1
            if delta < self.tol:
                converged = True
                break

        # Approximate log marginal via the same Laplace-style identity
        # at the converged ``q_mean`` (acceptable for small N; not
        # guaranteed PL-consistent — strategies expose their own marginal
        # when the user needs strategy-specific values).
        log_marg = _laplace_log_marginal(
            log_prob_per_point, q_mean, y, prior_mean, K, nat2
        )

        return NonGaussConditionedGP(
            prior=prior,
            y=y,
            site_nat1=nat1,
            site_nat2=nat2,
            q_mean=q_mean,
            q_var=q_var,
            log_marginal_approx=log_marg,
            n_iter=n_iter,
            converged=converged,
        )

ExpectationPropagation

Bases: Module

Parallel expectation propagation (Minka 2001) with damping.

Per outer iteration: for every site simultaneously, form the cavity, compute the tilted moments E_{q^c(f) p(y|f)}[1, f, f^2] via the integrator, match them to a Gaussian, and update site naturals (damped). Parallel-EP is embarrassingly vectorizable and converges for well-behaved log-concave likelihoods; for problematic models reduce damping.

Attributes:

Name Type Description
integrator AbstractIntegrator

Cavity integrator. Default gaussx.GaussHermiteIntegrator(order=20). EP only reads the integrator's order attribute when delegating to gaussx.ep_tilted_moments; non-Gauss-Hermite integrators fall back to order=20.

max_iter int

Iterations. Default 40.

damping float

Damping in (0, 1]. Default 0.5.

tol float

inf-norm convergence on the posterior mean.

precision_floor float

Lower bound on the diagonal precision.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss.py
class ExpectationPropagation(eqx.Module):
    r"""Parallel expectation propagation (Minka 2001) with damping.

    Per outer iteration: for every site simultaneously, form the
    cavity, compute the *tilted* moments
    ``E_{q^c(f) p(y|f)}[1, f, f^2]`` via the integrator, match them to
    a Gaussian, and update site naturals (damped). Parallel-EP is
    embarrassingly vectorizable and converges for well-behaved
    log-concave likelihoods; for problematic models reduce ``damping``.

    Attributes:
        integrator: Cavity integrator. Default
            ``gaussx.GaussHermiteIntegrator(order=20)``. EP only reads
            the integrator's ``order`` attribute when delegating to
            `gaussx.ep_tilted_moments`; non-Gauss-Hermite
            integrators fall back to ``order=20``.
        max_iter: Iterations. Default ``40``.
        damping: Damping in (0, 1]. Default ``0.5``.
        tol: ``inf``-norm convergence on the posterior mean.
        precision_floor: Lower bound on the diagonal precision.
    """

    integrator: AbstractIntegrator = eqx.field(
        default_factory=lambda: GaussHermiteIntegrator(order=20)
    )
    max_iter: int = eqx.field(static=True, default=40)
    damping: float = eqx.field(static=True, default=0.5)
    tol: float = eqx.field(static=True, default=1e-6)
    precision_floor: float = eqx.field(static=True, default=1e-6)

    def fit(
        self,
        prior: GPPrior,
        likelihood: Likelihood,
        y: Float[Array, " N"],
    ) -> NonGaussConditionedGP:
        _check_scalar_latent(likelihood)
        K = _prior_K(prior)
        prior_mean = prior.mean(prior.X)
        log_prob_per_point = _log_prob_per_point_factory(likelihood)

        N = K.shape[0]
        nat1 = jnp.zeros(N, dtype=K.dtype)
        nat2 = jnp.full((N,), self.precision_floor, dtype=K.dtype)
        q_mean = jnp.asarray(prior_mean)
        q_var = jnp.diag(K)

        def lp(f_n: Float[Array, ""], y_n: Float[Array, ""]) -> Float[Array, ""]:
            return log_prob_per_point(f_n[None], y_n[None])[0]

        order = getattr(self.integrator, "order", 20)

        def _per_site_tilted(
            m_n: Float[Array, ""],
            v_n: Float[Array, ""],
            y_n: Float[Array, ""],
        ) -> tuple[Float[Array, ""], Float[Array, ""]]:
            return ep_tilted_moments(lambda f: lp(f, y_n), m_n, v_n, order=order)

        converged = False
        n_iter = 0
        for it in range(self.max_iter):
            cav_mean, cav_var = cavity_distribution(
                q_mean, q_var, nat1, nat2, precision_floor=self.precision_floor
            )
            assert isinstance(cav_var, jax.Array)  # diagonal path in, diagonal out

            # Tilted moments via `gaussx.ep_tilted_moments`. The
            # gaussx API expects a ``log_lik_fn(f)`` with the per-site
            # target baked in, so ``_per_site_tilted`` closes over
            # ``y_n`` and ``vmap`` recovers the ``(N,)`` shape contract.
            tilted_mean, tilted_var = jax.vmap(_per_site_tilted)(cav_mean, cav_var, y)

            # New site naturals from matched moments minus the cavity.
            new_prec = jnp.reciprocal(tilted_var) - jnp.reciprocal(cav_var)
            new_prec = jnp.maximum(new_prec, self.precision_floor)
            new_nat1 = tilted_mean / tilted_var - cav_mean / cav_var
            nat1, nat2 = damped_natural_update(
                nat1, nat2, new_nat1, new_prec, lr=self.damping
            )
            assert isinstance(nat2, jax.Array)  # pyrox sites are diagonal arrays
            nat2 = jnp.maximum(nat2, self.precision_floor)

            q_mean_new, q_var = _posterior_from_diag_sites(K, nat1, nat2, prior_mean)
            delta = jnp.max(jnp.abs(q_mean_new - q_mean))
            q_mean = q_mean_new
            n_iter = it + 1
            if delta < self.tol:
                converged = True
                break

        # Approximate marginal via Laplace identity at the EP mean.
        log_marg = _laplace_log_marginal(
            log_prob_per_point, q_mean, y, prior_mean, K, nat2
        )

        return NonGaussConditionedGP(
            prior=prior,
            y=y,
            site_nat1=nat1,
            site_nat2=nat2,
            q_mean=q_mean,
            q_var=q_var,
            log_marginal_approx=log_marg,
            n_iter=n_iter,
            converged=converged,
        )

QuasiNewtonInference

Bases: Module

MAP optimization via L-BFGS, Laplace covariance at convergence.

Optimizes the unnormalized log posterior \(\log p(y | f) - \tfrac12 (f - \mu)^\top K^{-1} (f - \mu)\) with optax.lbfgs, then forms a Laplace-style Gaussian approximation centered at the optimum using the exact per-point Hessian. The optimization is the cheap path for very high N where you cannot afford dense Hessians per iteration; the final Laplace covariance is computed once.

For full low-rank posterior covariance from L-BFGS history (without the final Hessian solve) see the follow-up issue — this class delivers the simpler "QN optimization, Laplace covariance" contract.

Attributes:

Name Type Description
max_iter int

L-BFGS iterations. Default 50.

tol float

Gradient-norm tolerance. Default 1e-6.

precision_floor float

Lower bound on the diagonal precision.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss.py
class QuasiNewtonInference(eqx.Module):
    r"""MAP optimization via L-BFGS, Laplace covariance at convergence.

    Optimizes the unnormalized log posterior
    $\log p(y | f) - \tfrac12 (f - \mu)^\top K^{-1} (f - \mu)$
    with ``optax.lbfgs``, then forms a Laplace-style Gaussian
    approximation centered at the optimum using the exact per-point
    Hessian. The optimization is the cheap path for very high N where
    you cannot afford dense Hessians per iteration; the final Laplace
    covariance is computed once.

    For full low-rank posterior covariance from L-BFGS history (without
    the final Hessian solve) see the follow-up issue — this class
    delivers the simpler "QN optimization, Laplace covariance"
    contract.

    Attributes:
        max_iter: L-BFGS iterations. Default ``50``.
        tol: Gradient-norm tolerance. Default ``1e-6``.
        precision_floor: Lower bound on the diagonal precision.
    """

    max_iter: int = eqx.field(static=True, default=50)
    tol: float = eqx.field(static=True, default=1e-6)
    precision_floor: float = eqx.field(static=True, default=1e-6)

    def fit(
        self,
        prior: GPPrior,
        likelihood: Likelihood,
        y: Float[Array, " N"],
    ) -> NonGaussConditionedGP:
        import optax

        _check_scalar_latent(likelihood)
        K = _prior_K(prior)
        prior_mean = prior.mean(prior.X)
        log_prob_per_point = _log_prob_per_point_factory(likelihood)
        L_K = _psd_safe_cholesky(K)

        def neg_log_post(f: Float[Array, " N"]) -> Float[Array, ""]:
            ll = log_prob_per_point(f, y).sum()
            r = f - prior_mean
            alpha = jax.scipy.linalg.cho_solve((L_K, True), r)
            return -(ll - 0.5 * jnp.dot(r, alpha))

        opt = optax.lbfgs()
        f = jnp.asarray(prior_mean)
        opt_state = opt.init(f)
        value_and_grad = optax.value_and_grad_from_state(neg_log_post)
        n_iter = 0
        converged = False
        for it in range(self.max_iter):
            v, g = value_and_grad(f, state=opt_state)
            updates, opt_state = opt.update(
                g, opt_state, f, value=v, grad=g, value_fn=neg_log_post
            )
            f = optax.apply_updates(f, updates)
            n_iter = it + 1
            if jnp.linalg.norm(g) < self.tol:
                converged = True
                break

        # Laplace covariance at the optimum. Cast ``f`` from optax's
        # generic ``Array`` to a typed JAX array so downstream typing
        # narrows correctly.
        f_opt = jnp.asarray(f)
        g, h = _per_point_grad_hess(log_prob_per_point, f_opt, y)
        nat1, Lam = newton_update(f_opt, g, h, precision_floor=self.precision_floor)
        q_mean, q_var = _posterior_from_diag_sites(K, nat1, Lam, prior_mean)

        log_marg = _laplace_log_marginal(
            log_prob_per_point, f_opt, y, prior_mean, K, Lam
        )

        return NonGaussConditionedGP(
            prior=prior,
            y=y,
            site_nat1=nat1,
            site_nat2=Lam,
            q_mean=q_mean,
            q_var=q_var,
            log_marginal_approx=log_marg,
            n_iter=n_iter,
            converged=converged,
        )

NonGaussConditionedGP

Bases: Module

GP conditioned on a non-Gaussian likelihood via an advanced strategy.

Equivalent role to pyrox_gp.ConditionedGP but the posterior over training latents is a generic q(f) = N(q_mean, q_cov) rather than a Gaussian-likelihood closed form. Predictions at test inputs use the standard site-as-pseudo-observation trick: any site-based Gaussian approximation of p(f | y) looks identical to a Gaussian- likelihood regression with synthetic per-point noise variance \(\sigma_n^2 = 1/\Lambda_n^{(2)}\) and synthetic targets \(\tilde y_n = \lambda_n^{(1)} / \Lambda_n^{(2)}\) (in the zero-mean prior frame). Predictions reconstruct K_reg = K + diag(1 / Lambda) and Cholesky-factorize it on each call (the full-Cholesky cost is O(N^3) per predict invocation; the cross-covariance contributions are O(M N)). Caching the solve on the module would require freezing the kernel hyperparameters at fit time and is intentionally not done here so prior'd-kernel workflows keep resampling correctly.

Attributes:

Name Type Description
prior GPPrior

The GPPrior.

y Float[Array, ' N']

Training targets (kept for round-trip / diagnostics).

site_nat1 Float[Array, ' N']

Diagonal site naturals \(\lambda^{(1)} \in \mathbb{R}^N\).

site_nat2 Float[Array, ' N']

Diagonal site precisions \(\Lambda^{(2)} \in \mathbb{R}^N\) (positive).

q_mean Float[Array, ' N']

Posterior mean over training latents.

q_var Float[Array, ' N']

Marginal posterior variance per training point.

log_marginal_approx Float[Array, '']

Approximate log marginal likelihood (the scalar each strategy reports — interpretation is strategy-specific; see the per-class docstring).

n_iter int

Iterations used by the strategy.

converged bool

Whether convergence tolerance was met.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss.py
class NonGaussConditionedGP(eqx.Module):
    """GP conditioned on a non-Gaussian likelihood via an advanced strategy.

    Equivalent role to `pyrox_gp.ConditionedGP` but the
    posterior over training latents is a generic
    ``q(f) = N(q_mean, q_cov)`` rather than a Gaussian-likelihood
    closed form. Predictions at test inputs use the standard
    *site-as-pseudo-observation* trick: any site-based Gaussian
    approximation of ``p(f | y)`` looks identical to a Gaussian-
    likelihood regression with synthetic per-point noise variance
    $\\sigma_n^2 = 1/\\Lambda_n^{(2)}$ and synthetic targets
    $\\tilde y_n = \\lambda_n^{(1)} / \\Lambda_n^{(2)}$ (in the
    zero-mean prior frame). Predictions reconstruct ``K_reg = K +
    diag(1 / Lambda)`` and Cholesky-factorize it on each call (the
    full-Cholesky cost is ``O(N^3)`` per ``predict`` invocation; the
    cross-covariance contributions are ``O(M N)``). Caching the solve
    on the module would require freezing the kernel hyperparameters at
    fit time and is intentionally not done here so prior'd-kernel
    workflows keep resampling correctly.

    Attributes:
        prior: The `GPPrior`.
        y: Training targets (kept for round-trip / diagnostics).
        site_nat1: Diagonal site naturals
            $\\lambda^{(1)} \\in \\mathbb{R}^N$.
        site_nat2: Diagonal site precisions
            $\\Lambda^{(2)} \\in \\mathbb{R}^N$ (positive).
        q_mean: Posterior mean over training latents.
        q_var: Marginal posterior variance per training point.
        log_marginal_approx: Approximate log marginal likelihood
            (the scalar each strategy reports — interpretation is
            strategy-specific; see the per-class docstring).
        n_iter: Iterations used by the strategy.
        converged: Whether convergence tolerance was met.
    """

    prior: GPPrior
    y: Float[Array, " N"]
    site_nat1: Float[Array, " N"]
    site_nat2: Float[Array, " N"]
    q_mean: Float[Array, " N"]
    q_var: Float[Array, " N"]
    log_marginal_approx: Float[Array, ""]
    n_iter: int = eqx.field(static=True)
    converged: bool = eqx.field(static=True)

    def _pseudo_factor(self) -> Float[Array, "N N"]:
        """Cholesky of ``K + diag(1/Λ + jitter)`` via `_psd_safe_cholesky`."""
        K = self.prior.kernel(self.prior.X, self.prior.X)
        diag_reg = jnp.reciprocal(self.site_nat2) + self.prior.jitter
        K_reg = K.at[jnp.diag_indices_from(K)].add(diag_reg)
        return _psd_safe_cholesky(K_reg)

    def predict_mean(self, X_star: Float[Array, "M D"]) -> Float[Array, " M"]:
        r"""$\mu_* = \mu(X_*) + K_{*f}\,\alpha$ with
        ``alpha`` derived from the site naturals."""
        with _kernel_context(self.prior.kernel):
            L = self._pseudo_factor()
            K_cross = self.prior.kernel(X_star, self.prior.X)
            prior_mean_train = self.prior.mean(self.prior.X)
        # Effective synthetic targets in the centered (zero-mean) frame.
        y_tilde = self.site_nat1 / self.site_nat2
        residual = y_tilde - prior_mean_train
        alpha = jax.scipy.linalg.cho_solve((L, True), residual)
        # μ_* = μ(X_*) + K_{*f} α,  K_{*f}: (M, N), α: (N,) → (M,).
        return self.prior.mean(X_star) + einx.dot("m n, n -> m", K_cross, alpha)

    def predict_var(self, X_star: Float[Array, "M D"]) -> Float[Array, " M"]:
        r"""$\Sigma_{**} - K_{*f} (K + \mathrm{diag}(1/\Lambda))^{-1} K_{f*}$."""
        with _kernel_context(self.prior.kernel):
            L = self._pseudo_factor()
            K_cross = self.prior.kernel(X_star, self.prior.X)
            K_diag = self.prior.kernel.diag(X_star)
        # v = L⁻¹ K_{f*}, shape (N, M); the predictive variance reduction is
        # the column-wise squared norm Σ_n v_{n,*}².
        v = jax.scipy.linalg.solve_triangular(
            L, einx.id("m n -> n m", K_cross), lower=True
        )
        return jnp.maximum(K_diag - einx.sum("[n] m", v * v), 0.0)

    def predict(
        self, X_star: Float[Array, "M D"]
    ) -> tuple[Float[Array, " M"], Float[Array, " M"]]:
        """Joint mean / marginal-variance prediction at ``X_star``.

        Both kernel evaluations share a single kernel context so
        Pattern B / C kernels with prior'd hyperparameters resample
        once and produce a self-consistent ``(mean, var)`` pair.
        """
        with _kernel_context(self.prior.kernel):
            L = self._pseudo_factor()
            K_cross = self.prior.kernel(X_star, self.prior.X)
            K_diag = self.prior.kernel.diag(X_star)
            prior_mean_train = self.prior.mean(self.prior.X)
            prior_mean_test = self.prior.mean(X_star)
        y_tilde = self.site_nat1 / self.site_nat2
        residual = y_tilde - prior_mean_train
        alpha = jax.scipy.linalg.cho_solve((L, True), residual)
        mean = prior_mean_test + einx.dot("m n, n -> m", K_cross, alpha)
        v = jax.scipy.linalg.solve_triangular(
            L, einx.id("m n -> n m", K_cross), lower=True
        )
        var = jnp.maximum(K_diag - einx.sum("[n] m", v * v), 0.0)
        return mean, var

predict_mean(X_star: Float[Array, 'M D']) -> Float[Array, ' M']

\(\mu_* = \mu(X_*) + K_{*f}\,\alpha\) with alpha derived from the site naturals.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss.py
def predict_mean(self, X_star: Float[Array, "M D"]) -> Float[Array, " M"]:
    r"""$\mu_* = \mu(X_*) + K_{*f}\,\alpha$ with
    ``alpha`` derived from the site naturals."""
    with _kernel_context(self.prior.kernel):
        L = self._pseudo_factor()
        K_cross = self.prior.kernel(X_star, self.prior.X)
        prior_mean_train = self.prior.mean(self.prior.X)
    # Effective synthetic targets in the centered (zero-mean) frame.
    y_tilde = self.site_nat1 / self.site_nat2
    residual = y_tilde - prior_mean_train
    alpha = jax.scipy.linalg.cho_solve((L, True), residual)
    # μ_* = μ(X_*) + K_{*f} α,  K_{*f}: (M, N), α: (N,) → (M,).
    return self.prior.mean(X_star) + einx.dot("m n, n -> m", K_cross, alpha)

predict_var(X_star: Float[Array, 'M D']) -> Float[Array, ' M']

\(\Sigma_{**} - K_{*f} (K + \mathrm{diag}(1/\Lambda))^{-1} K_{f*}\).

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss.py
def predict_var(self, X_star: Float[Array, "M D"]) -> Float[Array, " M"]:
    r"""$\Sigma_{**} - K_{*f} (K + \mathrm{diag}(1/\Lambda))^{-1} K_{f*}$."""
    with _kernel_context(self.prior.kernel):
        L = self._pseudo_factor()
        K_cross = self.prior.kernel(X_star, self.prior.X)
        K_diag = self.prior.kernel.diag(X_star)
    # v = L⁻¹ K_{f*}, shape (N, M); the predictive variance reduction is
    # the column-wise squared norm Σ_n v_{n,*}².
    v = jax.scipy.linalg.solve_triangular(
        L, einx.id("m n -> n m", K_cross), lower=True
    )
    return jnp.maximum(K_diag - einx.sum("[n] m", v * v), 0.0)

predict(X_star: Float[Array, 'M D']) -> tuple[Float[Array, ' M'], Float[Array, ' M']]

Joint mean / marginal-variance prediction at X_star.

Both kernel evaluations share a single kernel context so Pattern B / C kernels with prior'd hyperparameters resample once and produce a self-consistent (mean, var) pair.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss.py
def predict(
    self, X_star: Float[Array, "M D"]
) -> tuple[Float[Array, " M"], Float[Array, " M"]]:
    """Joint mean / marginal-variance prediction at ``X_star``.

    Both kernel evaluations share a single kernel context so
    Pattern B / C kernels with prior'd hyperparameters resample
    once and produce a self-consistent ``(mean, var)`` pair.
    """
    with _kernel_context(self.prior.kernel):
        L = self._pseudo_factor()
        K_cross = self.prior.kernel(X_star, self.prior.X)
        K_diag = self.prior.kernel.diag(X_star)
        prior_mean_train = self.prior.mean(self.prior.X)
        prior_mean_test = self.prior.mean(X_star)
    y_tilde = self.site_nat1 / self.site_nat2
    residual = y_tilde - prior_mean_train
    alpha = jax.scipy.linalg.cho_solve((L, True), residual)
    mean = prior_mean_test + einx.dot("m n, n -> m", K_cross, alpha)
    v = jax.scipy.linalg.solve_triangular(
        L, einx.id("m n -> n m", K_cross), lower=True
    )
    var = jnp.maximum(K_diag - einx.sum("[n] m", v * v), 0.0)
    return mean, var

Multi-output GPs

Vector-valued GPs via coregionalization: the linear model of coregionalization (LMC) mixes Q independent latent processes through a learned matrix, the intrinsic coregionalization model (ICM) shares one latent kernel, and the orthogonal instantaneous linear mixing model (OILMM) constrains the mixing to be orthogonal so inference decouples per latent process (the projections delegate to gaussx.oilmm_project / gaussx.oilmm_back_project).

LMCKernel

Bases: Module

Linear model of coregionalization for vector-valued GPs.

Each output is a linear combination of latent scalar GPs: f_p(x) = sum_q W[p, q] g_q(x). The cross-output covariance is Cov[f_p(x), f_{p'}(x')] = sum_q (w_q w_q^T)[p, p'] k_q(x, x').

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
class LMCKernel(eqx.Module):
    """Linear model of coregionalization for vector-valued GPs.

    Each output is a linear combination of latent scalar GPs:
    ``f_p(x) = sum_q W[p, q] g_q(x)``. The cross-output covariance is
    ``Cov[f_p(x), f_{p'}(x')] = sum_q (w_q w_q^T)[p, p'] k_q(x, x')``.
    """

    kernels: tuple[Kernel, ...]
    mixing: Float[Array, "P Q"]

    def __check_init__(self) -> None:
        _validate_mixing(self.mixing)
        _validate_kernel_count(self.kernels, self.mixing.shape[1])
        _validate_kernel_scopes_unique(self.kernels)

    @property
    def num_outputs(self) -> int:
        """Number of observed output channels ``P``."""
        return self.mixing.shape[0]

    @property
    def num_latents(self) -> int:
        """Number of latent scalar GPs ``Q``."""
        return self.mixing.shape[1]

    def coregionalization_matrix(self, q: int) -> Float[Array, "P P"]:
        """Return the rank-1 coregionalization matrix ``w_q w_q^T``."""
        if not 0 <= q < self.num_latents:
            raise IndexError(f"latent index out of range: {q}")
        column = self.mixing[:, q]
        return einx.dot("i, j -> i j", column, column)

    def kronecker_factors(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> tuple[tuple[Float[Array, "P P"], Float[Array, "N1 N2"]], ...]:
        """Return ``(B_q, K_q(X1, X2))`` factors for each latent process.

        All latent kernel evaluations share one per-call context per
        unique kernel instance, so reusing the same kernel across latents
        (for hyperparameter tying) registers each sample site exactly once.
        """
        with _kernel_contexts(self.kernels):
            return tuple(
                (self.coregionalization_matrix(q), kernel(X1, X2))
                for q, kernel in enumerate(self.kernels)
            )

    def output_covariance(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "P P N1 N2"]:
        """Return output-pair covariance blocks with shape ``(P, P, N1, N2)``."""
        # Per-latent outer product B_q ⊗ K_q → (P, P, N1, N2).
        terms = [
            einx.multiply("p1 p2, n1 n2 -> p1 p2 n1 n2", B_q, K_q)
            for B_q, K_q in self.kronecker_factors(X1, X2)
        ]
        return functools.reduce(jnp.add, terms)

    def cross_covariance_operator(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> lx.AbstractLinearOperator:
        """Return ``Cov[vec(F(X1)), vec(F(X2))]`` as a sum-of-Kroneckers operator.

        The returned operator preserves the per-latent Kronecker structure
        so structure-aware solvers can avoid materializing the full
        ``(P*N1, P*N2)`` matrix.
        """
        psd_K = X1 is X2
        terms = [
            _kron_block_op(B_q, K_q, psd_K=psd_K)
            for B_q, K_q in self.kronecker_factors(X1, X2)
        ]
        if len(terms) == 1:
            return terms[0]
        return SumOperator(*terms)

    def cross_covariance(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "PN1 PN2"]:
        """Return the dense covariance of ``vec(F(X1))`` and ``vec(F(X2))``."""
        return self.cross_covariance_operator(X1, X2).as_matrix()

    def full_covariance_operator(
        self, X: Float[Array, "N D"]
    ) -> lx.AbstractLinearOperator:
        """Return the Gram operator for isotopic multi-output observations."""
        return self.cross_covariance_operator(X, X)

    def full_covariance(self, X: Float[Array, "N D"]) -> Float[Array, "PN PN"]:
        """Return the dense Gram matrix for isotopic multi-output observations."""
        return self.full_covariance_operator(X).as_matrix()

    def diag(self, X: Float[Array, "N D"]) -> Float[Array, "N P"]:
        """Return per-input, per-output marginal variances with shape ``(N, P)``."""
        with _kernel_contexts(self.kernels):
            # Per-latent variance kₙ ⊗ wₚ² → (N, P), summed over latents.
            terms = [
                einx.multiply(
                    "n, p -> n p", kernel.diag(X), jnp.square(self.mixing[:, q])
                )
                for q, kernel in enumerate(self.kernels)
            ]
        return functools.reduce(jnp.add, terms)

num_outputs: int property

Number of observed output channels P.

num_latents: int property

Number of latent scalar GPs Q.

coregionalization_matrix(q: int) -> Float[Array, 'P P']

Return the rank-1 coregionalization matrix w_q w_q^T.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def coregionalization_matrix(self, q: int) -> Float[Array, "P P"]:
    """Return the rank-1 coregionalization matrix ``w_q w_q^T``."""
    if not 0 <= q < self.num_latents:
        raise IndexError(f"latent index out of range: {q}")
    column = self.mixing[:, q]
    return einx.dot("i, j -> i j", column, column)

kronecker_factors(X1: Float[Array, 'N1 D'], X2: Float[Array, 'N2 D']) -> tuple[tuple[Float[Array, 'P P'], Float[Array, 'N1 N2']], ...]

Return (B_q, K_q(X1, X2)) factors for each latent process.

All latent kernel evaluations share one per-call context per unique kernel instance, so reusing the same kernel across latents (for hyperparameter tying) registers each sample site exactly once.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def kronecker_factors(
    self,
    X1: Float[Array, "N1 D"],
    X2: Float[Array, "N2 D"],
) -> tuple[tuple[Float[Array, "P P"], Float[Array, "N1 N2"]], ...]:
    """Return ``(B_q, K_q(X1, X2))`` factors for each latent process.

    All latent kernel evaluations share one per-call context per
    unique kernel instance, so reusing the same kernel across latents
    (for hyperparameter tying) registers each sample site exactly once.
    """
    with _kernel_contexts(self.kernels):
        return tuple(
            (self.coregionalization_matrix(q), kernel(X1, X2))
            for q, kernel in enumerate(self.kernels)
        )

output_covariance(X1: Float[Array, 'N1 D'], X2: Float[Array, 'N2 D']) -> Float[Array, 'P P N1 N2']

Return output-pair covariance blocks with shape (P, P, N1, N2).

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def output_covariance(
    self,
    X1: Float[Array, "N1 D"],
    X2: Float[Array, "N2 D"],
) -> Float[Array, "P P N1 N2"]:
    """Return output-pair covariance blocks with shape ``(P, P, N1, N2)``."""
    # Per-latent outer product B_q ⊗ K_q → (P, P, N1, N2).
    terms = [
        einx.multiply("p1 p2, n1 n2 -> p1 p2 n1 n2", B_q, K_q)
        for B_q, K_q in self.kronecker_factors(X1, X2)
    ]
    return functools.reduce(jnp.add, terms)

cross_covariance_operator(X1: Float[Array, 'N1 D'], X2: Float[Array, 'N2 D']) -> lx.AbstractLinearOperator

Return Cov[vec(F(X1)), vec(F(X2))] as a sum-of-Kroneckers operator.

The returned operator preserves the per-latent Kronecker structure so structure-aware solvers can avoid materializing the full (P*N1, P*N2) matrix.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def cross_covariance_operator(
    self,
    X1: Float[Array, "N1 D"],
    X2: Float[Array, "N2 D"],
) -> lx.AbstractLinearOperator:
    """Return ``Cov[vec(F(X1)), vec(F(X2))]`` as a sum-of-Kroneckers operator.

    The returned operator preserves the per-latent Kronecker structure
    so structure-aware solvers can avoid materializing the full
    ``(P*N1, P*N2)`` matrix.
    """
    psd_K = X1 is X2
    terms = [
        _kron_block_op(B_q, K_q, psd_K=psd_K)
        for B_q, K_q in self.kronecker_factors(X1, X2)
    ]
    if len(terms) == 1:
        return terms[0]
    return SumOperator(*terms)

cross_covariance(X1: Float[Array, 'N1 D'], X2: Float[Array, 'N2 D']) -> Float[Array, 'PN1 PN2']

Return the dense covariance of vec(F(X1)) and vec(F(X2)).

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def cross_covariance(
    self,
    X1: Float[Array, "N1 D"],
    X2: Float[Array, "N2 D"],
) -> Float[Array, "PN1 PN2"]:
    """Return the dense covariance of ``vec(F(X1))`` and ``vec(F(X2))``."""
    return self.cross_covariance_operator(X1, X2).as_matrix()

full_covariance_operator(X: Float[Array, 'N D']) -> lx.AbstractLinearOperator

Return the Gram operator for isotopic multi-output observations.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def full_covariance_operator(
    self, X: Float[Array, "N D"]
) -> lx.AbstractLinearOperator:
    """Return the Gram operator for isotopic multi-output observations."""
    return self.cross_covariance_operator(X, X)

full_covariance(X: Float[Array, 'N D']) -> Float[Array, 'PN PN']

Return the dense Gram matrix for isotopic multi-output observations.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def full_covariance(self, X: Float[Array, "N D"]) -> Float[Array, "PN PN"]:
    """Return the dense Gram matrix for isotopic multi-output observations."""
    return self.full_covariance_operator(X).as_matrix()

diag(X: Float[Array, 'N D']) -> Float[Array, 'N P']

Return per-input, per-output marginal variances with shape (N, P).

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def diag(self, X: Float[Array, "N D"]) -> Float[Array, "N P"]:
    """Return per-input, per-output marginal variances with shape ``(N, P)``."""
    with _kernel_contexts(self.kernels):
        # Per-latent variance kₙ ⊗ wₚ² → (N, P), summed over latents.
        terms = [
            einx.multiply(
                "n, p -> n p", kernel.diag(X), jnp.square(self.mixing[:, q])
            )
            for q, kernel in enumerate(self.kernels)
        ]
    return functools.reduce(jnp.add, terms)

ICMKernel

Bases: Module

Intrinsic coregionalization model with one shared latent kernel.

The cross-output covariance is kron(B, k(X1, X2)) with B = W W^T + diag(kappa). When kappa is None the extra diagonal term is omitted.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
class ICMKernel(eqx.Module):
    """Intrinsic coregionalization model with one shared latent kernel.

    The cross-output covariance is ``kron(B, k(X1, X2))`` with
    ``B = W W^T + diag(kappa)``. When ``kappa is None`` the extra diagonal
    term is omitted.
    """

    kernel: Kernel
    mixing: Float[Array, "P Q"]
    kappa: Float[Array, " P"] | None = None

    def __check_init__(self) -> None:
        _validate_mixing(self.mixing)
        if self.kappa is not None:
            if self.kappa.shape != (self.mixing.shape[0],):
                raise ValueError(
                    "kappa must have shape (num_outputs,) when provided; "
                    f"got {self.kappa.shape}."
                )
            # ``B = W W^T + diag(kappa)`` is downstream wrapped with the
            # PSD tag (``_psd_matrix_op``). ``W W^T`` is PSD by
            # construction, but a negative ``kappa`` entry can pull
            # ``B`` out of PSD and silently break Cholesky-based solver
            # paths. Reject at construction when the value is concrete;
            # under ``jax.jit`` the check is a no-op (see
            # `_check_nonnegative_concrete`).
            _check_nonnegative_concrete(self.kappa, name="ICMKernel.kappa")

    @property
    def num_outputs(self) -> int:
        """Number of observed output channels ``P``."""
        return self.mixing.shape[0]

    @property
    def num_latents(self) -> int:
        """Number of latent scalar GPs ``Q``."""
        return self.mixing.shape[1]

    def coregionalization_matrix(self) -> Float[Array, "P P"]:
        """Return ``B = W W^T + diag(kappa)``."""
        # B = W Wᵀ: contract the shared latent axis q.
        B = einx.dot("p q, r q -> p r", self.mixing, self.mixing)
        if self.kappa is None:
            return B
        return B + jnp.diag(self.kappa)

    def kronecker_factors(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> tuple[Float[Array, "P P"], Float[Array, "N1 N2"]]:
        """Return the shared ``(B, K(X1, X2))`` Kronecker factors."""
        with _kernel_context(self.kernel):
            return self.coregionalization_matrix(), self.kernel(X1, X2)

    def output_covariance(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "P P N1 N2"]:
        """Return output-pair covariance blocks with shape ``(P, P, N1, N2)``."""
        B, K = self.kronecker_factors(X1, X2)
        # Outer product B ⊗ K → (P, P, N1, N2).
        return einx.multiply("p1 p2, n1 n2 -> p1 p2 n1 n2", B, K)

    def cross_covariance_operator(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Kronecker:
        """Return ``Cov[vec(F(X1)), vec(F(X2))]`` as a ``Kronecker`` operator.

        The ``kron(B, K)`` structure lets downstream solvers apply
        `gaussx.kronecker_mll` and related Kronecker-exact routines
        instead of materializing a ``(P*N1, P*N2)`` matrix.
        """
        B, K = self.kronecker_factors(X1, X2)
        return _kron_block_op(B, K, psd_K=X1 is X2)  # ty: ignore[invalid-return-type]

    def cross_covariance(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "PN1 PN2"]:
        """Return the dense covariance of ``vec(F(X1))`` and ``vec(F(X2))``."""
        return self.cross_covariance_operator(X1, X2).as_matrix()

    def full_covariance_operator(self, X: Float[Array, "N D"]) -> Kronecker:
        """Return the Gram operator for isotopic multi-output observations."""
        return self.cross_covariance_operator(X, X)

    def full_covariance(self, X: Float[Array, "N D"]) -> Float[Array, "PN PN"]:
        """Return the dense Gram matrix for isotopic multi-output observations."""
        return self.full_covariance_operator(X).as_matrix()

    def diag(self, X: Float[Array, "N D"]) -> Float[Array, "N P"]:
        """Return per-input, per-output marginal variances with shape ``(N, P)``."""
        with _kernel_context(self.kernel):
            # kₙ ⊗ diag(B)ₚ → (N, P) marginal variances.
            return einx.multiply(
                "n, p -> n p",
                self.kernel.diag(X),
                jnp.diag(self.coregionalization_matrix()),
            )

num_outputs: int property

Number of observed output channels P.

num_latents: int property

Number of latent scalar GPs Q.

coregionalization_matrix() -> Float[Array, 'P P']

Return B = W W^T + diag(kappa).

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def coregionalization_matrix(self) -> Float[Array, "P P"]:
    """Return ``B = W W^T + diag(kappa)``."""
    # B = W Wᵀ: contract the shared latent axis q.
    B = einx.dot("p q, r q -> p r", self.mixing, self.mixing)
    if self.kappa is None:
        return B
    return B + jnp.diag(self.kappa)

kronecker_factors(X1: Float[Array, 'N1 D'], X2: Float[Array, 'N2 D']) -> tuple[Float[Array, 'P P'], Float[Array, 'N1 N2']]

Return the shared (B, K(X1, X2)) Kronecker factors.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def kronecker_factors(
    self,
    X1: Float[Array, "N1 D"],
    X2: Float[Array, "N2 D"],
) -> tuple[Float[Array, "P P"], Float[Array, "N1 N2"]]:
    """Return the shared ``(B, K(X1, X2))`` Kronecker factors."""
    with _kernel_context(self.kernel):
        return self.coregionalization_matrix(), self.kernel(X1, X2)

output_covariance(X1: Float[Array, 'N1 D'], X2: Float[Array, 'N2 D']) -> Float[Array, 'P P N1 N2']

Return output-pair covariance blocks with shape (P, P, N1, N2).

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def output_covariance(
    self,
    X1: Float[Array, "N1 D"],
    X2: Float[Array, "N2 D"],
) -> Float[Array, "P P N1 N2"]:
    """Return output-pair covariance blocks with shape ``(P, P, N1, N2)``."""
    B, K = self.kronecker_factors(X1, X2)
    # Outer product B ⊗ K → (P, P, N1, N2).
    return einx.multiply("p1 p2, n1 n2 -> p1 p2 n1 n2", B, K)

cross_covariance_operator(X1: Float[Array, 'N1 D'], X2: Float[Array, 'N2 D']) -> Kronecker

Return Cov[vec(F(X1)), vec(F(X2))] as a Kronecker operator.

The kron(B, K) structure lets downstream solvers apply gaussx.kronecker_mll and related Kronecker-exact routines instead of materializing a (P*N1, P*N2) matrix.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def cross_covariance_operator(
    self,
    X1: Float[Array, "N1 D"],
    X2: Float[Array, "N2 D"],
) -> Kronecker:
    """Return ``Cov[vec(F(X1)), vec(F(X2))]`` as a ``Kronecker`` operator.

    The ``kron(B, K)`` structure lets downstream solvers apply
    `gaussx.kronecker_mll` and related Kronecker-exact routines
    instead of materializing a ``(P*N1, P*N2)`` matrix.
    """
    B, K = self.kronecker_factors(X1, X2)
    return _kron_block_op(B, K, psd_K=X1 is X2)  # ty: ignore[invalid-return-type]

cross_covariance(X1: Float[Array, 'N1 D'], X2: Float[Array, 'N2 D']) -> Float[Array, 'PN1 PN2']

Return the dense covariance of vec(F(X1)) and vec(F(X2)).

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def cross_covariance(
    self,
    X1: Float[Array, "N1 D"],
    X2: Float[Array, "N2 D"],
) -> Float[Array, "PN1 PN2"]:
    """Return the dense covariance of ``vec(F(X1))`` and ``vec(F(X2))``."""
    return self.cross_covariance_operator(X1, X2).as_matrix()

full_covariance_operator(X: Float[Array, 'N D']) -> Kronecker

Return the Gram operator for isotopic multi-output observations.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def full_covariance_operator(self, X: Float[Array, "N D"]) -> Kronecker:
    """Return the Gram operator for isotopic multi-output observations."""
    return self.cross_covariance_operator(X, X)

full_covariance(X: Float[Array, 'N D']) -> Float[Array, 'PN PN']

Return the dense Gram matrix for isotopic multi-output observations.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def full_covariance(self, X: Float[Array, "N D"]) -> Float[Array, "PN PN"]:
    """Return the dense Gram matrix for isotopic multi-output observations."""
    return self.full_covariance_operator(X).as_matrix()

diag(X: Float[Array, 'N D']) -> Float[Array, 'N P']

Return per-input, per-output marginal variances with shape (N, P).

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def diag(self, X: Float[Array, "N D"]) -> Float[Array, "N P"]:
    """Return per-input, per-output marginal variances with shape ``(N, P)``."""
    with _kernel_context(self.kernel):
        # kₙ ⊗ diag(B)ₚ → (N, P) marginal variances.
        return einx.multiply(
            "n, p -> n p",
            self.kernel.diag(X),
            jnp.diag(self.coregionalization_matrix()),
        )

OILMMKernel

Bases: Module

Orthogonal instantaneous linear mixing model.

The latent GP kernels stay independent. Orthogonal mixing makes it possible to project observations into latent space and run Q scalar GP problems instead of one monolithic multi-output solve. Observation noise lives in a separate pyrox_gp.Likelihood; this class returns noise-free signal covariance, matching the LMC/ICM convention.

Pass check_orthogonal=True to verify W^T W ≈ I at construction — useful as a defensive check when W comes from an external computation that may drift off the Stiefel manifold.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
class OILMMKernel(eqx.Module):
    """Orthogonal instantaneous linear mixing model.

    The latent GP kernels stay independent. Orthogonal mixing makes it
    possible to project observations into latent space and run ``Q`` scalar
    GP problems instead of one monolithic multi-output solve. Observation
    noise lives in a separate `pyrox_gp.Likelihood`; this class
    returns noise-free signal covariance, matching the LMC/ICM convention.

    Pass ``check_orthogonal=True`` to verify ``W^T W ≈ I`` at
    construction — useful as a defensive check when ``W`` comes from
    an external computation that may drift off the Stiefel manifold.
    """

    kernels: tuple[Kernel, ...]
    mixing: Float[Array, "P Q"]
    check_orthogonal: bool = eqx.field(static=True, default=False)

    def __check_init__(self) -> None:
        _validate_mixing(self.mixing)
        _validate_kernel_count(self.kernels, self.mixing.shape[1])
        if self.mixing.shape[1] > self.mixing.shape[0]:
            raise ValueError(
                "orthogonal mixing requires num_latents <= num_outputs to form "
                "a valid semi-orthogonal mixing matrix; "
                f"got {self.mixing.shape[1]} latents and "
                f"{self.mixing.shape[0]} outputs."
            )
        if self.check_orthogonal and not self.is_orthogonal():
            raise ValueError(
                "mixing must satisfy W^T W ≈ I when check_orthogonal=True. "
                "Project via `jnp.linalg.qr(W)[0]` before construction, or "
                "pass check_orthogonal=False to bypass."
            )
        _validate_kernel_scopes_unique(self.kernels)

    @property
    def num_outputs(self) -> int:
        """Number of observed output channels ``P``."""
        return self.mixing.shape[0]

    @property
    def num_latents(self) -> int:
        """Number of latent scalar GPs ``Q``."""
        return self.mixing.shape[1]

    def is_orthogonal(self, *, atol: float = 1e-6, rtol: float = 1e-6) -> bool:
        """Whether the current mixing matrix satisfies ``W^T W ≈ I``.

        Returns a Python ``bool`` via a host sync; not usable inside
        ``jax.jit`` / ``jax.vmap``.
        """
        # Wᵀ W: contract the shared output axis p.
        gram = einx.dot("p q, p r -> q r", self.mixing, self.mixing)
        eye = jnp.eye(self.num_latents, dtype=self.mixing.dtype)
        return bool(jnp.allclose(gram, eye, atol=atol, rtol=rtol))

    def project(
        self,
        Y: Float[Array, "N P"],
        noise_var: Float[Array, " P"] | float,
    ) -> tuple[Float[Array, "N Q"], Float[Array, " Q"]]:
        """Project observations to latent space + per-latent noise variances.

        Delegates to `gaussx.oilmm_project`. Returns
        ``(Y_latent, noise_latent)`` with shapes ``(N, Q)`` and ``(Q,)``;
        the per-latent noise is ``noise_latent = (W**2).T @ noise_var``.
        """
        if Y.ndim != 2 or Y.shape[1] != self.num_outputs:
            raise ValueError(
                f"Y must have shape (N, {self.num_outputs}); got {Y.shape}."
            )
        return oilmm_project(Y, self.mixing, noise_var)

    def back_project(
        self,
        f_means: Float[Array, "N Q"],
        f_vars: Float[Array, "N Q"],
    ) -> tuple[Float[Array, "N P"], Float[Array, "N P"]]:
        """Back-project latent GP predictive ``(means, vars)`` to output space.

        Delegates to `gaussx.oilmm_back_project`.
        """
        return oilmm_back_project(f_means, f_vars, self.mixing)

    def independent_gps(self) -> tuple[Kernel, ...]:
        """Return the latent scalar GP kernels used after projection."""
        return self.kernels

    def signal_factors(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> tuple[tuple[Float[Array, "P P"], Float[Array, "N1 N2"]], ...]:
        """Return the latent signal factors before any observation noise.

        All latent kernel evaluations share one per-call context per
        unique kernel instance, mirroring `LMCKernel.kronecker_factors`.
        """
        with _kernel_contexts(self.kernels):
            return tuple(
                (
                    einx.dot("i, j -> i j", self.mixing[:, q], self.mixing[:, q]),
                    kernel(X1, X2),
                )
                for q, kernel in enumerate(self.kernels)
            )

    def signal_covariance_operator(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> lx.AbstractLinearOperator:
        """Return the noise-free signal covariance as a structured operator."""
        psd_K = X1 is X2
        terms = [
            _kron_block_op(B_q, K_q, psd_K=psd_K)
            for B_q, K_q in self.signal_factors(X1, X2)
        ]
        if len(terms) == 1:
            return terms[0]
        return SumOperator(*terms)

    def signal_covariance(
        self,
        X1: Float[Array, "N1 D"],
        X2: Float[Array, "N2 D"],
    ) -> Float[Array, "PN1 PN2"]:
        """Return the dense noise-free signal covariance matrix."""
        return self.signal_covariance_operator(X1, X2).as_matrix()

    # LMC / ICM-parity aliases — OILMM's "signal covariance" is the
    # noise-free vec-cross covariance, same semantics as
    # ``LMCKernel.cross_covariance`` since noise is no longer kernel-side.
    cross_covariance_operator = signal_covariance_operator
    cross_covariance = signal_covariance

    def full_covariance_operator(
        self, X: Float[Array, "N D"]
    ) -> lx.AbstractLinearOperator:
        """Return the noise-free Gram operator for isotopic observations."""
        return self.signal_covariance_operator(X, X)

    def full_covariance(self, X: Float[Array, "N D"]) -> Float[Array, "PN PN"]:
        """Return the dense noise-free Gram matrix."""
        return self.full_covariance_operator(X).as_matrix()

    def diag(self, X: Float[Array, "N D"]) -> Float[Array, "N P"]:
        """Return per-input, per-output marginal signal variances."""
        with _kernel_contexts(self.kernels):
            # Per-latent variance kₙ ⊗ wₚ² → (N, P), summed over latents.
            terms = [
                einx.multiply(
                    "n, p -> n p", kernel.diag(X), jnp.square(self.mixing[:, q])
                )
                for q, kernel in enumerate(self.kernels)
            ]
        return functools.reduce(jnp.add, terms)

num_outputs: int property

Number of observed output channels P.

num_latents: int property

Number of latent scalar GPs Q.

is_orthogonal(*, atol: float = 1e-06, rtol: float = 1e-06) -> bool

Whether the current mixing matrix satisfies W^T W ≈ I.

Returns a Python bool via a host sync; not usable inside jax.jit / jax.vmap.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def is_orthogonal(self, *, atol: float = 1e-6, rtol: float = 1e-6) -> bool:
    """Whether the current mixing matrix satisfies ``W^T W ≈ I``.

    Returns a Python ``bool`` via a host sync; not usable inside
    ``jax.jit`` / ``jax.vmap``.
    """
    # Wᵀ W: contract the shared output axis p.
    gram = einx.dot("p q, p r -> q r", self.mixing, self.mixing)
    eye = jnp.eye(self.num_latents, dtype=self.mixing.dtype)
    return bool(jnp.allclose(gram, eye, atol=atol, rtol=rtol))

project(Y: Float[Array, 'N P'], noise_var: Float[Array, ' P'] | float) -> tuple[Float[Array, 'N Q'], Float[Array, ' Q']]

Project observations to latent space + per-latent noise variances.

Delegates to gaussx.oilmm_project. Returns (Y_latent, noise_latent) with shapes (N, Q) and (Q,); the per-latent noise is noise_latent = (W**2).T @ noise_var.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def project(
    self,
    Y: Float[Array, "N P"],
    noise_var: Float[Array, " P"] | float,
) -> tuple[Float[Array, "N Q"], Float[Array, " Q"]]:
    """Project observations to latent space + per-latent noise variances.

    Delegates to `gaussx.oilmm_project`. Returns
    ``(Y_latent, noise_latent)`` with shapes ``(N, Q)`` and ``(Q,)``;
    the per-latent noise is ``noise_latent = (W**2).T @ noise_var``.
    """
    if Y.ndim != 2 or Y.shape[1] != self.num_outputs:
        raise ValueError(
            f"Y must have shape (N, {self.num_outputs}); got {Y.shape}."
        )
    return oilmm_project(Y, self.mixing, noise_var)

back_project(f_means: Float[Array, 'N Q'], f_vars: Float[Array, 'N Q']) -> tuple[Float[Array, 'N P'], Float[Array, 'N P']]

Back-project latent GP predictive (means, vars) to output space.

Delegates to gaussx.oilmm_back_project.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def back_project(
    self,
    f_means: Float[Array, "N Q"],
    f_vars: Float[Array, "N Q"],
) -> tuple[Float[Array, "N P"], Float[Array, "N P"]]:
    """Back-project latent GP predictive ``(means, vars)`` to output space.

    Delegates to `gaussx.oilmm_back_project`.
    """
    return oilmm_back_project(f_means, f_vars, self.mixing)

independent_gps() -> tuple[Kernel, ...]

Return the latent scalar GP kernels used after projection.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def independent_gps(self) -> tuple[Kernel, ...]:
    """Return the latent scalar GP kernels used after projection."""
    return self.kernels

signal_factors(X1: Float[Array, 'N1 D'], X2: Float[Array, 'N2 D']) -> tuple[tuple[Float[Array, 'P P'], Float[Array, 'N1 N2']], ...]

Return the latent signal factors before any observation noise.

All latent kernel evaluations share one per-call context per unique kernel instance, mirroring LMCKernel.kronecker_factors.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def signal_factors(
    self,
    X1: Float[Array, "N1 D"],
    X2: Float[Array, "N2 D"],
) -> tuple[tuple[Float[Array, "P P"], Float[Array, "N1 N2"]], ...]:
    """Return the latent signal factors before any observation noise.

    All latent kernel evaluations share one per-call context per
    unique kernel instance, mirroring `LMCKernel.kronecker_factors`.
    """
    with _kernel_contexts(self.kernels):
        return tuple(
            (
                einx.dot("i, j -> i j", self.mixing[:, q], self.mixing[:, q]),
                kernel(X1, X2),
            )
            for q, kernel in enumerate(self.kernels)
        )

signal_covariance_operator(X1: Float[Array, 'N1 D'], X2: Float[Array, 'N2 D']) -> lx.AbstractLinearOperator

Return the noise-free signal covariance as a structured operator.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def signal_covariance_operator(
    self,
    X1: Float[Array, "N1 D"],
    X2: Float[Array, "N2 D"],
) -> lx.AbstractLinearOperator:
    """Return the noise-free signal covariance as a structured operator."""
    psd_K = X1 is X2
    terms = [
        _kron_block_op(B_q, K_q, psd_K=psd_K)
        for B_q, K_q in self.signal_factors(X1, X2)
    ]
    if len(terms) == 1:
        return terms[0]
    return SumOperator(*terms)

signal_covariance(X1: Float[Array, 'N1 D'], X2: Float[Array, 'N2 D']) -> Float[Array, 'PN1 PN2']

Return the dense noise-free signal covariance matrix.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def signal_covariance(
    self,
    X1: Float[Array, "N1 D"],
    X2: Float[Array, "N2 D"],
) -> Float[Array, "PN1 PN2"]:
    """Return the dense noise-free signal covariance matrix."""
    return self.signal_covariance_operator(X1, X2).as_matrix()

full_covariance_operator(X: Float[Array, 'N D']) -> lx.AbstractLinearOperator

Return the noise-free Gram operator for isotopic observations.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def full_covariance_operator(
    self, X: Float[Array, "N D"]
) -> lx.AbstractLinearOperator:
    """Return the noise-free Gram operator for isotopic observations."""
    return self.signal_covariance_operator(X, X)

full_covariance(X: Float[Array, 'N D']) -> Float[Array, 'PN PN']

Return the dense noise-free Gram matrix.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def full_covariance(self, X: Float[Array, "N D"]) -> Float[Array, "PN PN"]:
    """Return the dense noise-free Gram matrix."""
    return self.full_covariance_operator(X).as_matrix()

diag(X: Float[Array, 'N D']) -> Float[Array, 'N P']

Return per-input, per-output marginal signal variances.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def diag(self, X: Float[Array, "N D"]) -> Float[Array, "N P"]:
    """Return per-input, per-output marginal signal variances."""
    with _kernel_contexts(self.kernels):
        # Per-latent variance kₙ ⊗ wₚ² → (N, P), summed over latents.
        terms = [
            einx.multiply(
                "n, p -> n p", kernel.diag(X), jnp.square(self.mixing[:, q])
            )
            for q, kernel in enumerate(self.kernels)
        ]
    return functools.reduce(jnp.add, terms)

MultiOutputInducingVariables

Bases: Module

Shared inducing-point structure for LMC-style sparse workflows.

mixing[p, q] is the weight with which latent process q enters output p; the block layout of K_uf matches that convention. ICMKernel with non-zero kappa cannot be represented here — the extra diagonal does not fit the per-latent factorization.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
class MultiOutputInducingVariables(eqx.Module):
    """Shared inducing-point structure for LMC-style sparse workflows.

    ``mixing[p, q]`` is the weight with which latent process ``q`` enters
    output ``p``; the block layout of `K_uf` matches that convention.
    ``ICMKernel`` with non-zero ``kappa`` cannot be represented here —
    the extra diagonal does not fit the per-latent factorization.
    """

    inducing: SharedInducingPoints
    mixing: Float[Array, "P Q"]

    def __check_init__(self) -> None:
        _validate_mixing(self.mixing)

    @classmethod
    def from_kernel(
        cls,
        kernel: LMCKernel | ICMKernel,
        inducing: SharedInducingPoints,
    ) -> MultiOutputInducingVariables:
        """Construct from a kernel, sharing its mixing matrix.

        Avoids the footgun of maintaining two independent ``mixing``
        copies that can silently disagree between ``K_ff`` and ``K_uf``.
        Only `LMCKernel` and `ICMKernel` are accepted;
        `OILMMKernel` is rejected because the sparse inducing
        workflow does not currently exploit orthogonal projection.

        `ICMKernel` with non-zero ``kappa`` is rejected: the
        sparse blocks assembled downstream (`K_uu` / `K_uf`)
        do not carry the ``diag(kappa)`` contribution that is present
        in `ICMKernel.K_ff`, so accepting it here would silently
        produce inconsistent covariance terms and underestimate output
        variance. Users who need the ``kappa`` contribution should use
        the dense multi-output solve.
        """
        if isinstance(kernel, (LMCKernel, ICMKernel)):
            if isinstance(kernel, ICMKernel) and kernel.kappa is not None:
                _require_all_zero_concrete(kernel.kappa, name="ICMKernel.kappa")
            return cls(inducing=inducing, mixing=kernel.mixing)
        raise TypeError(
            "from_kernel only accepts LMCKernel or ICMKernel; "
            f"got {type(kernel).__name__}."
        )

    @property
    def num_outputs(self) -> int:
        """Number of observed output channels ``P``."""
        return self.mixing.shape[0]

    @property
    def num_latents(self) -> int:
        """Number of latent scalar GPs ``Q``."""
        return self.mixing.shape[1]

    def K_uu_operator(self, kernels: tuple[Kernel, ...]) -> BlockDiag:
        """Return the block-diagonal inducing covariance as a ``BlockDiag``."""
        _validate_kernel_count(kernels, self.num_latents)
        return self.inducing.K_uu_operator(kernels)

    def K_uu(self, kernels: tuple[Kernel, ...]) -> Float[Array, "QM QM"]:
        """Return the block-diagonal inducing covariance over latent processes."""
        _validate_kernel_count(kernels, self.num_latents)
        return self.inducing.K_uu(kernels)

    def K_uf(
        self,
        X: Float[Array, "N D"],
        kernels: tuple[Kernel, ...],
    ) -> Float[Array, "QM PN"]:
        """Return the inducing-to-output cross-covariance for isotopic outputs.

        The block layout is
        ``K_uf[q*M:(q+1)*M, p*N:(p+1)*N] = mixing[p, q] * k_q(Z, X)``.
        """
        _validate_kernel_count(kernels, self.num_latents)
        latent_blocks = self.inducing.cross_covariances(X, kernels)
        return self._assemble_K_uf(latent_blocks)

    def _assemble_K_uf(
        self, latent_blocks: tuple[Float[Array, "M N"], ...]
    ) -> Float[Array, "QM PN"]:
        """Stack per-latent ``K(Z, X)`` blocks into the ``(Q*M, P*N)`` ``K_uf``."""
        rows = []
        for q, K_zx in enumerate(latent_blocks):
            # Scale K(Z, X) by each output weight, then interleave the output
            # axis into the columns: (P,) ⊙ (M, N) → (M, P·N).
            scaled = einx.multiply("p, m n -> p m n", self.mixing[:, q], K_zx)
            row = einx.id("p m n -> m (p n)", scaled)
            rows.append(row)
        return jnp.concatenate(rows, axis=0)

    def inducing_blocks(
        self,
        X: Float[Array, "N D"],
        kernels: tuple[Kernel, ...],
    ) -> tuple[BlockDiag, Float[Array, "QM PN"]]:
        """Return ``(K_uu_op, K_uf)`` under one shared kernel context.

        Use this in place of separate `K_uu_operator` /
        `K_uf` calls when assembling an SVGP-style sparse
        predictive: it shares one per-call context per unique kernel
        instance across both blocks, so Pattern B/C kernels with priored
        hyperparameters register their `pyrox_sample` sites
        exactly once and the two blocks see the same hyperparameter
        draw. Sequential calls would close the per-call context between
        ``K_uu`` and ``K_uf``, which would re-register sites under a
        NumPyro trace (or resample inconsistent hyperparameters under
        `numpyro.handlers.seed`).
        """
        _validate_kernel_count(kernels, self.num_latents)
        K_uu_blocks, K_uf_blocks = self.inducing.inducing_blocks(X, kernels)
        K_uu_op = BlockDiag(*(_psd_matrix_op(B) for B in K_uu_blocks))
        K_uf = self._assemble_K_uf(K_uf_blocks)
        return K_uu_op, K_uf

num_outputs: int property

Number of observed output channels P.

num_latents: int property

Number of latent scalar GPs Q.

from_kernel(kernel: LMCKernel | ICMKernel, inducing: SharedInducingPoints) -> MultiOutputInducingVariables classmethod

Construct from a kernel, sharing its mixing matrix.

Avoids the footgun of maintaining two independent mixing copies that can silently disagree between K_ff and K_uf. Only LMCKernel and ICMKernel are accepted; OILMMKernel is rejected because the sparse inducing workflow does not currently exploit orthogonal projection.

ICMKernel with non-zero kappa is rejected: the sparse blocks assembled downstream (K_uu / K_uf) do not carry the diag(kappa) contribution that is present in ICMKernel.K_ff, so accepting it here would silently produce inconsistent covariance terms and underestimate output variance. Users who need the kappa contribution should use the dense multi-output solve.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
@classmethod
def from_kernel(
    cls,
    kernel: LMCKernel | ICMKernel,
    inducing: SharedInducingPoints,
) -> MultiOutputInducingVariables:
    """Construct from a kernel, sharing its mixing matrix.

    Avoids the footgun of maintaining two independent ``mixing``
    copies that can silently disagree between ``K_ff`` and ``K_uf``.
    Only `LMCKernel` and `ICMKernel` are accepted;
    `OILMMKernel` is rejected because the sparse inducing
    workflow does not currently exploit orthogonal projection.

    `ICMKernel` with non-zero ``kappa`` is rejected: the
    sparse blocks assembled downstream (`K_uu` / `K_uf`)
    do not carry the ``diag(kappa)`` contribution that is present
    in `ICMKernel.K_ff`, so accepting it here would silently
    produce inconsistent covariance terms and underestimate output
    variance. Users who need the ``kappa`` contribution should use
    the dense multi-output solve.
    """
    if isinstance(kernel, (LMCKernel, ICMKernel)):
        if isinstance(kernel, ICMKernel) and kernel.kappa is not None:
            _require_all_zero_concrete(kernel.kappa, name="ICMKernel.kappa")
        return cls(inducing=inducing, mixing=kernel.mixing)
    raise TypeError(
        "from_kernel only accepts LMCKernel or ICMKernel; "
        f"got {type(kernel).__name__}."
    )

K_uu_operator(kernels: tuple[Kernel, ...]) -> BlockDiag

Return the block-diagonal inducing covariance as a BlockDiag.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def K_uu_operator(self, kernels: tuple[Kernel, ...]) -> BlockDiag:
    """Return the block-diagonal inducing covariance as a ``BlockDiag``."""
    _validate_kernel_count(kernels, self.num_latents)
    return self.inducing.K_uu_operator(kernels)

K_uu(kernels: tuple[Kernel, ...]) -> Float[Array, 'QM QM']

Return the block-diagonal inducing covariance over latent processes.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def K_uu(self, kernels: tuple[Kernel, ...]) -> Float[Array, "QM QM"]:
    """Return the block-diagonal inducing covariance over latent processes."""
    _validate_kernel_count(kernels, self.num_latents)
    return self.inducing.K_uu(kernels)

K_uf(X: Float[Array, 'N D'], kernels: tuple[Kernel, ...]) -> Float[Array, 'QM PN']

Return the inducing-to-output cross-covariance for isotopic outputs.

The block layout is K_uf[q*M:(q+1)*M, p*N:(p+1)*N] = mixing[p, q] * k_q(Z, X).

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def K_uf(
    self,
    X: Float[Array, "N D"],
    kernels: tuple[Kernel, ...],
) -> Float[Array, "QM PN"]:
    """Return the inducing-to-output cross-covariance for isotopic outputs.

    The block layout is
    ``K_uf[q*M:(q+1)*M, p*N:(p+1)*N] = mixing[p, q] * k_q(Z, X)``.
    """
    _validate_kernel_count(kernels, self.num_latents)
    latent_blocks = self.inducing.cross_covariances(X, kernels)
    return self._assemble_K_uf(latent_blocks)

inducing_blocks(X: Float[Array, 'N D'], kernels: tuple[Kernel, ...]) -> tuple[BlockDiag, Float[Array, 'QM PN']]

Return (K_uu_op, K_uf) under one shared kernel context.

Use this in place of separate K_uu_operator / K_uf calls when assembling an SVGP-style sparse predictive: it shares one per-call context per unique kernel instance across both blocks, so Pattern B/C kernels with priored hyperparameters register their pyrox_sample sites exactly once and the two blocks see the same hyperparameter draw. Sequential calls would close the per-call context between K_uu and K_uf, which would re-register sites under a NumPyro trace (or resample inconsistent hyperparameters under numpyro.handlers.seed).

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def inducing_blocks(
    self,
    X: Float[Array, "N D"],
    kernels: tuple[Kernel, ...],
) -> tuple[BlockDiag, Float[Array, "QM PN"]]:
    """Return ``(K_uu_op, K_uf)`` under one shared kernel context.

    Use this in place of separate `K_uu_operator` /
    `K_uf` calls when assembling an SVGP-style sparse
    predictive: it shares one per-call context per unique kernel
    instance across both blocks, so Pattern B/C kernels with priored
    hyperparameters register their `pyrox_sample` sites
    exactly once and the two blocks see the same hyperparameter
    draw. Sequential calls would close the per-call context between
    ``K_uu`` and ``K_uf``, which would re-register sites under a
    NumPyro trace (or resample inconsistent hyperparameters under
    `numpyro.handlers.seed`).
    """
    _validate_kernel_count(kernels, self.num_latents)
    K_uu_blocks, K_uf_blocks = self.inducing.inducing_blocks(X, kernels)
    K_uu_op = BlockDiag(*(_psd_matrix_op(B) for B in K_uu_blocks))
    K_uf = self._assemble_K_uf(K_uf_blocks)
    return K_uu_op, K_uf

SharedInducingPoints

Bases: Module

Shared inducing inputs for multi-output latent GP collections.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
class SharedInducingPoints(eqx.Module):
    """Shared inducing inputs for multi-output latent GP collections."""

    locations: Float[Array, "M D"]

    def __check_init__(self) -> None:
        if self.locations.ndim != 2:
            raise ValueError("locations must have shape (num_inducing, input_dim).")

    @property
    def num_inducing(self) -> int:
        """Number of inducing inputs ``M`` shared by all latent processes."""
        return self.locations.shape[0]

    def latent_covariances(
        self, kernels: tuple[Kernel, ...]
    ) -> tuple[Float[Array, "M M"], ...]:
        """Return one inducing covariance block per latent kernel.

        Shares a per-call context per unique kernel instance so a kernel
        reused across latents registers its sample sites once.
        """
        if not kernels:
            raise ValueError("kernels must be non-empty.")
        _validate_kernel_scopes_unique(kernels)
        with _kernel_contexts(kernels):
            return tuple(kernel(self.locations, self.locations) for kernel in kernels)

    def K_uu_operator(self, kernels: tuple[Kernel, ...]) -> BlockDiag:
        """Return the block-diagonal inducing covariance as a ``BlockDiag``.

        Downstream solvers decompose a ``(Q*M, Q*M)`` solve into ``Q``
        independent ``(M, M)`` solves via the ``block_diagonal_tag``.
        """
        if not kernels:
            raise ValueError("kernels must be non-empty.")
        _validate_kernel_scopes_unique(kernels)
        with _kernel_contexts(kernels):
            blocks = tuple(
                _psd_matrix_op(kernel(self.locations, self.locations))
                for kernel in kernels
            )
        return BlockDiag(*blocks)

    def K_uu(self, kernels: tuple[Kernel, ...]) -> Float[Array, "QM QM"]:
        """Materialize the block-diagonal inducing covariance over all latents."""
        return self.K_uu_operator(kernels).as_matrix()

    def cross_covariances(
        self,
        X: Float[Array, "N D"],
        kernels: tuple[Kernel, ...],
    ) -> tuple[Float[Array, "M N"], ...]:
        """Return one ``K(Z, X)`` block per latent kernel."""
        if not kernels:
            raise ValueError("kernels must be non-empty.")
        _validate_kernel_scopes_unique(kernels)
        with _kernel_contexts(kernels):
            return tuple(kernel(self.locations, X) for kernel in kernels)

    def inducing_blocks(
        self,
        X: Float[Array, "N D"],
        kernels: tuple[Kernel, ...],
    ) -> tuple[tuple[Float[Array, "M M"], ...], tuple[Float[Array, "M N"], ...]]:
        """Return ``(K_uu_blocks, K_uf_blocks)`` under one shared context.

        Pairs the per-latent ``K(Z, Z)`` and ``K(Z, X)`` evaluations so
        Pattern B/C kernels with priored hyperparameters register their
        `pyrox_sample` sites once across the K_uu/K_uf pair.
        Calling `latent_covariances` and `cross_covariances`
        sequentially would close and reopen each kernel's per-call
        context between the two calls, clearing the sample-site cache
        and tripping duplicate-site registration under a NumPyro trace
        (or resampling inconsistent hyperparameters under
        `numpyro.handlers.seed`).
        """
        if not kernels:
            raise ValueError("kernels must be non-empty.")
        _validate_kernel_scopes_unique(kernels)
        with _kernel_contexts(kernels):
            K_uu_blocks = tuple(
                kernel(self.locations, self.locations) for kernel in kernels
            )
            K_uf_blocks = tuple(kernel(self.locations, X) for kernel in kernels)
        return K_uu_blocks, K_uf_blocks

num_inducing: int property

Number of inducing inputs M shared by all latent processes.

latent_covariances(kernels: tuple[Kernel, ...]) -> tuple[Float[Array, 'M M'], ...]

Return one inducing covariance block per latent kernel.

Shares a per-call context per unique kernel instance so a kernel reused across latents registers its sample sites once.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def latent_covariances(
    self, kernels: tuple[Kernel, ...]
) -> tuple[Float[Array, "M M"], ...]:
    """Return one inducing covariance block per latent kernel.

    Shares a per-call context per unique kernel instance so a kernel
    reused across latents registers its sample sites once.
    """
    if not kernels:
        raise ValueError("kernels must be non-empty.")
    _validate_kernel_scopes_unique(kernels)
    with _kernel_contexts(kernels):
        return tuple(kernel(self.locations, self.locations) for kernel in kernels)

K_uu_operator(kernels: tuple[Kernel, ...]) -> BlockDiag

Return the block-diagonal inducing covariance as a BlockDiag.

Downstream solvers decompose a (Q*M, Q*M) solve into Q independent (M, M) solves via the block_diagonal_tag.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def K_uu_operator(self, kernels: tuple[Kernel, ...]) -> BlockDiag:
    """Return the block-diagonal inducing covariance as a ``BlockDiag``.

    Downstream solvers decompose a ``(Q*M, Q*M)`` solve into ``Q``
    independent ``(M, M)`` solves via the ``block_diagonal_tag``.
    """
    if not kernels:
        raise ValueError("kernels must be non-empty.")
    _validate_kernel_scopes_unique(kernels)
    with _kernel_contexts(kernels):
        blocks = tuple(
            _psd_matrix_op(kernel(self.locations, self.locations))
            for kernel in kernels
        )
    return BlockDiag(*blocks)

K_uu(kernels: tuple[Kernel, ...]) -> Float[Array, 'QM QM']

Materialize the block-diagonal inducing covariance over all latents.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def K_uu(self, kernels: tuple[Kernel, ...]) -> Float[Array, "QM QM"]:
    """Materialize the block-diagonal inducing covariance over all latents."""
    return self.K_uu_operator(kernels).as_matrix()

cross_covariances(X: Float[Array, 'N D'], kernels: tuple[Kernel, ...]) -> tuple[Float[Array, 'M N'], ...]

Return one K(Z, X) block per latent kernel.

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def cross_covariances(
    self,
    X: Float[Array, "N D"],
    kernels: tuple[Kernel, ...],
) -> tuple[Float[Array, "M N"], ...]:
    """Return one ``K(Z, X)`` block per latent kernel."""
    if not kernels:
        raise ValueError("kernels must be non-empty.")
    _validate_kernel_scopes_unique(kernels)
    with _kernel_contexts(kernels):
        return tuple(kernel(self.locations, X) for kernel in kernels)

inducing_blocks(X: Float[Array, 'N D'], kernels: tuple[Kernel, ...]) -> tuple[tuple[Float[Array, 'M M'], ...], tuple[Float[Array, 'M N'], ...]]

Return (K_uu_blocks, K_uf_blocks) under one shared context.

Pairs the per-latent K(Z, Z) and K(Z, X) evaluations so Pattern B/C kernels with priored hyperparameters register their pyrox_sample sites once across the K_uu/K_uf pair. Calling latent_covariances and cross_covariances sequentially would close and reopen each kernel's per-call context between the two calls, clearing the sample-site cache and tripping duplicate-site registration under a NumPyro trace (or resampling inconsistent hyperparameters under numpyro.handlers.seed).

Source code in packages/pyrox-gp/src/pyrox_gp/_multi_output.py
def inducing_blocks(
    self,
    X: Float[Array, "N D"],
    kernels: tuple[Kernel, ...],
) -> tuple[tuple[Float[Array, "M M"], ...], tuple[Float[Array, "M N"], ...]]:
    """Return ``(K_uu_blocks, K_uf_blocks)`` under one shared context.

    Pairs the per-latent ``K(Z, Z)`` and ``K(Z, X)`` evaluations so
    Pattern B/C kernels with priored hyperparameters register their
    `pyrox_sample` sites once across the K_uu/K_uf pair.
    Calling `latent_covariances` and `cross_covariances`
    sequentially would close and reopen each kernel's per-call
    context between the two calls, clearing the sample-site cache
    and tripping duplicate-site registration under a NumPyro trace
    (or resampling inconsistent hyperparameters under
    `numpyro.handlers.seed`).
    """
    if not kernels:
        raise ValueError("kernels must be non-empty.")
    _validate_kernel_scopes_unique(kernels)
    with _kernel_contexts(kernels):
        K_uu_blocks = tuple(
            kernel(self.locations, self.locations) for kernel in kernels
        )
        K_uf_blocks = tuple(kernel(self.locations, X) for kernel in kernels)
    return K_uu_blocks, K_uf_blocks

Latent factor

Collapsed latent-factor regression: the linear decoder (mixing matrix) carries a fixed unit-normal prior and is marginalized analytically, so the likelihood factorizes only a Q x Q capacitance matrix and costs O(NQP) in the output dimension. Contrast the coregionalization kernels above, which hold the mixing matrix as a concrete array.

LatentFactorGPPrior

Bases: Module

GP latent factors with an analytically marginalized linear decoder.

Unlike pyrox_gp.OILMMKernel, the mixing matrix is not a field of this module. It is a random variable with a fixed \(\mathcal{N}(0,1)\) prior, integrated out in collapsed_lfr_log_prob and recovered in closed form by decoder_posterior. There is therefore no orthogonality requirement and no Q <= P constraint, and the output dimension P is not fixed at construction — it is read from Y at condition time.

Caveats to keep in mind:

  • Only the span of Z is identified. The \(\mathcal{N}(0, I)\) prior on the decoder is rotation-invariant, so any orthogonal \(Z \to ZA\), \(W \to A^\top W\) leaves the objective unchanged. Individual latent factors carry no physical meaning without an extra rotation criterion.
  • The kernel amplitude is identifiable (unlike the rotation above). \(Z \to cZ\), \(W \to W/c\) is degenerate only when \(W\) is a free parameter; here \(W\) carries a fixed \(\mathcal{N}(0, I)\) prior and is marginalized, so scaling \(Z\) changes the collapsed covariance \(ZZ^\top + \sigma^2 I\) and the amplitude has a finite, data-dependent optimum. Fixing variance = 1.0 per latent is a legitimate modelling choice (it makes the factors comparable), not a requirement for a valid fit.
  • predict_latents variance understates uncertainty. It treats the MAP Z as noiseless observations of the latent processes, so it captures input-space extrapolation but not uncertainty in the point estimate itself.
  • z_* and W are assumed independent in predict. They are not — both depend on the training data through the fitted Z. Reasonable when P is large; understates variance for small P.
  • This is a low-data model. Z is N x Q free parameters, does not amortize, does not minibatch over N, and carries O(Q N^3) through the latent priors. The reference experiments top out at N = 800.

Warp-specific caveats (ignore when warp is None):

  • Y must sit inside the warp's support. gauss_flows.RQSplineMarginal is linear outside [-interval, interval]; scale Y sensibly or the warp degenerates to affine.
  • Identifiability gets worse. A per-channel affine warp is degenerate with \(\sigma\) and with the columns of \(W\). Identity init plus a distinct learning-rate group (param_group_optimizer) is the mitigation.
  • The warp is fitted to the marginals of Y, not the residuals. A warp that Gaussianizes raw channel marginals is not necessarily the one that Gaussianizes the noise — a genuine modelling approximation, not a bug.
  • Invertibility is mandatory. The change of variables requires a bijection, unlike warped-likelihood GPs where a bound holds for any monotone link.

Attributes:

Name Type Description
kernels tuple[Kernel, ...]

One pyrox_gp.Kernel per latent process; len is Q. Distinct pyrox.PyroxModule kernels must carry distinct pyrox_name values or their NumPyro sites collide.

X Float[Array, 'N D']

Training inputs of shape (N, D).

warp AbstractBijection | None

Optional bijection with event shape (P,) — one marginal transform per output channel. The warped observations \(G^{-1}(Y)\) are modelled as the linear factor model, which handles skewed / heavy-tailed / positive channels without breaking the analytic decoder marginalization (the log-det Jacobian is free of \(Z\), \(W\), and \(\sigma\)). None keeps the plain Gaussian factor model. Prefer gauss_flows.RQSplineMarginal — closed-form in both directions and the exact identity at initialization, so a warped fit starts precisely at the unwarped fit. Conditional (input-dependent) warps are rejected.

latent_noise float

Model nugget added to each latent GP covariance. A modelling choice that controls how tightly the latent GP interpolates the MAP factors — distinct from jitter.

jitter float

Numerical diagonal regularization for Cholesky stability.

Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor_models.py
class LatentFactorGPPrior(eqx.Module):
    """GP latent factors with an analytically marginalized linear decoder.

    Unlike `pyrox_gp.OILMMKernel`, the mixing matrix is not a field of
    this module. It is a random variable with a fixed $\\mathcal{N}(0,1)$
    prior, integrated out in
    [`collapsed_lfr_log_prob`][pyrox_gp.collapsed_lfr_log_prob] and recovered
    in closed form by
    [`decoder_posterior`][pyrox_gp.decoder_posterior]. There is therefore no
    orthogonality requirement and no ``Q <= P`` constraint, and the output
    dimension ``P`` is not fixed at construction — it is read from ``Y`` at
    `condition` time.

    Caveats to keep in mind:

    - **Only the span of ``Z`` is identified.** The $\\mathcal{N}(0, I)$
      prior on the decoder is rotation-invariant, so any orthogonal
      $Z \\to ZA$, $W \\to A^\\top W$ leaves the objective unchanged.
      Individual latent factors carry no physical meaning without an extra
      rotation criterion.
    - **The kernel amplitude is identifiable** (unlike the rotation above).
      $Z \\to cZ$, $W \\to W/c$ is degenerate only when $W$ is a free
      parameter; here $W$ carries a fixed $\\mathcal{N}(0, I)$ prior and is
      marginalized, so scaling $Z$ changes the collapsed covariance
      $ZZ^\\top + \\sigma^2 I$ and the amplitude has a finite,
      data-dependent optimum. Fixing ``variance = 1.0`` per latent is a
      legitimate modelling choice (it makes the factors comparable), not a
      requirement for a valid fit.
    - **`predict_latents` variance understates uncertainty.** It treats the
      MAP ``Z`` as noiseless observations of the latent processes, so it
      captures input-space extrapolation but not uncertainty in the point
      estimate itself.
    - **``z_*`` and ``W`` are assumed independent** in `predict`. They are
      not — both depend on the training data through the fitted ``Z``.
      Reasonable when ``P`` is large; understates variance for small ``P``.
    - **This is a low-data model.** ``Z`` is ``N x Q`` free parameters, does
      not amortize, does not minibatch over ``N``, and carries ``O(Q N^3)``
      through the latent priors. The reference experiments top out at
      ``N = 800``.

    Warp-specific caveats (ignore when ``warp is None``):

    - **``Y`` must sit inside the warp's support.**
      ``gauss_flows.RQSplineMarginal`` is linear outside
      ``[-interval, interval]``; scale ``Y`` sensibly or the warp
      degenerates to affine.
    - **Identifiability gets worse.** A per-channel affine warp is
      degenerate with $\\sigma$ and with the columns of $W$. Identity
      init plus a distinct learning-rate group
      ([`param_group_optimizer`][pyrox.inference.param_group_optimizer])
      is the mitigation.
    - **The warp is fitted to the marginals of ``Y``, not the residuals.**
      A warp that Gaussianizes raw channel marginals is not necessarily
      the one that Gaussianizes the noise — a genuine modelling
      approximation, not a bug.
    - **Invertibility is mandatory.** The change of variables requires a
      bijection, unlike warped-*likelihood* GPs where a bound holds for
      any monotone link.

    Attributes:
        kernels: One `pyrox_gp.Kernel` per latent process; ``len`` is ``Q``.
            Distinct `pyrox.PyroxModule` kernels must carry distinct
            ``pyrox_name`` values or their NumPyro sites collide.
        X: Training inputs of shape ``(N, D)``.
        warp: Optional bijection with event shape ``(P,)`` — one marginal
            transform per output channel. The *warped* observations
            $G^{-1}(Y)$ are modelled as the linear factor model, which
            handles skewed / heavy-tailed / positive channels without
            breaking the analytic decoder marginalization (the log-det
            Jacobian is free of $Z$, $W$, and $\\sigma$). ``None`` keeps
            the plain Gaussian factor model. Prefer
            ``gauss_flows.RQSplineMarginal`` — closed-form in both
            directions and the exact identity at initialization, so a
            warped fit starts precisely at the unwarped fit. Conditional
            (input-dependent) warps are rejected.
        latent_noise: Model nugget added to each latent GP covariance. A
            modelling choice that controls how tightly the latent GP
            interpolates the MAP factors — distinct from ``jitter``.
        jitter: Numerical diagonal regularization for Cholesky stability.
    """

    kernels: tuple[Kernel, ...]
    X: Float[Array, "N D"]
    warp: AbstractBijection | None = None
    latent_noise: float = 1e-3
    jitter: float = 1e-6

    def __check_init__(self) -> None:
        if not self.kernels:
            raise ValueError(
                "kernels must contain at least one latent kernel; an empty "
                "tuple gives Q = 0 and fails inside latent_cholesky."
            )
        _validate_kernel_scopes_unique(self.kernels)
        if self.warp is not None and self.warp.cond_shape is not None:
            raise ValueError(
                "Conditional warps are not supported: the log-det Jacobian "
                "would depend on the inputs, and the per-channel marginal "
                "interpretation no longer holds."
            )

    @property
    def num_latents(self) -> int:
        """Number of latent scalar GPs ``Q``."""
        return len(self.kernels)

    def latent_priors(self) -> tuple[GPPrior, ...]:
        """Return one scalar `pyrox_gp.GPPrior` per latent process."""
        return tuple(
            GPPrior(kernel=k, X=self.X, jitter=self.jitter) for k in self.kernels
        )

    def latent_cholesky(self) -> Float[Array, "Q N N"]:
        """Batched Cholesky factors of the ``Q`` latent GP covariances."""
        n = self.X.shape[0]
        eye = jnp.eye(n, dtype=self.X.dtype)
        with _kernel_contexts(self.kernels):
            covs = jnp.stack(
                [
                    k(self.X, self.X) + (self.latent_noise + self.jitter) * eye
                    for k in self.kernels
                ]
            )
        return jnp.linalg.cholesky(covs)

    def condition(
        self,
        Y: Float[Array, "N P"],
        Z: Float[Array, "N Q"],
        noise_var: Float[Array, ""],
    ) -> ConditionedLatentFactorGP:
        """Recover the decoder posterior for MAP latents ``Z``.

        Returns a `pyrox_gp.ConditionedLatentFactorGP` holding the
        closed-form matrix-normal decoder posterior alongside the latents,
        with each latent GP conditioned once so repeated `predict` calls
        do not repeat the ``O(Q N^3)`` training solve.

        !!! warning "Call this under the same context you fitted in"
            The kernels resolve their hyperparameters when this method
            evaluates them. Outside a NumPyro substitution context they
            resolve to their *initial* values, so predictions would come
            from different kernels than the fit. Wrap the call in
            ``numpyro.handlers.substitute(fn, result.params)`` — see the
            module docstring for the end-to-end pattern.
        """
        if Y.shape[0] != self.X.shape[0]:
            raise ValueError(
                f"Y must have {self.X.shape[0]} rows to match X; got {Y.shape[0]}."
            )
        if Z.shape != (self.X.shape[0], self.num_latents):
            raise ValueError(
                f"Z must have shape {(self.X.shape[0], self.num_latents)}; "
                f"got {Z.shape}."
            )
        if self.warp is None:
            mu_W, Sigma_W = decoder_posterior(Y, Z, noise_var)
        else:
            mu_W, Sigma_W = warped_decoder_posterior(Y, Z, noise_var, self.warp)
        # Condition each latent GP once, here, rather than on every predict
        # call: the training solve is O(N^3) per factor and does not depend
        # on the test inputs. Doing it here also captures the kernel
        # hyperparameters under whatever context the caller conditions in.
        with _kernel_contexts(self.kernels):
            latents = tuple(
                prior.condition(Z[:, q], noise_var=jnp.asarray(self.latent_noise))
                for q, prior in enumerate(self.latent_priors())
            )
        return ConditionedLatentFactorGP(
            prior=self,
            Z=Z,
            mu_W=mu_W,
            Sigma_W=Sigma_W,
            noise_var=noise_var,
            latents=latents,
        )

num_latents: int property

Number of latent scalar GPs Q.

latent_priors() -> tuple[GPPrior, ...]

Return one scalar pyrox_gp.GPPrior per latent process.

Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor_models.py
def latent_priors(self) -> tuple[GPPrior, ...]:
    """Return one scalar `pyrox_gp.GPPrior` per latent process."""
    return tuple(
        GPPrior(kernel=k, X=self.X, jitter=self.jitter) for k in self.kernels
    )

latent_cholesky() -> Float[Array, 'Q N N']

Batched Cholesky factors of the Q latent GP covariances.

Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor_models.py
def latent_cholesky(self) -> Float[Array, "Q N N"]:
    """Batched Cholesky factors of the ``Q`` latent GP covariances."""
    n = self.X.shape[0]
    eye = jnp.eye(n, dtype=self.X.dtype)
    with _kernel_contexts(self.kernels):
        covs = jnp.stack(
            [
                k(self.X, self.X) + (self.latent_noise + self.jitter) * eye
                for k in self.kernels
            ]
        )
    return jnp.linalg.cholesky(covs)

condition(Y: Float[Array, 'N P'], Z: Float[Array, 'N Q'], noise_var: Float[Array, '']) -> ConditionedLatentFactorGP

Recover the decoder posterior for MAP latents Z.

Returns a pyrox_gp.ConditionedLatentFactorGP holding the closed-form matrix-normal decoder posterior alongside the latents, with each latent GP conditioned once so repeated predict calls do not repeat the O(Q N^3) training solve.

Call this under the same context you fitted in

The kernels resolve their hyperparameters when this method evaluates them. Outside a NumPyro substitution context they resolve to their initial values, so predictions would come from different kernels than the fit. Wrap the call in numpyro.handlers.substitute(fn, result.params) — see the module docstring for the end-to-end pattern.

Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor_models.py
def condition(
    self,
    Y: Float[Array, "N P"],
    Z: Float[Array, "N Q"],
    noise_var: Float[Array, ""],
) -> ConditionedLatentFactorGP:
    """Recover the decoder posterior for MAP latents ``Z``.

    Returns a `pyrox_gp.ConditionedLatentFactorGP` holding the
    closed-form matrix-normal decoder posterior alongside the latents,
    with each latent GP conditioned once so repeated `predict` calls
    do not repeat the ``O(Q N^3)`` training solve.

    !!! warning "Call this under the same context you fitted in"
        The kernels resolve their hyperparameters when this method
        evaluates them. Outside a NumPyro substitution context they
        resolve to their *initial* values, so predictions would come
        from different kernels than the fit. Wrap the call in
        ``numpyro.handlers.substitute(fn, result.params)`` — see the
        module docstring for the end-to-end pattern.
    """
    if Y.shape[0] != self.X.shape[0]:
        raise ValueError(
            f"Y must have {self.X.shape[0]} rows to match X; got {Y.shape[0]}."
        )
    if Z.shape != (self.X.shape[0], self.num_latents):
        raise ValueError(
            f"Z must have shape {(self.X.shape[0], self.num_latents)}; "
            f"got {Z.shape}."
        )
    if self.warp is None:
        mu_W, Sigma_W = decoder_posterior(Y, Z, noise_var)
    else:
        mu_W, Sigma_W = warped_decoder_posterior(Y, Z, noise_var, self.warp)
    # Condition each latent GP once, here, rather than on every predict
    # call: the training solve is O(N^3) per factor and does not depend
    # on the test inputs. Doing it here also captures the kernel
    # hyperparameters under whatever context the caller conditions in.
    with _kernel_contexts(self.kernels):
        latents = tuple(
            prior.condition(Z[:, q], noise_var=jnp.asarray(self.latent_noise))
            for q, prior in enumerate(self.latent_priors())
        )
    return ConditionedLatentFactorGP(
        prior=self,
        Z=Z,
        mu_W=mu_W,
        Sigma_W=Sigma_W,
        noise_var=noise_var,
        latents=latents,
    )

ConditionedLatentFactorGP

Bases: Module

Latent-factor posterior — MAP latents plus the decoder posterior.

Attributes:

Name Type Description
prior LatentFactorGPPrior

The pyrox_gp.LatentFactorGPPrior this was conditioned from.

Z Float[Array, 'N Q']

MAP latent factor values at the training inputs, (N, Q).

mu_W Float[Array, 'Q P']

Decoder posterior mean, (Q, P).

Sigma_W Float[Array, 'Q Q']

Decoder posterior row covariance, (Q, Q).

noise_var Float[Array, '']

Scalar isotropic observation noise variance.

latents tuple[ConditionedGP, ...]

One conditioned scalar GP per latent process, built once by LatentFactorGPPrior.condition so the O(Q N^3) training solve is not repeated on every prediction.

Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor_models.py
class ConditionedLatentFactorGP(eqx.Module):
    """Latent-factor posterior — MAP latents plus the decoder posterior.

    Attributes:
        prior: The `pyrox_gp.LatentFactorGPPrior` this was conditioned from.
        Z: MAP latent factor values at the training inputs, ``(N, Q)``.
        mu_W: Decoder posterior mean, ``(Q, P)``.
        Sigma_W: Decoder posterior row covariance, ``(Q, Q)``.
        noise_var: Scalar isotropic observation noise variance.
        latents: One conditioned scalar GP per latent process, built once
            by `LatentFactorGPPrior.condition` so the ``O(Q N^3)``
            training solve is not repeated on every prediction.
    """

    prior: LatentFactorGPPrior
    Z: Float[Array, "N Q"]
    mu_W: Float[Array, "Q P"]
    Sigma_W: Float[Array, "Q Q"]
    noise_var: Float[Array, ""]
    latents: tuple[ConditionedGP, ...]

    def predict_latents(
        self, X_new: Float[Array, "T D"]
    ) -> tuple[Float[Array, "T Q"], Float[Array, "T Q"]]:
        """Per-latent GP conditional means and marginal variances at ``X_new``.

        The MAP latents are treated as observations of the latent processes
        with ``latent_noise`` as their noise, so the returned variance
        reflects input-space extrapolation only — not uncertainty in the
        point estimate ``Z`` itself.
        """
        means, variances = [], []
        with _kernel_contexts(self.prior.kernels):
            for cond in self.latents:
                m, v = cond.predict(X_new)
                means.append(m)
                variances.append(v)
        return jnp.stack(means, -1), jnp.stack(variances, -1)

    def predict(
        self,
        X_new: Float[Array, "T D"],
        *,
        include_noise: bool = False,
        quad_order: int = 32,
    ) -> tuple[Float[Array, "T P"], Float[Array, "T P"]]:
        """Posterior predictive mean and variance over all ``P`` outputs.

        Composes the latent GP conditional with the decoder posterior via
        [`lfr_predictive_moments`][pyrox_gp.lfr_predictive_moments]. The
        variance decomposes into a decoder term, a latent term, and an
        interaction term; the first and third are output-independent.

        With a warp, the warped-space moments are pushed through $G$
        (``transform``) by Gauss-Hermite quadrature. The returned mean is
        $\\mathbb{E}[G(f)]$, **not** $G(\\mathbb{E}[f])$ — the latter is the
        pushforward median for a monotone warp and is badly biased for a
        skewed one.

        !!! warning "The warped predictive is an approximation"
            $f = z_*^\\top W$ is a *product* of two independent Gaussians,
            which is not itself Gaussian;
            [`lfr_predictive_moments`][pyrox_gp.lfr_predictive_moments]
            gives its exact first two moments, and the quadrature below
            builds Gaussian nodes from them. The observation-space moments
            are therefore moment-matched, not exact, and the error grows
            with the product of the latent and decoder variances relative
            to the mean. The unwarped path (``warp=None``) is unaffected —
            it returns the exact moments directly.

        The latent variance includes the model's ``latent_noise``
        nugget, matching the prior `lfr_model` fits. ``include_noise``
        additionally adds the *observation* noise, a separate quantity.

        Args:
            X_new: Test inputs, shape ``(T, D)``.
            include_noise: Add the observation noise variance. With a
                warp, the noise lives in the warped space, so it is added
                before the pushforward.
            quad_order: Gauss-Hermite order for the warped pushforward.
                Ignored when ``warp is None``. A moderate order is fine —
                the quadrature appears only here, never in training.

        """
        z_mean, z_var = self.predict_latents(X_new)
        # lfr_model gives each latent factor the covariance K + nugget*I,
        # so a factor value at a *new* input carries an independent nugget
        # too. predict_latents returns the smooth GP conditional (there the
        # nugget acts as the noise the MAP factors are observed with), so
        # add it back or the decoder propagates an understated variance.
        z_var = z_var + self.prior.latent_noise
        mean, var = lfr_predictive_moments(
            z_mean,
            z_var,
            self.mu_W,
            self.Sigma_W,
            self.noise_var if include_noise else None,
        )
        warp = self.prior.warp
        if warp is None:
            return mean, var
        nodes, weights = np.polynomial.hermite_e.hermegauss(quad_order)
        nodes = jnp.asarray(nodes)
        weights = jnp.asarray(weights) / np.sqrt(2.0 * np.pi)
        # fs: (order, T, P) — per-node evaluation points in the warped space.
        fs = mean[None] + jnp.sqrt(var)[None] * nodes[:, None, None]
        g = jax.vmap(jax.vmap(warp.transform))(fs)
        m1 = einx.dot("s, s t p -> t p", weights, g)
        # Accumulate the *centered* second moment. E[G^2] - E[G]^2 cancels
        # catastrophically when the transformed values carry a large offset
        # relative to their spread (a mean of 1e4 with variance 1 loses the
        # variance entirely in float32), and can even come out negative.
        centered = g - m1[None]
        var = einx.dot("s, s t p -> t p", weights, centered**2)
        return m1, var

predict_latents(X_new: Float[Array, 'T D']) -> tuple[Float[Array, 'T Q'], Float[Array, 'T Q']]

Per-latent GP conditional means and marginal variances at X_new.

The MAP latents are treated as observations of the latent processes with latent_noise as their noise, so the returned variance reflects input-space extrapolation only — not uncertainty in the point estimate Z itself.

Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor_models.py
def predict_latents(
    self, X_new: Float[Array, "T D"]
) -> tuple[Float[Array, "T Q"], Float[Array, "T Q"]]:
    """Per-latent GP conditional means and marginal variances at ``X_new``.

    The MAP latents are treated as observations of the latent processes
    with ``latent_noise`` as their noise, so the returned variance
    reflects input-space extrapolation only — not uncertainty in the
    point estimate ``Z`` itself.
    """
    means, variances = [], []
    with _kernel_contexts(self.prior.kernels):
        for cond in self.latents:
            m, v = cond.predict(X_new)
            means.append(m)
            variances.append(v)
    return jnp.stack(means, -1), jnp.stack(variances, -1)

predict(X_new: Float[Array, 'T D'], *, include_noise: bool = False, quad_order: int = 32) -> tuple[Float[Array, 'T P'], Float[Array, 'T P']]

Posterior predictive mean and variance over all P outputs.

Composes the latent GP conditional with the decoder posterior via lfr_predictive_moments. The variance decomposes into a decoder term, a latent term, and an interaction term; the first and third are output-independent.

With a warp, the warped-space moments are pushed through \(G\) (transform) by Gauss-Hermite quadrature. The returned mean is \(\mathbb{E}[G(f)]\), not \(G(\mathbb{E}[f])\) — the latter is the pushforward median for a monotone warp and is badly biased for a skewed one.

The warped predictive is an approximation

\(f = z_*^\top W\) is a product of two independent Gaussians, which is not itself Gaussian; lfr_predictive_moments gives its exact first two moments, and the quadrature below builds Gaussian nodes from them. The observation-space moments are therefore moment-matched, not exact, and the error grows with the product of the latent and decoder variances relative to the mean. The unwarped path (warp=None) is unaffected — it returns the exact moments directly.

The latent variance includes the model's latent_noise nugget, matching the prior lfr_model fits. include_noise additionally adds the observation noise, a separate quantity.

Parameters:

Name Type Description Default
X_new Float[Array, 'T D']

Test inputs, shape (T, D).

required
include_noise bool

Add the observation noise variance. With a warp, the noise lives in the warped space, so it is added before the pushforward.

False
quad_order int

Gauss-Hermite order for the warped pushforward. Ignored when warp is None. A moderate order is fine — the quadrature appears only here, never in training.

32
Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor_models.py
def predict(
    self,
    X_new: Float[Array, "T D"],
    *,
    include_noise: bool = False,
    quad_order: int = 32,
) -> tuple[Float[Array, "T P"], Float[Array, "T P"]]:
    """Posterior predictive mean and variance over all ``P`` outputs.

    Composes the latent GP conditional with the decoder posterior via
    [`lfr_predictive_moments`][pyrox_gp.lfr_predictive_moments]. The
    variance decomposes into a decoder term, a latent term, and an
    interaction term; the first and third are output-independent.

    With a warp, the warped-space moments are pushed through $G$
    (``transform``) by Gauss-Hermite quadrature. The returned mean is
    $\\mathbb{E}[G(f)]$, **not** $G(\\mathbb{E}[f])$ — the latter is the
    pushforward median for a monotone warp and is badly biased for a
    skewed one.

    !!! warning "The warped predictive is an approximation"
        $f = z_*^\\top W$ is a *product* of two independent Gaussians,
        which is not itself Gaussian;
        [`lfr_predictive_moments`][pyrox_gp.lfr_predictive_moments]
        gives its exact first two moments, and the quadrature below
        builds Gaussian nodes from them. The observation-space moments
        are therefore moment-matched, not exact, and the error grows
        with the product of the latent and decoder variances relative
        to the mean. The unwarped path (``warp=None``) is unaffected —
        it returns the exact moments directly.

    The latent variance includes the model's ``latent_noise``
    nugget, matching the prior `lfr_model` fits. ``include_noise``
    additionally adds the *observation* noise, a separate quantity.

    Args:
        X_new: Test inputs, shape ``(T, D)``.
        include_noise: Add the observation noise variance. With a
            warp, the noise lives in the warped space, so it is added
            before the pushforward.
        quad_order: Gauss-Hermite order for the warped pushforward.
            Ignored when ``warp is None``. A moderate order is fine —
            the quadrature appears only here, never in training.

    """
    z_mean, z_var = self.predict_latents(X_new)
    # lfr_model gives each latent factor the covariance K + nugget*I,
    # so a factor value at a *new* input carries an independent nugget
    # too. predict_latents returns the smooth GP conditional (there the
    # nugget acts as the noise the MAP factors are observed with), so
    # add it back or the decoder propagates an understated variance.
    z_var = z_var + self.prior.latent_noise
    mean, var = lfr_predictive_moments(
        z_mean,
        z_var,
        self.mu_W,
        self.Sigma_W,
        self.noise_var if include_noise else None,
    )
    warp = self.prior.warp
    if warp is None:
        return mean, var
    nodes, weights = np.polynomial.hermite_e.hermegauss(quad_order)
    nodes = jnp.asarray(nodes)
    weights = jnp.asarray(weights) / np.sqrt(2.0 * np.pi)
    # fs: (order, T, P) — per-node evaluation points in the warped space.
    fs = mean[None] + jnp.sqrt(var)[None] * nodes[:, None, None]
    g = jax.vmap(jax.vmap(warp.transform))(fs)
    m1 = einx.dot("s, s t p -> t p", weights, g)
    # Accumulate the *centered* second moment. E[G^2] - E[G]^2 cancels
    # catastrophically when the transformed values carry a large offset
    # relative to their spread (a mean of 1e4 with variance 1 loses the
    # variance entirely in float32), and can even come out negative.
    centered = g - m1[None]
    var = einx.dot("s, s t p -> t p", weights, centered**2)
    return m1, var

lfr_model(X: Float[Array, 'N D'], Y: Float[Array, 'N P'], prior: LatentFactorGPPrior, *, beta: float | None = None, noise_prior_scale: float = 0.5) -> None

MAP-over-latents collapsed latent-factor regression model.

Pair with numpyro.infer.autoguide.AutoDelta — the latents are a point estimate, not a marginalized quantity. Latents are stored transposed, as (Q, N), so the Q independent GP priors form a batch dimension of one batched multivariate normal.

Parameters:

Name Type Description Default
X Float[Array, 'N D']

Training inputs, (N, D). Must match prior.X — the latent covariance is built from the prior, so X contributes only shape and dtype and a mismatch is rejected rather than fitted against the wrong locations.

required
Y Float[Array, 'N P']

Centered observations, (N, P).

required
prior LatentFactorGPPrior

The pyrox_gp.LatentFactorGPPrior to fit.

required
beta float | None

Inverse temperature forwarded to lfr_factor; None means Q / P.

None
noise_prior_scale float

Scale of the half-normal prior on the noise standard deviation.

0.5

With a warp on the prior, the warp's array leaves are registered as a single pytree-valued numpyro.param site ("warp_params") so all four blocks — latents, noise, kernel hyperparameters, and the warp — fit jointly under SVI. After fitting, rebuild the warp for prediction with eqx.combine(result.params["warp_params"], eqx.partition(prior.warp, eqx.is_inexact_array)[1]) and pass it via eqx.tree_at (or reconstruct the prior) before calling condition.

Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor_models.py
def lfr_model(
    X: Float[Array, "N D"],
    Y: Float[Array, "N P"],
    prior: LatentFactorGPPrior,
    *,
    beta: float | None = None,
    noise_prior_scale: float = 0.5,
) -> None:
    """MAP-over-latents collapsed latent-factor regression model.

    Pair with ``numpyro.infer.autoguide.AutoDelta`` — the latents are a
    point estimate, not a marginalized quantity. Latents are stored
    transposed, as ``(Q, N)``, so the ``Q`` independent GP priors form a
    batch dimension of one batched multivariate normal.

    Args:
        X: Training inputs, ``(N, D)``. Must match ``prior.X`` — the latent
            covariance is built from the prior, so ``X`` contributes only
            shape and dtype and a mismatch is rejected rather than fitted
            against the wrong locations.
        Y: Centered observations, ``(N, P)``.
        prior: The `pyrox_gp.LatentFactorGPPrior` to fit.
        beta: Inverse temperature forwarded to
            [`lfr_factor`][pyrox_gp.lfr_factor]; ``None`` means ``Q / P``.
        noise_prior_scale: Scale of the half-normal prior on the noise
            standard deviation.

    With a warp on the prior, the warp's array leaves are registered as a
    single pytree-valued ``numpyro.param`` site (``"warp_params"``) so all
    four blocks — latents, noise, kernel hyperparameters, and the warp —
    fit jointly under SVI. After fitting, rebuild the warp for prediction
    with ``eqx.combine(result.params["warp_params"],
    eqx.partition(prior.warp, eqx.is_inexact_array)[1])`` and pass it via
    ``eqx.tree_at`` (or reconstruct the prior) before calling `condition`.
    """
    if X.shape != prior.X.shape:
        raise ValueError(
            f"X must be the prior's training inputs, shape {prior.X.shape}; "
            f"got {X.shape}. The latent covariance is built from prior.X, so "
            "a different X would silently pair Y with the wrong locations."
        )
    n = X.shape[0]
    q = prior.num_latents
    Z_T = jnp.asarray(
        numpyro.sample(
            "Z_T",
            dist.MultivariateNormal(
                loc=jnp.zeros((q, n), dtype=X.dtype),
                scale_tril=prior.latent_cholesky(),
            ).to_event(1),
        )
    )
    noise = jnp.asarray(numpyro.sample("noise", dist.HalfNormal(noise_prior_scale)))
    if prior.warp is None:
        warp = None
    else:
        # Register the warp's array leaves as one pytree-valued numpyro.param
        # (the same mechanism numpyro.contrib.module uses for flax/haiku
        # params), so SVI fits the warp jointly with Z, noise, and the
        # kernel hyperparameters.
        params, static = eqx.partition(prior.warp, eqx.is_inexact_array)
        params = numpyro.param("warp_params", params)
        warp = cast(AbstractBijection, eqx.combine(params, static))
    lfr_factor(Y, Z_T.T, noise**2, warp=warp, beta=beta)

lfr_factor(Y: Float[Array, 'N P'], Z: Float[Array, 'N Q'], noise_var: Float[Array, ''], *, warp: AbstractBijection | None = None, beta: float | None = None, name: str = 'collapsed_lfr') -> None

Register the collapsed latent-factor log-likelihood with NumPyro.

The likelihood term grows as \(O(NP)\) while the GP prior on \(Z\) grows as \(O(NQ)\), so for \(P \gg Q\) an untempered MAP drives \(Z\) to interpolate noise. The likelihood is therefore scaled by an inverse temperature beta; priors are left at unit weight.

Parameters:

Name Type Description Default
Y Float[Array, 'N P']

Observations, (N, P).

required
Z Float[Array, 'N Q']

Latent factor values, (N, Q).

required
noise_var Float[Array, '']

Scalar isotropic observation noise variance (in the warped space when warp is given).

required
warp AbstractBijection | None

Optional bijection with event shape (P,); when given, registers warped_lfr_log_prob instead of the plain collapsed likelihood.

None
beta float | None

Inverse temperature on the likelihood. None selects Q / P, which keeps the likelihood and the latent GP prior balanced as the output dimension grows. Tune by held-out likelihood when it matters.

None
name str

NumPyro factor site name.

'collapsed_lfr'
Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor_models.py
def lfr_factor(
    Y: Float[Array, "N P"],
    Z: Float[Array, "N Q"],
    noise_var: Float[Array, ""],
    *,
    warp: AbstractBijection | None = None,
    beta: float | None = None,
    name: str = "collapsed_lfr",
) -> None:
    """Register the collapsed latent-factor log-likelihood with NumPyro.

    The likelihood term grows as $O(NP)$ while the GP prior on $Z$ grows
    as $O(NQ)$, so for $P \\gg Q$ an untempered MAP drives $Z$ to
    interpolate noise. The likelihood is therefore scaled by an inverse
    temperature ``beta``; priors are left at unit weight.

    Args:
        Y: Observations, ``(N, P)``.
        Z: Latent factor values, ``(N, Q)``.
        noise_var: Scalar isotropic observation noise variance (in the
            warped space when ``warp`` is given).
        warp: Optional bijection with event shape ``(P,)``; when given,
            registers [`warped_lfr_log_prob`][pyrox_gp.warped_lfr_log_prob]
            instead of the plain collapsed likelihood.
        beta: Inverse temperature on the likelihood. ``None`` selects
            ``Q / P``, which keeps the likelihood and the latent GP prior
            balanced as the output dimension grows. Tune by held-out
            likelihood when it matters.
        name: NumPyro factor site name.
    """
    if beta is None:
        beta = Z.shape[1] / Y.shape[1]
    if warp is None:
        log_prob = collapsed_lfr_log_prob(Y, Z, noise_var)
    else:
        log_prob = warped_lfr_log_prob(Y, Z, noise_var, warp)
    with numpyro.handlers.scale(scale=beta):
        numpyro.factor(name, log_prob)

latent_total_correlation(Z: Float[Array, 'N Q']) -> Float[Array, '']

Total correlation of the fitted latent factors.

The model places an independent GP prior on each of the \(Q\) latent factors. A large value here means that assumption is violated in the fit — usually because \(Q\) is larger than the data supports, or because the optimizer landed in a badly rotated gauge (only the span of \(Z\) is identified; see pyrox_gp.LatentFactorGPPrior).

Diagnostic only — this does not enter the objective. Requires the flows optional dependency (pip install 'pyrox-gp[flows]').

Parameters:

Name Type Description Default
Z Float[Array, 'N Q']

MAP latent factor values, shape (N, Q).

required

Returns:

Type Description
Float[Array, '']

Scalar total correlation, zero for exactly independent factors.

Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor_models.py
def latent_total_correlation(Z: Float[Array, "N Q"]) -> Float[Array, ""]:
    """Total correlation of the fitted latent factors.

    The model places an independent GP prior on each of the $Q$ latent
    factors. A large value here means that assumption is violated in the
    fit — usually because $Q$ is larger than the data supports, or because
    the optimizer landed in a badly rotated gauge (only the span of $Z$ is
    identified; see `pyrox_gp.LatentFactorGPPrior`).

    Diagnostic only — this does not enter the objective. Requires the
    ``flows`` optional dependency (``pip install 'pyrox-gp[flows]'``).

    Args:
        Z: MAP latent factor values, shape ``(N, Q)``.

    Returns:
        Scalar total correlation, zero for exactly independent factors.
    """
    from gauss_flows import (  # ty: ignore[unresolved-import]
        gaussian_total_correlation,
    )

    # jnp.cov collapses to a scalar for a single factor; the total-
    # correlation routine needs a (Q, Q) matrix (and returns zero there).
    return gaussian_total_correlation(jnp.atleast_2d(jnp.cov(Z.T)))

collapsed_lfr_log_prob(Y: Float[Array, 'N P'], Z: Float[Array, 'N Q'], noise_var: Float[Array, ''], *, jitter: float = 1e-12) -> Float[Array, '']

Log-likelihood with the linear decoder marginalized analytically.

Evaluates

\[ p(Y \mid Z, \sigma^2) = \prod_{j=1}^{P} \mathcal{N}(y_j \mid 0,\; ZZ^\top + \sigma^2 I_N) \]

via the Woodbury identity in the \(Q \times Q\) capacitance matrix \(\Psi = I_Q + \sigma^{-2} Z^\top Z\), so no \((N, N)\) matrix is ever formed. Cost is \(O(NQ^2 + NQP + Q^3)\); the output dimension \(P\) enters only through \(\lVert Y \rVert_F^2\) and \(Z^\top Y\).

Parameters:

Name Type Description Default
Y Float[Array, 'N P']

Observations of shape (N, P). Assumed centered.

required
Z Float[Array, 'N Q']

Latent factor values at the training inputs, shape (N, Q).

required
noise_var Float[Array, '']

Scalar isotropic observation noise variance. Per-output noise is not supported — it breaks the shared-covariance identity the derivation rests on.

required
jitter float

Added to noise_var before inversion.

1e-12

Returns:

Type Description
Float[Array, '']

Scalar log-likelihood.

Examples:

>>> import jax.numpy as jnp
>>> Y = jnp.zeros((5, 100))
>>> Z = jnp.ones((5, 2))
>>> float(collapsed_lfr_log_prob(Y, Z, 0.1)) < 0.0
True
Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor.py
def collapsed_lfr_log_prob(
    Y: Float[Array, "N P"],
    Z: Float[Array, "N Q"],
    noise_var: Float[Array, ""],
    *,
    jitter: float = 1e-12,
) -> Float[Array, ""]:
    """Log-likelihood with the linear decoder marginalized analytically.

    Evaluates

    $$
    p(Y \\mid Z, \\sigma^2) = \\prod_{j=1}^{P}
        \\mathcal{N}(y_j \\mid 0,\\; ZZ^\\top + \\sigma^2 I_N)
    $$

    via the Woodbury identity in the $Q \\times Q$ capacitance matrix
    $\\Psi = I_Q + \\sigma^{-2} Z^\\top Z$, so no $(N, N)$ matrix is ever
    formed. Cost is $O(NQ^2 + NQP + Q^3)$; the output dimension $P$ enters
    only through $\\lVert Y \\rVert_F^2$ and $Z^\\top Y$.

    Args:
        Y: Observations of shape ``(N, P)``. Assumed centered.
        Z: Latent factor values at the training inputs, shape ``(N, Q)``.
        noise_var: Scalar isotropic observation noise variance. Per-output
            noise is not supported — it breaks the shared-covariance identity
            the derivation rests on.
        jitter: Added to ``noise_var`` before inversion.

    Returns:
        Scalar log-likelihood.

    Examples:
        >>> import jax.numpy as jnp
        >>> Y = jnp.zeros((5, 100))
        >>> Z = jnp.ones((5, 2))
        >>> float(collapsed_lfr_log_prob(Y, Z, 0.1)) < 0.0
        True
    """
    _check_rows_match(Y, Z)
    _check_scalar_noise(noise_var)
    N, P = Y.shape
    s2 = noise_var + jitter
    c, low = _psi_factor(Z, s2)
    logdet_psi = 2.0 * jnp.sum(jnp.log(jnp.diagonal(c)))
    ZTY = einx.dot("n q, n p -> q p", Z, Y)
    # Cancellation-free quadratic. The direct Woodbury form
    # ``||Y||^2 / s2 - ||Z^T Y||^2_{Psi^-1} / s2^2`` differences two terms of
    # order ``1 / s2`` whose true difference is order one whenever Y lies
    # close to the span of Z, losing most significant digits (and possibly
    # going negative) at small noise in float32. The equivalent penalized
    # residual form is a sum of non-negative terms:
    #   y^T C^-1 y = ||Y - Z W_map||_F^2 / s2 + ||W_map||_F^2,
    # with W_map the decoder posterior mean. Same O(NQP) cost.
    W_map = cho_solve((c, low), ZTY) / s2
    resid = Y - einx.dot("n q, q p -> n p", Z, W_map)
    quad = jnp.sum(resid * resid) / s2 + jnp.sum(W_map * W_map)
    return -0.5 * (
        N * P * jnp.log(2.0 * jnp.pi) + N * P * jnp.log(s2) + P * logdet_psi + quad
    )

decoder_posterior(Y: Float[Array, 'N P'], Z: Float[Array, 'N Q'], noise_var: Float[Array, ''], *, jitter: float = 1e-12) -> tuple[Float[Array, 'Q P'], Float[Array, 'Q Q']]

Closed-form matrix-normal posterior over the decoder.

\[ p(W \mid Z, Y, \sigma) = \mathcal{MN}_{Q \times P} \big(\sigma^{-2}\Psi^{-1}Z^\top Y,\; \Psi^{-1},\; I_P\big) \]

The row covariance is shared across all \(P\) columns and the column covariance is the identity, so the full \(QP \times QP\) posterior covariance is exactly \(\Psi^{-1} \otimes I_P\) and is never built.

Parameters:

Name Type Description Default
Y Float[Array, 'N P']

Observations of shape (N, P).

required
Z Float[Array, 'N Q']

Latent factor values, shape (N, Q).

required
noise_var Float[Array, '']

Scalar isotropic observation noise variance.

required
jitter float

Added to noise_var before inversion.

1e-12

Returns:

Type Description
tuple[Float[Array, 'Q P'], Float[Array, 'Q Q']]

Tuple of (mean, row_cov) with shapes (Q, P) and (Q, Q).

Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor.py
def decoder_posterior(
    Y: Float[Array, "N P"],
    Z: Float[Array, "N Q"],
    noise_var: Float[Array, ""],
    *,
    jitter: float = 1e-12,
) -> tuple[Float[Array, "Q P"], Float[Array, "Q Q"]]:
    """Closed-form matrix-normal posterior over the decoder.

    $$
    p(W \\mid Z, Y, \\sigma) = \\mathcal{MN}_{Q \\times P}
        \\big(\\sigma^{-2}\\Psi^{-1}Z^\\top Y,\\; \\Psi^{-1},\\; I_P\\big)
    $$

    The row covariance is shared across all $P$ columns and the column
    covariance is the identity, so the full $QP \\times QP$ posterior
    covariance is exactly $\\Psi^{-1} \\otimes I_P$ and is never built.

    Args:
        Y: Observations of shape ``(N, P)``.
        Z: Latent factor values, shape ``(N, Q)``.
        noise_var: Scalar isotropic observation noise variance.
        jitter: Added to ``noise_var`` before inversion.

    Returns:
        Tuple of ``(mean, row_cov)`` with shapes ``(Q, P)`` and ``(Q, Q)``.
    """
    _check_rows_match(Y, Z)
    _check_scalar_noise(noise_var)
    s2 = noise_var + jitter
    c, low = _psi_factor(Z, s2)
    mean = cho_solve((c, low), einx.dot("n q, n p -> q p", Z, Y)) / s2
    row_cov = cho_solve((c, low), jnp.eye(Z.shape[1], dtype=Z.dtype))
    return mean, row_cov

lfr_predictive_moments(z_mean: Float[Array, 'T Q'], z_var: Float[Array, 'T Q'], mu_W: Float[Array, 'Q P'], Sigma_W: Float[Array, 'Q Q'], noise_var: Float[Array, ''] | None = None) -> tuple[Float[Array, 'T P'], Float[Array, 'T P']]

Exact moments of the product of two independent Gaussians.

For \(z \sim \mathcal{N}(m, \mathrm{diag}(v))\) independent of \(W \sim \mathcal{MN}(\mu_W, \Sigma_W, I_P)\), the predictive \(f = z^\top W\) has

\[ \mathrm{Var}[f_j] = m^\top \Sigma_W m + \sum_q v_q \mu_{W,qj}^2 + \sum_q v_q \Sigma_{W,qq} \]

The first and third terms do not depend on the output index, so they are computed once as (T, 1) columns and broadcast over P. Predictive variance for very wide outputs therefore costs barely more than the mean.

Parameters:

Name Type Description Default
z_mean Float[Array, 'T Q']

Latent predictive means, shape (T, Q).

required
z_var Float[Array, 'T Q']

Latent predictive marginal variances, shape (T, Q).

required
mu_W Float[Array, 'Q P']

Decoder posterior mean, shape (Q, P).

required
Sigma_W Float[Array, 'Q Q']

Decoder posterior row covariance, shape (Q, Q).

required
noise_var Float[Array, ''] | None

If given, added to the variance for the observation-noise predictive rather than the signal predictive.

None

Returns:

Type Description
tuple[Float[Array, 'T P'], Float[Array, 'T P']]

Tuple of (mean, variance), both shape (T, P).

Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor.py
def lfr_predictive_moments(
    z_mean: Float[Array, "T Q"],
    z_var: Float[Array, "T Q"],
    mu_W: Float[Array, "Q P"],
    Sigma_W: Float[Array, "Q Q"],
    noise_var: Float[Array, ""] | None = None,
) -> tuple[Float[Array, "T P"], Float[Array, "T P"]]:
    """Exact moments of the product of two independent Gaussians.

    For $z \\sim \\mathcal{N}(m, \\mathrm{diag}(v))$ independent of
    $W \\sim \\mathcal{MN}(\\mu_W, \\Sigma_W, I_P)$, the predictive
    $f = z^\\top W$ has

    $$
    \\mathrm{Var}[f_j] = m^\\top \\Sigma_W m
        + \\sum_q v_q \\mu_{W,qj}^2
        + \\sum_q v_q \\Sigma_{W,qq}
    $$

    The first and third terms do not depend on the output index, so they are
    computed once as ``(T, 1)`` columns and broadcast over ``P``. Predictive
    variance for very wide outputs therefore costs barely more than the mean.

    Args:
        z_mean: Latent predictive means, shape ``(T, Q)``.
        z_var: Latent predictive marginal variances, shape ``(T, Q)``.
        mu_W: Decoder posterior mean, shape ``(Q, P)``.
        Sigma_W: Decoder posterior row covariance, shape ``(Q, Q)``.
        noise_var: If given, added to the variance for the observation-noise
            predictive rather than the signal predictive.

    Returns:
        Tuple of ``(mean, variance)``, both shape ``(T, P)``.
    """
    mean = einx.dot("t q, q p -> t p", z_mean, mu_W)
    decoder = jnp.sum((z_mean @ Sigma_W) * z_mean, axis=-1, keepdims=True)
    latent = einx.dot("t q, q p -> t p", z_var, mu_W**2)
    cross = z_var @ jnp.diagonal(Sigma_W)[:, None]
    var = decoder + latent + cross
    return mean, var if noise_var is None else var + noise_var

warp_to_base(warp: AbstractBijection, Y: Float[Array, 'N P']) -> tuple[Float[Array, 'N P'], Float[Array, '']]

Map observations into the space where the factor model is linear.

The warp has event shape (P,) -- one marginal transform per output channel -- and is applied independently to each of the N rows.

The warp must be elementwise

The event shape check below cannot tell a per-channel bijection from a coupled one (a triangular affine transform has the same (P,) event shape and passes). A coupled warp trains and conditions without complaint, but prediction keeps only per-channel marginal moments and evaluates every channel at the same scalar quadrature node, which silently imposes perfect standardized correlation. Pass a genuinely marginal transform such as gauss_flows.RQSplineMarginal; a coupled one would need the full joint covariance and multidimensional integration.

Parameters:

Name Type Description Default
warp AbstractBijection

Elementwise bijection with event shape (P,) — one marginal transform per output channel. flowjax convention: transform maps base to data, inverse maps data to base, so this uses inverse.

required
Y Float[Array, 'N P']

Observations of shape (N, P).

required

Returns:

Type Description
tuple[Float[Array, 'N P'], Float[Array, '']]

Tuple of (Y_tilde, total_log_det).

Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor.py
def warp_to_base(
    warp: AbstractBijection,
    Y: Float[Array, "N P"],
) -> tuple[Float[Array, "N P"], Float[Array, ""]]:
    """Map observations into the space where the factor model is linear.

    The warp has event shape ``(P,)`` -- one marginal transform per output
    channel -- and is applied independently to each of the ``N`` rows.

    !!! warning "The warp must be elementwise"
        The event shape check below cannot tell a per-channel bijection
        from a *coupled* one (a triangular affine transform has the same
        ``(P,)`` event shape and passes). A coupled warp trains and
        conditions without complaint, but prediction keeps only per-channel
        marginal moments and evaluates every channel at the same scalar
        quadrature node, which silently imposes perfect standardized
        correlation. Pass a genuinely marginal transform such as
        ``gauss_flows.RQSplineMarginal``; a coupled one would need the full
        joint covariance and multidimensional integration.

    Args:
        warp: Elementwise bijection with event shape ``(P,)`` — one marginal
            transform per output channel. flowjax convention:
            ``transform`` maps base to data, ``inverse`` maps data to base,
            so this uses ``inverse``.
        Y: Observations of shape ``(N, P)``.

    Returns:
        Tuple of ``(Y_tilde, total_log_det)``.
    """
    if warp.cond_shape is not None:
        raise ValueError(
            "Conditional warps are not supported here: the log-det Jacobian "
            "would depend on the inputs, and the per-channel marginal "
            "interpretation no longer holds."
        )
    if warp.shape != (Y.shape[1],):
        raise ValueError(
            f"Warp event shape must be (P,) = {(Y.shape[1],)} -- one marginal "
            f"transform per output channel; got {warp.shape}."
        )
    Ytil, log_det = jax.vmap(warp.inverse_and_log_det)(Y)
    return Ytil, jnp.sum(log_det)

warped_lfr_log_prob(Y: Float[Array, 'N P'], Z: Float[Array, 'N Q'], noise_var: Float[Array, ''], warp: AbstractBijection, *, jitter: float = 1e-12) -> Float[Array, '']

Collapsed latent-factor log-likelihood on warped observations.

Models \(G^{-1}(y_j) = Z w_j + \epsilon\), so

\[ \log p(Y) = \log p_{\mathrm{collapsed}}(G^{-1}(Y), Z, \sigma^2) + \sum_{n,j} \log \left| \partial G^{-1} / \partial y \right| \]

The Jacobian term is free of \(Z\), \(W\) and \(\sigma\), so the analytic decoder marginalization of collapsed_lfr_log_prob is unchanged and the cost stays linear in \(P\).

Warp direction is a performance cliff

This evaluates G^{-1} (inverse) on every step, so a warp that is cheap forward and expensive backward costs dearly: MixtureGaussianCDF.inverse runs a bisection solver and is ~40x its own forward cost, ~459x slower end-to-end than the alternative. Prefer gauss_flows.RQSplineMarginal — closed-form in both directions and the exact identity at initialization.

Wrapping a mixture-CDF warp in flowjax.bijections.Invert is cheap, but it is a different model, not a faster route to the same one: this function applies warp.inverse, which is M.inverse for M and M.transform for Invert(M), so the two map Y to different base values and define different likelihoods. Choose it because the flipped map is the warp you want, never as a drop-in speedup.

Parameters:

Name Type Description Default
Y Float[Array, 'N P']

Observations of shape (N, P). Must lie inside the warp's support.

required
Z Float[Array, 'N Q']

Latent factor values at the training inputs, shape (N, Q).

required
noise_var Float[Array, '']

Scalar isotropic observation noise variance, in the warped space.

required
warp AbstractBijection

Bijection with event shape (P,).

required
jitter float

Added to noise_var before inversion.

1e-12

Returns:

Type Description
Float[Array, '']

Scalar log-likelihood.

Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor.py
def warped_lfr_log_prob(
    Y: Float[Array, "N P"],
    Z: Float[Array, "N Q"],
    noise_var: Float[Array, ""],
    warp: AbstractBijection,
    *,
    jitter: float = 1e-12,
) -> Float[Array, ""]:
    r"""Collapsed latent-factor log-likelihood on warped observations.

    Models $G^{-1}(y_j) = Z w_j + \epsilon$, so

    $$
    \log p(Y) = \log p_{\mathrm{collapsed}}(G^{-1}(Y), Z, \sigma^2)
              + \sum_{n,j} \log \left| \partial G^{-1} / \partial y \right|
    $$

    The Jacobian term is free of $Z$, $W$ and $\sigma$, so the analytic
    decoder marginalization of
    [`collapsed_lfr_log_prob`][pyrox_gp.collapsed_lfr_log_prob] is
    unchanged and the cost stays linear in $P$.

    !!! warning "Warp direction is a performance cliff"
        This evaluates ``G^{-1}`` (``inverse``) on every step, so a warp
        that is cheap forward and expensive backward costs dearly:
        ``MixtureGaussianCDF.inverse`` runs a bisection solver and is ~40x
        its own forward cost, ~459x slower end-to-end than the
        alternative. Prefer ``gauss_flows.RQSplineMarginal`` — closed-form
        in both directions and the exact identity at initialization.

        Wrapping a mixture-CDF warp in ``flowjax.bijections.Invert`` is
        cheap, but it is **a different model, not a faster route to the
        same one**: this function applies ``warp.inverse``, which is
        ``M.inverse`` for ``M`` and ``M.transform`` for ``Invert(M)``, so
        the two map ``Y`` to different base values and define different
        likelihoods. Choose it because the flipped map is the warp you
        want, never as a drop-in speedup.

    Args:
        Y: Observations of shape ``(N, P)``. Must lie inside the warp's
            support.
        Z: Latent factor values at the training inputs, shape ``(N, Q)``.
        noise_var: Scalar isotropic observation noise variance, in the
            warped space.
        warp: Bijection with event shape ``(P,)``.
        jitter: Added to ``noise_var`` before inversion.

    Returns:
        Scalar log-likelihood.
    """
    Ytil, log_det = warp_to_base(warp, Y)
    return collapsed_lfr_log_prob(Ytil, Z, noise_var, jitter=jitter) + log_det

warped_decoder_posterior(Y: Float[Array, 'N P'], Z: Float[Array, 'N Q'], noise_var: Float[Array, ''], warp: AbstractBijection, *, jitter: float = 1e-12) -> tuple[Float[Array, 'Q P'], Float[Array, 'Q Q']]

Matrix-normal decoder posterior, read in the warped space.

Identical to decoder_posterior applied to G^{-1}(Y) -- the warp does not change the conjugacy.

Parameters:

Name Type Description Default
Y Float[Array, 'N P']

Observations of shape (N, P).

required
Z Float[Array, 'N Q']

Latent factor values, shape (N, Q).

required
noise_var Float[Array, '']

Scalar isotropic observation noise variance, in the warped space.

required
warp AbstractBijection

Bijection with event shape (P,).

required
jitter float

Added to noise_var before inversion.

1e-12

Returns:

Type Description
tuple[Float[Array, 'Q P'], Float[Array, 'Q Q']]

Tuple of (mean, row_cov) with shapes (Q, P) and (Q, Q).

Source code in packages/pyrox-gp/src/pyrox_gp/_latent_factor.py
def warped_decoder_posterior(
    Y: Float[Array, "N P"],
    Z: Float[Array, "N Q"],
    noise_var: Float[Array, ""],
    warp: AbstractBijection,
    *,
    jitter: float = 1e-12,
) -> tuple[Float[Array, "Q P"], Float[Array, "Q Q"]]:
    """Matrix-normal decoder posterior, read in the warped space.

    Identical to [`decoder_posterior`][pyrox_gp.decoder_posterior] applied
    to ``G^{-1}(Y)`` -- the warp does not change the conjugacy.

    Args:
        Y: Observations of shape ``(N, P)``.
        Z: Latent factor values, shape ``(N, Q)``.
        noise_var: Scalar isotropic observation noise variance, in the
            warped space.
        warp: Bijection with event shape ``(P,)``.
        jitter: Added to ``noise_var`` before inversion.

    Returns:
        Tuple of ``(mean, row_cov)`` with shapes ``(Q, P)`` and ``(Q, Q)``.
    """
    Ytil, _ = warp_to_base(warp, Y)
    return decoder_posterior(Ytil, Z, noise_var, jitter=jitter)

Pathwise posterior samplers (#39)

Callable posterior function draws via Matheron's rule. Each sampled path is a PathwiseFunction that evaluates in O(N_* · F · D + N_* · N_corr) per path — N_* · F · D for the RFF prior draw and N_* · N_corr for the kernel correction against the N_corr training (exact) or inducing (decoupled) points — so the same draw can be reused at arbitrary test sets without rebuilding a test-set covariance. Standard enabler for Thompson sampling, Bayesian optimization, and posterior visualization.

from pyrox_gp import (
    RBF,
    GPPrior,
    PathwiseSampler,
    DecoupledPathwiseSampler,
    FullRankGuide,
    SparseGPPrior,
)
import jax
import jax.numpy as jnp

# Exact GP:
posterior = GPPrior(kernel=RBF(), X=X).condition(y, jnp.array(0.05))
paths = PathwiseSampler(posterior, n_features=512).sample_paths(
    jax.random.PRNGKey(0), n_paths=32
)
draws = paths(X_star)            # (32, N_star)

# Sparse / decoupled:
sparse  = SparseGPPrior(kernel=RBF(), Z=Z)
guide   = FullRankGuide.init(Z.shape[0])
paths   = DecoupledPathwiseSampler(sparse, guide).sample_paths(key, n_paths=16)
samples = paths(X_star)

Currently supports RBF and Matern kernels. Point-inducing SparseGPPrior only — inducing-feature priors raise at construction.

PathwiseSampler

Bases: Module

Exact-GP pathwise posterior sampler using Matheron's rule.

Given a ConditionedGP, draws a zero-mean RFF prior path f_tilde and an iid noise draw eps_tilde at the training inputs, forms the residual y - mu(X) - f_tilde(X) - eps_tilde, solves it against the cached noisy operator K + (jitter + sigma^2)I, and stores the result as posterior correction weights. The returned PathwiseFunction is callable at any X_* in \(\mathcal{O}(N_* \cdot F \cdot D + N_* \cdot N)\) per path, where N is the number of training (correction) points: the RFF prior term recomputes features over X_* each call (N_* · F · D), and the correction term forms a fresh K(X_*, X) block (N_* · N).

Examples:

>>> posterior = GPPrior(kernel=RBF(), X=X).condition(y, jnp.array(0.05))
>>> sampler = PathwiseSampler(posterior, n_features=512)
>>> paths = sampler.sample_paths(key, n_paths=32)
>>> draws = paths(X_star)

Examples:

>>> sampler = PathwiseSampler(posterior, n_features=1024)
>>> thompson = sampler.sample_paths(key, n_paths=1)
>>> values = thompson(X_candidates)
Source code in packages/pyrox-gp/src/pyrox_gp/_pathwise.py
class PathwiseSampler(eqx.Module):
    """Exact-GP pathwise posterior sampler using Matheron's rule.

    Given a `ConditionedGP`, draws a zero-mean RFF prior path
    ``f_tilde`` and an iid noise draw ``eps_tilde`` at the training
    inputs, forms the residual ``y - mu(X) - f_tilde(X) - eps_tilde``,
    solves it against the cached noisy operator ``K + (jitter + sigma^2)I``,
    and stores the result as posterior correction weights. The returned
    `PathwiseFunction` is callable at any ``X_*`` in
    $\\mathcal{O}(N_* \\cdot F \\cdot D + N_* \\cdot N)$ per path,
    where ``N`` is the number of training (correction) points: the RFF
    prior term recomputes features over ``X_*`` each call
    (``N_* · F · D``), and the correction term forms a fresh
    ``K(X_*, X)`` block (``N_* · N``).

    Examples:
        >>> posterior = GPPrior(kernel=RBF(), X=X).condition(y, jnp.array(0.05))
        >>> sampler = PathwiseSampler(posterior, n_features=512)
        >>> paths = sampler.sample_paths(key, n_paths=32)
        >>> draws = paths(X_star)

    Examples:
        >>> sampler = PathwiseSampler(posterior, n_features=1024)
        >>> thompson = sampler.sample_paths(key, n_paths=1)
        >>> values = thompson(X_candidates)
    """

    conditioned_gp: ConditionedGP
    n_features: int = eqx.field(static=True, default=512)

    def sample_paths(self, key: Array, n_paths: int = 1) -> PathwiseFunction:
        """Sample callable posterior paths.

        ``key`` is split into three subkeys: one for the RFF basis,
        one for the iid training-noise draw, and one reserved for
        future extensions.
        """
        rff_key, noise_key, _reserved = jax.random.split(key, 3)

        X = self.conditioned_gp.prior.X
        kernel = self.conditioned_gp.prior.kernel
        # Reuse the resolved (variance, lengthscale) captured by
        # GPPrior.condition under its kernel context. For Pattern B/C
        # kernels the cached operator was built with these exact
        # values; resampling here would put the RFF basis in a
        # different posterior than the cached training solve.
        cached = self.conditioned_gp.resolved_hyperparams
        cached_variance, cached_lengthscale = (
            cached if cached is not None else (None, None)
        )
        variance, lengthscale, omega, phase, feature_weights = draw_rff_cosine_basis(
            kernel,
            rff_key,
            n_paths=n_paths,
            n_features=self.n_features,
            in_features=X.shape[1],
            dtype=X.dtype,
            variance=cached_variance,
            lengthscale=cached_lengthscale,
        )
        prior_train = evaluate_rff_cosine_paths(
            X,
            variance=variance,
            lengthscale=lengthscale,
            omega=omega,
            phase=phase,
            weights=feature_weights,
        )
        mean_train = _broadcast_mean(self.conditioned_gp.prior.mean_fn, X)
        # Matheron requires Cov(eps_tilde) to match the diagonal added to
        # the cached operator. _noisy_operator uses (noise_var + jitter) I,
        # so eps_tilde must have the same variance — otherwise the
        # correction solve is inconsistent and paths are under-dispersed
        # (pronounced when jitter is bumped up for stability).
        noise_var = jnp.asarray(self.conditioned_gp.noise_var, dtype=X.dtype)
        jitter = jnp.asarray(self.conditioned_gp.prior.jitter, dtype=X.dtype)
        eps_var = noise_var + jitter
        noise = jnp.sqrt(eps_var) * jax.random.normal(
            noise_key, shape=(n_paths, X.shape[0]), dtype=X.dtype
        )
        residual = (
            self.conditioned_gp.y[None, :] - (mean_train[None, :] + prior_train) - noise
        )
        # Per-path Matheron correction weights alpha = (K + eps I)^{-1} r
        # via gaussx.solve_rows — structured dispatch on the cached operator.
        correction_weights = solve_rows(self.conditioned_gp.operator, residual)
        return PathwiseFunction(
            kernel_fn=_frozen_kernel_fn(kernel, variance, lengthscale),
            correction_points=X,
            correction_weights=correction_weights,
            omega=omega,
            phase=phase,
            feature_weights=feature_weights,
            variance=variance,
            lengthscale=lengthscale,
            mean_fn=self.conditioned_gp.prior.mean_fn,
        )

    def __call__(
        self,
        key: Array,
        X_star: Float[Array, "N D"],
        n_paths: int = 1,
    ) -> Float[Array, "S N"]:
        """Convenience wrapper for ``sample_paths(key, n_paths)(X_star)``."""
        return self.sample_paths(key, n_paths=n_paths)(X_star)

sample_paths(key: Array, n_paths: int = 1) -> PathwiseFunction

Sample callable posterior paths.

key is split into three subkeys: one for the RFF basis, one for the iid training-noise draw, and one reserved for future extensions.

Source code in packages/pyrox-gp/src/pyrox_gp/_pathwise.py
def sample_paths(self, key: Array, n_paths: int = 1) -> PathwiseFunction:
    """Sample callable posterior paths.

    ``key`` is split into three subkeys: one for the RFF basis,
    one for the iid training-noise draw, and one reserved for
    future extensions.
    """
    rff_key, noise_key, _reserved = jax.random.split(key, 3)

    X = self.conditioned_gp.prior.X
    kernel = self.conditioned_gp.prior.kernel
    # Reuse the resolved (variance, lengthscale) captured by
    # GPPrior.condition under its kernel context. For Pattern B/C
    # kernels the cached operator was built with these exact
    # values; resampling here would put the RFF basis in a
    # different posterior than the cached training solve.
    cached = self.conditioned_gp.resolved_hyperparams
    cached_variance, cached_lengthscale = (
        cached if cached is not None else (None, None)
    )
    variance, lengthscale, omega, phase, feature_weights = draw_rff_cosine_basis(
        kernel,
        rff_key,
        n_paths=n_paths,
        n_features=self.n_features,
        in_features=X.shape[1],
        dtype=X.dtype,
        variance=cached_variance,
        lengthscale=cached_lengthscale,
    )
    prior_train = evaluate_rff_cosine_paths(
        X,
        variance=variance,
        lengthscale=lengthscale,
        omega=omega,
        phase=phase,
        weights=feature_weights,
    )
    mean_train = _broadcast_mean(self.conditioned_gp.prior.mean_fn, X)
    # Matheron requires Cov(eps_tilde) to match the diagonal added to
    # the cached operator. _noisy_operator uses (noise_var + jitter) I,
    # so eps_tilde must have the same variance — otherwise the
    # correction solve is inconsistent and paths are under-dispersed
    # (pronounced when jitter is bumped up for stability).
    noise_var = jnp.asarray(self.conditioned_gp.noise_var, dtype=X.dtype)
    jitter = jnp.asarray(self.conditioned_gp.prior.jitter, dtype=X.dtype)
    eps_var = noise_var + jitter
    noise = jnp.sqrt(eps_var) * jax.random.normal(
        noise_key, shape=(n_paths, X.shape[0]), dtype=X.dtype
    )
    residual = (
        self.conditioned_gp.y[None, :] - (mean_train[None, :] + prior_train) - noise
    )
    # Per-path Matheron correction weights alpha = (K + eps I)^{-1} r
    # via gaussx.solve_rows — structured dispatch on the cached operator.
    correction_weights = solve_rows(self.conditioned_gp.operator, residual)
    return PathwiseFunction(
        kernel_fn=_frozen_kernel_fn(kernel, variance, lengthscale),
        correction_points=X,
        correction_weights=correction_weights,
        omega=omega,
        phase=phase,
        feature_weights=feature_weights,
        variance=variance,
        lengthscale=lengthscale,
        mean_fn=self.conditioned_gp.prior.mean_fn,
    )

DecoupledPathwiseSampler

Bases: Module

Sparse/decoupled pathwise sampler with RFF prior + inducing update.

The prior draw uses random features while the correction is represented in the inducing-point basis, so each sampled path stays callable at arbitrary inputs after a one-time inducing solve.

Supported for point-inducing SparseGPPrior (Z=...); inducing-feature priors (inducing=...) are rejected at construction with a clear error.

Handles WhitenedGuide automatically: whitened guide draws v ~ q(v) are unwhitened to inducing values u = L_ZZ v via gaussx.unwhiten before forming the inducing-space residual.

Examples:

>>> prior = SparseGPPrior(kernel=RBF(), Z=Z)
>>> guide = FullRankGuide.init(Z.shape[0])
>>> sampler = DecoupledPathwiseSampler(prior, guide, n_features=512)
>>> paths = sampler.sample_paths(key, n_paths=16)
>>> draws = paths(X_star)
Source code in packages/pyrox-gp/src/pyrox_gp/_pathwise.py
class DecoupledPathwiseSampler(eqx.Module):
    """Sparse/decoupled pathwise sampler with RFF prior + inducing update.

    The prior draw uses random features while the correction is represented in
    the inducing-point basis, so each sampled path stays callable at arbitrary
    inputs after a one-time inducing solve.

    Supported for point-inducing `SparseGPPrior` (``Z=...``);
    inducing-feature priors (``inducing=...``) are rejected at
    construction with a clear error.

    Handles `WhitenedGuide` automatically: whitened guide draws
    ``v ~ q(v)`` are unwhitened to inducing values ``u = L_ZZ v`` via
    `gaussx.unwhiten` before forming the inducing-space residual.

    Examples:
        >>> prior = SparseGPPrior(kernel=RBF(), Z=Z)
        >>> guide = FullRankGuide.init(Z.shape[0])
        >>> sampler = DecoupledPathwiseSampler(prior, guide, n_features=512)
        >>> paths = sampler.sample_paths(key, n_paths=16)
        >>> draws = paths(X_star)
    """

    prior: SparseGPPrior
    guide: Guide
    n_features: int = eqx.field(static=True, default=512)

    def __check_init__(self) -> None:
        if self.prior.Z is None:
            raise ValueError(
                "DecoupledPathwiseSampler currently requires a point-inducing "
                "SparseGPPrior constructed with `Z=...`. Inducing-feature "
                "priors (FourierInducingFeatures, SphericalHarmonicInducingFeatures, "
                "LaplacianInducingFeatures) are not yet supported."
            )

    def sample_paths(self, key: Array, n_paths: int = 1) -> PathwiseFunction:
        """Sample callable sparse posterior paths.

        ``key`` is split into three subkeys: one for the RFF basis, one
        for ``n_paths`` independent guide draws, and one for the
        jitter-augmentation of the prior inducing draw. The RFF basis
        draw and the $K_{zz}$ assembly share a single
        ``_kernel_context`` so kernels with hyperparameter priors
        (Pattern B / C) sample ``(variance, lengthscale)`` once.

        The Matheron correction needs ``Cov(u_tilde) = K_{zz} + \
        \\text{jitter}\\,I`` so it matches the operator that the
        correction is solved against. The bare RFF draw at ``Z``
        produces only the ``K_{zz}`` part; we add an iid Gaussian
        with variance ``jitter`` per inducing index to close the gap —
        without this, paths are under-dispersed when jitter is bumped
        up for stability.
        """
        rff_key, guide_key, jitter_key = jax.random.split(key, 3)

        Z = self.prior.Z
        assert Z is not None  # __check_init__ guarantees
        with _kernel_context(self.prior.kernel):
            basis = draw_rff_cosine_basis(
                self.prior.kernel,
                rff_key,
                n_paths=n_paths,
                n_features=self.n_features,
                in_features=Z.shape[1],
                dtype=Z.dtype,
            )
            variance, lengthscale, omega, phase, feature_weights = basis
            prior_inducing = evaluate_rff_cosine_paths(
                Z,
                variance=variance,
                lengthscale=lengthscale,
                omega=omega,
                phase=phase,
                weights=feature_weights,
            )
            inducing_op = self.prior.inducing_operator()

        # See docstring: u_tilde must have covariance K_zz + jitter I.
        jitter = jnp.asarray(self.prior.jitter, dtype=Z.dtype)
        prior_inducing = prior_inducing + jnp.sqrt(jitter) * jax.random.normal(
            jitter_key, shape=prior_inducing.shape, dtype=Z.dtype
        )

        guide_keys = jax.random.split(guide_key, n_paths)
        guide_samples = jax.vmap(self.guide.sample)(guide_keys)
        if isinstance(self.guide, WhitenedGuide):
            inducing_chol = cholesky(inducing_op)
            inducing_samples = jax.vmap(lambda sample: unwhiten(sample, inducing_chol))(
                guide_samples
            )
        else:
            inducing_samples = guide_samples

        # Per-path Matheron correction weights via gaussx.solve_rows.
        correction_weights = solve_rows(
            inducing_op,
            inducing_samples - prior_inducing,
        )
        return PathwiseFunction(
            kernel_fn=_frozen_kernel_fn(self.prior.kernel, variance, lengthscale),
            correction_points=Z,
            correction_weights=correction_weights,
            omega=omega,
            phase=phase,
            feature_weights=feature_weights,
            variance=variance,
            lengthscale=lengthscale,
            mean_fn=self.prior.mean_fn,
        )

    def __call__(
        self,
        key: Array,
        X_star: Float[Array, "N D"],
        n_paths: int = 1,
    ) -> Float[Array, "S N"]:
        """Convenience wrapper for ``sample_paths(key, n_paths)(X_star)``."""
        return self.sample_paths(key, n_paths=n_paths)(X_star)

sample_paths(key: Array, n_paths: int = 1) -> PathwiseFunction

Sample callable sparse posterior paths.

key is split into three subkeys: one for the RFF basis, one for n_paths independent guide draws, and one for the jitter-augmentation of the prior inducing draw. The RFF basis draw and the \(K_{zz}\) assembly share a single _kernel_context so kernels with hyperparameter priors (Pattern B / C) sample (variance, lengthscale) once.

The Matheron correction needs Cov(u_tilde) = K_{zz} + \text{jitter}\,I so it matches the operator that the correction is solved against. The bare RFF draw at Z produces only the K_{zz} part; we add an iid Gaussian with variance jitter per inducing index to close the gap — without this, paths are under-dispersed when jitter is bumped up for stability.

Source code in packages/pyrox-gp/src/pyrox_gp/_pathwise.py
def sample_paths(self, key: Array, n_paths: int = 1) -> PathwiseFunction:
    """Sample callable sparse posterior paths.

    ``key`` is split into three subkeys: one for the RFF basis, one
    for ``n_paths`` independent guide draws, and one for the
    jitter-augmentation of the prior inducing draw. The RFF basis
    draw and the $K_{zz}$ assembly share a single
    ``_kernel_context`` so kernels with hyperparameter priors
    (Pattern B / C) sample ``(variance, lengthscale)`` once.

    The Matheron correction needs ``Cov(u_tilde) = K_{zz} + \
    \\text{jitter}\\,I`` so it matches the operator that the
    correction is solved against. The bare RFF draw at ``Z``
    produces only the ``K_{zz}`` part; we add an iid Gaussian
    with variance ``jitter`` per inducing index to close the gap —
    without this, paths are under-dispersed when jitter is bumped
    up for stability.
    """
    rff_key, guide_key, jitter_key = jax.random.split(key, 3)

    Z = self.prior.Z
    assert Z is not None  # __check_init__ guarantees
    with _kernel_context(self.prior.kernel):
        basis = draw_rff_cosine_basis(
            self.prior.kernel,
            rff_key,
            n_paths=n_paths,
            n_features=self.n_features,
            in_features=Z.shape[1],
            dtype=Z.dtype,
        )
        variance, lengthscale, omega, phase, feature_weights = basis
        prior_inducing = evaluate_rff_cosine_paths(
            Z,
            variance=variance,
            lengthscale=lengthscale,
            omega=omega,
            phase=phase,
            weights=feature_weights,
        )
        inducing_op = self.prior.inducing_operator()

    # See docstring: u_tilde must have covariance K_zz + jitter I.
    jitter = jnp.asarray(self.prior.jitter, dtype=Z.dtype)
    prior_inducing = prior_inducing + jnp.sqrt(jitter) * jax.random.normal(
        jitter_key, shape=prior_inducing.shape, dtype=Z.dtype
    )

    guide_keys = jax.random.split(guide_key, n_paths)
    guide_samples = jax.vmap(self.guide.sample)(guide_keys)
    if isinstance(self.guide, WhitenedGuide):
        inducing_chol = cholesky(inducing_op)
        inducing_samples = jax.vmap(lambda sample: unwhiten(sample, inducing_chol))(
            guide_samples
        )
    else:
        inducing_samples = guide_samples

    # Per-path Matheron correction weights via gaussx.solve_rows.
    correction_weights = solve_rows(
        inducing_op,
        inducing_samples - prior_inducing,
    )
    return PathwiseFunction(
        kernel_fn=_frozen_kernel_fn(self.prior.kernel, variance, lengthscale),
        correction_points=Z,
        correction_weights=correction_weights,
        omega=omega,
        phase=phase,
        feature_weights=feature_weights,
        variance=variance,
        lengthscale=lengthscale,
        mean_fn=self.prior.mean_fn,
    )

PathwiseFunction

Bases: Module

Callable posterior function draw(s) produced by a pathwise sampler.

Carries the random-feature prior basis (omega, phase, feature_weights) and the posterior correction weights evaluated against either the training inputs (exact) or the inducing inputs (sparse). Calling the instance on test points X_star evaluates

\[ f_{\text{post}}(x_*) = \tilde{f}(x_*) + K(x_*,\, X_{\mathrm{corr}})\,\alpha + \mu(x_*), \]

where \(\tilde f\) is the stored RFF prior draw and \(X_{\mathrm{corr}}\) is either the training set (exact) or the inducing set (sparse).

The kernel enters only as a frozen (X1, X2) -> K callable with the sample-time variance and lengthscale baked in, so repeated evaluations stay consistent with the original RFF draw even for Pattern B/C kernels that register hyperparameter priors.

Examples:

>>> prior = GPPrior(kernel=RBF(), X=X)
>>> posterior = prior.condition(y, noise_var=jnp.array(0.05))
>>> sampler = PathwiseSampler(posterior, n_features=512)
>>> paths = sampler.sample_paths(key, n_paths=8)
>>> samples = paths(X_star)

Examples:

>>> sparse_prior = SparseGPPrior(kernel=RBF(), Z=Z)
>>> guide = FullRankGuide.init(Z.shape[0])
>>> paths = DecoupledPathwiseSampler(sparse_prior, guide).sample_paths(key)
>>> thompson_values = paths(X_candidates)
Source code in packages/pyrox-gp/src/pyrox_gp/_pathwise.py
class PathwiseFunction(eqx.Module):
    """Callable posterior function draw(s) produced by a pathwise sampler.

    Carries the random-feature prior basis (``omega``, ``phase``,
    ``feature_weights``) and the posterior correction weights evaluated
    against either the training inputs (exact) or the inducing inputs
    (sparse). Calling the instance on test points ``X_star`` evaluates

    $$
    f_{\\text{post}}(x_*) =
        \\tilde{f}(x_*)
        + K(x_*,\\, X_{\\mathrm{corr}})\\,\\alpha
        + \\mu(x_*),
    $$

    where $\\tilde f$ is the stored RFF prior draw and
    $X_{\\mathrm{corr}}$ is either the training set (exact) or
    the inducing set (sparse).

    The kernel enters only as a frozen ``(X1, X2) -> K`` callable with
    the sample-time ``variance`` and ``lengthscale`` baked in, so
    repeated evaluations stay consistent with the original RFF draw
    even for Pattern B/C kernels that register hyperparameter priors.

    Examples:
        >>> prior = GPPrior(kernel=RBF(), X=X)
        >>> posterior = prior.condition(y, noise_var=jnp.array(0.05))
        >>> sampler = PathwiseSampler(posterior, n_features=512)
        >>> paths = sampler.sample_paths(key, n_paths=8)
        >>> samples = paths(X_star)

    Examples:
        >>> sparse_prior = SparseGPPrior(kernel=RBF(), Z=Z)
        >>> guide = FullRankGuide.init(Z.shape[0])
        >>> paths = DecoupledPathwiseSampler(sparse_prior, guide).sample_paths(key)
        >>> thompson_values = paths(X_candidates)
    """

    kernel_fn: Callable[
        [Float[Array, "N1 D"], Float[Array, "N2 D"]], Float[Array, "N1 N2"]
    ]
    correction_points: Float[Array, "R D"]
    correction_weights: Float[Array, "S R"]
    omega: Float[Array, "S D F"]
    phase: Float[Array, "S F"]
    feature_weights: Float[Array, "S F"]
    variance: Float[Array, ""]
    lengthscale: Float[Array, ""]
    mean_fn: Callable[[Float[Array, "N D"]], Float[Array, " N"]] | None = None

    def __call__(self, X_star: Float[Array, "N D"]) -> Float[Array, "S N"]:
        """Evaluate the sampled function(s) at arbitrary inputs ``X_star``."""
        prior = evaluate_rff_cosine_paths(
            X_star,
            variance=self.variance,
            lengthscale=self.lengthscale,
            omega=self.omega,
            phase=self.phase,
            weights=self.feature_weights,
        )
        K_cross = self.kernel_fn(X_star, self.correction_points)
        # Matheron correction: contract the correction-point axis r,
        # broadcasting over sample paths s → (S, N).
        update = einx.dot("n r, s r -> s n", K_cross, self.correction_weights)
        mean = _broadcast_mean(self.mean_fn, X_star)
        return prior + update + mean[None, :]

State-space (SDE) kernels

Stationary 1-D kernels expressed as linear time-invariant SDEs. Once in state-space form, GP inference on a 1-D grid reduces to Kalman filtering in O(N d^3) instead of O(N^3) Cholesky. The protocol exposes sde_params() -> (F, L, H, Q_c, P_inf) and discretise(dt) -> (A_k, Q_k) for downstream Kalman / RTS use.

import jax.numpy as jnp
from pyrox_gp import (
    ConstantSDE, CosineSDE, MaternSDE, PeriodicSDE,
    ProductSDE, QuasiPeriodicSDE, SumSDE,
)

# Primitive kernels
matern = MaternSDE(variance=1.0, lengthscale=0.5, order=1)  # nu = 3/2
cos    = CosineSDE(variance=1.0, frequency=2.0)
const  = ConstantSDE(variance=0.3)
per    = PeriodicSDE(variance=1.0, lengthscale=1.0, period=2.0, n_harmonics=7)

# Composition: trend + offset
trend = SumSDE((matern, const))                   # state dim = 2 + 1 = 3

# Composition: damped oscillation (Matern x Cosine)
damped = ProductSDE(matern, cos)                  # state dim = 2 * 2 = 4

# Quasi-periodic (Matern x Periodic) — convenience wrapper around ProductSDE
qp = QuasiPeriodicSDE(matern, per)                # state dim = 2 * 15 = 30

Non-stationary priors

Every kernel above is stationary: the process is started in its stationary distribution, so the filter seeds from P_inf and sde_autocovariance reports K(tau).

IntegratedWienerSDE is not. At order=1 it is the local linear trend — smooth, and linear unless the data push back, with no lengthscale to choose — and its marginal variance grows without bound, so there is no P_inf at all. It reports stationary = False and sde_params().P_inf = None, and the filter starts from an explicit initial covariance instead:

import jax.numpy as jnp
from pyrox_gp import IntegratedWienerSDE, MarkovGPPrior

trend = IntegratedWienerSDE(diffusion=1e-4)       # state = [level, slope]
prior = MarkovGPPrior(trend, times)               # diffuse default P_0
# ... or state what is known before the first observation:
prior = MarkovGPPrior(trend, times, init_cov=jnp.diag(jnp.array([1e2, 1.0])))

Filtering, smoothing, prediction and the site-based non-Gaussian strategies all work; the dense paths (log_prob and markov_gp_sample) raise, because K_ij = H exp(F |t_i - t_j|) P_inf H^T is a function of the lag alone, while a non-stationary covariance depends on both times.

A SumSDE may mix the two — a trend plus a stationary seasonal — and starts from the block-diagonal of its components' initial covariances.

SDEKernel

Bases: Module

Abstract base class for state-space kernel representations.

Subclasses implement sde_params to provide the continuous-time SDE matrices (F, L, H, Q_c, P_inf). The default discretise uses the matrix exponential for discretization; subclasses may override with closed-form solutions.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_sde_kernel.py
class SDEKernel(eqx.Module):
    """Abstract base class for state-space kernel representations.

    Subclasses implement `sde_params` to provide the continuous-time
    SDE matrices ``(F, L, H, Q_c, P_inf)``. The default `discretise`
    uses the matrix exponential for discretization; subclasses may override
    with closed-form solutions.
    """

    @property
    @abc.abstractmethod
    def state_dim(self) -> int:
        """Dimension of the latent state vector."""
        ...

    @property
    def stationary(self) -> bool:
        """Whether the process has a stationary distribution.

        ``True`` for every kernel in the zoo, which is why that is the
        default; a non-stationary kernel such as
        `gaussx.IntegratedWienerSDE` overrides it. Consumers should
        branch on this rather than on ``sde_params().P_inf is None``:
        the two are different questions, since a stationary kernel may
        report ``P_inf=None`` when it has no *closed form* for it (a
        learned drift, say).
        """
        return True

    @abc.abstractmethod
    def sde_params(self) -> SDEParams:
        """Return continuous-time SDE parameters."""
        ...

    def initial_covariance(self) -> Float[Array, "d d"]:
        r"""Return the covariance of the state at the first time point.

        For a stationary kernel the process is assumed started in its
        stationary distribution, so this is $P_\infty$ — the default
        implementation returns it and existing kernels need no change.
        A non-stationary kernel has no such limit to start from and must
        override this with an explicit choice.

        Returns:
            Initial state covariance, shape ``(d, d)``.

        Raises:
            ValueError: If ``sde_params().P_inf`` is ``None``, so there
                is no stationary covariance to fall back on.
        """
        P_inf = self.sde_params().P_inf
        if P_inf is None:
            msg = (
                f"{type(self).__name__} has no initial covariance: it "
                f"reports P_inf=None, and the default initial covariance "
                f"is the stationary one. Override initial_covariance() "
                f"with an explicit choice (a diffuse prior, typically), "
                f"or give the kernel a P_inf."
            )
            raise ValueError(msg)
        return P_inf

    def discretise(
        self,
        dt: Float[Array, ""],
    ) -> tuple[Float[Array, "d d"], Float[Array, "d d"]]:
        """Discretise the SDE at time step ``dt``.

        Default implementation computes:

            A = expm(F * dt)
            Q = P_inf - A @ P_inf @ A^T

        When ``sde_params()`` returns ``P_inf=None`` — a kernel with no
        closed-form stationary covariance, such as one whose drift is a
        learned parameter — this falls back to `gaussx.discretise_mfd`,
        which recovers both ``A`` and ``Q`` from one matrix exponential and
        is well defined for every ``F``. The fallback is chosen at trace
        time from a static ``None`` check, so kernels that do supply
        ``P_inf`` keep the stationary route and its precision exactly.

        Subclasses may override with closed-form expressions.

        Args:
            dt: Time step (scalar, non-negative).

        Returns:
            Tuple ``(A, Q)`` where A is the transition matrix and
            Q is the process noise covariance.
        """
        params = self.sde_params()
        if params.P_inf is None:
            diffusion = params.L @ params.Q_c @ params.L.T
            return discretise_mfd(params.F, diffusion, dt)
        A = jsl.expm(params.F * dt)
        Q = symmetrize(process_noise_covariance(A, params.P_inf))
        return A, Q

    def discretise_sequence(
        self,
        dt: Float[Array, " N"],
    ) -> tuple[Float[Array, "N d d"], Float[Array, "N d d"]]:
        """Discretise the SDE at multiple time steps.

        Args:
            dt: Time steps, shape ``(N,)``.

        Returns:
            Tuple ``(A_seq, Q_seq)`` with shapes ``(N, d, d)``.
        """
        return jax.vmap(self.discretise)(dt)

state_dim: int abstractmethod property

Dimension of the latent state vector.

stationary: bool property

Whether the process has a stationary distribution.

True for every kernel in the zoo, which is why that is the default; a non-stationary kernel such as gaussx.IntegratedWienerSDE overrides it. Consumers should branch on this rather than on sde_params().P_inf is None: the two are different questions, since a stationary kernel may report P_inf=None when it has no closed form for it (a learned drift, say).

sde_params() -> SDEParams abstractmethod

Return continuous-time SDE parameters.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_sde_kernel.py
@abc.abstractmethod
def sde_params(self) -> SDEParams:
    """Return continuous-time SDE parameters."""
    ...

initial_covariance() -> Float[Array, 'd d']

Return the covariance of the state at the first time point.

For a stationary kernel the process is assumed started in its stationary distribution, so this is \(P_\infty\) — the default implementation returns it and existing kernels need no change. A non-stationary kernel has no such limit to start from and must override this with an explicit choice.

Returns:

Type Description
Float[Array, 'd d']

Initial state covariance, shape (d, d).

Raises:

Type Description
ValueError

If sde_params().P_inf is None, so there is no stationary covariance to fall back on.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_sde_kernel.py
def initial_covariance(self) -> Float[Array, "d d"]:
    r"""Return the covariance of the state at the first time point.

    For a stationary kernel the process is assumed started in its
    stationary distribution, so this is $P_\infty$ — the default
    implementation returns it and existing kernels need no change.
    A non-stationary kernel has no such limit to start from and must
    override this with an explicit choice.

    Returns:
        Initial state covariance, shape ``(d, d)``.

    Raises:
        ValueError: If ``sde_params().P_inf`` is ``None``, so there
            is no stationary covariance to fall back on.
    """
    P_inf = self.sde_params().P_inf
    if P_inf is None:
        msg = (
            f"{type(self).__name__} has no initial covariance: it "
            f"reports P_inf=None, and the default initial covariance "
            f"is the stationary one. Override initial_covariance() "
            f"with an explicit choice (a diffuse prior, typically), "
            f"or give the kernel a P_inf."
        )
        raise ValueError(msg)
    return P_inf

discretise(dt: Float[Array, '']) -> tuple[Float[Array, 'd d'], Float[Array, 'd d']]

Discretise the SDE at time step dt.

Default implementation computes:

A = expm(F * dt)
Q = P_inf - A @ P_inf @ A^T

When sde_params() returns P_inf=None — a kernel with no closed-form stationary covariance, such as one whose drift is a learned parameter — this falls back to gaussx.discretise_mfd, which recovers both A and Q from one matrix exponential and is well defined for every F. The fallback is chosen at trace time from a static None check, so kernels that do supply P_inf keep the stationary route and its precision exactly.

Subclasses may override with closed-form expressions.

Parameters:

Name Type Description Default
dt Float[Array, '']

Time step (scalar, non-negative).

required

Returns:

Type Description
Float[Array, 'd d']

Tuple (A, Q) where A is the transition matrix and

Float[Array, 'd d']

Q is the process noise covariance.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_sde_kernel.py
def discretise(
    self,
    dt: Float[Array, ""],
) -> tuple[Float[Array, "d d"], Float[Array, "d d"]]:
    """Discretise the SDE at time step ``dt``.

    Default implementation computes:

        A = expm(F * dt)
        Q = P_inf - A @ P_inf @ A^T

    When ``sde_params()`` returns ``P_inf=None`` — a kernel with no
    closed-form stationary covariance, such as one whose drift is a
    learned parameter — this falls back to `gaussx.discretise_mfd`,
    which recovers both ``A`` and ``Q`` from one matrix exponential and
    is well defined for every ``F``. The fallback is chosen at trace
    time from a static ``None`` check, so kernels that do supply
    ``P_inf`` keep the stationary route and its precision exactly.

    Subclasses may override with closed-form expressions.

    Args:
        dt: Time step (scalar, non-negative).

    Returns:
        Tuple ``(A, Q)`` where A is the transition matrix and
        Q is the process noise covariance.
    """
    params = self.sde_params()
    if params.P_inf is None:
        diffusion = params.L @ params.Q_c @ params.L.T
        return discretise_mfd(params.F, diffusion, dt)
    A = jsl.expm(params.F * dt)
    Q = symmetrize(process_noise_covariance(A, params.P_inf))
    return A, Q

discretise_sequence(dt: Float[Array, ' N']) -> tuple[Float[Array, 'N d d'], Float[Array, 'N d d']]

Discretise the SDE at multiple time steps.

Parameters:

Name Type Description Default
dt Float[Array, ' N']

Time steps, shape (N,).

required

Returns:

Type Description
tuple[Float[Array, 'N d d'], Float[Array, 'N d d']]

Tuple (A_seq, Q_seq) with shapes (N, d, d).

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_sde_kernel.py
def discretise_sequence(
    self,
    dt: Float[Array, " N"],
) -> tuple[Float[Array, "N d d"], Float[Array, "N d d"]]:
    """Discretise the SDE at multiple time steps.

    Args:
        dt: Time steps, shape ``(N,)``.

    Returns:
        Tuple ``(A_seq, Q_seq)`` with shapes ``(N, d, d)``.
    """
    return jax.vmap(self.discretise)(dt)

SDEParams

Bases: NamedTuple

Continuous-time SDE parameters for a linear SDE.

Defines the linear time-invariant SDE:

dx = F x dt + L dW,   W ~ N(0, Q_c dt)

with observation model y = H x.

Attributes:

Name Type Description
F Float[Array, 'd d']

Drift matrix, shape (d, d).

L Float[Array, 'd s']

Diffusion matrix, shape (d, s).

H Float[Array, '1 d']

Observation matrix, shape (1, d).

Q_c Float[Array, 's s']

Spectral density, shape (s, s).

P_inf Float[Array, 'd d'] | None

Stationary covariance, shape (d, d), or None when the kernel has no closed-form stationary covariance — as for a learned drift matrix, or a non-stationary kernel such as gaussx.IntegratedWienerSDE, which has no stationary covariance at all. SDEKernel.discretise then falls back to gaussx.discretise_mfd, which needs no P_inf, and the filter is started from SDEKernel.initial_covariance instead.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_sde_kernel.py
class SDEParams(NamedTuple):
    """Continuous-time SDE parameters for a linear SDE.

    Defines the linear time-invariant SDE:

        dx = F x dt + L dW,   W ~ N(0, Q_c dt)

    with observation model ``y = H x``.

    Attributes:
        F: Drift matrix, shape ``(d, d)``.
        L: Diffusion matrix, shape ``(d, s)``.
        H: Observation matrix, shape ``(1, d)``.
        Q_c: Spectral density, shape ``(s, s)``.
        P_inf: Stationary covariance, shape ``(d, d)``, or ``None`` when
            the kernel has no closed-form stationary covariance — as for a
            learned drift matrix, or a non-stationary kernel such as
            `gaussx.IntegratedWienerSDE`, which has no stationary
            covariance at all. `SDEKernel.discretise` then falls back to
            `gaussx.discretise_mfd`, which needs no ``P_inf``, and the
            filter is started from `SDEKernel.initial_covariance`
            instead.
    """

    F: Float[Array, "d d"]
    L: Float[Array, "d s"]
    H: Float[Array, "1 d"]
    Q_c: Float[Array, "s s"]
    P_inf: Float[Array, "d d"] | None = None

MaternSDE

Bases: SDEKernel

State-space representation of the Matern kernel.

Supports orders 0 (Matern-1/2), 1 (Matern-3/2), and 2 (Matern-5/2). The state dimension is order + 1.

Attributes:

Name Type Description
variance Float[Array, '']

Signal variance \(\sigma^2\).

lengthscale Float[Array, '']

Lengthscale \(\ell\).

order int

Matern order (0, 1, or 2).

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_matern.py
class MaternSDE(SDEKernel):
    r"""State-space representation of the Matern kernel.

    Supports orders 0 (Matern-1/2), 1 (Matern-3/2), and 2 (Matern-5/2).
    The state dimension is ``order + 1``.

    Attributes:
        variance: Signal variance $\sigma^2$.
        lengthscale: Lengthscale $\ell$.
        order: Matern order (0, 1, or 2).
    """

    variance: Float[Array, ""]
    lengthscale: Float[Array, ""]
    order: int = eqx.field(static=True)

    @property
    def state_dim(self) -> int:
        return self.order + 1

    def sde_params(self) -> SDEParams:
        """Compute SDE parameters for the Matern kernel."""
        if self.order == 0:
            return self._matern12()
        elif self.order == 1:
            return self._matern32()
        elif self.order == 2:
            return self._matern52()
        else:
            msg = f"Unsupported Matern order {self.order}; must be 0, 1, or 2"
            raise ValueError(msg)

    def _matern12(self) -> SDEParams:
        # Constant blocks follow the hyperparameter dtype; an untyped
        # ``jnp.array`` of Python floats is float64 under x64 (gh-224).
        dtype = jnp.result_type(self.variance, self.lengthscale)
        lam = 1.0 / self.lengthscale
        F = jnp.array([[-lam]])
        L = jnp.array([[1.0]], dtype=dtype)
        H = jnp.array([[1.0]], dtype=dtype)
        Q_c = jnp.array([[2.0 * lam * self.variance]])
        P_inf = jnp.array([[self.variance]])
        return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=P_inf)

    def _matern32(self) -> SDEParams:
        dtype = jnp.result_type(self.variance, self.lengthscale)
        lam = jnp.sqrt(3.0) / self.lengthscale
        F = jnp.array([[0.0, 1.0], [-(lam**2), -2.0 * lam]])
        L = jnp.array([[0.0], [1.0]], dtype=dtype)
        H = jnp.array([[1.0, 0.0]], dtype=dtype)
        q = 4.0 * lam**3 * self.variance
        Q_c = jnp.array([[q]])
        P_inf = jnp.array([[self.variance, 0.0], [0.0, lam**2 * self.variance]])
        return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=P_inf)

    def _matern52(self) -> SDEParams:
        dtype = jnp.result_type(self.variance, self.lengthscale)
        lam = jnp.sqrt(5.0) / self.lengthscale
        F = jnp.array(
            [
                [0.0, 1.0, 0.0],
                [0.0, 0.0, 1.0],
                [-(lam**3), -3.0 * lam**2, -3.0 * lam],
            ]
        )
        L = jnp.array([[0.0], [0.0], [1.0]], dtype=dtype)
        H = jnp.array([[1.0, 0.0, 0.0]], dtype=dtype)
        kappa = 5.0 / 3.0 * self.variance / self.lengthscale**2
        q = 16.0 / 3.0 * lam**5 * self.variance
        Q_c = jnp.array([[q]])
        P_inf = jnp.array(
            [
                [self.variance, 0.0, -kappa],
                [0.0, kappa, 0.0],
                [-kappa, 0.0, lam**4 * self.variance],
            ]
        )
        return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=P_inf)

sde_params() -> SDEParams

Compute SDE parameters for the Matern kernel.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_matern.py
def sde_params(self) -> SDEParams:
    """Compute SDE parameters for the Matern kernel."""
    if self.order == 0:
        return self._matern12()
    elif self.order == 1:
        return self._matern32()
    elif self.order == 2:
        return self._matern52()
    else:
        msg = f"Unsupported Matern order {self.order}; must be 0, 1, or 2"
        raise ValueError(msg)

ConstantSDE

Bases: SDEKernel

State-space representation of a constant kernel.

Models \(k(\tau) = \sigma^2\) — a degenerate kernel with zero dynamics and zero diffusion. State dimension is 1.

Attributes:

Name Type Description
variance Float[Array, '']

Signal variance \(\sigma^2\).

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_constant.py
class ConstantSDE(SDEKernel):
    r"""State-space representation of a constant kernel.

    Models $k(\tau) = \sigma^2$ — a degenerate kernel with zero
    dynamics and zero diffusion. State dimension is 1.

    Attributes:
        variance: Signal variance $\sigma^2$.
    """

    variance: Float[Array, ""]

    @property
    def state_dim(self) -> int:
        return 1

    def sde_params(self) -> SDEParams:
        """Return SDE parameters for the constant kernel."""
        # Constant blocks follow the hyperparameter dtype; untyped
        # ``jnp.zeros``/``jnp.array`` are float64 under x64 (gh-224).
        dtype = jnp.result_type(self.variance)
        F = jnp.zeros((1, 1), dtype=dtype)
        L = jnp.zeros((1, 1), dtype=dtype)
        H = jnp.array([[1.0]], dtype=dtype)
        Q_c = jnp.zeros((1, 1), dtype=dtype)
        P_inf = jnp.array([[self.variance]])
        return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=P_inf)

    def discretise(
        self,
        dt: Float[Array, ""],
    ) -> tuple[Float[Array, "d d"], Float[Array, "d d"]]:
        """Closed-form: A = I, Q = 0 (no dynamics)."""
        dtype = jnp.result_type(self.variance)
        A = jnp.eye(1, dtype=dtype)
        Q = jnp.zeros((1, 1), dtype=dtype)
        return A, Q

sde_params() -> SDEParams

Return SDE parameters for the constant kernel.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_constant.py
def sde_params(self) -> SDEParams:
    """Return SDE parameters for the constant kernel."""
    # Constant blocks follow the hyperparameter dtype; untyped
    # ``jnp.zeros``/``jnp.array`` are float64 under x64 (gh-224).
    dtype = jnp.result_type(self.variance)
    F = jnp.zeros((1, 1), dtype=dtype)
    L = jnp.zeros((1, 1), dtype=dtype)
    H = jnp.array([[1.0]], dtype=dtype)
    Q_c = jnp.zeros((1, 1), dtype=dtype)
    P_inf = jnp.array([[self.variance]])
    return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=P_inf)

discretise(dt: Float[Array, '']) -> tuple[Float[Array, 'd d'], Float[Array, 'd d']]

Closed-form: A = I, Q = 0 (no dynamics).

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_constant.py
def discretise(
    self,
    dt: Float[Array, ""],
) -> tuple[Float[Array, "d d"], Float[Array, "d d"]]:
    """Closed-form: A = I, Q = 0 (no dynamics)."""
    dtype = jnp.result_type(self.variance)
    A = jnp.eye(1, dtype=dtype)
    Q = jnp.zeros((1, 1), dtype=dtype)
    return A, Q

CosineSDE

Bases: SDEKernel

State-space representation of the cosine kernel.

Models \(k(\tau) = \sigma^2 \cos(\omega_0 \tau)\) via a 2-D rotation SDE. State dimension is 2.

Attributes:

Name Type Description
variance Float[Array, '']

Signal variance \(\sigma^2\).

frequency Float[Array, '']

Angular frequency \(\omega_0\).

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_periodic.py
class CosineSDE(SDEKernel):
    r"""State-space representation of the cosine kernel.

    Models $k(\tau) = \sigma^2 \cos(\omega_0 \tau)$ via a 2-D
    rotation SDE. State dimension is 2.

    Attributes:
        variance: Signal variance $\sigma^2$.
        frequency: Angular frequency $\omega_0$.
    """

    variance: Float[Array, ""]
    frequency: Float[Array, ""]

    @property
    def state_dim(self) -> int:
        return 2

    def sde_params(self) -> SDEParams:
        """Return SDE parameters for the cosine kernel."""
        # Constant blocks follow the hyperparameter dtype; untyped
        # ``jnp.zeros``/``jnp.eye`` are float64 under x64 (gh-224).
        dtype = jnp.result_type(self.variance, self.frequency)
        w = self.frequency
        F = jnp.array([[0.0, -w], [w, 0.0]])
        L = jnp.zeros((2, 1), dtype=dtype)
        H = jnp.array([[1.0, 0.0]], dtype=dtype)
        Q_c = jnp.zeros((1, 1), dtype=dtype)
        P_inf = self.variance * jnp.eye(2, dtype=dtype)
        return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=P_inf)

    def discretise(
        self,
        dt: Float[Array, ""],
    ) -> tuple[Float[Array, "d d"], Float[Array, "d d"]]:
        """Closed-form rotation matrix discretization."""
        w = self.frequency
        cos_wdt = jnp.cos(w * dt)
        sin_wdt = jnp.sin(w * dt)
        A = jnp.array([[cos_wdt, -sin_wdt], [sin_wdt, cos_wdt]])
        Q = jnp.zeros((2, 2), dtype=A.dtype)
        return A, Q

sde_params() -> SDEParams

Return SDE parameters for the cosine kernel.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_periodic.py
def sde_params(self) -> SDEParams:
    """Return SDE parameters for the cosine kernel."""
    # Constant blocks follow the hyperparameter dtype; untyped
    # ``jnp.zeros``/``jnp.eye`` are float64 under x64 (gh-224).
    dtype = jnp.result_type(self.variance, self.frequency)
    w = self.frequency
    F = jnp.array([[0.0, -w], [w, 0.0]])
    L = jnp.zeros((2, 1), dtype=dtype)
    H = jnp.array([[1.0, 0.0]], dtype=dtype)
    Q_c = jnp.zeros((1, 1), dtype=dtype)
    P_inf = self.variance * jnp.eye(2, dtype=dtype)
    return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=P_inf)

discretise(dt: Float[Array, '']) -> tuple[Float[Array, 'd d'], Float[Array, 'd d']]

Closed-form rotation matrix discretization.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_periodic.py
def discretise(
    self,
    dt: Float[Array, ""],
) -> tuple[Float[Array, "d d"], Float[Array, "d d"]]:
    """Closed-form rotation matrix discretization."""
    w = self.frequency
    cos_wdt = jnp.cos(w * dt)
    sin_wdt = jnp.sin(w * dt)
    A = jnp.array([[cos_wdt, -sin_wdt], [sin_wdt, cos_wdt]])
    Q = jnp.zeros((2, 2), dtype=A.dtype)
    return A, Q

PeriodicSDE

Bases: SDEKernel

State-space representation of the periodic (MacKay) kernel.

Approximates the periodic kernel via Fourier series truncation to n_harmonics terms. State dimension is 2 * n_harmonics.

Attributes:

Name Type Description
variance Float[Array, '']

Signal variance \(\sigma^2\).

lengthscale Float[Array, '']

Lengthscale \(\ell\).

period Float[Array, '']

Period \(T\).

n_harmonics int

Number of Fourier harmonics (truncation order).

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_periodic.py
class PeriodicSDE(SDEKernel):
    r"""State-space representation of the periodic (MacKay) kernel.

    Approximates the periodic kernel via Fourier series truncation
    to ``n_harmonics`` terms. State dimension is ``2 * n_harmonics``.

    Attributes:
        variance: Signal variance $\sigma^2$.
        lengthscale: Lengthscale $\ell$.
        period: Period $T$.
        n_harmonics: Number of Fourier harmonics (truncation order).
    """

    variance: Float[Array, ""]
    lengthscale: Float[Array, ""]
    period: Float[Array, ""]
    n_harmonics: int = eqx.field(static=True, default=6)

    @property
    def state_dim(self) -> int:
        return 2 * self.n_harmonics

    def sde_params(self) -> SDEParams:
        """Return SDE parameters for the periodic kernel."""
        dtype = jnp.result_type(self.variance, self.lengthscale, self.period)
        J = self.n_harmonics
        d = 2 * J
        w0 = 2.0 * jnp.pi / self.period

        inv_ell_sq = 1.0 / self.lengthscale**2
        # ``js`` feeds the Bessel series; an integer arange would promote it
        # to float64 under x64 (gh-224).
        js = jnp.arange(1, J + 1, dtype=dtype)
        log_ij = self._log_bessel_i(js, inv_ell_sq)
        log_q = jnp.log(2.0) + log_ij - inv_ell_sq
        q_j = self.variance * jnp.exp(log_q)

        F = jnp.zeros((d, d), dtype=dtype)
        P_inf = jnp.zeros((d, d), dtype=dtype)
        for j_idx in range(J):
            freq = (j_idx + 1) * w0
            block_start = 2 * j_idx
            F = F.at[block_start, block_start + 1].set(-freq)
            F = F.at[block_start + 1, block_start].set(freq)
            P_inf = P_inf.at[block_start, block_start].set(q_j[j_idx])
            P_inf = P_inf.at[block_start + 1, block_start + 1].set(q_j[j_idx])

        L = jnp.zeros((d, 1), dtype=dtype)
        H = jnp.zeros((1, d), dtype=dtype)
        for j_idx in range(J):
            H = H.at[0, 2 * j_idx].set(1.0)

        Q_c = jnp.zeros((1, 1), dtype=dtype)
        return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=P_inf)

    def discretise(
        self,
        dt: Float[Array, ""],
    ) -> tuple[Float[Array, "d d"], Float[Array, "d d"]]:
        """Closed-form: block-diagonal rotation matrices."""
        J = self.n_harmonics
        d = 2 * J
        w0 = 2.0 * jnp.pi / self.period

        A = jnp.zeros((d, d), dtype=jnp.result_type(w0, dt))
        for j_idx in range(J):
            freq = (j_idx + 1) * w0
            cos_val = jnp.cos(freq * dt)
            sin_val = jnp.sin(freq * dt)
            block_start = 2 * j_idx
            A = A.at[block_start, block_start].set(cos_val)
            A = A.at[block_start, block_start + 1].set(-sin_val)
            A = A.at[block_start + 1, block_start].set(sin_val)
            A = A.at[block_start + 1, block_start + 1].set(cos_val)

        Q = jnp.zeros((d, d), dtype=A.dtype)
        return A, Q

    @staticmethod
    def _log_bessel_i(
        order: Float[Array, " J"],
        x: Float[Array, ""],
    ) -> Float[Array, " J"]:
        """Log of modified Bessel function I_n(x) via series."""
        half_x = x / 2.0
        log_half_x = jnp.log(half_x)

        log_leading = order * log_half_x - jss.gammaln(order + 1.0)

        x2_over_4 = x**2 / 4.0
        K = 20
        log_sum = jnp.zeros_like(order)
        log_term = jnp.zeros_like(order)
        for k in range(1, K + 1):
            log_term = (
                log_term
                + jnp.log(x2_over_4)
                - jnp.log(jnp.array(k, dtype=order.dtype))
                - jnp.log(order + k)
            )
            log_sum = jnp.logaddexp(log_sum, log_term)

        return log_leading + log_sum

sde_params() -> SDEParams

Return SDE parameters for the periodic kernel.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_periodic.py
def sde_params(self) -> SDEParams:
    """Return SDE parameters for the periodic kernel."""
    dtype = jnp.result_type(self.variance, self.lengthscale, self.period)
    J = self.n_harmonics
    d = 2 * J
    w0 = 2.0 * jnp.pi / self.period

    inv_ell_sq = 1.0 / self.lengthscale**2
    # ``js`` feeds the Bessel series; an integer arange would promote it
    # to float64 under x64 (gh-224).
    js = jnp.arange(1, J + 1, dtype=dtype)
    log_ij = self._log_bessel_i(js, inv_ell_sq)
    log_q = jnp.log(2.0) + log_ij - inv_ell_sq
    q_j = self.variance * jnp.exp(log_q)

    F = jnp.zeros((d, d), dtype=dtype)
    P_inf = jnp.zeros((d, d), dtype=dtype)
    for j_idx in range(J):
        freq = (j_idx + 1) * w0
        block_start = 2 * j_idx
        F = F.at[block_start, block_start + 1].set(-freq)
        F = F.at[block_start + 1, block_start].set(freq)
        P_inf = P_inf.at[block_start, block_start].set(q_j[j_idx])
        P_inf = P_inf.at[block_start + 1, block_start + 1].set(q_j[j_idx])

    L = jnp.zeros((d, 1), dtype=dtype)
    H = jnp.zeros((1, d), dtype=dtype)
    for j_idx in range(J):
        H = H.at[0, 2 * j_idx].set(1.0)

    Q_c = jnp.zeros((1, 1), dtype=dtype)
    return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=P_inf)

discretise(dt: Float[Array, '']) -> tuple[Float[Array, 'd d'], Float[Array, 'd d']]

Closed-form: block-diagonal rotation matrices.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_periodic.py
def discretise(
    self,
    dt: Float[Array, ""],
) -> tuple[Float[Array, "d d"], Float[Array, "d d"]]:
    """Closed-form: block-diagonal rotation matrices."""
    J = self.n_harmonics
    d = 2 * J
    w0 = 2.0 * jnp.pi / self.period

    A = jnp.zeros((d, d), dtype=jnp.result_type(w0, dt))
    for j_idx in range(J):
        freq = (j_idx + 1) * w0
        cos_val = jnp.cos(freq * dt)
        sin_val = jnp.sin(freq * dt)
        block_start = 2 * j_idx
        A = A.at[block_start, block_start].set(cos_val)
        A = A.at[block_start, block_start + 1].set(-sin_val)
        A = A.at[block_start + 1, block_start].set(sin_val)
        A = A.at[block_start + 1, block_start + 1].set(cos_val)

    Q = jnp.zeros((d, d), dtype=A.dtype)
    return A, Q

IntegratedWienerSDE

Bases: SDEKernel

State-space representation of an integrated Wiener process.

The \(p\)-times integrated Wiener process — the local linear trend prior at order=1 — with state \(x(t) = [f(t), f'(t), \dots, f^{(p)}(t)]\) and

\[ \mathrm{d} f^{(p)}(t) = \sqrt{q} \, \mathrm{d} W(t), \]

so the top derivative is white noise and every lower component is its integral. State dimension is order + 1. The SDE matrices are

\[ F = \begin{bmatrix} 0 & I_p \\ 0 & 0 \end{bmatrix}, \quad L = e_p, \quad H = e_0^\top, \quad Q_c = q , \]

with \(F\) nilpotent, which is what makes the process non-stationary: its marginal variance grows without bound, so no \(P_\infty\) exists and sde_params reports P_inf=None. discretise is overridden with the exact closed form (below) and never needs one, but the filter has to be started from an explicit initial_covariance instead.

That non-stationarity is the point: unlike a Matérn prior, the local linear trend commits to no lengthscale — it is smooth and linear unless the data push back — so a long record may drift without the model asserting a scale on which it must revert.

Attributes:

Name Type Description
diffusion Float[Array, '']

Diffusion intensity \(q\), the spectral density of the white noise driving the top derivative. Must be non-negative — \(Q\) is linear in it, so a negative value returns something that is not a covariance. Like every hyperparameter in the kernel zoo this is assumed rather than checked; constrain it with a positive transform when it is learned.

order int

Number of integrations \(p\). 0 is a Brownian random walk; 1 (the default) is the local linear trend, which is also the cubic-spline-equivalent prior — it is \(f''\) that is white noise there, and the smoothed posterior mean is the cubic smoothing spline. Each further order raises the spline by two degrees, so 2 is the quintic-spline prior.

P_0 Float[Array, 'd d'] | None

Initial state covariance, shape (order + 1, order + 1). A modelling choice rather than something the kernel can derive; None (the default) means a diffuse _default_diffuse_variance(dtype) * I.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_wiener.py
class IntegratedWienerSDE(SDEKernel):
    r"""State-space representation of an integrated Wiener process.

    The $p$-times integrated Wiener process — the local linear trend
    prior at ``order=1`` — with state
    $x(t) = [f(t), f'(t), \dots, f^{(p)}(t)]$ and

    $$
    \mathrm{d} f^{(p)}(t) = \sqrt{q} \, \mathrm{d} W(t),
    $$

    so the top derivative is white noise and every lower component is its
    integral. State dimension is ``order + 1``. The SDE matrices are

    $$
    F = \begin{bmatrix} 0 & I_p \\ 0 & 0 \end{bmatrix}, \quad
    L = e_p, \quad
    H = e_0^\top, \quad
    Q_c = q ,
    $$

    with $F$ nilpotent, which is what makes the process **non-stationary**:
    its marginal variance grows without bound, so no $P_\infty$ exists and
    `sde_params` reports ``P_inf=None``. `discretise` is overridden with
    the exact closed form (below) and never needs one, but the filter has
    to be started from an explicit `initial_covariance` instead.

    That non-stationarity is the point: unlike a Matérn prior, the local
    linear trend commits to no lengthscale — it is smooth and linear
    unless the data push back — so a long record may drift without the
    model asserting a scale on which it must revert.

    Attributes:
        diffusion: Diffusion intensity $q$, the spectral density of the
            white noise driving the top derivative. Must be
            non-negative — $Q$ is linear in it, so a negative value
            returns something that is not a covariance. Like every
            hyperparameter in the kernel zoo this is assumed rather
            than checked; constrain it with a positive transform when
            it is learned.
        order: Number of integrations $p$. ``0`` is a Brownian random
            walk; ``1`` (the default) is the local linear trend, which
            is also the cubic-spline-equivalent prior — it is $f''$ that
            is white noise there, and the smoothed posterior mean is the
            cubic smoothing spline. Each further order raises the spline
            by two degrees, so ``2`` is the quintic-spline prior.
        P_0: Initial state covariance, shape ``(order + 1, order + 1)``.
            A modelling choice rather than something the kernel can
            derive; ``None`` (the default) means a diffuse
            ``_default_diffuse_variance(dtype) * I``.
    """

    diffusion: Float[Array, ""]
    order: int = eqx.field(static=True, default=1)
    P_0: Float[Array, "d d"] | None = None

    def __check_init__(self) -> None:
        """Reject a negative order at construction.

        The state dimension is ``order + 1``, so a negative order gives
        an empty state and fails later, deep inside whichever method is
        called first, with an index error about an axis of size zero.
        """
        if self.order < 0:
            msg = (
                f"IntegratedWienerSDE order must be non-negative "
                f"(the state is [f, ..., f^(order)], of dimension "
                f"order + 1), got {self.order}."
            )
            raise ValueError(msg)

    @property
    def state_dim(self) -> int:
        return self.order + 1

    @property
    def stationary(self) -> bool:
        """``False`` — the marginal variance grows without bound."""
        return False

    def sde_params(self) -> SDEParams:
        """Return SDE parameters, with ``P_inf=None``.

        The drift is nilpotent, so no stationary covariance exists; see
        `initial_covariance` for what starts the filter instead.
        """
        # Constant blocks follow the hyperparameter dtype; untyped
        # ``jnp.zeros``/``jnp.eye`` are float64 under x64 (gh-224).
        dtype = _inexact_dtype(self.diffusion)
        d = self.state_dim
        F = jnp.eye(d, k=1, dtype=dtype)
        L = jnp.zeros((d, 1), dtype=dtype).at[d - 1, 0].set(1.0)
        H = jnp.zeros((1, d), dtype=dtype).at[0, 0].set(1.0)
        Q_c = jnp.reshape(self.diffusion, (1, 1)).astype(dtype)
        return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=None)

    def initial_covariance(self) -> Float[Array, "d d"]:
        r"""Return the initial state covariance $P_0$.

        Defaults to a diffuse ``kappa * I`` when the ``P_0`` field is
        ``None``, with ``kappa`` set from the dtype's precision by
        `_default_diffuse_variance` — a vaguer prior than that is not
        merely wasteful but actively wrong, since the Kalman update
        cancels it to a zero-variance first estimate.

        Pass a ``P_0`` to encode what is actually known about the level
        and its derivatives at the first time point — e.g.
        ``diag(kappa, s^2)`` for a vague level and a slope of scale
        ``s``. Do so in particular when the observation noise is large:
        the default is diffuse relative to a noise variance of order
        one, not relative to every scale.
        """
        dtype = _inexact_dtype(self.diffusion)
        if self.P_0 is None:
            kappa = _default_diffuse_variance(dtype)
            return kappa * jnp.eye(self.state_dim, dtype=dtype)
        return jnp.asarray(self.P_0, dtype=dtype)

    def discretise(
        self,
        dt: Float[Array, ""],
    ) -> tuple[Float[Array, "d d"], Float[Array, "d d"]]:
        r"""Closed-form discretisation — no ``expm``, no $P_\infty$.

        The drift is nilpotent, so its exponential terminates:

        $$
        A(\Delta t)_{ij} = \frac{\Delta t^{\,j-i}}{(j-i)!} \;\; (j \ge i),
        \qquad
        Q(\Delta t)_{ij} = q \,
            \frac{\Delta t^{\,2p+1-i-j}}
                 {(p-i)!\,(p-j)!\,(2p+1-i-j)} ,
        $$

        which at ``order=1`` is the familiar

        $$
        A = \begin{bmatrix} 1 & \Delta t \\ 0 & 1 \end{bmatrix},
        \qquad
        Q = q \begin{bmatrix}
            \Delta t^3/3 & \Delta t^2/2 \\
            \Delta t^2/2 & \Delta t
        \end{bmatrix}.
        $$

        Both are exact, so this route is cheaper *and* more accurate than
        the `gaussx.discretise_mfd` fallback the ``P_inf=None`` base
        implementation would otherwise take.

        Note:
            The coefficients carry $1/(p-i)!$, which outruns float64
            from roughly ``order=90`` up and underflows to zero there
            (XLA flushes the subnormals in between). Nothing is
            silently wrong — the entries are zero rather than garbage —
            but a prior that integrates white noise ninety times is not
            a numerically meaningful object in any dtype.

        Args:
            dt: Time step. Must be non-negative — a negative step would
                return a ``Q`` that is not a covariance (at ``order=1``
                it is negative definite), rather than the harmless
                reverse-time transition the sign might suggest. Checked
                with `equinox.error_if`, matching
                `gaussx.discretise_mfd`, so under ``jit`` the error
                fires at evaluation rather than trace time.

        Returns:
            Tuple ``(A, Q)``, both shape ``(order + 1, order + 1)``.
        """
        dtype = _inexact_dtype(self.diffusion, dt)
        p = self.order
        d = self.state_dim

        # The closed form below is a polynomial in dt with no guard of
        # its own, unlike the ``expm`` routes; a negative step would run
        # it happily and hand back an indefinite Q.
        dt = eqx.error_if(
            dt, dt < 0, "IntegratedWienerSDE.discretise requires dt >= 0."
        )

        # Powers are taken with *static* Python exponents so JAX lowers
        # them to ``integer_pow``, whose derivative is exact at zero. A
        # traced exponent would go through the generic ``y * x**(y-1)``
        # rule instead, and dt = 0 -- which the natural
        # ``diff(times, prepend=times[0])`` produces at the first step --
        # would differentiate to 0 * inf = NaN.
        def power(exponent: int) -> Float[Array, ""]:
            if exponent == 0:
                return jnp.ones_like(dt)
            return dt**exponent

        # Only 2p+2 distinct powers appear across both matrices, so they
        # are formed once and gathered by a static index table. That is
        # O(d) traced operations rather than one per entry: at order 40
        # the per-entry version took over five seconds to trace.
        powers = jnp.stack([power(k) for k in range(2 * p + 2)]).astype(dtype)

        i = np.arange(d)[:, None]
        j = np.arange(d)[None, :]
        upper = j >= i
        # Zeroed through the *coefficient* rather than by masking a
        # negative power, so no entry raises dt to a negative exponent.
        a_index = np.maximum(j - i, 0)
        q_index = 2 * p + 1 - i - j

        # ``1 / n`` keeps both operands Python ints, which are unbounded
        # and divide to a correctly rounded float. ``1.0 / n`` would
        # convert first and raise OverflowError from order 98 up, where
        # the denominator exceeds the float range even though its
        # reciprocal is perfectly representable (down to a subnormal,
        # and to zero beyond that).
        factorial = [math.factorial(n) for n in range(d)]
        a_coeff = np.array(
            [
                [1 / factorial[j - i] if j >= i else 0.0 for j in range(d)]
                for i in range(d)
            ]
        )
        q_coeff = np.array(
            [
                [
                    1 / (factorial[p - i] * factorial[p - j] * (2 * p + 1 - i - j))
                    for j in range(d)
                ]
                for i in range(d)
            ]
        )

        A = jnp.asarray(a_coeff * upper, dtype=dtype) * powers[a_index]
        Q = self.diffusion * jnp.asarray(q_coeff, dtype=dtype) * powers[q_index]
        return A.astype(dtype), Q.astype(dtype)

stationary: bool property

False — the marginal variance grows without bound.

sde_params() -> SDEParams

Return SDE parameters, with P_inf=None.

The drift is nilpotent, so no stationary covariance exists; see initial_covariance for what starts the filter instead.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_wiener.py
def sde_params(self) -> SDEParams:
    """Return SDE parameters, with ``P_inf=None``.

    The drift is nilpotent, so no stationary covariance exists; see
    `initial_covariance` for what starts the filter instead.
    """
    # Constant blocks follow the hyperparameter dtype; untyped
    # ``jnp.zeros``/``jnp.eye`` are float64 under x64 (gh-224).
    dtype = _inexact_dtype(self.diffusion)
    d = self.state_dim
    F = jnp.eye(d, k=1, dtype=dtype)
    L = jnp.zeros((d, 1), dtype=dtype).at[d - 1, 0].set(1.0)
    H = jnp.zeros((1, d), dtype=dtype).at[0, 0].set(1.0)
    Q_c = jnp.reshape(self.diffusion, (1, 1)).astype(dtype)
    return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=None)

initial_covariance() -> Float[Array, 'd d']

Return the initial state covariance \(P_0\).

Defaults to a diffuse kappa * I when the P_0 field is None, with kappa set from the dtype's precision by _default_diffuse_variance — a vaguer prior than that is not merely wasteful but actively wrong, since the Kalman update cancels it to a zero-variance first estimate.

Pass a P_0 to encode what is actually known about the level and its derivatives at the first time point — e.g. diag(kappa, s^2) for a vague level and a slope of scale s. Do so in particular when the observation noise is large: the default is diffuse relative to a noise variance of order one, not relative to every scale.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_wiener.py
def initial_covariance(self) -> Float[Array, "d d"]:
    r"""Return the initial state covariance $P_0$.

    Defaults to a diffuse ``kappa * I`` when the ``P_0`` field is
    ``None``, with ``kappa`` set from the dtype's precision by
    `_default_diffuse_variance` — a vaguer prior than that is not
    merely wasteful but actively wrong, since the Kalman update
    cancels it to a zero-variance first estimate.

    Pass a ``P_0`` to encode what is actually known about the level
    and its derivatives at the first time point — e.g.
    ``diag(kappa, s^2)`` for a vague level and a slope of scale
    ``s``. Do so in particular when the observation noise is large:
    the default is diffuse relative to a noise variance of order
    one, not relative to every scale.
    """
    dtype = _inexact_dtype(self.diffusion)
    if self.P_0 is None:
        kappa = _default_diffuse_variance(dtype)
        return kappa * jnp.eye(self.state_dim, dtype=dtype)
    return jnp.asarray(self.P_0, dtype=dtype)

discretise(dt: Float[Array, '']) -> tuple[Float[Array, 'd d'], Float[Array, 'd d']]

Closed-form discretisation — no expm, no \(P_\infty\).

The drift is nilpotent, so its exponential terminates:

\[ A(\Delta t)_{ij} = \frac{\Delta t^{\,j-i}}{(j-i)!} \;\; (j \ge i), \qquad Q(\Delta t)_{ij} = q \, \frac{\Delta t^{\,2p+1-i-j}} {(p-i)!\,(p-j)!\,(2p+1-i-j)} , \]

which at order=1 is the familiar

\[ A = \begin{bmatrix} 1 & \Delta t \\ 0 & 1 \end{bmatrix}, \qquad Q = q \begin{bmatrix} \Delta t^3/3 & \Delta t^2/2 \\ \Delta t^2/2 & \Delta t \end{bmatrix}. \]

Both are exact, so this route is cheaper and more accurate than the gaussx.discretise_mfd fallback the P_inf=None base implementation would otherwise take.

Note

The coefficients carry \(1/(p-i)!\), which outruns float64 from roughly order=90 up and underflows to zero there (XLA flushes the subnormals in between). Nothing is silently wrong — the entries are zero rather than garbage — but a prior that integrates white noise ninety times is not a numerically meaningful object in any dtype.

Parameters:

Name Type Description Default
dt Float[Array, '']

Time step. Must be non-negative — a negative step would return a Q that is not a covariance (at order=1 it is negative definite), rather than the harmless reverse-time transition the sign might suggest. Checked with equinox.error_if, matching gaussx.discretise_mfd, so under jit the error fires at evaluation rather than trace time.

required

Returns:

Type Description
tuple[Float[Array, 'd d'], Float[Array, 'd d']]

Tuple (A, Q), both shape (order + 1, order + 1).

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_wiener.py
def discretise(
    self,
    dt: Float[Array, ""],
) -> tuple[Float[Array, "d d"], Float[Array, "d d"]]:
    r"""Closed-form discretisation — no ``expm``, no $P_\infty$.

    The drift is nilpotent, so its exponential terminates:

    $$
    A(\Delta t)_{ij} = \frac{\Delta t^{\,j-i}}{(j-i)!} \;\; (j \ge i),
    \qquad
    Q(\Delta t)_{ij} = q \,
        \frac{\Delta t^{\,2p+1-i-j}}
             {(p-i)!\,(p-j)!\,(2p+1-i-j)} ,
    $$

    which at ``order=1`` is the familiar

    $$
    A = \begin{bmatrix} 1 & \Delta t \\ 0 & 1 \end{bmatrix},
    \qquad
    Q = q \begin{bmatrix}
        \Delta t^3/3 & \Delta t^2/2 \\
        \Delta t^2/2 & \Delta t
    \end{bmatrix}.
    $$

    Both are exact, so this route is cheaper *and* more accurate than
    the `gaussx.discretise_mfd` fallback the ``P_inf=None`` base
    implementation would otherwise take.

    Note:
        The coefficients carry $1/(p-i)!$, which outruns float64
        from roughly ``order=90`` up and underflows to zero there
        (XLA flushes the subnormals in between). Nothing is
        silently wrong — the entries are zero rather than garbage —
        but a prior that integrates white noise ninety times is not
        a numerically meaningful object in any dtype.

    Args:
        dt: Time step. Must be non-negative — a negative step would
            return a ``Q`` that is not a covariance (at ``order=1``
            it is negative definite), rather than the harmless
            reverse-time transition the sign might suggest. Checked
            with `equinox.error_if`, matching
            `gaussx.discretise_mfd`, so under ``jit`` the error
            fires at evaluation rather than trace time.

    Returns:
        Tuple ``(A, Q)``, both shape ``(order + 1, order + 1)``.
    """
    dtype = _inexact_dtype(self.diffusion, dt)
    p = self.order
    d = self.state_dim

    # The closed form below is a polynomial in dt with no guard of
    # its own, unlike the ``expm`` routes; a negative step would run
    # it happily and hand back an indefinite Q.
    dt = eqx.error_if(
        dt, dt < 0, "IntegratedWienerSDE.discretise requires dt >= 0."
    )

    # Powers are taken with *static* Python exponents so JAX lowers
    # them to ``integer_pow``, whose derivative is exact at zero. A
    # traced exponent would go through the generic ``y * x**(y-1)``
    # rule instead, and dt = 0 -- which the natural
    # ``diff(times, prepend=times[0])`` produces at the first step --
    # would differentiate to 0 * inf = NaN.
    def power(exponent: int) -> Float[Array, ""]:
        if exponent == 0:
            return jnp.ones_like(dt)
        return dt**exponent

    # Only 2p+2 distinct powers appear across both matrices, so they
    # are formed once and gathered by a static index table. That is
    # O(d) traced operations rather than one per entry: at order 40
    # the per-entry version took over five seconds to trace.
    powers = jnp.stack([power(k) for k in range(2 * p + 2)]).astype(dtype)

    i = np.arange(d)[:, None]
    j = np.arange(d)[None, :]
    upper = j >= i
    # Zeroed through the *coefficient* rather than by masking a
    # negative power, so no entry raises dt to a negative exponent.
    a_index = np.maximum(j - i, 0)
    q_index = 2 * p + 1 - i - j

    # ``1 / n`` keeps both operands Python ints, which are unbounded
    # and divide to a correctly rounded float. ``1.0 / n`` would
    # convert first and raise OverflowError from order 98 up, where
    # the denominator exceeds the float range even though its
    # reciprocal is perfectly representable (down to a subnormal,
    # and to zero beyond that).
    factorial = [math.factorial(n) for n in range(d)]
    a_coeff = np.array(
        [
            [1 / factorial[j - i] if j >= i else 0.0 for j in range(d)]
            for i in range(d)
        ]
    )
    q_coeff = np.array(
        [
            [
                1 / (factorial[p - i] * factorial[p - j] * (2 * p + 1 - i - j))
                for j in range(d)
            ]
            for i in range(d)
        ]
    )

    A = jnp.asarray(a_coeff * upper, dtype=dtype) * powers[a_index]
    Q = self.diffusion * jnp.asarray(q_coeff, dtype=dtype) * powers[q_index]
    return A.astype(dtype), Q.astype(dtype)

SumSDE

Bases: SDEKernel

Sum of SDE kernels via block-diagonal composition.

Attributes:

Name Type Description
kernels tuple[SDEKernel, ...]

Tuple of component SDE kernels.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_composition.py
class SumSDE(SDEKernel):
    """Sum of SDE kernels via block-diagonal composition.

    Attributes:
        kernels: Tuple of component SDE kernels.
    """

    kernels: tuple[SDEKernel, ...] = eqx.field()

    @property
    def state_dim(self) -> int:
        return sum(k.state_dim for k in self.kernels)

    @property
    def stationary(self) -> bool:
        """Stationary only if every component is."""
        return all(k.stationary for k in self.kernels)

    def initial_covariance(self) -> Float[Array, "d d"]:
        """Return the block-diagonal initial covariance.

        The components are independent, so their initial covariances
        stack block-diagonally — which lets a sum mix stationary and
        non-stationary components (a local linear trend plus a Matern
        seasonal, say), each started from its own.
        """
        return jsl.block_diag(*[k.initial_covariance() for k in self.kernels])

    def sde_params(self) -> SDEParams:
        """Return block-diagonal SDE parameters."""
        params_list = [k.sde_params() for k in self.kernels]

        F = jsl.block_diag(*[p.F for p in params_list])
        # A component with no closed-form stationary covariance leaves the
        # sum without one either; propagating ``None`` routes the composite
        # through ``discretise_mfd`` rather than fabricating a ``P_inf``.
        component_p_inf = [p.P_inf for p in params_list]
        P_inf = (
            None
            if any(block is None for block in component_p_inf)
            else jsl.block_diag(*component_p_inf)
        )

        L_blocks = [p.L for p in params_list]
        total_rows = sum(b.shape[0] for b in L_blocks)
        total_cols = sum(b.shape[1] for b in L_blocks)
        L = jnp.zeros((total_rows, total_cols))
        row_offset = 0
        col_offset = 0
        for block in L_blocks:
            r, c = block.shape
            L = L.at[row_offset : row_offset + r, col_offset : col_offset + c].set(
                block
            )
            row_offset += r
            col_offset += c

        Q_c = jsl.block_diag(*[p.Q_c for p in params_list])
        H = jnp.concatenate([p.H for p in params_list], axis=1)

        return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=P_inf)

stationary: bool property

Stationary only if every component is.

initial_covariance() -> Float[Array, 'd d']

Return the block-diagonal initial covariance.

The components are independent, so their initial covariances stack block-diagonally — which lets a sum mix stationary and non-stationary components (a local linear trend plus a Matern seasonal, say), each started from its own.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_composition.py
def initial_covariance(self) -> Float[Array, "d d"]:
    """Return the block-diagonal initial covariance.

    The components are independent, so their initial covariances
    stack block-diagonally — which lets a sum mix stationary and
    non-stationary components (a local linear trend plus a Matern
    seasonal, say), each started from its own.
    """
    return jsl.block_diag(*[k.initial_covariance() for k in self.kernels])

sde_params() -> SDEParams

Return block-diagonal SDE parameters.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_composition.py
def sde_params(self) -> SDEParams:
    """Return block-diagonal SDE parameters."""
    params_list = [k.sde_params() for k in self.kernels]

    F = jsl.block_diag(*[p.F for p in params_list])
    # A component with no closed-form stationary covariance leaves the
    # sum without one either; propagating ``None`` routes the composite
    # through ``discretise_mfd`` rather than fabricating a ``P_inf``.
    component_p_inf = [p.P_inf for p in params_list]
    P_inf = (
        None
        if any(block is None for block in component_p_inf)
        else jsl.block_diag(*component_p_inf)
    )

    L_blocks = [p.L for p in params_list]
    total_rows = sum(b.shape[0] for b in L_blocks)
    total_cols = sum(b.shape[1] for b in L_blocks)
    L = jnp.zeros((total_rows, total_cols))
    row_offset = 0
    col_offset = 0
    for block in L_blocks:
        r, c = block.shape
        L = L.at[row_offset : row_offset + r, col_offset : col_offset + c].set(
            block
        )
        row_offset += r
        col_offset += c

    Q_c = jsl.block_diag(*[p.Q_c for p in params_list])
    H = jnp.concatenate([p.H for p in params_list], axis=1)

    return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=P_inf)

ProductSDE

Bases: SDEKernel

Product of two SDE kernels via Kronecker composition.

Attributes:

Name Type Description
kernel1 SDEKernel

First component kernel.

kernel2 SDEKernel

Second component kernel.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_composition.py
class ProductSDE(SDEKernel):
    """Product of two SDE kernels via Kronecker composition.

    Attributes:
        kernel1: First component kernel.
        kernel2: Second component kernel.
    """

    kernel1: SDEKernel
    kernel2: SDEKernel

    @property
    def state_dim(self) -> int:
        return self.kernel1.state_dim * self.kernel2.state_dim

    @property
    def stationary(self) -> bool:
        """Stationary only if both factors are."""
        return self.kernel1.stationary and self.kernel2.stationary

    def sde_params(self) -> SDEParams:
        r"""Return Kronecker-structured SDE parameters.

        The drift is the Kronecker **sum** $F_1 \oplus F_2$ and the
        stationary covariance the Kronecker **product**
        $P_1 \otimes P_2$. Substituting those into the Lyapunov equation
        $F P + P F^\top + B = 0$ fixes the composite diffusion at

        $$
        B \;=\; B_1 \otimes P_2 \;+\; P_1 \otimes B_2,
        \qquad B_i = L_i Q_{c,i} L_i^\top,
        $$

        which is **not** $B_1 \otimes B_2$ — the value a naive
        $L_1 \otimes L_2$, $Q_{c,1} \otimes Q_{c,2}$ pair would imply.
        Reporting the latter used to hand out a tuple that failed its own
        Lyapunov equation (gh-219); for a Matérn ⊗ Cosine product it was
        identically zero, since `CosineSDE` has $Q_c = 0$.

        The sum of two Kronecker products is still expressible in the
        ``(L, Q_c)`` form, by widening the noise dimension and carrying
        each factor's stationary covariance in the spectral density:

        $$
        L = \bigl[\, L_1 \otimes I_{d_2} \;\;\big|\;\; I_{d_1} \otimes L_2 \,\bigr],
        \qquad
        Q_c = \operatorname{blockdiag}\!\bigl(
            Q_{c,1} \otimes P_2,\; P_1 \otimes Q_{c,2}
        \bigr),
        $$

        so that $L Q_c L^\top$ telescopes to exactly the $B$ above by the
        mixed-product property. Writing it this way rather than through
        square roots $S_i S_i^\top = P_i$ keeps the result exact for
        singular or zero $P_\infty$ (where a Cholesky would need jitter,
        and would then violate the very Lyapunov equation this enforces)
        and keeps ``sde_params`` reverse-mode differentiable.

        Note:
            ``SDEParams`` currently types its fields as dense
            ``jaxtyping.Float[Array, ...]``. The Kronecker products
            below are dense materializations of size
            ``(state_dim, state_dim)``, where ``state_dim`` is
            ``kernel1.state_dim * kernel2.state_dim`` — for typical SSM
            kernels (Matérn-3/2, periodic) this is ≤ 32, so the
            materialization is bounded and cheap. A future refactor
            could expose a parallel ``sde_operators()`` method that
            returns `gaussx.Kronecker` operators for downstream
            filters that can exploit the structure (issue #153).

        Raises:
            NotImplementedError: If either factor lacks a stationary
                covariance. The composite diffusion needs both, so the
                tuple cannot be built — and reporting the inconsistent
                Kronecker product instead is what gh-219 was about.
        """
        p1 = self.kernel1.sde_params()
        p2 = self.kernel2.sde_params()

        d1 = self.kernel1.state_dim
        d2 = self.kernel2.state_dim

        if p1.P_inf is None or p2.P_inf is None:
            msg = (
                f"ProductSDE cannot report SDE parameters when a factor has "
                f"no stationary covariance: "
                f"{type(self.kernel1).__name__}.P_inf is "
                f"{'None' if p1.P_inf is None else 'set'} and "
                f"{type(self.kernel2).__name__}.P_inf is "
                f"{'None' if p2.P_inf is None else 'set'}. The composite "
                f"diffusion of a product kernel is B1 (x) P2 + P1 (x) B2, "
                f"which needs both. Use the factors' own parameters, or "
                f"give the factor a P_inf."
            )
            raise NotImplementedError(msg)

        # Identities carry the factor dtypes: an untyped ``jnp.eye`` is
        # float64 under x64 and would promote float32 kernels.
        eye1 = jnp.eye(d1, dtype=p1.F.dtype)
        eye2 = jnp.eye(d2, dtype=p2.F.dtype)

        F = jnp.kron(p1.F, eye2) + jnp.kron(eye1, p2.F)
        H = jnp.kron(p1.H, p2.H)
        P_inf = jnp.kron(p1.P_inf, p2.P_inf)

        # B = B1 (x) P2 + P1 (x) B2, kept in the (L, Q_c) pair by putting
        # each factor's P_inf in the *spectral density* rather than taking
        # its square root:
        #
        #     B1 (x) P2 = (L1 (x) I) (Q_c1 (x) P2) (L1 (x) I)^T
        #     P1 (x) B2 = (I (x) L2) (P1 (x) Q_c2) (I (x) L2)^T
        #
        # by the mixed-product property. Exact for any PSD P_inf --
        # including singular or zero ones, where a Cholesky would need
        # jitter and stop satisfying the Lyapunov equation this is here to
        # enforce. The noise dimension widens from s1*s2 to s1*d2 + d1*s2.
        L = jnp.concatenate(
            [jnp.kron(p1.L, eye2), jnp.kron(eye1, p2.L)],
            axis=1,
        )
        Q_c = jsl.block_diag(
            jnp.kron(p1.Q_c, p2.P_inf),
            jnp.kron(p1.P_inf, p2.Q_c),
        )

        return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=P_inf)

    def discretise(
        self,
        dt: Float[Array, ""],
    ) -> tuple[Float[Array, "d d"], Float[Array, "d d"]]:
        r"""Discretise via the Kronecker matrix-exponential identity.

        For a product kernel ``F = F_1 \oplus F_2 = F_1 \otimes I + I \otimes F_2``,
        the factors ``F_1 \otimes I`` and ``I \otimes F_2`` commute, so

        $$
        \exp(F \, dt) = \exp(F_1 \, dt) \otimes \exp(F_2 \, dt).
        $$

        This computes two ``expm`` calls of size ``d_1`` and ``d_2``
        each, plus one Kronecker product, instead of one ``expm`` of
        size ``d_1 \cdot d_2``. Numerically equivalent to the dense
        ``expm`` on ``F`` but cheaper for moderate factor sizes.

        ``Q = P_\infty - A P_\infty A^T`` exploits the same factorisation.
        By the mixed-product property,

        $$
        (A_1 \otimes A_2)(P_1 \otimes P_2)(A_1 \otimes A_2)^\top
            = (A_1 P_1 A_1^\top) \otimes (A_2 P_2 A_2^\top),
        $$

        so the congruence is evaluated per factor via
        `gaussx.process_noise_covariance` on `gaussx.Kronecker`
        operands — ``O(d_1^3 + d_2^3)`` instead of the
        ``O((d_1 d_2)^3)`` triple product on the full matrix. Only the
        final ``Q`` is materialised, to keep the consumer-facing
        ``(A, Q)`` interface unchanged.

        Args:
            dt: Time step (scalar, positive).

        Returns:
            Tuple ``(A, Q)`` matching `SDEKernel.discretise`.
        """
        p1 = self.kernel1.sde_params()
        p2 = self.kernel2.sde_params()
        A1 = jsl.expm(p1.F * dt)
        A2 = jsl.expm(p2.F * dt)
        A = jnp.kron(A1, A2)

        # Use the per-factor stationary covariances directly; building
        # the full ``F`` via ``self.sde_params()`` would defeat the
        # whole point of this override. Keeping both operands as
        # ``Kronecker`` lets the shared helper contract each factor
        # separately instead of forming the (d1 d2)-square triple product.
        if p1.P_inf is None or p2.P_inf is None:
            # Deferring to the base MFD path here would be *wrong*, not
            # merely slower: the composite diffusion
            #
            #     B = B_1 (x) P_2  +  P_1 (x) B_2
            #
            # needs both factor stationary covariances -- exactly what is
            # missing -- so MFD has nothing correct to consume. Checked
            # against the factors directly rather than via
            # ``self.sde_params()`` (which now raises for the same reason)
            # so the message names the offending factor, and so the
            # happy path never builds the full-size drift it would
            # otherwise have to discard.
            msg = (
                f"ProductSDE cannot be discretised when a factor has no "
                f"stationary covariance: "
                f"{type(self.kernel1).__name__}.P_inf is "
                f"{'None' if p1.P_inf is None else 'set'} and "
                f"{type(self.kernel2).__name__}.P_inf is "
                f"{'None' if p2.P_inf is None else 'set'}. The composite "
                f"diffusion of a product kernel is B1 (x) P2 + P1 (x) B2, "
                f"which needs both. Discretise the factors separately, or "
                f"give the factor a P_inf."
            )
            raise NotImplementedError(msg)

        A_op = Kronecker(
            lx.MatrixLinearOperator(A1),
            lx.MatrixLinearOperator(A2),
        )
        P_op = Kronecker(
            lx.MatrixLinearOperator(p1.P_inf, lx.symmetric_tag),
            lx.MatrixLinearOperator(p2.P_inf, lx.symmetric_tag),
        )
        Q = symmetrize(process_noise_covariance(A_op, P_op).as_matrix())
        return A, Q

stationary: bool property

Stationary only if both factors are.

sde_params() -> SDEParams

Return Kronecker-structured SDE parameters.

The drift is the Kronecker sum \(F_1 \oplus F_2\) and the stationary covariance the Kronecker product \(P_1 \otimes P_2\). Substituting those into the Lyapunov equation \(F P + P F^\top + B = 0\) fixes the composite diffusion at

\[ B \;=\; B_1 \otimes P_2 \;+\; P_1 \otimes B_2, \qquad B_i = L_i Q_{c,i} L_i^\top, \]

which is not \(B_1 \otimes B_2\) — the value a naive \(L_1 \otimes L_2\), \(Q_{c,1} \otimes Q_{c,2}\) pair would imply. Reporting the latter used to hand out a tuple that failed its own Lyapunov equation (gh-219); for a Matérn ⊗ Cosine product it was identically zero, since CosineSDE has \(Q_c = 0\).

The sum of two Kronecker products is still expressible in the (L, Q_c) form, by widening the noise dimension and carrying each factor's stationary covariance in the spectral density:

\[ L = \bigl[\, L_1 \otimes I_{d_2} \;\;\big|\;\; I_{d_1} \otimes L_2 \,\bigr], \qquad Q_c = \operatorname{blockdiag}\!\bigl( Q_{c,1} \otimes P_2,\; P_1 \otimes Q_{c,2} \bigr), \]

so that \(L Q_c L^\top\) telescopes to exactly the \(B\) above by the mixed-product property. Writing it this way rather than through square roots \(S_i S_i^\top = P_i\) keeps the result exact for singular or zero \(P_\infty\) (where a Cholesky would need jitter, and would then violate the very Lyapunov equation this enforces) and keeps sde_params reverse-mode differentiable.

Note

SDEParams currently types its fields as dense jaxtyping.Float[Array, ...]. The Kronecker products below are dense materializations of size (state_dim, state_dim), where state_dim is kernel1.state_dim * kernel2.state_dim — for typical SSM kernels (Matérn-3/2, periodic) this is ≤ 32, so the materialization is bounded and cheap. A future refactor could expose a parallel sde_operators() method that returns gaussx.Kronecker operators for downstream filters that can exploit the structure (issue #153).

Raises:

Type Description
NotImplementedError

If either factor lacks a stationary covariance. The composite diffusion needs both, so the tuple cannot be built — and reporting the inconsistent Kronecker product instead is what gh-219 was about.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_composition.py
def sde_params(self) -> SDEParams:
    r"""Return Kronecker-structured SDE parameters.

    The drift is the Kronecker **sum** $F_1 \oplus F_2$ and the
    stationary covariance the Kronecker **product**
    $P_1 \otimes P_2$. Substituting those into the Lyapunov equation
    $F P + P F^\top + B = 0$ fixes the composite diffusion at

    $$
    B \;=\; B_1 \otimes P_2 \;+\; P_1 \otimes B_2,
    \qquad B_i = L_i Q_{c,i} L_i^\top,
    $$

    which is **not** $B_1 \otimes B_2$ — the value a naive
    $L_1 \otimes L_2$, $Q_{c,1} \otimes Q_{c,2}$ pair would imply.
    Reporting the latter used to hand out a tuple that failed its own
    Lyapunov equation (gh-219); for a Matérn ⊗ Cosine product it was
    identically zero, since `CosineSDE` has $Q_c = 0$.

    The sum of two Kronecker products is still expressible in the
    ``(L, Q_c)`` form, by widening the noise dimension and carrying
    each factor's stationary covariance in the spectral density:

    $$
    L = \bigl[\, L_1 \otimes I_{d_2} \;\;\big|\;\; I_{d_1} \otimes L_2 \,\bigr],
    \qquad
    Q_c = \operatorname{blockdiag}\!\bigl(
        Q_{c,1} \otimes P_2,\; P_1 \otimes Q_{c,2}
    \bigr),
    $$

    so that $L Q_c L^\top$ telescopes to exactly the $B$ above by the
    mixed-product property. Writing it this way rather than through
    square roots $S_i S_i^\top = P_i$ keeps the result exact for
    singular or zero $P_\infty$ (where a Cholesky would need jitter,
    and would then violate the very Lyapunov equation this enforces)
    and keeps ``sde_params`` reverse-mode differentiable.

    Note:
        ``SDEParams`` currently types its fields as dense
        ``jaxtyping.Float[Array, ...]``. The Kronecker products
        below are dense materializations of size
        ``(state_dim, state_dim)``, where ``state_dim`` is
        ``kernel1.state_dim * kernel2.state_dim`` — for typical SSM
        kernels (Matérn-3/2, periodic) this is ≤ 32, so the
        materialization is bounded and cheap. A future refactor
        could expose a parallel ``sde_operators()`` method that
        returns `gaussx.Kronecker` operators for downstream
        filters that can exploit the structure (issue #153).

    Raises:
        NotImplementedError: If either factor lacks a stationary
            covariance. The composite diffusion needs both, so the
            tuple cannot be built — and reporting the inconsistent
            Kronecker product instead is what gh-219 was about.
    """
    p1 = self.kernel1.sde_params()
    p2 = self.kernel2.sde_params()

    d1 = self.kernel1.state_dim
    d2 = self.kernel2.state_dim

    if p1.P_inf is None or p2.P_inf is None:
        msg = (
            f"ProductSDE cannot report SDE parameters when a factor has "
            f"no stationary covariance: "
            f"{type(self.kernel1).__name__}.P_inf is "
            f"{'None' if p1.P_inf is None else 'set'} and "
            f"{type(self.kernel2).__name__}.P_inf is "
            f"{'None' if p2.P_inf is None else 'set'}. The composite "
            f"diffusion of a product kernel is B1 (x) P2 + P1 (x) B2, "
            f"which needs both. Use the factors' own parameters, or "
            f"give the factor a P_inf."
        )
        raise NotImplementedError(msg)

    # Identities carry the factor dtypes: an untyped ``jnp.eye`` is
    # float64 under x64 and would promote float32 kernels.
    eye1 = jnp.eye(d1, dtype=p1.F.dtype)
    eye2 = jnp.eye(d2, dtype=p2.F.dtype)

    F = jnp.kron(p1.F, eye2) + jnp.kron(eye1, p2.F)
    H = jnp.kron(p1.H, p2.H)
    P_inf = jnp.kron(p1.P_inf, p2.P_inf)

    # B = B1 (x) P2 + P1 (x) B2, kept in the (L, Q_c) pair by putting
    # each factor's P_inf in the *spectral density* rather than taking
    # its square root:
    #
    #     B1 (x) P2 = (L1 (x) I) (Q_c1 (x) P2) (L1 (x) I)^T
    #     P1 (x) B2 = (I (x) L2) (P1 (x) Q_c2) (I (x) L2)^T
    #
    # by the mixed-product property. Exact for any PSD P_inf --
    # including singular or zero ones, where a Cholesky would need
    # jitter and stop satisfying the Lyapunov equation this is here to
    # enforce. The noise dimension widens from s1*s2 to s1*d2 + d1*s2.
    L = jnp.concatenate(
        [jnp.kron(p1.L, eye2), jnp.kron(eye1, p2.L)],
        axis=1,
    )
    Q_c = jsl.block_diag(
        jnp.kron(p1.Q_c, p2.P_inf),
        jnp.kron(p1.P_inf, p2.Q_c),
    )

    return SDEParams(F=F, L=L, H=H, Q_c=Q_c, P_inf=P_inf)

discretise(dt: Float[Array, '']) -> tuple[Float[Array, 'd d'], Float[Array, 'd d']]

Discretise via the Kronecker matrix-exponential identity.

For a product kernel F = F_1 \oplus F_2 = F_1 \otimes I + I \otimes F_2, the factors F_1 \otimes I and I \otimes F_2 commute, so

\[ \exp(F \, dt) = \exp(F_1 \, dt) \otimes \exp(F_2 \, dt). \]

This computes two expm calls of size d_1 and d_2 each, plus one Kronecker product, instead of one expm of size d_1 \cdot d_2. Numerically equivalent to the dense expm on F but cheaper for moderate factor sizes.

Q = P_\infty - A P_\infty A^T exploits the same factorisation. By the mixed-product property,

\[ (A_1 \otimes A_2)(P_1 \otimes P_2)(A_1 \otimes A_2)^\top = (A_1 P_1 A_1^\top) \otimes (A_2 P_2 A_2^\top), \]

so the congruence is evaluated per factor via gaussx.process_noise_covariance on gaussx.Kronecker operands — O(d_1^3 + d_2^3) instead of the O((d_1 d_2)^3) triple product on the full matrix. Only the final Q is materialised, to keep the consumer-facing (A, Q) interface unchanged.

Parameters:

Name Type Description Default
dt Float[Array, '']

Time step (scalar, positive).

required

Returns:

Type Description
tuple[Float[Array, 'd d'], Float[Array, 'd d']]

Tuple (A, Q) matching SDEKernel.discretise.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_composition.py
def discretise(
    self,
    dt: Float[Array, ""],
) -> tuple[Float[Array, "d d"], Float[Array, "d d"]]:
    r"""Discretise via the Kronecker matrix-exponential identity.

    For a product kernel ``F = F_1 \oplus F_2 = F_1 \otimes I + I \otimes F_2``,
    the factors ``F_1 \otimes I`` and ``I \otimes F_2`` commute, so

    $$
    \exp(F \, dt) = \exp(F_1 \, dt) \otimes \exp(F_2 \, dt).
    $$

    This computes two ``expm`` calls of size ``d_1`` and ``d_2``
    each, plus one Kronecker product, instead of one ``expm`` of
    size ``d_1 \cdot d_2``. Numerically equivalent to the dense
    ``expm`` on ``F`` but cheaper for moderate factor sizes.

    ``Q = P_\infty - A P_\infty A^T`` exploits the same factorisation.
    By the mixed-product property,

    $$
    (A_1 \otimes A_2)(P_1 \otimes P_2)(A_1 \otimes A_2)^\top
        = (A_1 P_1 A_1^\top) \otimes (A_2 P_2 A_2^\top),
    $$

    so the congruence is evaluated per factor via
    `gaussx.process_noise_covariance` on `gaussx.Kronecker`
    operands — ``O(d_1^3 + d_2^3)`` instead of the
    ``O((d_1 d_2)^3)`` triple product on the full matrix. Only the
    final ``Q`` is materialised, to keep the consumer-facing
    ``(A, Q)`` interface unchanged.

    Args:
        dt: Time step (scalar, positive).

    Returns:
        Tuple ``(A, Q)`` matching `SDEKernel.discretise`.
    """
    p1 = self.kernel1.sde_params()
    p2 = self.kernel2.sde_params()
    A1 = jsl.expm(p1.F * dt)
    A2 = jsl.expm(p2.F * dt)
    A = jnp.kron(A1, A2)

    # Use the per-factor stationary covariances directly; building
    # the full ``F`` via ``self.sde_params()`` would defeat the
    # whole point of this override. Keeping both operands as
    # ``Kronecker`` lets the shared helper contract each factor
    # separately instead of forming the (d1 d2)-square triple product.
    if p1.P_inf is None or p2.P_inf is None:
        # Deferring to the base MFD path here would be *wrong*, not
        # merely slower: the composite diffusion
        #
        #     B = B_1 (x) P_2  +  P_1 (x) B_2
        #
        # needs both factor stationary covariances -- exactly what is
        # missing -- so MFD has nothing correct to consume. Checked
        # against the factors directly rather than via
        # ``self.sde_params()`` (which now raises for the same reason)
        # so the message names the offending factor, and so the
        # happy path never builds the full-size drift it would
        # otherwise have to discard.
        msg = (
            f"ProductSDE cannot be discretised when a factor has no "
            f"stationary covariance: "
            f"{type(self.kernel1).__name__}.P_inf is "
            f"{'None' if p1.P_inf is None else 'set'} and "
            f"{type(self.kernel2).__name__}.P_inf is "
            f"{'None' if p2.P_inf is None else 'set'}. The composite "
            f"diffusion of a product kernel is B1 (x) P2 + P1 (x) B2, "
            f"which needs both. Discretise the factors separately, or "
            f"give the factor a P_inf."
        )
        raise NotImplementedError(msg)

    A_op = Kronecker(
        lx.MatrixLinearOperator(A1),
        lx.MatrixLinearOperator(A2),
    )
    P_op = Kronecker(
        lx.MatrixLinearOperator(p1.P_inf, lx.symmetric_tag),
        lx.MatrixLinearOperator(p2.P_inf, lx.symmetric_tag),
    )
    Q = symmetrize(process_noise_covariance(A_op, P_op).as_matrix())
    return A, Q

QuasiPeriodicSDE

Bases: ProductSDE

Quasi-periodic kernel: product of Matern and Periodic SDE.

Attributes:

Name Type Description
kernel1 SDEKernel

Modulating kernel (typically Matern).

kernel2 SDEKernel

Periodic kernel.

Source code in .venv/lib/python3.12/site-packages/gaussx/_ssm/_composition.py
class QuasiPeriodicSDE(ProductSDE):
    """Quasi-periodic kernel: product of Matern and Periodic SDE.

    Attributes:
        kernel1: Modulating kernel (typically Matern).
        kernel2: Periodic kernel.
    """

    pass

Markov GP — Kalman / RTS workflow

MarkovGPPrior consumes any SDEKernel over a sorted 1-D grid and gives O(N d^3) marginal likelihood (forward Kalman filter) and posterior smoothing (backward RTS), where d is the SDE state dimension. Use it for temporal GP regression / forecasting when the training grid lives on a single time axis. Predictions at arbitrary test times — including forecasting, backcasting, and within-window interpolation — re-run the filter+smoother over the merged grid with the test points masked out of the update step.

import jax.numpy as jnp
from pyrox_gp import MaternSDE, MarkovGPPrior, markov_gp_factor

times = jnp.linspace(0.0, 5.0, 200)
y     = jnp.sin(times) + 0.05 * jnp.cos(7.0 * times)

prior = MarkovGPPrior(
    MaternSDE(variance=1.0, lengthscale=0.5, order=1),  # Matern-3/2
    times,
)
log_marg = prior.log_marginal(y, jnp.asarray(0.01))     # Kalman forward
cond     = prior.condition(y, jnp.asarray(0.01))        # filter + RTS smoother
mean, var = cond.predict(jnp.linspace(-0.5, 6.0, 50))   # arbitrary test times

Inside a NumPyro model, swap gp_factor for markov_gp_factor:

import jax.numpy as jnp
import numpyro
from numpyro import distributions as dist
from pyrox_gp import MarkovGPPrior, MaternSDE, markov_gp_factor

def temporal_model(times, y):
    sigma2 = numpyro.sample("variance",  dist.LogNormal(0.0, 1.0))
    ell    = numpyro.sample("lengthscale", dist.LogNormal(0.0, 1.0))
    sde    = MaternSDE(variance=sigma2, lengthscale=ell, order=1)
    prior  = MarkovGPPrior(sde, times)
    markov_gp_factor("obs", prior, y, jnp.array(0.01))

For non-Gaussian likelihoods on the Markov path, see the Markov non-Gaussian strategies below; for inducing-grid scalability, the sparse Markov GP.

MarkovGPPrior

Bases: Module

Linear-time temporal GP prior over a sorted 1-D grid.

Wraps any pyrox_gp.SDEKernel (e.g. pyrox_gp.MaternSDE, pyrox_gp.SumSDE, pyrox_gp.PeriodicSDE) to give Kalman filtering for the marginal log-likelihood and RTS smoothing for the posterior on the training grid. Supports an optional mean function and a small observation-noise floor for numerical stability.

Attributes:

Name Type Description
sde_kernel SDEKernel

Any SDEKernel. Provides (F, L, H, Q_c, P_inf) via sde_params() and the discrete transition sequence via discretise_sequence(dt).

times Float[Array, ' N']

Sorted, strictly increasing observation times of shape (N,). Concrete (non-traced) times arrays are validated for monotonicity at construction time; under jax.jit / SVI / MCMC the input is a tracer and the check is silently skipped — callers must guarantee monotonicity in that case.

mean_fn Callable[[Float[Array, ' N']], Float[Array, ' N']] | None

Optional callable mapping times -> (N,) mean values. Defaults to the zero mean. The mean is subtracted from observations before filtering and added back at predict time.

obs_noise_floor float

Small extra diagonal added to the observation variance R = noise_var + obs_noise_floor for stability when noise_var is near zero. Defaults to 0.0.

init_cov Float[Array, 'd d'] | None

Optional (d, d) covariance of the state at the first time point. None (the default) asks the kernel via initial_covariance(), which is \(P_\infty\) for a stationary kernel — so nothing changes for the kernels that have one — and an explicit diffuse covariance for a non-stationary one such as gaussx.IntegratedWienerSDE. Supply it to state what is known about the level (and its derivatives) before the first observation. Worth supplying for the site-based non-Gaussian strategies in particular: pyrox_gp.PosteriorLinearizationMarkov and pyrox_gp.ExpectationPropagationMarkov seed their cavities from the prior marginal variance, and a very diffuse prior can overflow their quadrature under an exponential-link likelihood — as a stationary kernel with an equally large variance does.

Examples:

>>> import jax.numpy as jnp
>>> from pyrox_gp import MaternSDE, MarkovGPPrior
>>> times = jnp.linspace(0.0, 5.0, 50)
>>> sde = MaternSDE(variance=1.0, lengthscale=0.5, order=1)
>>> prior = MarkovGPPrior(sde, times)
>>> y = jnp.sin(times) + 0.05 * jnp.cos(3.0 * times)
>>> log_marg = prior.log_marginal(y, jnp.asarray(0.01))
Notes

The solver-strategy plumbing used by pyrox_gp.GPPrior does not apply here — Kalman filtering is its own linear-algebra path and does not factor through gaussx.AbstractSolverStrategy.

Source code in packages/pyrox-gp/src/pyrox_gp/_markov.py
class MarkovGPPrior(eqx.Module):
    r"""Linear-time temporal GP prior over a sorted 1-D grid.

    Wraps any `pyrox_gp.SDEKernel` (e.g. `pyrox_gp.MaternSDE`,
    `pyrox_gp.SumSDE`, `pyrox_gp.PeriodicSDE`) to give Kalman
    filtering for the marginal log-likelihood and RTS smoothing for the
    posterior on the training grid. Supports an optional mean function and a
    small observation-noise floor for numerical stability.

    Attributes:
        sde_kernel: Any `SDEKernel`. Provides ``(F, L, H, Q_c, P_inf)``
            via ``sde_params()`` and the discrete transition sequence via
            ``discretise_sequence(dt)``.
        times: Sorted, strictly increasing observation times of shape
            ``(N,)``. Concrete (non-traced) ``times`` arrays are validated
            for monotonicity at construction time; under `jax.jit` /
            SVI / MCMC the input is a tracer and the check is silently
            skipped — callers must guarantee monotonicity in that case.
        mean_fn: Optional callable mapping ``times -> (N,)`` mean values.
            Defaults to the zero mean. The mean is subtracted from
            observations before filtering and added back at predict time.
        obs_noise_floor: Small extra diagonal added to the observation
            variance ``R = noise_var + obs_noise_floor`` for stability when
            ``noise_var`` is near zero. Defaults to ``0.0``.
        init_cov: Optional ``(d, d)`` covariance of the state at the first
            time point. ``None`` (the default) asks the kernel via
            ``initial_covariance()``, which is $P_\infty$ for a stationary
            kernel — so nothing changes for the kernels that have one — and
            an explicit diffuse covariance for a non-stationary one such as
            `gaussx.IntegratedWienerSDE`. Supply it to state what is known
            about the level (and its derivatives) before the first
            observation. Worth supplying for the site-based non-Gaussian
            strategies in particular: `pyrox_gp.PosteriorLinearizationMarkov`
            and `pyrox_gp.ExpectationPropagationMarkov` seed their cavities
            from the prior marginal variance, and a very diffuse prior can
            overflow their quadrature under an exponential-link likelihood
            — as a stationary kernel with an equally large variance does.

    Examples:
        >>> import jax.numpy as jnp
        >>> from pyrox_gp import MaternSDE, MarkovGPPrior
        >>> times = jnp.linspace(0.0, 5.0, 50)
        >>> sde = MaternSDE(variance=1.0, lengthscale=0.5, order=1)
        >>> prior = MarkovGPPrior(sde, times)
        >>> y = jnp.sin(times) + 0.05 * jnp.cos(3.0 * times)
        >>> log_marg = prior.log_marginal(y, jnp.asarray(0.01))

    Notes:
        The solver-strategy plumbing used by `pyrox_gp.GPPrior` does
        not apply here — Kalman filtering is its own linear-algebra path
        and does not factor through ``gaussx.AbstractSolverStrategy``.
    """

    sde_kernel: SDEKernel
    times: Float[Array, " N"]
    mean_fn: Callable[[Float[Array, " N"]], Float[Array, " N"]] | None = None
    obs_noise_floor: float = eqx.field(static=True, default=0.0)
    init_cov: Float[Array, "d d"] | None = None

    def __init__(
        self,
        sde_kernel: SDEKernel,
        times: Float[Array, " N"],
        mean_fn: Callable[[Float[Array, " N"]], Float[Array, " N"]] | None = None,
        obs_noise_floor: float = 0.0,
        init_cov: Float[Array, "d d"] | None = None,
    ) -> None:
        if obs_noise_floor < 0:
            raise ValueError(
                f"obs_noise_floor must be non-negative, got {obs_noise_floor!r}"
            )
        times_arr = jnp.asarray(times, dtype=jnp.result_type(times, 0.0))
        if times_arr.ndim != 1:
            raise ValueError(f"times must be 1-D, got shape {tuple(times_arr.shape)!r}")
        # Eager monotonicity check for concrete (non-traced) inputs only.
        # Under ``jax.jit`` / SVI / similar transforms ``times`` may arrive as
        # a tracer; the ``bool`` conversion would raise, so we silence that
        # path and let downstream Kalman steps trust the contract.
        if times_arr.shape[0] >= 2:
            try:
                if not bool(jnp.all(jnp.diff(times_arr) > 0)):
                    raise ValueError("times must be strictly increasing")
            except jax.errors.TracerBoolConversionError:
                pass
        if init_cov is not None:
            # Promoted like ``times``: an integer covariance would seed the
            # scan with an int carry and fail inside it on a dtype mismatch.
            init_cov = jnp.asarray(init_cov, dtype=jnp.result_type(init_cov, 0.0))
            d = sde_kernel.state_dim
            if init_cov.shape != (d, d):
                raise ValueError(
                    f"init_cov must be ({d}, {d}) for a state dimension of "
                    f"{d}, got shape {tuple(init_cov.shape)!r}"
                )
        self.sde_kernel = sde_kernel
        self.times = times_arr
        self.mean_fn = mean_fn
        self.obs_noise_floor = float(obs_noise_floor)
        self.init_cov = init_cov

    @property
    def state_dim(self) -> int:
        """SDE state dimension $d$ for this kernel."""
        return self.sde_kernel.state_dim

    def initial_covariance(self) -> Float[Array, "d d"]:
        r"""Covariance the Kalman recursions start from.

        The ``init_cov`` given at construction if there is one, otherwise
        whatever the kernel reports — $P_\infty$ for a stationary kernel,
        an explicit diffuse covariance for a non-stationary one. This is
        the single seed every filtering surface here uses, so a kernel
        with no $P_\infty$ is only a problem for the paths that are
        *defined* in terms of it (`log_prob`, the dense Gram).
        """
        if self.init_cov is not None:
            return self.init_cov
        return _kernel_initial_covariance(self.sde_kernel)

    def mean(self, times: Float[Array, " M"]) -> Float[Array, " M"]:
        """Evaluate the mean function at ``times``; zero by default."""
        if self.mean_fn is None:
            return jnp.zeros_like(times)
        return self.mean_fn(times)

    def _residual(self, y: Float[Array, " N"]) -> Float[Array, " N"]:
        return y - self.mean(self.times)

    def _R(self, noise_var: Float[Array, ""]) -> Float[Array, ""]:
        return jnp.asarray(noise_var) + jnp.asarray(self.obs_noise_floor)

    def filter(
        self,
        y: Float[Array, " N"],
        noise_var: Float[Array, ""],
    ) -> tuple[
        Float[Array, "N d"],
        Float[Array, "N d d"],
        Float[Array, "N d"],
        Float[Array, "N d d"],
        Float[Array, ""],
    ]:
        """Run the forward Kalman filter on the training grid.

        Returns:
            Tuple ``(m_pred, P_pred, m_filt, P_filt, log_marginal)`` where
            each ``*_pred`` / ``*_filt`` is shaped ``(N, d)`` or
            ``(N, d, d)`` and ``log_marginal`` is the scalar log-likelihood
            ``log p(y | theta)``.
        """
        F, _L, H, _Qc, _P_inf = self.sde_kernel.sde_params()
        init_cov = self.initial_covariance()
        dt_full = _build_dt_full(self.times)
        A_seq, Q_seq = self.sde_kernel.discretise_sequence(dt_full)
        residual = self._residual(y)
        mask = jnp.ones_like(self.times)
        R_seq = jnp.broadcast_to(self._R(noise_var), self.times.shape)
        return _kalman_filter(F, H, init_cov, A_seq, Q_seq, residual, mask, R_seq)

    def log_marginal(
        self,
        y: Float[Array, " N"],
        noise_var: Float[Array, ""],
    ) -> Float[Array, ""]:
        r"""Marginal log-likelihood ``log p(y | theta)`` via Kalman filtering."""
        *_, log_marg = self.filter(y, noise_var)
        return log_marg

    def smooth(
        self,
        y: Float[Array, " N"],
        noise_var: Float[Array, ""],
    ) -> tuple[Float[Array, "N d"], Float[Array, "N d d"], Float[Array, ""]]:
        """Run filter + RTS smoother on the training grid.

        Returns ``(m_smooth, P_smooth, log_marginal)`` over the training
        times.
        """
        F, _L, H, _Qc, _P_inf = self.sde_kernel.sde_params()
        init_cov = self.initial_covariance()
        dt_full = _build_dt_full(self.times)
        A_seq, Q_seq = self.sde_kernel.discretise_sequence(dt_full)
        residual = self._residual(y)
        mask = jnp.ones_like(self.times)
        R_seq = jnp.broadcast_to(self._R(noise_var), self.times.shape)
        m_pred, P_pred, m_filt, P_filt, log_marg = _kalman_filter(
            F, H, init_cov, A_seq, Q_seq, residual, mask, R_seq
        )
        m_smooth, P_smooth = _rts_smoother(m_pred, P_pred, m_filt, P_filt, A_seq)
        return m_smooth, P_smooth, log_marg

    def condition(
        self,
        y: Float[Array, " N"],
        noise_var: Float[Array, ""],
    ) -> ConditionedMarkovGP:
        """Condition on Gaussian-likelihood observations via filter + smoother."""
        m_smooth, P_smooth, log_marg = self.smooth(y, noise_var)
        return ConditionedMarkovGP(
            prior=self,
            y=y,
            noise_var=jnp.asarray(noise_var),
            smoothed_means=m_smooth,
            smoothed_covs=P_smooth,
            log_marginal=log_marg,
        )

    def condition_nongauss(
        self,
        likelihood: Likelihood,
        y: Float[Array, " N"],
        *,
        strategy: _NonGaussMarkovStrategy,
    ) -> NonGaussConditionedMarkovGP:
        """Condition on a non-Gaussian likelihood via a site-based strategy.

        Convenience that forwards to ``strategy.fit(self, likelihood, y)``.
        Pick any of the Markov-aware site-based strategies in
        `pyrox_gp._inference_nongauss_markov`:
        `pyrox_gp.LaplaceMarkovInference`,
        `pyrox_gp.GaussNewtonMarkovInference`,
        `pyrox_gp.PosteriorLinearizationMarkov`, or
        `pyrox_gp.ExpectationPropagationMarkov`. Returns a
        `pyrox_gp.NonGaussConditionedMarkovGP` with the same
        ``predict`` API as the Gaussian-likelihood
        `ConditionedMarkovGP`.
        """
        return strategy.fit(self, likelihood, y)

    def log_prob(self, f: Float[Array, " N"]) -> Float[Array, ""]:
        r"""Log density of an exact-state path $f(t_n) = H x_n$ under the prior.

        Evaluates ``log N(f | mu(times), K_NN)`` where ``K_NN`` is the dense
        Gram of the kernel encoded by ``sde_kernel`` on ``self.times``.
        Computes the dense covariance via ``H exp(F |t_i - t_j|) P_inf H^T``
        — one ``expm`` per pairwise lag, costing $O(N^2 d^3)$ for the
        Gram plus $O(N^3)$ for the Cholesky solve — intended for
        sanity checks and small-grid use rather than scalable inference.
        For training, prefer `log_marginal`.
        """
        K = _dense_sde_gram(self.sde_kernel, self.times)
        cov_op = lx.MatrixLinearOperator(K, lx.positive_semidefinite_tag)
        return gaussx.gaussian_log_prob(self.mean(self.times), cov_op, f)

state_dim: int property

SDE state dimension \(d\) for this kernel.

initial_covariance() -> Float[Array, 'd d']

Covariance the Kalman recursions start from.

The init_cov given at construction if there is one, otherwise whatever the kernel reports — \(P_\infty\) for a stationary kernel, an explicit diffuse covariance for a non-stationary one. This is the single seed every filtering surface here uses, so a kernel with no \(P_\infty\) is only a problem for the paths that are defined in terms of it (log_prob, the dense Gram).

Source code in packages/pyrox-gp/src/pyrox_gp/_markov.py
def initial_covariance(self) -> Float[Array, "d d"]:
    r"""Covariance the Kalman recursions start from.

    The ``init_cov`` given at construction if there is one, otherwise
    whatever the kernel reports — $P_\infty$ for a stationary kernel,
    an explicit diffuse covariance for a non-stationary one. This is
    the single seed every filtering surface here uses, so a kernel
    with no $P_\infty$ is only a problem for the paths that are
    *defined* in terms of it (`log_prob`, the dense Gram).
    """
    if self.init_cov is not None:
        return self.init_cov
    return _kernel_initial_covariance(self.sde_kernel)

mean(times: Float[Array, ' M']) -> Float[Array, ' M']

Evaluate the mean function at times; zero by default.

Source code in packages/pyrox-gp/src/pyrox_gp/_markov.py
def mean(self, times: Float[Array, " M"]) -> Float[Array, " M"]:
    """Evaluate the mean function at ``times``; zero by default."""
    if self.mean_fn is None:
        return jnp.zeros_like(times)
    return self.mean_fn(times)

filter(y: Float[Array, ' N'], noise_var: Float[Array, '']) -> tuple[Float[Array, 'N d'], Float[Array, 'N d d'], Float[Array, 'N d'], Float[Array, 'N d d'], Float[Array, '']]

Run the forward Kalman filter on the training grid.

Returns:

Type Description
Float[Array, 'N d']

Tuple (m_pred, P_pred, m_filt, P_filt, log_marginal) where

Float[Array, 'N d d']

each *_pred / *_filt is shaped (N, d) or

Float[Array, 'N d']

(N, d, d) and log_marginal is the scalar log-likelihood

Float[Array, 'N d d']

log p(y | theta).

Source code in packages/pyrox-gp/src/pyrox_gp/_markov.py
def filter(
    self,
    y: Float[Array, " N"],
    noise_var: Float[Array, ""],
) -> tuple[
    Float[Array, "N d"],
    Float[Array, "N d d"],
    Float[Array, "N d"],
    Float[Array, "N d d"],
    Float[Array, ""],
]:
    """Run the forward Kalman filter on the training grid.

    Returns:
        Tuple ``(m_pred, P_pred, m_filt, P_filt, log_marginal)`` where
        each ``*_pred`` / ``*_filt`` is shaped ``(N, d)`` or
        ``(N, d, d)`` and ``log_marginal`` is the scalar log-likelihood
        ``log p(y | theta)``.
    """
    F, _L, H, _Qc, _P_inf = self.sde_kernel.sde_params()
    init_cov = self.initial_covariance()
    dt_full = _build_dt_full(self.times)
    A_seq, Q_seq = self.sde_kernel.discretise_sequence(dt_full)
    residual = self._residual(y)
    mask = jnp.ones_like(self.times)
    R_seq = jnp.broadcast_to(self._R(noise_var), self.times.shape)
    return _kalman_filter(F, H, init_cov, A_seq, Q_seq, residual, mask, R_seq)

log_marginal(y: Float[Array, ' N'], noise_var: Float[Array, '']) -> Float[Array, '']

Marginal log-likelihood log p(y | theta) via Kalman filtering.

Source code in packages/pyrox-gp/src/pyrox_gp/_markov.py
def log_marginal(
    self,
    y: Float[Array, " N"],
    noise_var: Float[Array, ""],
) -> Float[Array, ""]:
    r"""Marginal log-likelihood ``log p(y | theta)`` via Kalman filtering."""
    *_, log_marg = self.filter(y, noise_var)
    return log_marg

smooth(y: Float[Array, ' N'], noise_var: Float[Array, '']) -> tuple[Float[Array, 'N d'], Float[Array, 'N d d'], Float[Array, '']]

Run filter + RTS smoother on the training grid.

Returns (m_smooth, P_smooth, log_marginal) over the training times.

Source code in packages/pyrox-gp/src/pyrox_gp/_markov.py
def smooth(
    self,
    y: Float[Array, " N"],
    noise_var: Float[Array, ""],
) -> tuple[Float[Array, "N d"], Float[Array, "N d d"], Float[Array, ""]]:
    """Run filter + RTS smoother on the training grid.

    Returns ``(m_smooth, P_smooth, log_marginal)`` over the training
    times.
    """
    F, _L, H, _Qc, _P_inf = self.sde_kernel.sde_params()
    init_cov = self.initial_covariance()
    dt_full = _build_dt_full(self.times)
    A_seq, Q_seq = self.sde_kernel.discretise_sequence(dt_full)
    residual = self._residual(y)
    mask = jnp.ones_like(self.times)
    R_seq = jnp.broadcast_to(self._R(noise_var), self.times.shape)
    m_pred, P_pred, m_filt, P_filt, log_marg = _kalman_filter(
        F, H, init_cov, A_seq, Q_seq, residual, mask, R_seq
    )
    m_smooth, P_smooth = _rts_smoother(m_pred, P_pred, m_filt, P_filt, A_seq)
    return m_smooth, P_smooth, log_marg

condition(y: Float[Array, ' N'], noise_var: Float[Array, '']) -> ConditionedMarkovGP

Condition on Gaussian-likelihood observations via filter + smoother.

Source code in packages/pyrox-gp/src/pyrox_gp/_markov.py
def condition(
    self,
    y: Float[Array, " N"],
    noise_var: Float[Array, ""],
) -> ConditionedMarkovGP:
    """Condition on Gaussian-likelihood observations via filter + smoother."""
    m_smooth, P_smooth, log_marg = self.smooth(y, noise_var)
    return ConditionedMarkovGP(
        prior=self,
        y=y,
        noise_var=jnp.asarray(noise_var),
        smoothed_means=m_smooth,
        smoothed_covs=P_smooth,
        log_marginal=log_marg,
    )

condition_nongauss(likelihood: Likelihood, y: Float[Array, ' N'], *, strategy: _NonGaussMarkovStrategy) -> NonGaussConditionedMarkovGP

Condition on a non-Gaussian likelihood via a site-based strategy.

Convenience that forwards to strategy.fit(self, likelihood, y). Pick any of the Markov-aware site-based strategies in pyrox_gp._inference_nongauss_markov: pyrox_gp.LaplaceMarkovInference, pyrox_gp.GaussNewtonMarkovInference, pyrox_gp.PosteriorLinearizationMarkov, or pyrox_gp.ExpectationPropagationMarkov. Returns a pyrox_gp.NonGaussConditionedMarkovGP with the same predict API as the Gaussian-likelihood ConditionedMarkovGP.

Source code in packages/pyrox-gp/src/pyrox_gp/_markov.py
def condition_nongauss(
    self,
    likelihood: Likelihood,
    y: Float[Array, " N"],
    *,
    strategy: _NonGaussMarkovStrategy,
) -> NonGaussConditionedMarkovGP:
    """Condition on a non-Gaussian likelihood via a site-based strategy.

    Convenience that forwards to ``strategy.fit(self, likelihood, y)``.
    Pick any of the Markov-aware site-based strategies in
    `pyrox_gp._inference_nongauss_markov`:
    `pyrox_gp.LaplaceMarkovInference`,
    `pyrox_gp.GaussNewtonMarkovInference`,
    `pyrox_gp.PosteriorLinearizationMarkov`, or
    `pyrox_gp.ExpectationPropagationMarkov`. Returns a
    `pyrox_gp.NonGaussConditionedMarkovGP` with the same
    ``predict`` API as the Gaussian-likelihood
    `ConditionedMarkovGP`.
    """
    return strategy.fit(self, likelihood, y)

log_prob(f: Float[Array, ' N']) -> Float[Array, '']

Log density of an exact-state path \(f(t_n) = H x_n\) under the prior.

Evaluates log N(f | mu(times), K_NN) where K_NN is the dense Gram of the kernel encoded by sde_kernel on self.times. Computes the dense covariance via H exp(F |t_i - t_j|) P_inf H^T — one expm per pairwise lag, costing \(O(N^2 d^3)\) for the Gram plus \(O(N^3)\) for the Cholesky solve — intended for sanity checks and small-grid use rather than scalable inference. For training, prefer log_marginal.

Source code in packages/pyrox-gp/src/pyrox_gp/_markov.py
def log_prob(self, f: Float[Array, " N"]) -> Float[Array, ""]:
    r"""Log density of an exact-state path $f(t_n) = H x_n$ under the prior.

    Evaluates ``log N(f | mu(times), K_NN)`` where ``K_NN`` is the dense
    Gram of the kernel encoded by ``sde_kernel`` on ``self.times``.
    Computes the dense covariance via ``H exp(F |t_i - t_j|) P_inf H^T``
    — one ``expm`` per pairwise lag, costing $O(N^2 d^3)$ for the
    Gram plus $O(N^3)$ for the Cholesky solve — intended for
    sanity checks and small-grid use rather than scalable inference.
    For training, prefer `log_marginal`.
    """
    K = _dense_sde_gram(self.sde_kernel, self.times)
    cov_op = lx.MatrixLinearOperator(K, lx.positive_semidefinite_tag)
    return gaussx.gaussian_log_prob(self.mean(self.times), cov_op, f)

ConditionedMarkovGP

Bases: Module

Markov GP conditioned on Gaussian-likelihood observations.

Holds the smoothed posterior on the training grid plus the marginal log-likelihood. Use predict for marginal posterior mean / variance at arbitrary test times.

Attributes:

Name Type Description
prior MarkovGPPrior

The originating MarkovGPPrior.

y Float[Array, ' N']

Observations of shape (N,).

noise_var Float[Array, '']

Observation variance used for conditioning.

smoothed_means Float[Array, 'N d']

(N, d) smoothed state means at training times.

smoothed_covs Float[Array, 'N d d']

(N, d, d) smoothed state covariances at training times.

log_marginal Float[Array, '']

Scalar \(\log p(y \mid \theta)\).

Source code in packages/pyrox-gp/src/pyrox_gp/_markov.py
class ConditionedMarkovGP(eqx.Module):
    """Markov GP conditioned on Gaussian-likelihood observations.

    Holds the smoothed posterior on the training grid plus the marginal
    log-likelihood. Use `predict` for marginal posterior mean / variance
    at arbitrary test times.

    Attributes:
        prior: The originating `MarkovGPPrior`.
        y: Observations of shape ``(N,)``.
        noise_var: Observation variance used for conditioning.
        smoothed_means: ``(N, d)`` smoothed state means at training times.
        smoothed_covs: ``(N, d, d)`` smoothed state covariances at training
            times.
        log_marginal: Scalar $\\log p(y \\mid \\theta)$.
    """

    prior: MarkovGPPrior
    y: Float[Array, " N"]
    noise_var: Float[Array, ""]
    smoothed_means: Float[Array, "N d"]
    smoothed_covs: Float[Array, "N d d"]
    log_marginal: Float[Array, ""]

    def predict(
        self,
        t_star: Float[Array, " M"],
    ) -> tuple[Float[Array, " M"], Float[Array, " M"]]:
        r"""Predictive marginals ``(mean, var)`` at arbitrary test times.

        Implementation: re-run the filter+smoother over the merged grid
        ``sort(times \\cup t_star)`` with the test points masked out of the
        update step, then read off the smoothed marginals at the test
        positions via ``H @ m`` and ``H @ P @ H^T``. Cost is
        $O((N + M)\\,d^3)$. Handles training-grid lookups, forecasting,
        backcasting, and within-window interpolation under one code path.
        """
        F, _L, H, _Qc, _P_inf = self.prior.sde_kernel.sde_params()
        init_cov = self.prior.initial_covariance()
        times = self.prior.times
        t_star = jnp.asarray(t_star)

        N = times.shape[0]
        M = t_star.shape[0]
        merged = jnp.concatenate([times, t_star], axis=0)
        # Stable sort so the relative ordering of identical times is preserved
        # (training point sorts before a duplicate test point, so the test
        # point still sees the observation update earlier in the grid).
        order = jnp.argsort(merged, stable=True)
        merged_sorted = merged[order]

        is_obs = jnp.concatenate(
            [jnp.ones(N, dtype=times.dtype), jnp.zeros(M, dtype=times.dtype)]
        )[order]
        residual_full = jnp.concatenate(
            [self.y - self.prior.mean(times), jnp.zeros(M, dtype=self.y.dtype)]
        )[order]

        dt_full = _build_dt_full(merged_sorted)
        A_seq, Q_seq = self.prior.sde_kernel.discretise_sequence(dt_full)
        R_seq = jnp.broadcast_to(self.prior._R(self.noise_var), merged_sorted.shape)
        m_pred, P_pred, m_filt, P_filt, _ = _kalman_filter(
            F, H, init_cov, A_seq, Q_seq, residual_full, is_obs, R_seq
        )
        m_smooth, P_smooth = _rts_smoother(m_pred, P_pred, m_filt, P_filt, A_seq)

        # Inverse permutation: position in the sorted grid for each original
        # entry, then slice off the trailing M test entries.
        inv_order = jnp.argsort(order, stable=True)
        test_positions = inv_order[N:]
        m_test_state = m_smooth[test_positions]  # (M, d)
        P_test_state = P_smooth[test_positions]  # (M, d, d)
        means = (m_test_state @ H.T)[:, 0] + self.prior.mean(t_star)
        # var = H P H^T per test point — vmap over axis 0
        vars_ = jax.vmap(lambda P: (H @ P @ H.T)[0, 0])(P_test_state)
        return means, vars_

predict(t_star: Float[Array, ' M']) -> tuple[Float[Array, ' M'], Float[Array, ' M']]

Predictive marginals (mean, var) at arbitrary test times.

Implementation: re-run the filter+smoother over the merged grid sort(times \\cup t_star) with the test points masked out of the update step, then read off the smoothed marginals at the test positions via H @ m and H @ P @ H^T. Cost is \(O((N + M)\\,d^3)\). Handles training-grid lookups, forecasting, backcasting, and within-window interpolation under one code path.

Source code in packages/pyrox-gp/src/pyrox_gp/_markov.py
def predict(
    self,
    t_star: Float[Array, " M"],
) -> tuple[Float[Array, " M"], Float[Array, " M"]]:
    r"""Predictive marginals ``(mean, var)`` at arbitrary test times.

    Implementation: re-run the filter+smoother over the merged grid
    ``sort(times \\cup t_star)`` with the test points masked out of the
    update step, then read off the smoothed marginals at the test
    positions via ``H @ m`` and ``H @ P @ H^T``. Cost is
    $O((N + M)\\,d^3)$. Handles training-grid lookups, forecasting,
    backcasting, and within-window interpolation under one code path.
    """
    F, _L, H, _Qc, _P_inf = self.prior.sde_kernel.sde_params()
    init_cov = self.prior.initial_covariance()
    times = self.prior.times
    t_star = jnp.asarray(t_star)

    N = times.shape[0]
    M = t_star.shape[0]
    merged = jnp.concatenate([times, t_star], axis=0)
    # Stable sort so the relative ordering of identical times is preserved
    # (training point sorts before a duplicate test point, so the test
    # point still sees the observation update earlier in the grid).
    order = jnp.argsort(merged, stable=True)
    merged_sorted = merged[order]

    is_obs = jnp.concatenate(
        [jnp.ones(N, dtype=times.dtype), jnp.zeros(M, dtype=times.dtype)]
    )[order]
    residual_full = jnp.concatenate(
        [self.y - self.prior.mean(times), jnp.zeros(M, dtype=self.y.dtype)]
    )[order]

    dt_full = _build_dt_full(merged_sorted)
    A_seq, Q_seq = self.prior.sde_kernel.discretise_sequence(dt_full)
    R_seq = jnp.broadcast_to(self.prior._R(self.noise_var), merged_sorted.shape)
    m_pred, P_pred, m_filt, P_filt, _ = _kalman_filter(
        F, H, init_cov, A_seq, Q_seq, residual_full, is_obs, R_seq
    )
    m_smooth, P_smooth = _rts_smoother(m_pred, P_pred, m_filt, P_filt, A_seq)

    # Inverse permutation: position in the sorted grid for each original
    # entry, then slice off the trailing M test entries.
    inv_order = jnp.argsort(order, stable=True)
    test_positions = inv_order[N:]
    m_test_state = m_smooth[test_positions]  # (M, d)
    P_test_state = P_smooth[test_positions]  # (M, d, d)
    means = (m_test_state @ H.T)[:, 0] + self.prior.mean(t_star)
    # var = H P H^T per test point — vmap over axis 0
    vars_ = jax.vmap(lambda P: (H @ P @ H.T)[0, 0])(P_test_state)
    return means, vars_

markov_gp_factor(name: str, prior: MarkovGPPrior, y: Float[Array, ' N'], noise_var: Float[Array, '']) -> None

Register the collapsed Markov-GP marginal log-likelihood with NumPyro.

Computes log p(y | times, theta) via Kalman filtering and adds it as numpyro.factor(name, ...). Use this inside a NumPyro model for Gaussian-likelihood temporal GP regression — the latent function is marginalized analytically.

Source code in packages/pyrox-gp/src/pyrox_gp/_markov.py
def markov_gp_factor(
    name: str,
    prior: MarkovGPPrior,
    y: Float[Array, " N"],
    noise_var: Float[Array, ""],
) -> None:
    """Register the collapsed Markov-GP marginal log-likelihood with NumPyro.

    Computes ``log p(y | times, theta)`` via Kalman filtering and adds it as
    ``numpyro.factor(name, ...)``. Use this inside a NumPyro model for
    Gaussian-likelihood temporal GP regression — the latent function is
    marginalized analytically.
    """
    numpyro.factor(name, prior.log_marginal(y, noise_var))

markov_gp_sample(name: str, prior: MarkovGPPrior) -> Float[Array, ' N']

Sample a latent function f at the prior's training times.

Registers a single numpyro.sample(name, MVN(mu, K)) site where K is the dense Gram derived from the SDE autocovariance H exp(F|tau|) P_inf H^T. This is the simple, dense path — use it when N is small. Scalable Markov-aware sample sites land in a later wave alongside non-Gaussian likelihood support.

Source code in packages/pyrox-gp/src/pyrox_gp/_markov.py
def markov_gp_sample(
    name: str,
    prior: MarkovGPPrior,
) -> Float[Array, " N"]:
    """Sample a latent function ``f`` at the prior's training times.

    Registers a single ``numpyro.sample(name, MVN(mu, K))`` site where ``K``
    is the dense Gram derived from the SDE autocovariance
    ``H exp(F|tau|) P_inf H^T``. This is the simple, dense path — use it
    when ``N`` is small. Scalable Markov-aware sample sites land in a
    later wave alongside non-Gaussian likelihood support.
    """
    K = _dense_sde_gram(prior.sde_kernel, prior.times)
    mu = prior.mean(prior.times)
    return numpyro.sample(  # ty: ignore[invalid-return-type]
        name, dist.MultivariateNormal(mu, covariance_matrix=K)
    )

Normalizing Kalman Filter

NormalizingKalmanPrior wraps a gaussx.LGSSM (or gaussx.MaskedLGSSM) base and an optional per-timestep observation warp into the same model surface: exact collapsed marginal via log_marginal, NumPyro registration via normalizing_kalman_factor, and observation-space predictive moments via predict (RTS smoothing followed by a Gauss-Hermite pushforward through the warp — the mean is E[G(z)], not G(E[z])).

The unwarped model (warp=None) works with the base install — LGSSM is a hard dependency. Passing a warp requires the flows extra: pip install 'pyrox-gp[flows]'. Because the warp acts on observations, the log-det term is independent of the latent state, the Kalman recursion stays exact, and none of the non-Gaussian Markov strategies below are involved.

import jax.numpy as jnp
import numpyro
from numpyro import distributions as dist
from gaussx import LGSSM
from pyrox_gp import NormalizingKalmanPrior, normalizing_kalman_factor

def nkf_model(y, warp=None, mask=None):
    T, M = y.shape
    log_q = numpyro.sample("log_q", dist.Normal(0.0, 1.0).expand([M]).to_event(1))
    log_r = numpyro.sample("log_r", dist.Normal(0.0, 1.0).expand([M]).to_event(1))
    base = LGSSM(0.9 * jnp.eye(M), jnp.eye(M),
                 jnp.diag(jnp.exp(jnp.asarray(log_q))),
                 jnp.diag(jnp.exp(jnp.asarray(log_r))),
                 jnp.zeros(M), jnp.eye(M), n_steps=T)
    prior = NormalizingKalmanPrior(base, warp=warp)
    normalizing_kalman_factor("nkf", prior, y, mask)

NormalizingKalmanPrior

Bases: Module

Normalizing Kalman Filter prior over a multivariate time grid.

Wraps a gaussx.LGSSM (or gaussx.MaskedLGSSM) base and an optional per-timestep observation warp into the pyrox_gp model surface, so the state-space parameters and the warp can carry NumPyro priors and be fitted with any pyrox.inference driver (pyrox.inference.EnsembleMAP, pyrox.inference.EnsembleVI, SVI, MCMC).

Because the warp acts on observations rather than on the latent state, the log-determinant term is independent of \(x_t\) and the Kalman recursion stays exact — the marginal likelihood is closed-form and no non-Gaussian inference strategy is required. With warp=None this reduces exactly to the base LGSSM marginal likelihood, with no flow dependencies at runtime.

The state-space parameters \((A, H, Q, R)\) are rotationally non-identifiable in the same way a latent-factor mixing matrix is, so ensembling over seeds (pyrox.inference.EnsembleMAP) is the recommended fitting pattern rather than a convenience.

Attributes:

Name Type Description
base LGSSM

gaussx.LGSSM or gaussx.MaskedLGSSM with event shape (T, M). Available without the flows extra. A gaussx.MaskedLGSSM base contributes its obs_mask whenever no call-time mask is given.

warp AbstractBijection | None

Optional bijection with event shape (M,), applied independently at each step (requires the flows extra). None means the identity, in which case this is exactly the base LGSSM. Conditional warps (non-None cond_shape) are rejected. A channel-mixing warp is usable with log_marginal on an unmasked base, but nothing else: gauss_flows refuses it over a masked base (masking and a non-diagonal Jacobian do not commute), and predict refuses it because its per-channel quadrature would be wrong (see there). Put cross-channel structure in H and R instead.

mean_fn Callable[[Float[Array, ' T']], Float[Array, 'T M']] | None

Optional callable times -> (T, M) evaluated on the integer grid 0, ..., T-1 of the base. The mean acts in observation space — it is subtracted from y before the inverse warp and filtering, and added back at predict time — so the model is \(y_t = m(t) + G(H x_t + r_t)\) and conjugacy is untouched.

Examples:

>>> import jax.numpy as jnp, jax.random as jr
>>> from gaussx import LGSSM
>>> from pyrox_gp import NormalizingKalmanPrior
>>> T, M = 12, 2
>>> base = LGSSM(0.9 * jnp.eye(M), jnp.eye(M), 0.1 * jnp.eye(M),
...              0.2 * jnp.eye(M), jnp.zeros(M), jnp.eye(M), n_steps=T)
>>> prior = NormalizingKalmanPrior(base)   # unwarped: no flow deps
>>> y = base.sample(jr.key(0))
>>> prior.log_marginal(y).shape
()
>>> mean, var = prior.predict(y, n_ahead=3)
>>> mean.shape
(15, 2)
Source code in packages/pyrox-gp/src/pyrox_gp/_markov_flow.py
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
class NormalizingKalmanPrior(eqx.Module):
    r"""Normalizing Kalman Filter prior over a multivariate time grid.

    Wraps a `gaussx.LGSSM` (or `gaussx.MaskedLGSSM`) base and an
    optional per-timestep observation warp into the ``pyrox_gp`` model
    surface, so the state-space parameters and the warp can carry NumPyro
    priors and be fitted with any ``pyrox.inference`` driver
    (`pyrox.inference.EnsembleMAP`, `pyrox.inference.EnsembleVI`, SVI,
    MCMC).

    Because the warp acts on observations rather than on the latent
    state, the log-determinant term is independent of $x_t$ and the
    Kalman recursion stays exact — the marginal likelihood is
    closed-form and no non-Gaussian inference strategy is required.
    With ``warp=None`` this reduces *exactly* to the base LGSSM marginal
    likelihood, with no flow dependencies at runtime.

    The state-space parameters $(A, H, Q, R)$ are rotationally
    non-identifiable in the same way a latent-factor mixing matrix is,
    so ensembling over seeds (`pyrox.inference.EnsembleMAP`) is the
    recommended fitting pattern rather than a convenience.

    Attributes:
        base: `gaussx.LGSSM` or `gaussx.MaskedLGSSM` with event shape
            ``(T, M)``. Available without the ``flows`` extra. A
            `gaussx.MaskedLGSSM` base contributes its ``obs_mask``
            whenever no call-time ``mask`` is given.
        warp: Optional bijection with event shape ``(M,)``, applied
            independently at each step (requires the ``flows`` extra).
            ``None`` means the identity, in which case this is exactly
            the base LGSSM. Conditional warps (non-``None``
            ``cond_shape``) are rejected. A channel-mixing warp is
            usable with `log_marginal` on an unmasked base, but nothing
            else: ``gauss_flows`` refuses it over a masked base (masking
            and a non-diagonal Jacobian do not commute), and `predict`
            refuses it because its per-channel quadrature would be wrong
            (see there). Put cross-channel structure in ``H`` and ``R``
            instead.
        mean_fn: Optional callable ``times -> (T, M)`` evaluated on the
            integer grid ``0, ..., T-1`` of the base. The mean acts in
            **observation space** — it is subtracted from ``y`` before
            the inverse warp and filtering, and added back at predict
            time — so the model is $y_t = m(t) + G(H x_t + r_t)$ and
            conjugacy is untouched.

    Examples:
        >>> import jax.numpy as jnp, jax.random as jr
        >>> from gaussx import LGSSM
        >>> from pyrox_gp import NormalizingKalmanPrior
        >>> T, M = 12, 2
        >>> base = LGSSM(0.9 * jnp.eye(M), jnp.eye(M), 0.1 * jnp.eye(M),
        ...              0.2 * jnp.eye(M), jnp.zeros(M), jnp.eye(M), n_steps=T)
        >>> prior = NormalizingKalmanPrior(base)   # unwarped: no flow deps
        >>> y = base.sample(jr.key(0))
        >>> prior.log_marginal(y).shape
        ()
        >>> mean, var = prior.predict(y, n_ahead=3)
        >>> mean.shape
        (15, 2)
    """

    base: LGSSM
    warp: AbstractBijection | None = None
    mean_fn: Callable[[Float[Array, " T"]], Float[Array, "T M"]] | None = None

    def __init__(
        self,
        base: LGSSM,
        warp: AbstractBijection | None = None,
        mean_fn: Callable[[Float[Array, " T"]], Float[Array, "T M"]] | None = None,
    ) -> None:
        if warp is not None:
            _require_flows()
            if warp.cond_shape is not None:
                raise ValueError(
                    "Conditional warps are not supported: the collapsed "
                    "marginal assumes one fixed observation warp per channel. "
                    f"Got cond_shape={warp.cond_shape!r}."
                )
            n_channels = base.event_shape[1]
            if len(warp.shape) != 1 or warp.shape[0] != n_channels:
                raise ValueError(
                    f"warp must have event shape ({n_channels},) to match the "
                    f"base's channel count; got {warp.shape!r}. Lift scalar "
                    "bijections over the channel axis first, e.g. "
                    "Vmap(RationalQuadraticSpline(...), in_axes=None, "
                    f"axis_size={n_channels})."
                )
        self.base = base
        self.warp = warp
        self.mean_fn = mean_fn

    @property
    def n_steps(self) -> int:
        """Sequence length ``T`` of the base."""
        return self.base.event_shape[0]

    @property
    def n_channels(self) -> int:
        """Observation dimension ``M`` of the base."""
        return self.base.event_shape[1]

    def _mean_grid(self, n_steps: int, dtype: np.dtype) -> Float[Array, "T M"]:
        """Mean values on the integer grid ``0, ..., n_steps - 1``."""
        if self.mean_fn is None:
            return jnp.zeros((n_steps, self.n_channels), dtype=dtype)
        return jnp.asarray(self.mean_fn(jnp.arange(n_steps, dtype=dtype)))

    def _effective_mask(
        self, mask: Bool[Array, "T M"] | None
    ) -> Bool[Array, "T M"] | None:
        """Call-time mask if given, else a `MaskedLGSSM` base's own mask."""
        if mask is not None:
            return jnp.asarray(mask, dtype=bool)
        if isinstance(self.base, MaskedLGSSM):
            return self.base.obs_mask
        return None

    def _check_shapes(
        self,
        y: Float[Array, "T M"],
        mask: Bool[Array, "T M"] | None,
    ) -> None:
        """Reject ``y`` / ``mask`` that do not match the base event shape.

        ``y - mean_grid`` and the filter's own broadcasting would other-
        wise accept a ``(T, 1)`` series against an ``(T, M)`` base,
        silently replicating the one observed channel across all ``M``
        and returning a finite log-likelihood for data the caller never
        supplied. Shapes are static under ``jit``, so this costs nothing
        at trace time.
        """
        expected = (self.n_steps, self.n_channels)
        if y.shape != expected:
            raise ValueError(
                f"y has shape {y.shape}, but the base event shape is "
                f"{expected}. Observations must match the base exactly — "
                "a mismatched channel or time axis would broadcast rather "
                "than raise."
            )
        if mask is not None and mask.shape != expected:
            raise ValueError(
                f"mask has shape {mask.shape}, but the base event shape "
                f"is {expected}. The mask marks observed entries of y, so "
                "it must have the same shape."
            )

    def _require_elementwise_warp(self, flows: ModuleType) -> None:
        """Raise unless ``gauss_flows`` classifies the warp as elementwise.

        Asked through the public surface rather than a private helper:
        `gauss_flows.normalizing_kalman_filter` documents that it
        refuses a channel-mixing warp over a **conditional** (mask-
        consuming) base, because masking and a non-diagonal Jacobian do
        not commute. Building that form is construction-only and cheap,
        so it doubles as the classifier — and it stays correct if
        ``gauss_flows`` refines what counts as elementwise. Shapes were
        already validated in ``__init__``, so a ``ValueError`` out of
        this constructor is that rejection.
        """
        try:
            self._nkf(flows, masked=True)
        except ValueError as exc:
            raise ValueError(
                "predict requires an elementwise warp: it pushes the "
                "per-channel marginal moments through the warp with a "
                "scalar Gauss-Hermite rule, which is only the right "
                "integral when each output channel depends on its own "
                "input channel alone. This warp mixes channels (or "
                "cannot be shown not to), so those moments would be "
                "silently wrong — log_marginal is unaffected and stays "
                "exact. Put cross-channel structure in H and R instead."
            ) from exc

    def _nkf(self, flows: ModuleType, *, masked: bool):
        """Build the ``gauss_flows`` NKF density for the current warp.

        Unmasked: wrap the base directly. Masked: rebuild the base's
        parameters into a `gaussx.LGSSMFactory` so the mask arrives as
        a flowjax ``condition`` — that is what routes the log-det through
        ``gauss_flows``' masked change-of-variables (observed channels
        only) and triggers its rejection of channel-mixing warps.
        """
        if not masked:
            return flows.normalizing_kalman_filter(
                flows.NumpyroBase(dist=self.base), self.warp
            )
        factory = gaussx.LGSSMFactory(
            self.base.A,
            self.base.H,
            self.base.Q,
            self.base.R,
            self.base.m0,
            self.base.P0,
            self.n_steps,
        )
        conditional_base = flows.NumpyroBase(
            dist_factory=factory,
            event_shape=tuple(self.base.event_shape),
            cond_shape=tuple(self.base.event_shape),
        )
        return flows.normalizing_kalman_filter(conditional_base, self.warp)

    def log_marginal(
        self,
        y: Float[Array, "T M"],
        mask: Bool[Array, "T M"] | None = None,
    ) -> Float[Array, ""]:
        r"""Exact log marginal likelihood, warp included.

        Computes $\log p(y) = \log p_{\mathrm{LGSSM}}(G^{-1}(y - m)) +
        \sum_t \log|\det \partial G^{-1}/\partial y_t|$ — a Kalman
        forward pass on the inverse-warped observations plus the summed
        log-determinant. The log-det does not depend on the latent path,
        so the result is exact, not a bound.

        Args:
            y: Observations, shape ``(T, M)``. Masked entries are never
                read and may be ``NaN``.
            mask: Optional observation mask, shape ``(T, M)``; ``True``
                marks an observed entry. Overrides a
                `gaussx.MaskedLGSSM` base's own mask when both are
                present. Unobserved channels are marginalised exactly;
                with a warp this requires an elementwise warp (see
                the class docstring).

        Returns:
            Scalar $\log p(y_{\mathrm{obs}} \mid \theta)$.

        Raises:
            ValueError: If ``y`` or ``mask`` does not match the base
                event shape.
        """
        y = jnp.asarray(y)
        mask_eff = self._effective_mask(mask)
        self._check_shapes(y, mask_eff)
        residual = y - self._mean_grid(y.shape[0], y.dtype)
        if self.warp is None:
            return gaussx.kalman_filter(
                self.base.A,
                self.base.H,
                self.base.Q,
                self.base.R,
                residual,
                self.base.m0,
                self.base.P0,
                mask=mask_eff,
            ).log_likelihood
        flows = _require_flows()
        nkf = self._nkf(flows, masked=mask_eff is not None)
        if mask_eff is None:
            return nkf.log_prob(residual)
        return nkf.log_prob(residual, condition=mask_eff)

    def predict(
        self,
        y: Float[Array, "T M"],
        mask: Bool[Array, "T M"] | None = None,
        *,
        n_ahead: int = 0,
        order: int = 32,
    ) -> tuple[Float[Array, "Tp M"], Float[Array, "Tp M"]]:
        r"""Predictive moments in **observation** space.

        Two stages. RTS smoothing gives Gaussian moments of the warped
        observation $z_t = H x_t + r_t \sim N(\mu_t, \sigma_t^2)$
        (per channel, noise included) on the training grid, extended
        ``n_ahead`` steps by open-loop propagation through $(A, Q)$.
        Those moments are then pushed through the warp by Gauss-Hermite
        quadrature:

        $$
        \mathbb{E}[y_t] = \int G(z)\, N(z; \mu_t, \sigma_t^2)\, dz,
        \qquad
        \mathrm{Var}[y_t] = \int G(z)^2 N(z; \mu_t, \sigma_t^2)\, dz
          - \mathbb{E}[y_t]^2 .
        $$

        The returned mean is $\mathbb{E}[G(z)]$, **not**
        $G(\mathbb{E}[z])$ — the latter is the pushforward median for a
        monotone warp and is badly biased for a skewed one.

        The quadrature is per-channel over the *marginal* warped-space
        moments, so it is the right integral only when each output
        channel depends on its own input channel alone. **This method
        therefore requires an elementwise warp** and raises otherwise —
        a channel-mixing warp's outputs depend on the full joint
        Gaussian, cross-channel covariances included, and per-channel
        moments would be silently wrong rather than merely imprecise.
        `log_marginal` carries no such restriction: it stays exact for
        any unmasked warp.

        Gauss-Hermite converges spectrally only for analytic
        integrands. A piecewise rational-quadratic spline warp is not
        analytic at its knots, so its quadrature error plateaus around
        ``~3e-3`` and can get *worse* with increasing order — the
        default ``order=32`` is the sweet spot for splines, and raising
        the order is not a convergence diagnostic.

        Args:
            y: Observations, shape ``(T, M)``. Masked entries are never
                read and may be ``NaN``.
            mask: Optional observation mask, shape ``(T, M)``; ``True``
                marks an observed entry. Same semantics as in
                `log_marginal`.
            n_ahead: Number of forecast steps appended after the
                training grid.
            order: Gauss-Hermite order for the warped pushforward.
                Ignored when ``warp is None``.

        Returns:
            Tuple ``(mean, var)`` of observation-space marginal moments,
            each of shape ``(T + n_ahead, M)``. Variances are clamped at
            zero, so ``sqrt`` on them is always safe.

        Raises:
            ValueError: If ``y`` or ``mask`` does not match the base
                event shape, or if the warp is not elementwise.
        """
        y = jnp.asarray(y)
        n_steps, n_channels = self.n_steps, self.n_channels
        mask_eff = self._effective_mask(mask)
        self._check_shapes(y, mask_eff)
        mean_grid = self._mean_grid(n_steps + n_ahead, y.dtype)
        residual = y - mean_grid[:n_steps]

        if self.warp is None:
            z = residual
        else:
            flows = _require_flows()
            # The scalar per-channel quadrature below is only the right
            # integral for an elementwise warp, so predict requires one
            # even when log_marginal would not.
            self._require_elementwise_warp(flows)
            if mask_eff is not None:
                # Unobserved slots hold junk (often NaN); substitute an
                # in-support reference before the inverse warp. The
                # filter never reads those entries afterwards.
                reference = self.warp.transform(jnp.zeros(n_channels, dtype=y.dtype))
                residual = jnp.where(mask_eff, residual, reference)
            z = jax.vmap(self.warp.inverse)(residual)  # (T, M)

        state = gaussx.kalman_filter(
            self.base.A,
            self.base.H,
            self.base.Q,
            self.base.R,
            z,
            self.base.m0,
            self.base.P0,
            mask=mask_eff,
        )
        m_smooth, P_smooth = gaussx.rts_smoother(state, self.base.A, self.base.Q)

        if n_ahead > 0:
            A_dense = _dense(self.base.A)
            Q_dense = _dense(self.base.Q)

            def step(carry, _):
                m, P = carry
                m_next = A_dense @ m
                P_next = A_dense @ P @ A_dense.T + Q_dense
                return (m_next, P_next), (m_next, P_next)

            last = (state.filtered_means[-1], state.filtered_covs[-1])
            _, (m_ahead, P_ahead) = jax.lax.scan(step, last, None, length=n_ahead)
            m_smooth = jnp.concatenate([m_smooth, m_ahead], axis=0)
            P_smooth = jnp.concatenate([P_smooth, P_ahead], axis=0)

        H_dense = _dense(self.base.H)
        R_diag = jnp.diagonal(_dense(self.base.R))
        # z-space marginals: mean H m_t, variance diag(H P_t H^T + R).
        mz = m_smooth @ H_dense.T  # (T', M)
        # einx.dot is typed as possibly returning a tuple (multi-output
        # patterns); narrow back to a single array for the typechecker.
        # Clamped at zero: the quadratic form is PSD in exact arithmetic,
        # but a rounding-negative entry would become NaN under the sqrt
        # taken for the quadrature nodes below.
        vz = jnp.maximum(
            jnp.asarray(einx.dot("m n, t n k, m k -> t m", H_dense, P_smooth, H_dense))
            + R_diag,
            0.0,
        )

        if self.warp is None:
            return mz + mean_grid, vz

        nodes, weights = np.polynomial.hermite_e.hermegauss(order)
        nodes = jnp.asarray(nodes, dtype=y.dtype)
        weights = jnp.asarray(weights, dtype=y.dtype) / np.sqrt(2.0 * np.pi)
        # fs: (order, T', M) — per-node evaluation points in the warped space.
        fs = mz[None] + jnp.sqrt(vz)[None] * nodes[:, None, None]
        g = jax.vmap(jax.vmap(self.warp.transform))(fs)
        m1 = jnp.asarray(einx.dot("s, s t m -> t m", weights, g))
        m2 = jnp.asarray(einx.dot("s, s t m -> t m", weights, g**2))
        # E[G^2] - E[G]^2 cancels catastrophically for a concentrated
        # predictive: the two moments agree to working precision and the
        # difference can land just below zero, which callers would turn
        # into NaN on the first sqrt.
        return m1 + mean_grid, jnp.maximum(m2 - m1**2, 0.0)

n_steps: int property

Sequence length T of the base.

n_channels: int property

Observation dimension M of the base.

log_marginal(y: Float[Array, 'T M'], mask: Bool[Array, 'T M'] | None = None) -> Float[Array, '']

Exact log marginal likelihood, warp included.

Computes \(\log p(y) = \log p_{\mathrm{LGSSM}}(G^{-1}(y - m)) + \sum_t \log|\det \partial G^{-1}/\partial y_t|\) — a Kalman forward pass on the inverse-warped observations plus the summed log-determinant. The log-det does not depend on the latent path, so the result is exact, not a bound.

Parameters:

Name Type Description Default
y Float[Array, 'T M']

Observations, shape (T, M). Masked entries are never read and may be NaN.

required
mask Bool[Array, 'T M'] | None

Optional observation mask, shape (T, M); True marks an observed entry. Overrides a gaussx.MaskedLGSSM base's own mask when both are present. Unobserved channels are marginalised exactly; with a warp this requires an elementwise warp (see the class docstring).

None

Returns:

Type Description
Float[Array, '']

Scalar \(\log p(y_{\mathrm{obs}} \mid \theta)\).

Raises:

Type Description
ValueError

If y or mask does not match the base event shape.

Source code in packages/pyrox-gp/src/pyrox_gp/_markov_flow.py
def log_marginal(
    self,
    y: Float[Array, "T M"],
    mask: Bool[Array, "T M"] | None = None,
) -> Float[Array, ""]:
    r"""Exact log marginal likelihood, warp included.

    Computes $\log p(y) = \log p_{\mathrm{LGSSM}}(G^{-1}(y - m)) +
    \sum_t \log|\det \partial G^{-1}/\partial y_t|$ — a Kalman
    forward pass on the inverse-warped observations plus the summed
    log-determinant. The log-det does not depend on the latent path,
    so the result is exact, not a bound.

    Args:
        y: Observations, shape ``(T, M)``. Masked entries are never
            read and may be ``NaN``.
        mask: Optional observation mask, shape ``(T, M)``; ``True``
            marks an observed entry. Overrides a
            `gaussx.MaskedLGSSM` base's own mask when both are
            present. Unobserved channels are marginalised exactly;
            with a warp this requires an elementwise warp (see
            the class docstring).

    Returns:
        Scalar $\log p(y_{\mathrm{obs}} \mid \theta)$.

    Raises:
        ValueError: If ``y`` or ``mask`` does not match the base
            event shape.
    """
    y = jnp.asarray(y)
    mask_eff = self._effective_mask(mask)
    self._check_shapes(y, mask_eff)
    residual = y - self._mean_grid(y.shape[0], y.dtype)
    if self.warp is None:
        return gaussx.kalman_filter(
            self.base.A,
            self.base.H,
            self.base.Q,
            self.base.R,
            residual,
            self.base.m0,
            self.base.P0,
            mask=mask_eff,
        ).log_likelihood
    flows = _require_flows()
    nkf = self._nkf(flows, masked=mask_eff is not None)
    if mask_eff is None:
        return nkf.log_prob(residual)
    return nkf.log_prob(residual, condition=mask_eff)

predict(y: Float[Array, 'T M'], mask: Bool[Array, 'T M'] | None = None, *, n_ahead: int = 0, order: int = 32) -> tuple[Float[Array, 'Tp M'], Float[Array, 'Tp M']]

Predictive moments in observation space.

Two stages. RTS smoothing gives Gaussian moments of the warped observation \(z_t = H x_t + r_t \sim N(\mu_t, \sigma_t^2)\) (per channel, noise included) on the training grid, extended n_ahead steps by open-loop propagation through \((A, Q)\). Those moments are then pushed through the warp by Gauss-Hermite quadrature:

\[ \mathbb{E}[y_t] = \int G(z)\, N(z; \mu_t, \sigma_t^2)\, dz, \qquad \mathrm{Var}[y_t] = \int G(z)^2 N(z; \mu_t, \sigma_t^2)\, dz - \mathbb{E}[y_t]^2 . \]

The returned mean is \(\mathbb{E}[G(z)]\), not \(G(\mathbb{E}[z])\) — the latter is the pushforward median for a monotone warp and is badly biased for a skewed one.

The quadrature is per-channel over the marginal warped-space moments, so it is the right integral only when each output channel depends on its own input channel alone. This method therefore requires an elementwise warp and raises otherwise — a channel-mixing warp's outputs depend on the full joint Gaussian, cross-channel covariances included, and per-channel moments would be silently wrong rather than merely imprecise. log_marginal carries no such restriction: it stays exact for any unmasked warp.

Gauss-Hermite converges spectrally only for analytic integrands. A piecewise rational-quadratic spline warp is not analytic at its knots, so its quadrature error plateaus around ~3e-3 and can get worse with increasing order — the default order=32 is the sweet spot for splines, and raising the order is not a convergence diagnostic.

Parameters:

Name Type Description Default
y Float[Array, 'T M']

Observations, shape (T, M). Masked entries are never read and may be NaN.

required
mask Bool[Array, 'T M'] | None

Optional observation mask, shape (T, M); True marks an observed entry. Same semantics as in log_marginal.

None
n_ahead int

Number of forecast steps appended after the training grid.

0
order int

Gauss-Hermite order for the warped pushforward. Ignored when warp is None.

32

Returns:

Type Description
Float[Array, 'Tp M']

Tuple (mean, var) of observation-space marginal moments,

Float[Array, 'Tp M']

each of shape (T + n_ahead, M). Variances are clamped at

tuple[Float[Array, 'Tp M'], Float[Array, 'Tp M']]

zero, so sqrt on them is always safe.

Raises:

Type Description
ValueError

If y or mask does not match the base event shape, or if the warp is not elementwise.

Source code in packages/pyrox-gp/src/pyrox_gp/_markov_flow.py
def predict(
    self,
    y: Float[Array, "T M"],
    mask: Bool[Array, "T M"] | None = None,
    *,
    n_ahead: int = 0,
    order: int = 32,
) -> tuple[Float[Array, "Tp M"], Float[Array, "Tp M"]]:
    r"""Predictive moments in **observation** space.

    Two stages. RTS smoothing gives Gaussian moments of the warped
    observation $z_t = H x_t + r_t \sim N(\mu_t, \sigma_t^2)$
    (per channel, noise included) on the training grid, extended
    ``n_ahead`` steps by open-loop propagation through $(A, Q)$.
    Those moments are then pushed through the warp by Gauss-Hermite
    quadrature:

    $$
    \mathbb{E}[y_t] = \int G(z)\, N(z; \mu_t, \sigma_t^2)\, dz,
    \qquad
    \mathrm{Var}[y_t] = \int G(z)^2 N(z; \mu_t, \sigma_t^2)\, dz
      - \mathbb{E}[y_t]^2 .
    $$

    The returned mean is $\mathbb{E}[G(z)]$, **not**
    $G(\mathbb{E}[z])$ — the latter is the pushforward median for a
    monotone warp and is badly biased for a skewed one.

    The quadrature is per-channel over the *marginal* warped-space
    moments, so it is the right integral only when each output
    channel depends on its own input channel alone. **This method
    therefore requires an elementwise warp** and raises otherwise —
    a channel-mixing warp's outputs depend on the full joint
    Gaussian, cross-channel covariances included, and per-channel
    moments would be silently wrong rather than merely imprecise.
    `log_marginal` carries no such restriction: it stays exact for
    any unmasked warp.

    Gauss-Hermite converges spectrally only for analytic
    integrands. A piecewise rational-quadratic spline warp is not
    analytic at its knots, so its quadrature error plateaus around
    ``~3e-3`` and can get *worse* with increasing order — the
    default ``order=32`` is the sweet spot for splines, and raising
    the order is not a convergence diagnostic.

    Args:
        y: Observations, shape ``(T, M)``. Masked entries are never
            read and may be ``NaN``.
        mask: Optional observation mask, shape ``(T, M)``; ``True``
            marks an observed entry. Same semantics as in
            `log_marginal`.
        n_ahead: Number of forecast steps appended after the
            training grid.
        order: Gauss-Hermite order for the warped pushforward.
            Ignored when ``warp is None``.

    Returns:
        Tuple ``(mean, var)`` of observation-space marginal moments,
        each of shape ``(T + n_ahead, M)``. Variances are clamped at
        zero, so ``sqrt`` on them is always safe.

    Raises:
        ValueError: If ``y`` or ``mask`` does not match the base
            event shape, or if the warp is not elementwise.
    """
    y = jnp.asarray(y)
    n_steps, n_channels = self.n_steps, self.n_channels
    mask_eff = self._effective_mask(mask)
    self._check_shapes(y, mask_eff)
    mean_grid = self._mean_grid(n_steps + n_ahead, y.dtype)
    residual = y - mean_grid[:n_steps]

    if self.warp is None:
        z = residual
    else:
        flows = _require_flows()
        # The scalar per-channel quadrature below is only the right
        # integral for an elementwise warp, so predict requires one
        # even when log_marginal would not.
        self._require_elementwise_warp(flows)
        if mask_eff is not None:
            # Unobserved slots hold junk (often NaN); substitute an
            # in-support reference before the inverse warp. The
            # filter never reads those entries afterwards.
            reference = self.warp.transform(jnp.zeros(n_channels, dtype=y.dtype))
            residual = jnp.where(mask_eff, residual, reference)
        z = jax.vmap(self.warp.inverse)(residual)  # (T, M)

    state = gaussx.kalman_filter(
        self.base.A,
        self.base.H,
        self.base.Q,
        self.base.R,
        z,
        self.base.m0,
        self.base.P0,
        mask=mask_eff,
    )
    m_smooth, P_smooth = gaussx.rts_smoother(state, self.base.A, self.base.Q)

    if n_ahead > 0:
        A_dense = _dense(self.base.A)
        Q_dense = _dense(self.base.Q)

        def step(carry, _):
            m, P = carry
            m_next = A_dense @ m
            P_next = A_dense @ P @ A_dense.T + Q_dense
            return (m_next, P_next), (m_next, P_next)

        last = (state.filtered_means[-1], state.filtered_covs[-1])
        _, (m_ahead, P_ahead) = jax.lax.scan(step, last, None, length=n_ahead)
        m_smooth = jnp.concatenate([m_smooth, m_ahead], axis=0)
        P_smooth = jnp.concatenate([P_smooth, P_ahead], axis=0)

    H_dense = _dense(self.base.H)
    R_diag = jnp.diagonal(_dense(self.base.R))
    # z-space marginals: mean H m_t, variance diag(H P_t H^T + R).
    mz = m_smooth @ H_dense.T  # (T', M)
    # einx.dot is typed as possibly returning a tuple (multi-output
    # patterns); narrow back to a single array for the typechecker.
    # Clamped at zero: the quadratic form is PSD in exact arithmetic,
    # but a rounding-negative entry would become NaN under the sqrt
    # taken for the quadrature nodes below.
    vz = jnp.maximum(
        jnp.asarray(einx.dot("m n, t n k, m k -> t m", H_dense, P_smooth, H_dense))
        + R_diag,
        0.0,
    )

    if self.warp is None:
        return mz + mean_grid, vz

    nodes, weights = np.polynomial.hermite_e.hermegauss(order)
    nodes = jnp.asarray(nodes, dtype=y.dtype)
    weights = jnp.asarray(weights, dtype=y.dtype) / np.sqrt(2.0 * np.pi)
    # fs: (order, T', M) — per-node evaluation points in the warped space.
    fs = mz[None] + jnp.sqrt(vz)[None] * nodes[:, None, None]
    g = jax.vmap(jax.vmap(self.warp.transform))(fs)
    m1 = jnp.asarray(einx.dot("s, s t m -> t m", weights, g))
    m2 = jnp.asarray(einx.dot("s, s t m -> t m", weights, g**2))
    # E[G^2] - E[G]^2 cancels catastrophically for a concentrated
    # predictive: the two moments agree to working precision and the
    # difference can land just below zero, which callers would turn
    # into NaN on the first sqrt.
    return m1 + mean_grid, jnp.maximum(m2 - m1**2, 0.0)

normalizing_kalman_factor(name: str, prior: NormalizingKalmanPrior, y: Float[Array, 'T M'], mask: Bool[Array, 'T M'] | None = None) -> None

Register the NKF marginal log-likelihood with NumPyro.

Computes the exact collapsed marginal via NormalizingKalmanPrior.log_marginal and adds it as numpyro.factor(name, ...). Mirrors pyrox_gp.markov_gp_factor — the latent state path is marginalised analytically, so the model only carries sample sites for the state-space (and warp) hyperparameters.

Parameters:

Name Type Description Default
name str

NumPyro factor site name.

required
prior NormalizingKalmanPrior

The NormalizingKalmanPrior.

required
y Float[Array, 'T M']

Observations, shape (T, M).

required
mask Bool[Array, 'T M'] | None

Optional observation mask, shape (T, M).

None
Source code in packages/pyrox-gp/src/pyrox_gp/_markov_flow.py
def normalizing_kalman_factor(
    name: str,
    prior: NormalizingKalmanPrior,
    y: Float[Array, "T M"],
    mask: Bool[Array, "T M"] | None = None,
) -> None:
    """Register the NKF marginal log-likelihood with NumPyro.

    Computes the exact collapsed marginal via `NormalizingKalmanPrior.log_marginal`
    and adds it as ``numpyro.factor(name, ...)``. Mirrors
    `pyrox_gp.markov_gp_factor` — the latent state path is marginalised
    analytically, so the model only carries sample sites for the
    state-space (and warp) hyperparameters.

    Args:
        name: NumPyro factor site name.
        prior: The `NormalizingKalmanPrior`.
        y: Observations, shape ``(T, M)``.
        mask: Optional observation mask, shape ``(T, M)``.
    """
    numpyro.factor(name, prior.log_marginal(y, mask))

Non-Gaussian inference (Markov)

The Markov-aware counterparts of the site-based strategies above: same diagonal-site math, but the global posterior recomputation runs through the Kalman filter / RTS smoother in O(N d^3) instead of a dense O(N^3) solve. Each fit(prior, likelihood, y) returns a NonGaussConditionedMarkovGP with the same predict API as the Gaussian-likelihood ConditionedMarkovGP.

LaplaceMarkovInference

Bases: Module

Laplace approximation for MarkovGPPrior.

Fixed-point Newton on the smoothed posterior: at each iteration run filter + smoother with the current sites to obtain marginals (m_n, V_n), then update site naturals from per-point gradient / Hessian of log p(y | f) evaluated at f = m. Site precision \(\Lambda = -h\) is clipped to precision_floor.

Attributes:

Name Type Description
max_iter int

Newton iterations. Default 20.

tol float

inf-norm convergence on the posterior mean. Default 1e-6.

damping float

Step-size in (0, 1] applied to each Newton update. Default 1.0 (full Newton step). Drop below 1 for non-log-concave likelihoods.

precision_floor float

Lower bound on the diagonal precision. Default 1e-6.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss_markov.py
class LaplaceMarkovInference(eqx.Module):
    r"""Laplace approximation for `MarkovGPPrior`.

    Fixed-point Newton on the smoothed posterior: at each iteration run
    filter + smoother with the current sites to obtain marginals
    ``(m_n, V_n)``, then update site naturals from per-point gradient /
    Hessian of ``log p(y | f)`` evaluated at ``f = m``. Site precision
    $\Lambda = -h$ is clipped to ``precision_floor``.

    Attributes:
        max_iter: Newton iterations. Default ``20``.
        tol: ``inf``-norm convergence on the posterior mean. Default ``1e-6``.
        damping: Step-size in (0, 1] applied to each Newton update.
            Default ``1.0`` (full Newton step). Drop below 1 for
            non-log-concave likelihoods.
        precision_floor: Lower bound on the diagonal precision. Default
            ``1e-6``.
    """

    max_iter: int = eqx.field(static=True, default=20)
    tol: float = eqx.field(static=True, default=1e-6)
    damping: float = eqx.field(static=True, default=1.0)
    precision_floor: float = eqx.field(static=True, default=1e-6)

    def fit(
        self,
        prior: MarkovGPPrior,
        likelihood: Likelihood,
        y: Float[Array, " N"],
    ) -> NonGaussConditionedMarkovGP:
        _check_scalar_latent(likelihood)
        log_prob_per_point = _log_prob_per_point_factory(likelihood)

        N = prior.times.shape[0]
        prior_mean = prior.mean(prior.times)
        nat1 = jnp.asarray(prior_mean) * self.precision_floor
        nat2 = jnp.full((N,), self.precision_floor, dtype=prior_mean.dtype)
        f = jnp.asarray(prior_mean)

        converged = False
        n_iter = 0
        for it in range(self.max_iter):
            g, h = _per_point_grad_hess(log_prob_per_point, f, y)
            new_nat1, Lam = newton_update(f, g, h, precision_floor=self.precision_floor)
            nat1, nat2 = damped_natural_update(
                nat1, nat2, new_nat1, Lam, lr=self.damping
            )
            assert isinstance(nat2, jax.Array)  # pyrox sites are diagonal arrays
            nat2 = jnp.maximum(nat2, self.precision_floor)
            f_new, _, _ = _markov_smoothed_posterior(prior, nat1, nat2)
            delta = jnp.max(jnp.abs(f_new - f))
            f = f_new
            n_iter = it + 1
            if delta < self.tol:
                converged = True
                break

        # Final site naturals at convergence.
        g, h = _per_point_grad_hess(log_prob_per_point, f, y)
        nat1, Lam = newton_update(f, g, h, precision_floor=self.precision_floor)
        nat2 = jnp.maximum(Lam, self.precision_floor)
        q_mean, q_var, log_marg = _markov_smoothed_posterior(prior, nat1, nat2)
        # Add the data-fit term ``log p(y | hat f)`` and subtract the
        # pseudo-data fit so the reported scalar is the standard Laplace
        # log-marginal approximation rather than the Kalman pseudo-data
        # likelihood.
        ll_data = log_prob_per_point(q_mean, y).sum()
        pseudo_targets = nat1 / nat2
        ll_pseudo = -0.5 * jnp.sum(jnp.log(2.0 * jnp.pi / nat2)) - 0.5 * jnp.sum(
            nat2 * (q_mean - pseudo_targets) ** 2
        )
        log_marg_corrected = log_marg + ll_data - ll_pseudo

        return NonGaussConditionedMarkovGP(
            prior=prior,
            y=y,
            site_nat1=nat1,
            site_nat2=nat2,
            q_mean=q_mean,
            q_var=q_var,
            log_marginal_approx=log_marg_corrected,
            n_iter=n_iter,
            converged=converged,
        )

GaussNewtonMarkovInference

Bases: Module

Gauss-Newton (Markov) inference: Newton with a strict PSD floor.

Identical to LaplaceMarkovInference for log-concave likelihoods (Bernoulli, Poisson). For non-log-concave likelihoods (StudentT) the larger precision_floor keeps the Newton step PSD-stable. Same fixed-point loop and damping behaviour as LaplaceMarkovInference.

Attributes:

Name Type Description
max_iter int

Newton iterations. Default 20.

tol float

inf-norm tolerance on the posterior mean. Default 1e-6.

damping float

Step-size in (0, 1]. Default 1.0.

precision_floor float

Strict positive floor. Default 1e-3.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss_markov.py
class GaussNewtonMarkovInference(eqx.Module):
    r"""Gauss-Newton (Markov) inference: Newton with a strict PSD floor.

    Identical to `LaplaceMarkovInference` for log-concave
    likelihoods (Bernoulli, Poisson). For non-log-concave likelihoods
    (StudentT) the larger ``precision_floor`` keeps the Newton step
    PSD-stable. Same fixed-point loop and damping behaviour as
    `LaplaceMarkovInference`.

    Attributes:
        max_iter: Newton iterations. Default ``20``.
        tol: ``inf``-norm tolerance on the posterior mean. Default ``1e-6``.
        damping: Step-size in (0, 1]. Default ``1.0``.
        precision_floor: Strict positive floor. Default ``1e-3``.
    """

    max_iter: int = eqx.field(static=True, default=20)
    tol: float = eqx.field(static=True, default=1e-6)
    damping: float = eqx.field(static=True, default=1.0)
    precision_floor: float = eqx.field(static=True, default=1e-3)

    def fit(
        self,
        prior: MarkovGPPrior,
        likelihood: Likelihood,
        y: Float[Array, " N"],
    ) -> NonGaussConditionedMarkovGP:
        inner = LaplaceMarkovInference(
            max_iter=self.max_iter,
            tol=self.tol,
            damping=self.damping,
            precision_floor=self.precision_floor,
        )
        return inner.fit(prior, likelihood, y)

PosteriorLinearizationMarkov

Bases: Module

Posterior linearization for MarkovGPPrior (Markov IPLF).

Iterates filter + smoother + cavity-averaged statistical linearization. At each iteration, with current sites producing smoothed marginals (m_n, V_n), form the cavity q_{\\setminus n}(f) = N(m_{c,n}, V_{c,n}), evaluate the expected per-point gradient and Hessian under the cavity (via Gauss-Hermite), and update the sites.

Attributes:

Name Type Description
max_iter int

Iterations. Default 20.

tol float

inf-norm tolerance on the posterior mean. Default 1e-6.

damping float

Step-size in (0, 1]. Default 0.5.

precision_floor float

Floor on the diagonal precision. Default 1e-6.

deg int

Gauss-Hermite degree for the cavity expectations. Default 20.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss_markov.py
class PosteriorLinearizationMarkov(eqx.Module):
    r"""Posterior linearization for `MarkovGPPrior` (Markov IPLF).

    Iterates filter + smoother + cavity-averaged statistical
    linearization. At each iteration, with current sites producing
    smoothed marginals ``(m_n, V_n)``, form the cavity
    ``q_{\\setminus n}(f) = N(m_{c,n}, V_{c,n})``, evaluate the *expected*
    per-point gradient and Hessian under the cavity (via Gauss-Hermite),
    and update the sites.

    Attributes:
        max_iter: Iterations. Default ``20``.
        tol: ``inf``-norm tolerance on the posterior mean. Default ``1e-6``.
        damping: Step-size in (0, 1]. Default ``0.5``.
        precision_floor: Floor on the diagonal precision. Default ``1e-6``.
        deg: Gauss-Hermite degree for the cavity expectations. Default ``20``.
    """

    max_iter: int = eqx.field(static=True, default=20)
    tol: float = eqx.field(static=True, default=1e-6)
    damping: float = eqx.field(static=True, default=0.5)
    precision_floor: float = eqx.field(static=True, default=1e-6)
    deg: int = eqx.field(static=True, default=20)

    def fit(
        self,
        prior: MarkovGPPrior,
        likelihood: Likelihood,
        y: Float[Array, " N"],
    ) -> NonGaussConditionedMarkovGP:
        from gaussx import gauss_hermite_points

        _check_scalar_latent(likelihood)
        log_prob_per_point = _log_prob_per_point_factory(likelihood)

        def lp(f_n: Float[Array, ""], y_n: Float[Array, ""]) -> Float[Array, ""]:
            return log_prob_per_point(f_n[None], y_n[None])[0]

        grad_fn = jax.vmap(jax.grad(lp))
        hess_fn = jax.vmap(jax.grad(jax.grad(lp)))

        nodes, weights = gauss_hermite_points(self.deg, dim=1)
        x_nodes = nodes[:, 0]
        w_nodes = weights / jnp.sqrt(2.0 * jnp.pi)

        N = prior.times.shape[0]
        prior_mean = prior.mean(prior.times)
        nat1 = jnp.zeros(N, dtype=prior_mean.dtype)
        nat2 = jnp.full((N,), self.precision_floor, dtype=prior_mean.dtype)
        q_mean = jnp.asarray(prior_mean)
        # Seed cavity computation with the prior marginal variance ``H P_inf H^T``
        # so the first cavity precision ``1/q_var - nat2`` is on the correct
        # scale for kernels with variance != 1, and so a zero-iteration return
        # produces a posterior variance that matches the prior.
        q_var = _prior_marginal_variance(prior).astype(prior_mean.dtype)

        converged = False
        n_iter = 0
        for it in range(self.max_iter):
            cav_mean, cav_var = cavity_distribution(
                q_mean, q_var, nat1, nat2, precision_floor=self.precision_floor
            )
            assert isinstance(cav_var, jax.Array)  # diagonal path in, diagonal out

            std = jnp.sqrt(cav_var)
            # Gauss-Hermite grid fₙ + σₙ ξ_q over (Q nodes, N sites).
            f_grid = einx.add(
                "q n, n -> q n", einx.multiply("q, n -> q n", x_nodes, std), cav_mean
            )
            g_grid = jax.vmap(lambda f_row: grad_fn(f_row, y))(f_grid)
            h_grid = jax.vmap(lambda f_row: hess_fn(f_row, y))(f_grid)
            # Quadrature average = weighted sum over the node axis q → (N,).
            g_avg = einx.dot("q, q n -> n", w_nodes, g_grid)
            h_avg = einx.dot("q, q n -> n", w_nodes, h_grid)

            new_nat1, new_prec = newton_update(
                cav_mean, g_avg, h_avg, precision_floor=self.precision_floor
            )
            nat1, nat2 = damped_natural_update(
                nat1, nat2, new_nat1, new_prec, lr=self.damping
            )
            assert isinstance(nat2, jax.Array)  # pyrox sites are diagonal arrays
            nat2 = jnp.maximum(nat2, self.precision_floor)

            q_mean_new, q_var, _ = _markov_smoothed_posterior(prior, nat1, nat2)
            delta = jnp.max(jnp.abs(q_mean_new - q_mean))
            q_mean = q_mean_new
            n_iter = it + 1
            if delta < self.tol:
                converged = True
                break

        _, _, log_marg = _markov_smoothed_posterior(prior, nat1, nat2)

        return NonGaussConditionedMarkovGP(
            prior=prior,
            y=y,
            site_nat1=nat1,
            site_nat2=nat2,
            q_mean=q_mean,
            q_var=q_var,
            log_marginal_approx=log_marg,
            n_iter=n_iter,
            converged=converged,
        )

ExpectationPropagationMarkov

Bases: Module

Parallel Expectation Propagation for MarkovGPPrior.

Each iteration: run filter + smoother, form cavities at every site, match tilted-distribution moments via Gauss-Hermite (with log-space stabilisation), and update sites with damping.

Attributes:

Name Type Description
max_iter int

EP iterations. Default 40.

tol float

inf-norm tolerance on the posterior mean. Default 1e-5.

damping float

Damping in (0, 1]. Default 0.5.

precision_floor float

Floor on the diagonal precision. Default 1e-6.

deg int

Gauss-Hermite degree for the tilted-moment integrals. Default 20.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss_markov.py
class ExpectationPropagationMarkov(eqx.Module):
    r"""Parallel Expectation Propagation for `MarkovGPPrior`.

    Each iteration: run filter + smoother, form cavities at every site,
    match tilted-distribution moments via Gauss-Hermite (with log-space
    stabilisation), and update sites with damping.

    Attributes:
        max_iter: EP iterations. Default ``40``.
        tol: ``inf``-norm tolerance on the posterior mean. Default ``1e-5``.
        damping: Damping in (0, 1]. Default ``0.5``.
        precision_floor: Floor on the diagonal precision. Default ``1e-6``.
        deg: Gauss-Hermite degree for the tilted-moment integrals.
            Default ``20``.
    """

    max_iter: int = eqx.field(static=True, default=40)
    tol: float = eqx.field(static=True, default=1e-5)
    damping: float = eqx.field(static=True, default=0.5)
    precision_floor: float = eqx.field(static=True, default=1e-6)
    deg: int = eqx.field(static=True, default=20)

    def fit(
        self,
        prior: MarkovGPPrior,
        likelihood: Likelihood,
        y: Float[Array, " N"],
    ) -> NonGaussConditionedMarkovGP:
        _check_scalar_latent(likelihood)
        log_prob_per_point = _log_prob_per_point_factory(likelihood)

        def lp(f_n: Float[Array, ""], y_n: Float[Array, ""]) -> Float[Array, ""]:
            return log_prob_per_point(f_n[None], y_n[None])[0]

        N = prior.times.shape[0]
        prior_mean = prior.mean(prior.times)
        nat1 = jnp.zeros(N, dtype=prior_mean.dtype)
        nat2 = jnp.full((N,), self.precision_floor, dtype=prior_mean.dtype)
        q_mean = jnp.asarray(prior_mean)
        # Seed cavity computation with the prior marginal variance ``H P_inf H^T``
        # so the first cavity precision ``1/q_var - nat2`` is on the correct
        # scale for kernels with variance != 1, and so a zero-iteration return
        # produces a posterior variance that matches the prior.
        q_var = _prior_marginal_variance(prior).astype(prior_mean.dtype)

        # ``gaussx.ep_tilted_moments`` requires ``log_lik_fn(f)`` with the
        # per-site target baked in; close over ``y_n`` per site via vmap
        # to recover the (N,) shape contract.
        deg = self.deg

        def _per_site_tilted(
            m_n: Float[Array, ""],
            v_n: Float[Array, ""],
            y_n: Float[Array, ""],
        ) -> tuple[Float[Array, ""], Float[Array, ""]]:
            return ep_tilted_moments(lambda f: lp(f, y_n), m_n, v_n, order=deg)

        converged = False
        n_iter = 0
        for it in range(self.max_iter):
            cav_mean, cav_var = cavity_distribution(
                q_mean, q_var, nat1, nat2, precision_floor=self.precision_floor
            )
            assert isinstance(cav_var, jax.Array)  # diagonal path in, diagonal out

            tilted_mean, tilted_var = jax.vmap(_per_site_tilted)(cav_mean, cav_var, y)

            new_prec = jnp.reciprocal(tilted_var) - jnp.reciprocal(cav_var)
            new_prec = jnp.maximum(new_prec, self.precision_floor)
            new_nat1 = tilted_mean / tilted_var - cav_mean / cav_var
            nat1, nat2 = damped_natural_update(
                nat1, nat2, new_nat1, new_prec, lr=self.damping
            )
            assert isinstance(nat2, jax.Array)  # pyrox sites are diagonal arrays
            nat2 = jnp.maximum(nat2, self.precision_floor)

            q_mean_new, q_var, _ = _markov_smoothed_posterior(prior, nat1, nat2)
            delta = jnp.max(jnp.abs(q_mean_new - q_mean))
            q_mean = q_mean_new
            n_iter = it + 1
            if delta < self.tol:
                converged = True
                break

        _, _, log_marg = _markov_smoothed_posterior(prior, nat1, nat2)

        return NonGaussConditionedMarkovGP(
            prior=prior,
            y=y,
            site_nat1=nat1,
            site_nat2=nat2,
            q_mean=q_mean,
            q_var=q_var,
            log_marginal_approx=log_marg,
            n_iter=n_iter,
            converged=converged,
        )

NonGaussConditionedMarkovGP

Bases: Module

Markov GP conditioned on a non-Gaussian likelihood via a site-based strategy.

Equivalent role to pyrox_gp.ConditionedMarkovGP but the posterior over training latents is a generic Gaussian approximation produced by one of the strategies in this module rather than the closed-form Gaussian-likelihood smoother. Predictions at arbitrary test times use the standard site-as-pseudo-observation trick on the merged grid sort(times | t_star) with the test points masked.

Attributes:

Name Type Description
prior MarkovGPPrior

The originating MarkovGPPrior.

y Float[Array, ' N']

Training targets (kept for round-trip / diagnostics).

site_nat1 Float[Array, ' N']

Diagonal site naturals \(\lambda \in \mathbb{R}^N\).

site_nat2 Float[Array, ' N']

Diagonal site precisions \(\Lambda \in \mathbb{R}^N\) (positive).

q_mean Float[Array, ' N']

Posterior mean over training latents.

q_var Float[Array, ' N']

Marginal posterior variance per training point.

log_marginal_approx Float[Array, '']

Approximate log marginal likelihood (the scalar each strategy reports — interpretation is strategy-specific).

n_iter int

Iterations used by the strategy.

converged bool

Whether convergence tolerance was met.

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss_markov.py
class NonGaussConditionedMarkovGP(eqx.Module):
    """Markov GP conditioned on a non-Gaussian likelihood via a site-based strategy.

    Equivalent role to `pyrox_gp.ConditionedMarkovGP` but the
    posterior over training latents is a generic Gaussian
    approximation produced by one of the strategies in this module
    rather than the closed-form Gaussian-likelihood smoother. Predictions
    at arbitrary test times use the standard *site-as-pseudo-observation*
    trick on the merged grid ``sort(times | t_star)`` with the test
    points masked.

    Attributes:
        prior: The originating `MarkovGPPrior`.
        y: Training targets (kept for round-trip / diagnostics).
        site_nat1: Diagonal site naturals $\\lambda \\in \\mathbb{R}^N$.
        site_nat2: Diagonal site precisions $\\Lambda \\in \\mathbb{R}^N$
            (positive).
        q_mean: Posterior mean over training latents.
        q_var: Marginal posterior variance per training point.
        log_marginal_approx: Approximate log marginal likelihood (the
            scalar each strategy reports — interpretation is
            strategy-specific).
        n_iter: Iterations used by the strategy.
        converged: Whether convergence tolerance was met.
    """

    prior: MarkovGPPrior
    y: Float[Array, " N"]
    site_nat1: Float[Array, " N"]
    site_nat2: Float[Array, " N"]
    q_mean: Float[Array, " N"]
    q_var: Float[Array, " N"]
    log_marginal_approx: Float[Array, ""]
    n_iter: int = eqx.field(static=True)
    converged: bool = eqx.field(static=True)

    def predict(
        self,
        t_star: Float[Array, " M"],
    ) -> tuple[Float[Array, " M"], Float[Array, " M"]]:
        r"""Predictive marginals ``(mean, var)`` at arbitrary test times.

        Re-runs filter + smoother over the merged grid
        ``sort(times \cup t_star)`` with per-step pseudo-observation
        variances ``1/Λ_n`` on the training points and the test points
        masked out of the update step. Cost is $O((N + M)\,d^3)$.
        """
        F, _L, H, _Qc, _P_inf = self.prior.sde_kernel.sde_params()
        init_cov = self.prior.initial_covariance()
        times = self.prior.times
        t_star = jnp.asarray(t_star)

        N = times.shape[0]
        M = t_star.shape[0]
        merged = jnp.concatenate([times, t_star], axis=0)
        order = jnp.argsort(merged, stable=True)
        merged_sorted = merged[order]

        is_obs = jnp.concatenate(
            [jnp.ones(N, dtype=times.dtype), jnp.zeros(M, dtype=times.dtype)]
        )[order]
        pseudo_targets = self.site_nat1 / self.site_nat2
        residual_full = jnp.concatenate(
            [
                pseudo_targets - self.prior.mean(times),
                jnp.zeros(M, dtype=self.y.dtype),
            ]
        )[order]
        # ``R_seq`` for masked steps is irrelevant (the update is
        # skipped) but must be finite to avoid NaN propagation through
        # the filter; use 1.0 as a safe placeholder.
        R_train = jnp.reciprocal(self.site_nat2)
        R_full = jnp.concatenate([R_train, jnp.ones(M, dtype=self.y.dtype)])[order]

        dt_full = _build_dt_full(merged_sorted)
        A_seq, Q_seq = self.prior.sde_kernel.discretise_sequence(dt_full)
        m_pred, P_pred, m_filt, P_filt, _ = _kalman_filter(
            F, H, init_cov, A_seq, Q_seq, residual_full, is_obs, R_full
        )
        m_smooth, P_smooth = _rts_smoother(m_pred, P_pred, m_filt, P_filt, A_seq)

        inv_order = jnp.argsort(order, stable=True)
        test_positions = inv_order[N:]
        m_test_state = m_smooth[test_positions]
        P_test_state = P_smooth[test_positions]
        means = (m_test_state @ H.T)[:, 0] + self.prior.mean(t_star)
        vars_ = jax.vmap(lambda P: (H @ P @ H.T)[0, 0])(P_test_state)
        return means, jnp.maximum(vars_, 0.0)

predict(t_star: Float[Array, ' M']) -> tuple[Float[Array, ' M'], Float[Array, ' M']]

Predictive marginals (mean, var) at arbitrary test times.

Re-runs filter + smoother over the merged grid sort(times \cup t_star) with per-step pseudo-observation variances 1/Λ_n on the training points and the test points masked out of the update step. Cost is \(O((N + M)\,d^3)\).

Source code in packages/pyrox-gp/src/pyrox_gp/_inference_nongauss_markov.py
def predict(
    self,
    t_star: Float[Array, " M"],
) -> tuple[Float[Array, " M"], Float[Array, " M"]]:
    r"""Predictive marginals ``(mean, var)`` at arbitrary test times.

    Re-runs filter + smoother over the merged grid
    ``sort(times \cup t_star)`` with per-step pseudo-observation
    variances ``1/Λ_n`` on the training points and the test points
    masked out of the update step. Cost is $O((N + M)\,d^3)$.
    """
    F, _L, H, _Qc, _P_inf = self.prior.sde_kernel.sde_params()
    init_cov = self.prior.initial_covariance()
    times = self.prior.times
    t_star = jnp.asarray(t_star)

    N = times.shape[0]
    M = t_star.shape[0]
    merged = jnp.concatenate([times, t_star], axis=0)
    order = jnp.argsort(merged, stable=True)
    merged_sorted = merged[order]

    is_obs = jnp.concatenate(
        [jnp.ones(N, dtype=times.dtype), jnp.zeros(M, dtype=times.dtype)]
    )[order]
    pseudo_targets = self.site_nat1 / self.site_nat2
    residual_full = jnp.concatenate(
        [
            pseudo_targets - self.prior.mean(times),
            jnp.zeros(M, dtype=self.y.dtype),
        ]
    )[order]
    # ``R_seq`` for masked steps is irrelevant (the update is
    # skipped) but must be finite to avoid NaN propagation through
    # the filter; use 1.0 as a safe placeholder.
    R_train = jnp.reciprocal(self.site_nat2)
    R_full = jnp.concatenate([R_train, jnp.ones(M, dtype=self.y.dtype)])[order]

    dt_full = _build_dt_full(merged_sorted)
    A_seq, Q_seq = self.prior.sde_kernel.discretise_sequence(dt_full)
    m_pred, P_pred, m_filt, P_filt, _ = _kalman_filter(
        F, H, init_cov, A_seq, Q_seq, residual_full, is_obs, R_full
    )
    m_smooth, P_smooth = _rts_smoother(m_pred, P_pred, m_filt, P_filt, A_seq)

    inv_order = jnp.argsort(order, stable=True)
    test_positions = inv_order[N:]
    m_test_state = m_smooth[test_positions]
    P_test_state = P_smooth[test_positions]
    means = (m_test_state @ H.T)[:, 0] + self.prior.mean(t_star)
    vars_ = jax.vmap(lambda P: (H @ P @ H.T)[0, 0])(P_test_state)
    return means, jnp.maximum(vars_, 0.0)

Sparse Markov GP

Sparse variational GP over an SDE kernel and an inducing time grid: the variational family lives on the inducing times while predictions exploit the Markov structure between them.

SparseMarkovGPPrior

Bases: Module

Sparse variational GP prior over an SDE kernel and inducing time grid.

Equivalent role to pyrox_gp.SparseGPPrior but the covariance is derived from the state-space autocovariance \(k(\tau) = H \exp(F\tau) P_\infty H^\top\) of an SDEKernel. Predictions and the SVGP ELBO go through the same predictive_blocks contract as the dense SparseGPPrior, so the existing variational guides (pyrox_gp.FullRankGuide, pyrox_gp.MeanFieldGuide, pyrox_gp.WhitenedGuide) and pyrox_gp.svgp_elbo work as-is.

Attributes:

Name Type Description
sde_kernel SDEKernel

Any SDEKernel (Matern, Periodic, Cosine, Sum/Product compositions, ...).

Z Float[Array, ' M']

Sorted, strictly increasing inducing times of shape (M,).

mean_fn Callable[[Float[Array, ' N']], Float[Array, ' N']] | None

Optional callable times -> (N,) global mean. Convenience accessor; not folded into the inducing prior.

solver AbstractSolverStrategy | None

Any gaussx.AbstractSolverStrategy. Defaults to gaussx.DenseSolver().

jitter float

Diagonal regularisation added to K_{ZZ} for numerical stability.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse_markov.py
class SparseMarkovGPPrior(eqx.Module):
    r"""Sparse variational GP prior over an SDE kernel and inducing time grid.

    Equivalent role to `pyrox_gp.SparseGPPrior` but the
    covariance is derived from the state-space autocovariance
    $k(\tau) = H \exp(F\tau) P_\infty H^\top$ of an
    `SDEKernel`. Predictions and the SVGP ELBO go through the
    same `predictive_blocks` contract as the dense
    `SparseGPPrior`, so the existing variational guides
    (`pyrox_gp.FullRankGuide`, `pyrox_gp.MeanFieldGuide`,
    `pyrox_gp.WhitenedGuide`) and `pyrox_gp.svgp_elbo`
    work as-is.

    Attributes:
        sde_kernel: Any `SDEKernel` (Matern, Periodic, Cosine,
            Sum/Product compositions, ...).
        Z: Sorted, strictly increasing inducing times of shape ``(M,)``.
        mean_fn: Optional callable ``times -> (N,)`` global mean.
            Convenience accessor; not folded into the inducing prior.
        solver: Any ``gaussx.AbstractSolverStrategy``. Defaults to
            ``gaussx.DenseSolver()``.
        jitter: Diagonal regularisation added to ``K_{ZZ}`` for numerical
            stability.
    """

    sde_kernel: SDEKernel
    Z: Float[Array, " M"]
    mean_fn: Callable[[Float[Array, " N"]], Float[Array, " N"]] | None = None
    solver: AbstractSolverStrategy | None = None
    jitter: float = eqx.field(static=True, default=1e-6)

    def __init__(
        self,
        sde_kernel: SDEKernel,
        Z: Float[Array, " M"],
        mean_fn: Callable[[Float[Array, " N"]], Float[Array, " N"]] | None = None,
        solver: AbstractSolverStrategy | None = None,
        jitter: float = 1e-6,
    ) -> None:
        Z_arr = jnp.asarray(Z, dtype=jnp.result_type(Z, 0.0))
        if Z_arr.ndim != 1:
            raise ValueError(f"Z must be 1-D, got shape {tuple(Z_arr.shape)!r}")
        if Z_arr.shape[0] >= 2:
            try:
                if not bool(jnp.all(jnp.diff(Z_arr) > 0)):
                    raise ValueError("Z must be strictly increasing")
            except jax.errors.TracerBoolConversionError:
                pass
        self.sde_kernel = sde_kernel
        self.Z = Z_arr
        self.mean_fn = mean_fn
        self.solver = solver
        self.jitter = float(jitter)

    @property
    def num_inducing(self) -> int:
        """Number of inducing time points $M$."""
        return self.Z.shape[0]

    def mean(self, times: Float[Array, " N"]) -> Float[Array, " N"]:
        """Evaluate the mean function at ``times``; zero by default."""
        if self.mean_fn is None:
            return jnp.zeros_like(times)
        return self.mean_fn(times)

    def _stationary_variance(self) -> Float[Array, ""]:
        r"""Stationary marginal variance ``H P_inf H^T`` (= kernel variance)."""
        _F, _L, H, _Qc, P_inf = self.sde_kernel.sde_params()
        P_inf = _require_stationary(self.sde_kernel, P_inf)
        return (H @ P_inf @ H.T)[0, 0]

    def inducing_operator(self) -> lx.AbstractLinearOperator:
        r"""Return ``K_{ZZ} + \text{jitter}\,I`` as a PSD ``lineax`` operator."""
        # Pairwise |Zᵢ - Zⱼ| lags between inducing times → (M, M).
        diffs = jnp.abs(einx.subtract("i, j -> i j", self.Z, self.Z))
        K_zz = sde_autocovariance(self.sde_kernel, diffs)
        K_zz = symmetrize(K_zz)
        K_zz = K_zz.at[jnp.diag_indices_from(K_zz)].add(self.jitter)
        return _psd_operator(K_zz)

    def cross_covariance(self, times: Float[Array, " N"]) -> Float[Array, "N M"]:
        """``K_{XZ}`` — pairwise SDE autocov between training and inducing times."""
        t = jnp.asarray(times)
        # Pairwise |tₙ - Zₘ| lags between training and inducing times → (N, M).
        diffs = jnp.abs(einx.subtract("n, m -> n m", t, self.Z))
        return sde_autocovariance(self.sde_kernel, diffs)

    def kernel_diag(self, times: Float[Array, " N"]) -> Float[Array, " N"]:
        r"""Prior diagonal $\mathrm{diag}\,K(X, X)$.

        Constant for stationary SDE kernels — equal to ``H P_inf H^T``.
        """
        var0 = self._stationary_variance()
        return jnp.broadcast_to(var0, jnp.asarray(times).shape).astype(self.Z.dtype)

    def predictive_blocks(
        self, times: Float[Array, " N"]
    ) -> tuple[
        lx.AbstractLinearOperator,
        Float[Array, "N M"],
        Float[Array, " N"],
    ]:
        r"""Return ``(K_zz_op, K_xz, K_xx_diag)`` for the SVGP predictive math.

        Mirrors `SparseGPPrior.predictive_blocks`. Delegates to the
        three independent accessors:

        * `inducing_operator` — ``K_{ZZ} + \\text{jitter}\\,I``
          wrapped as a PSD `lineax` operator.
        * `cross_covariance` — pairwise SDE autocov ``K_{XZ}``.
        * `kernel_diag` — prior marginal variance, constant for
          stationary SDE kernels.

        Each accessor calls `SDEKernel.sde_params` independently;
        these calls are cheap (parameter unpacking, not a kernel build)
        so no shared-state caching is performed.
        """
        K_zz_op = self.inducing_operator()
        K_xz = self.cross_covariance(times)
        K_xx_diag = self.kernel_diag(times)
        return K_zz_op, K_xz, K_xx_diag

    def _resolved_solver(self) -> AbstractSolverStrategy:
        return DenseSolver() if self.solver is None else self.solver

    def log_prob(self, u: Float[Array, " M"]) -> Float[Array, ""]:
        r"""Log-density under the inducing prior.

        $p(u) = \mathcal{N}(0,\, K_{ZZ} + \text{jitter}\,I)$.
        """
        m = jnp.zeros(self.num_inducing, dtype=u.dtype)
        return gaussian_log_prob(
            m, self.inducing_operator(), u, solver=self._resolved_solver()
        )

    def sample(self, key: Array) -> Float[Array, " M"]:
        r"""Draw $u \sim p(u)$ from the inducing prior."""
        op = self.inducing_operator()
        loc = jnp.zeros(self.num_inducing, dtype=op.out_structure().dtype)
        mvn = MultivariateNormal(loc, op, solver=self._resolved_solver())
        return mvn.sample(key)

num_inducing: int property

Number of inducing time points \(M\).

mean(times: Float[Array, ' N']) -> Float[Array, ' N']

Evaluate the mean function at times; zero by default.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse_markov.py
def mean(self, times: Float[Array, " N"]) -> Float[Array, " N"]:
    """Evaluate the mean function at ``times``; zero by default."""
    if self.mean_fn is None:
        return jnp.zeros_like(times)
    return self.mean_fn(times)

inducing_operator() -> lx.AbstractLinearOperator

Return K_{ZZ} + \text{jitter}\,I as a PSD lineax operator.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse_markov.py
def inducing_operator(self) -> lx.AbstractLinearOperator:
    r"""Return ``K_{ZZ} + \text{jitter}\,I`` as a PSD ``lineax`` operator."""
    # Pairwise |Zᵢ - Zⱼ| lags between inducing times → (M, M).
    diffs = jnp.abs(einx.subtract("i, j -> i j", self.Z, self.Z))
    K_zz = sde_autocovariance(self.sde_kernel, diffs)
    K_zz = symmetrize(K_zz)
    K_zz = K_zz.at[jnp.diag_indices_from(K_zz)].add(self.jitter)
    return _psd_operator(K_zz)

cross_covariance(times: Float[Array, ' N']) -> Float[Array, 'N M']

K_{XZ} — pairwise SDE autocov between training and inducing times.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse_markov.py
def cross_covariance(self, times: Float[Array, " N"]) -> Float[Array, "N M"]:
    """``K_{XZ}`` — pairwise SDE autocov between training and inducing times."""
    t = jnp.asarray(times)
    # Pairwise |tₙ - Zₘ| lags between training and inducing times → (N, M).
    diffs = jnp.abs(einx.subtract("n, m -> n m", t, self.Z))
    return sde_autocovariance(self.sde_kernel, diffs)

kernel_diag(times: Float[Array, ' N']) -> Float[Array, ' N']

Prior diagonal \(\mathrm{diag}\,K(X, X)\).

Constant for stationary SDE kernels — equal to H P_inf H^T.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse_markov.py
def kernel_diag(self, times: Float[Array, " N"]) -> Float[Array, " N"]:
    r"""Prior diagonal $\mathrm{diag}\,K(X, X)$.

    Constant for stationary SDE kernels — equal to ``H P_inf H^T``.
    """
    var0 = self._stationary_variance()
    return jnp.broadcast_to(var0, jnp.asarray(times).shape).astype(self.Z.dtype)

predictive_blocks(times: Float[Array, ' N']) -> tuple[lx.AbstractLinearOperator, Float[Array, 'N M'], Float[Array, ' N']]

Return (K_zz_op, K_xz, K_xx_diag) for the SVGP predictive math.

Mirrors SparseGPPrior.predictive_blocks. Delegates to the three independent accessors:

  • inducing_operator — K_{ZZ} + \\text{jitter}\\,I wrapped as a PSD lineax operator.
  • cross_covariance — pairwise SDE autocov K_{XZ}.
  • kernel_diag — prior marginal variance, constant for stationary SDE kernels.

Each accessor calls SDEKernel.sde_params independently; these calls are cheap (parameter unpacking, not a kernel build) so no shared-state caching is performed.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse_markov.py
def predictive_blocks(
    self, times: Float[Array, " N"]
) -> tuple[
    lx.AbstractLinearOperator,
    Float[Array, "N M"],
    Float[Array, " N"],
]:
    r"""Return ``(K_zz_op, K_xz, K_xx_diag)`` for the SVGP predictive math.

    Mirrors `SparseGPPrior.predictive_blocks`. Delegates to the
    three independent accessors:

    * `inducing_operator` — ``K_{ZZ} + \\text{jitter}\\,I``
      wrapped as a PSD `lineax` operator.
    * `cross_covariance` — pairwise SDE autocov ``K_{XZ}``.
    * `kernel_diag` — prior marginal variance, constant for
      stationary SDE kernels.

    Each accessor calls `SDEKernel.sde_params` independently;
    these calls are cheap (parameter unpacking, not a kernel build)
    so no shared-state caching is performed.
    """
    K_zz_op = self.inducing_operator()
    K_xz = self.cross_covariance(times)
    K_xx_diag = self.kernel_diag(times)
    return K_zz_op, K_xz, K_xx_diag

log_prob(u: Float[Array, ' M']) -> Float[Array, '']

Log-density under the inducing prior.

\(p(u) = \mathcal{N}(0,\, K_{ZZ} + \text{jitter}\,I)\).

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse_markov.py
def log_prob(self, u: Float[Array, " M"]) -> Float[Array, ""]:
    r"""Log-density under the inducing prior.

    $p(u) = \mathcal{N}(0,\, K_{ZZ} + \text{jitter}\,I)$.
    """
    m = jnp.zeros(self.num_inducing, dtype=u.dtype)
    return gaussian_log_prob(
        m, self.inducing_operator(), u, solver=self._resolved_solver()
    )

sample(key: Array) -> Float[Array, ' M']

Draw \(u \sim p(u)\) from the inducing prior.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse_markov.py
def sample(self, key: Array) -> Float[Array, " M"]:
    r"""Draw $u \sim p(u)$ from the inducing prior."""
    op = self.inducing_operator()
    loc = jnp.zeros(self.num_inducing, dtype=op.out_structure().dtype)
    mvn = MultivariateNormal(loc, op, solver=self._resolved_solver())
    return mvn.sample(key)

SparseConditionedMarkovGP

Bases: Module

Sparse Markov GP fitted to a guide.

Bundles the SparseMarkovGPPrior and a fitted pyrox_gp.Guide so that predict(t_star) can be called against arbitrary test times. The math is the standard SVGP predictive

\[ \mu_*(t) = K_{*Z} K_{ZZ}^{-1} m_q + \mu(t),\qquad \sigma_*^2(t) = k(t, t) - K_{*Z} K_{ZZ}^{-1} K_{Z*} + K_{*Z} K_{ZZ}^{-1} S_q K_{ZZ}^{-1} K_{Z*}. \]

Cost is \(O(M^3 + |t_*|\,M)\) per call.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse_markov.py
class SparseConditionedMarkovGP(eqx.Module):
    """Sparse Markov GP fitted to a guide.

    Bundles the `SparseMarkovGPPrior` and a fitted
    `pyrox_gp.Guide` so that ``predict(t_star)`` can be called
    against arbitrary test times. The math is the standard SVGP
    predictive

    $$
    \\mu_*(t) = K_{*Z} K_{ZZ}^{-1} m_q + \\mu(t),\\qquad
    \\sigma_*^2(t) = k(t, t) - K_{*Z} K_{ZZ}^{-1} K_{Z*}
                    + K_{*Z} K_{ZZ}^{-1} S_q K_{ZZ}^{-1} K_{Z*}.
    $$

    Cost is $O(M^3 + |t_*|\\,M)$ per call.
    """

    prior: SparseMarkovGPPrior
    guide: Guide

    def predict(
        self, t_star: Float[Array, " M_star"]
    ) -> tuple[Float[Array, " M_star"], Float[Array, " M_star"]]:
        """Predictive ``(mean, var)`` at arbitrary test times."""
        K_zz_op, K_xz, K_xx_diag = self.prior.predictive_blocks(t_star)
        mean, var = self.guide.predict(K_xz, K_zz_op, K_xx_diag)  # ty: ignore[unresolved-attribute]
        return mean + self.prior.mean(t_star), var

predict(t_star: Float[Array, ' M_star']) -> tuple[Float[Array, ' M_star'], Float[Array, ' M_star']]

Predictive (mean, var) at arbitrary test times.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse_markov.py
def predict(
    self, t_star: Float[Array, " M_star"]
) -> tuple[Float[Array, " M_star"], Float[Array, " M_star"]]:
    """Predictive ``(mean, var)`` at arbitrary test times."""
    K_zz_op, K_xz, K_xx_diag = self.prior.predictive_blocks(t_star)
    mean, var = self.guide.predict(K_xz, K_zz_op, K_xx_diag)  # ty: ignore[unresolved-attribute]
    return mean + self.prior.mean(t_star), var

sparse_markov_elbo(prior: SparseMarkovGPPrior, guide: Guide, likelihood: Likelihood, times: Float[Array, ' N'], y: Float[Array, ' N'], *, integrator: AbstractIntegrator | None = None) -> Float[Array, '']

Sparse variational ELBO for SparseMarkovGPPrior.

Mirrors pyrox_gp.svgp_elbo for the SDE-derived sparse Markov prior. Builds the SVGP predictive blocks \((K_{ZZ}, K_{XZ}, \mathrm{diag}\,K_{XX})\) from the prior, asks the guide for the predictive marginals \((\mu_n, \sigma_n^2) = q(f_n)\), and combines them with a closed-form Gaussian or quadrature-based expected log-likelihood and the inducing KL term:

\[ \mathcal{L} = \sum_n \mathbb{E}_{q(f_n)}[\log p(y_n \mid f_n)] - \mathrm{KL}[q(u) \\| p(u)]. \]

Unlike pyrox_gp.svgp_elbo, times stays as a 1-D vector of shape (N,) — the SDE-pair autocov works on raw 1-D times and has no need for a feature dimension.

Parameters:

Name Type Description Default
prior SparseMarkovGPPrior

SparseMarkovGPPrior over an SDE kernel and inducing time grid.

required
guide Guide

Variational guide over inducing values.

required
likelihood Likelihood

Observation model.

required
times Float[Array, ' N']

Training times of shape (N,).

required
y Float[Array, ' N']

Observations of shape (N,).

required
integrator AbstractIntegrator | None

gaussx integrator for non-conjugate likelihoods. None is fine for pyrox_gp.GaussianLikelihood.

None

Returns:

Type Description
Float[Array, '']

Scalar ELBO value (higher is better).

Raises:

Type Description
ValueError

If a non-conjugate likelihood is used without an integrator.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse_markov.py
def sparse_markov_elbo(
    prior: SparseMarkovGPPrior,
    guide: Guide,
    likelihood: Likelihood,
    times: Float[Array, " N"],
    y: Float[Array, " N"],
    *,
    integrator: AbstractIntegrator | None = None,
) -> Float[Array, ""]:
    r"""Sparse variational ELBO for `SparseMarkovGPPrior`.

    Mirrors `pyrox_gp.svgp_elbo` for the SDE-derived sparse
    Markov prior. Builds the SVGP predictive blocks
    $(K_{ZZ}, K_{XZ}, \mathrm{diag}\,K_{XX})$ from the prior, asks
    the guide for the predictive marginals
    $(\mu_n, \sigma_n^2) = q(f_n)$, and combines them with a closed-form
    Gaussian or quadrature-based expected log-likelihood and the
    inducing KL term:

    $$
    \mathcal{L} = \sum_n \mathbb{E}_{q(f_n)}[\log p(y_n \mid f_n)]
                - \mathrm{KL}[q(u) \\| p(u)].
    $$

    Unlike `pyrox_gp.svgp_elbo`, ``times`` stays as a 1-D vector
    of shape ``(N,)`` — the SDE-pair autocov works on raw 1-D times and
    has no need for a feature dimension.

    Args:
        prior: `SparseMarkovGPPrior` over an SDE kernel and
            inducing time grid.
        guide: Variational guide over inducing values.
        likelihood: Observation model.
        times: Training times of shape ``(N,)``.
        y: Observations of shape ``(N,)``.
        integrator: ``gaussx`` integrator for non-conjugate
            likelihoods. ``None`` is fine for
            `pyrox_gp.GaussianLikelihood`.

    Returns:
        Scalar ELBO value (higher is better).

    Raises:
        ValueError: If a non-conjugate likelihood is used without an
            integrator.
    """
    from gaussx import variational_elbo_gaussian

    times_arr = jnp.asarray(times)
    K_zz_op, K_xz, K_xx_diag = prior.predictive_blocks(times_arr)
    f_loc, f_var = guide.predict(K_xz, K_zz_op, K_xx_diag)  # ty: ignore[unresolved-attribute]
    f_loc = f_loc + prior.mean(times_arr)
    kl = guide.kl_divergence(K_zz_op)  # ty: ignore[unresolved-attribute]

    if isinstance(likelihood, GaussianLikelihood):
        return variational_elbo_gaussian(
            y,
            f_loc,
            f_var,
            likelihood.noise_var,  # ty: ignore[invalid-argument-type]
            kl,
        )

    if integrator is None:
        raise ValueError(
            "Non-conjugate likelihoods require an integrator "
            "(e.g. gaussx.GaussHermiteIntegrator). "
            "Pass integrator=GaussHermiteIntegrator(order=20)."
        )
    ell = _ell_numerical(likelihood, y, f_loc, f_var, integrator)
    return ell - kl

sparse_markov_factor(name: str, prior: SparseMarkovGPPrior, guide: Guide, likelihood: Likelihood, times: Float[Array, ' N'], y: Float[Array, ' N'], *, integrator: AbstractIntegrator | None = None) -> None

Register sparse_markov_elbo as a NumPyro factor site.

Source code in packages/pyrox-gp/src/pyrox_gp/_sparse_markov.py
def sparse_markov_factor(
    name: str,
    prior: SparseMarkovGPPrior,
    guide: Guide,
    likelihood: Likelihood,
    times: Float[Array, " N"],
    y: Float[Array, " N"],
    *,
    integrator: AbstractIntegrator | None = None,
) -> None:
    """Register `sparse_markov_elbo` as a NumPyro factor site."""
    numpyro.factor(
        name,
        sparse_markov_elbo(prior, guide, likelihood, times, y, integrator=integrator),
    )

Component protocols

Abstract pyrox-local bases for the orthogonal component stack — the contracts that the concrete kernels, guides, and likelihoods above implement. Cubature integrators (Gauss-Hermite, Monte Carlo) come from gaussx.AbstractIntegrator and its concrete subclasses; solver strategies live in gaussx.

Kernel = AbstractKernel module-attribute

Guide

Bases: Module

Abstract base for variational posterior families.

Concrete guides (DeltaGuide, MeanFieldGuide, LowRankGuide, FullRankGuide, etc.) land in the dedicated guide waves (#28, #29). The whitening principle keeps optimization geometry well-conditioned — sample from a unit-scale latent and unwhiten with the prior Cholesky.

Two distinct entry points:

  • sample / log_prob — pure variational draws and densities. sample(self, key) returns a draw from q(f); log_prob(self, f) evaluates log q(f). Neither touches the NumPyro trace.
  • register(name, prior) (optional) — the NumPyro-integration hook invoked by pyrox_gp.gp_sample when a guide is supplied. Use it to register a sample / param site (or compose one out of guide state) under name and return the latent function value. Concrete guides that participate in gp_sample should implement this; the protocol leaves it unspecified so guides usable purely outside NumPyro stay valid.
Source code in packages/pyrox-gp/src/pyrox_gp/_protocols.py
class Guide(eqx.Module):
    """Abstract base for variational posterior families.

    Concrete guides (``DeltaGuide``, ``MeanFieldGuide``, ``LowRankGuide``,
    ``FullRankGuide``, etc.) land in the dedicated guide waves (#28, #29).
    The whitening principle keeps optimization geometry well-conditioned —
    sample from a unit-scale latent and unwhiten with the prior Cholesky.

    Two distinct entry points:

    * `sample` / `log_prob` — pure variational draws and
      densities. ``sample(self, key)`` returns a draw from ``q(f)``;
      ``log_prob(self, f)`` evaluates ``log q(f)``. Neither touches the
      NumPyro trace.
    * ``register(name, prior)`` (optional) — the NumPyro-integration hook
      invoked by `pyrox_gp.gp_sample` when a guide is supplied. Use
      it to register a sample / param site (or compose one out of guide
      state) under ``name`` and return the latent function value. Concrete
      guides that participate in `gp_sample` should implement this;
      the protocol leaves it unspecified so guides usable purely outside
      NumPyro stay valid.
    """

    @abstractmethod
    def sample(self, key: Any) -> Float[Array, " ..."]:
        raise NotImplementedError

    @abstractmethod
    def log_prob(self, f: Float[Array, " ..."]) -> Float[Array, ""]:
        raise NotImplementedError

Likelihood

Bases: Module

Abstract base for observation models.

Implements the conditional p(y | f). The advanced inference strategies in pyrox_gp._inference_nongauss integrate log p(y | f) against a Gaussian cavity via any gaussx.AbstractIntegrator. Concrete scalar-latent likelihoods (GaussianLikelihood, BernoulliLikelihood, PoissonLikelihood, StudentTLikelihood) and multi-latent ones (SoftmaxLikelihood, HeteroscedasticGaussianLikelihood) live in pyrox_gp._likelihoods.

Multi-latent likelihoods declare latent_dim: int as a static field (e.g. latent_dim = num_classes for softmax). Scalar likelihoods may omit the field; consumers should read getattr(lik, "latent_dim", 1).

Source code in packages/pyrox-gp/src/pyrox_gp/_protocols.py
class Likelihood(eqx.Module):
    """Abstract base for observation models.

    Implements the conditional ``p(y | f)``. The advanced inference
    strategies in `pyrox_gp._inference_nongauss` integrate
    ``log p(y | f)`` against a Gaussian cavity via any
    `gaussx.AbstractIntegrator`. Concrete scalar-latent likelihoods
    (`GaussianLikelihood`, `BernoulliLikelihood`,
    `PoissonLikelihood`, `StudentTLikelihood`) and
    multi-latent ones (`SoftmaxLikelihood`,
    `HeteroscedasticGaussianLikelihood`) live in
    `pyrox_gp._likelihoods`.

    Multi-latent likelihoods declare ``latent_dim: int`` as a static
    field (e.g. ``latent_dim = num_classes`` for softmax). Scalar
    likelihoods may omit the field; consumers should read
    ``getattr(lik, "latent_dim", 1)``.
    """

    @abstractmethod
    def log_prob(
        self,
        f: Float[Array, " ..."],
        y: Float[Array, " ..."],
        X: Float[Array, " ..."] | None = None,
    ) -> Float[Array, ""]:
        """Conditional ``p(y | f)``, optionally conditioned on inputs.

        ``X`` is consumed only by likelihoods whose parameters are
        functions of the input (see `pyrox_gp.WarpedGaussianLikelihood`
        with a conditional warp). Scalar observation models ignore it.
        """
        raise NotImplementedError

log_prob(f: Float[Array, ' ...'], y: Float[Array, ' ...'], X: Float[Array, ' ...'] | None = None) -> Float[Array, ''] abstractmethod

Conditional p(y | f), optionally conditioned on inputs.

X is consumed only by likelihoods whose parameters are functions of the input (see pyrox_gp.WarpedGaussianLikelihood with a conditional warp). Scalar observation models ignore it.

Source code in packages/pyrox-gp/src/pyrox_gp/_protocols.py
@abstractmethod
def log_prob(
    self,
    f: Float[Array, " ..."],
    y: Float[Array, " ..."],
    X: Float[Array, " ..."] | None = None,
) -> Float[Array, ""]:
    """Conditional ``p(y | f)``, optionally conditioned on inputs.

    ``X`` is consumed only by likelihoods whose parameters are
    functions of the input (see `pyrox_gp.WarpedGaussianLikelihood`
    with a conditional warp). Scalar observation models ignore it.
    """
    raise NotImplementedError

Math primitives

Pure JAX kernel functions. Stateless, differentiable, composable — (Array, ..., hyperparams) -> Gram. No NumPyro, no protocols.

Deprecated: the pure kernel functions moved to kernellib.functional.

This module re-exports them unchanged so existing imports keep working, and warns on import. The math, and every kernel value, is identical: kernellib ported this module and its tests verbatim.

Replace

from pyrox_gp._src.kernels import rbf_kernel

with

from kernellib.functional import rbf_kernel