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
|
|
num_basis_per_dim |
tuple[int, ...]
|
Per-axis number of 1D eigenfunctions; total
count is prod(num_basis_per_dim).
|
L |
tuple[float, ...]
|
|
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()
|
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
|
|
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
|
|
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.
- Predict
(f_loc, f_var) = guide.predict(...)
- Compute per-point natural-gradient targets
\((\lambda_n^{(1)}, \Lambda_n^{(2)})\).
- 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}
\]
- Damped update via
NaturalGuide.natural_update.
Parameters:
| Name |
Type |
Description |
Default |
prior
|
SparseGPPrior
|
|
required
|
guide
|
NaturalGuide
|
Current natural-parameter guide.
|
required
|
likelihood
|
Likelihood
|
|
required
|
X
|
Float[Array, 'N D']
|
Training inputs, shape (N, D).
|
required
|
y
|
Float[Array, ' N']
|
Training targets, shape (N,).
|
required
|
Returns:
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
|
|
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
|
|
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
|
|
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
|
|
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']
|
|
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, '']
|
|
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, '']
|
|
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']
|
|
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, '']
|
|
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, '']
|
|
period |
Float[Array, '']
|
|
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:
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
|
|
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']
|
|
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
| 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
|
|
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
|
|
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