cuthbertlib 0.0.8__tar.gz → 0.0.10__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/PKG-INFO +1 -1
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/__init__.py +1 -0
- cuthbertlib-0.0.10/cuthbertlib/enkf/README.md +14 -0
- cuthbertlib-0.0.10/cuthbertlib/enkf/__init__.py +1 -0
- cuthbertlib-0.0.10/cuthbertlib/enkf/filtering.py +119 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/kalman/filtering.py +37 -1
- cuthbertlib-0.0.10/cuthbertlib/linalg/tria.py +97 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linearize/log_density.py +2 -0
- cuthbertlib-0.0.10/cuthbertlib/resampling/README.md +53 -0
- cuthbertlib-0.0.10/cuthbertlib/resampling/__init__.py +11 -0
- cuthbertlib-0.0.10/cuthbertlib/resampling/adaptive.py +72 -0
- cuthbertlib-0.0.10/cuthbertlib/resampling/autodiff.py +64 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/resampling/killing.py +24 -10
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/resampling/multinomial.py +18 -7
- cuthbertlib-0.0.10/cuthbertlib/resampling/protocols.py +98 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/resampling/systematic.py +20 -9
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/resampling/utils.py +12 -3
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/smoothing/exact_sampling.py +2 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/smoothing/mcmc.py +5 -3
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/smoothing/protocols.py +1 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/types.py +3 -1
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/pyproject.toml +1 -1
- cuthbertlib-0.0.8/cuthbertlib/linalg/tria.py +0 -21
- cuthbertlib-0.0.8/cuthbertlib/resampling/README.md +0 -26
- cuthbertlib-0.0.8/cuthbertlib/resampling/__init__.py +0 -3
- cuthbertlib-0.0.8/cuthbertlib/resampling/protocols.py +0 -92
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/.gitignore +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/README.md +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/README.md +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/discrete/__init__.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/discrete/filtering.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/discrete/smoothing.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/kalman/README.md +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/kalman/__init__.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/kalman/generate.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/kalman/sampling.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/kalman/smoothing.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linalg/README.md +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linalg/__init__.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linalg/collect_nans_chol.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linalg/marginal_sqrt_cov.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linalg/symmetric_inv_sqrt.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linearize/README.md +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linearize/__init__.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linearize/moments.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linearize/taylor.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/README.md +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/__init__.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/common.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/cubature.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/gauss_hermite.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/linearize.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/unscented.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/utils.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/README.md +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/__init__.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/ess.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/smoothing/__init__.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/smoothing/tracing.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/stats/README.md +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/stats/__init__.py +0 -0
- {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/stats/multivariate_normal.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: cuthbertlib
|
|
3
|
-
Version: 0.0.
|
|
3
|
+
Version: 0.0.10
|
|
4
4
|
Summary: Atomic building blocks for state-space model inference with JAX
|
|
5
5
|
Author-email: Sam Duffield <s@mduffield.com>, Sahel Iqbal <sahel13miqbal@proton.me>, Adrien Corenflos <adrien.corenflos.stats@gmail.com>
|
|
6
6
|
License: Apache-2.0
|
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
# Ensemble Kalman Filter (EnKF)
|
|
2
|
+
|
|
3
|
+
This sub-repository provides modular functions for the Ensemble Kalman Filter.
|
|
4
|
+
|
|
5
|
+
The core functions are:
|
|
6
|
+
|
|
7
|
+
- `predict`: Propagate ensemble members through nonlinear dynamics with additive Gaussian noise.
|
|
8
|
+
- `update`: Update ensemble members with an observation using the EnKF update equation.
|
|
9
|
+
|
|
10
|
+
Together, `predict` and `update` can be used to perform an online EnKF filtering step.
|
|
11
|
+
|
|
12
|
+
The EnKF uses an ensemble of particles with a Kalman-style measurement update based on
|
|
13
|
+
empirical covariances. Unlike the EKF, it does not require Jacobians, while naturally
|
|
14
|
+
handling nonlinear dynamics.
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
from cuthbertlib.enkf.filtering import predict, update
|
|
@@ -0,0 +1,119 @@
|
|
|
1
|
+
"""Implements the Ensemble Kalman Filter (EnKF) predict and update steps.
|
|
2
|
+
|
|
3
|
+
See Algorithm 10.2, [Sanz-Alonso et al., Inverse Problems and Data Assimilation](https://arxiv.org/abs/1810.06191).
|
|
4
|
+
Based in part on the [CD-Dynamax implementation](https://github.com/hd-UQ/cd_dynamax/blob/public/cd_dynamax/src/continuous_discrete_nonlinear_gaussian_ssm/inference_enkf.py).
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from typing import Callable
|
|
8
|
+
|
|
9
|
+
import jax
|
|
10
|
+
import jax.numpy as jnp
|
|
11
|
+
from jax import random
|
|
12
|
+
from jax.scipy.linalg import cho_solve
|
|
13
|
+
|
|
14
|
+
from cuthbertlib.linalg import collect_nans_chol, tria
|
|
15
|
+
from cuthbertlib.stats import multivariate_normal
|
|
16
|
+
from cuthbertlib.types import Array, ArrayTreeLike, KeyArray, ScalarArray
|
|
17
|
+
|
|
18
|
+
ObservationFn = Callable[[Array], Array]
|
|
19
|
+
DynamicsFn = Callable[[Array, KeyArray], Array]
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def predict(
|
|
23
|
+
key: KeyArray,
|
|
24
|
+
ensemble: Array,
|
|
25
|
+
dynamics_fn: DynamicsFn,
|
|
26
|
+
inflation: float = 0.0,
|
|
27
|
+
) -> Array:
|
|
28
|
+
"""Propagate ensemble members through an arbitrary simulator p(x_{t+1} | x_t).
|
|
29
|
+
|
|
30
|
+
Args:
|
|
31
|
+
key: JAX PRNG key.
|
|
32
|
+
ensemble: Ensemble of state vectors, shape (N, x_dim).
|
|
33
|
+
dynamics_fn: Dynamics function mapping (state, key) -> state.
|
|
34
|
+
inflation: Multiplicative inflation factor applied to ensemble deviations.
|
|
35
|
+
|
|
36
|
+
Returns:
|
|
37
|
+
Predicted ensemble, shape (N, x_dim).
|
|
38
|
+
"""
|
|
39
|
+
N, x_dim = ensemble.shape
|
|
40
|
+
|
|
41
|
+
# Propagate each member through the dynamics
|
|
42
|
+
keys = random.split(key, N)
|
|
43
|
+
propagated = jax.vmap(dynamics_fn, (0, 0))(ensemble, keys)
|
|
44
|
+
|
|
45
|
+
# Apply multiplicative inflation
|
|
46
|
+
mean = jnp.mean(propagated, axis=0)
|
|
47
|
+
propagated = mean + (1 + inflation) * (propagated - mean)
|
|
48
|
+
|
|
49
|
+
return propagated
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def update(
|
|
53
|
+
key: KeyArray,
|
|
54
|
+
predicted_ensemble: Array,
|
|
55
|
+
observation_fn: ObservationFn,
|
|
56
|
+
chol_R: Array,
|
|
57
|
+
y: Array,
|
|
58
|
+
perturbed_obs: bool = True,
|
|
59
|
+
) -> tuple[Array, ScalarArray]:
|
|
60
|
+
"""Update ensemble members with an observation using the EnKF update.
|
|
61
|
+
|
|
62
|
+
NaNs in ``y`` are treated as missing dimensions and are excluded from the
|
|
63
|
+
update. When ``y`` is entirely NaN, the update is a no-op: the predicted
|
|
64
|
+
ensemble is returned unchanged with zero log-likelihood contribution.
|
|
65
|
+
|
|
66
|
+
Args:
|
|
67
|
+
key: JAX PRNG key.
|
|
68
|
+
predicted_ensemble: Predicted ensemble, shape (N, x_dim).
|
|
69
|
+
observation_fn: Observation function mapping state -> obs.
|
|
70
|
+
chol_R: Cholesky factor of the observation noise covariance, shape (y_dim, y_dim).
|
|
71
|
+
y: Observation vector, shape (y_dim,). NaNs indicate missing dimensions.
|
|
72
|
+
perturbed_obs: If True, use perturbed observations (stochastic EnKF).
|
|
73
|
+
If False, use deterministic update.
|
|
74
|
+
|
|
75
|
+
Returns:
|
|
76
|
+
Tuple of (updated_ensemble, log_likelihood).
|
|
77
|
+
"""
|
|
78
|
+
N, x_dim = predicted_ensemble.shape
|
|
79
|
+
|
|
80
|
+
# Map ensemble to observation space
|
|
81
|
+
y_pred = jax.vmap(observation_fn, (0,))(predicted_ensemble)
|
|
82
|
+
|
|
83
|
+
# Handle partially-missing observations by reordering and zeroing missing dims.
|
|
84
|
+
# Use y_pred.T because y_pred is (N, y_dim) and we want to reorder along axis 0.
|
|
85
|
+
flag = jnp.isnan(y)
|
|
86
|
+
flag, chol_R, y, y_pred = collect_nans_chol(flag, chol_R, y, y_pred.T)
|
|
87
|
+
y_pred = y_pred.T
|
|
88
|
+
y_dim = y.shape[0]
|
|
89
|
+
|
|
90
|
+
# Ensemble means
|
|
91
|
+
x_mean = jnp.mean(predicted_ensemble, axis=0)
|
|
92
|
+
y_mean = jnp.mean(y_pred, axis=0)
|
|
93
|
+
|
|
94
|
+
# Deviations from ensemble mean
|
|
95
|
+
x_dev = predicted_ensemble - x_mean
|
|
96
|
+
y_dev = y_pred - y_mean
|
|
97
|
+
|
|
98
|
+
# Square-root innovation covariance via tria
|
|
99
|
+
chol_S = tria(jnp.concatenate([y_dev.T / jnp.sqrt(N - 1), chol_R], axis=1))
|
|
100
|
+
|
|
101
|
+
# Cross-covariance
|
|
102
|
+
C_xy = x_dev.T @ y_dev / (N - 1)
|
|
103
|
+
|
|
104
|
+
# Kalman gain: K = C_xy @ S^{-1} = C_xy @ cho_solve(chol_S, I)
|
|
105
|
+
K = cho_solve((chol_S, True), C_xy.T).T
|
|
106
|
+
|
|
107
|
+
# Innovation per member
|
|
108
|
+
if perturbed_obs:
|
|
109
|
+
y_n = y[None, :] + (chol_R @ random.normal(key, (y_dim, N))).T
|
|
110
|
+
else:
|
|
111
|
+
y_n = jnp.broadcast_to(y[None, :], (N, y_dim))
|
|
112
|
+
|
|
113
|
+
# Update ensemble
|
|
114
|
+
updated = predicted_ensemble + (y_n - y_pred) @ K.T
|
|
115
|
+
|
|
116
|
+
# Log-likelihood
|
|
117
|
+
ll = multivariate_normal.logpdf(y, y_mean, chol_S, nan_support=False)
|
|
118
|
+
|
|
119
|
+
return updated, jnp.asarray(ll)
|
|
@@ -208,6 +208,42 @@ def filtering_operator(
|
|
|
208
208
|
|
|
209
209
|
mu = cho_solve((U1, True), b1)
|
|
210
210
|
t1 = b1 @ mu - (eta2 + mu) @ tmp_2
|
|
211
|
-
|
|
211
|
+
|
|
212
|
+
# Derivation for O(nx) log-determinant computation of D_inv:
|
|
213
|
+
# This is a long comment but I wanted to include the full derivation for clarity and future reference.
|
|
214
|
+
# The key idea is to express D_inv in terms of the blocks of Xi and then apply Sylvester's determinant theorem
|
|
215
|
+
# to compute its determinant efficiently.
|
|
216
|
+
#
|
|
217
|
+
# 1. Expand the blocks of Xi @ Xi.T:
|
|
218
|
+
# (Xi @ Xi.T)[1,1] = I + U1.T @ Z2 @ Z2.T @ U1
|
|
219
|
+
# (Xi @ Xi.T)[2,1] = Z2 @ Z2.T @ U1
|
|
220
|
+
#
|
|
221
|
+
# 2. Equate to the corresponding blocks of L @ L.T:
|
|
222
|
+
# Xi11 @ Xi11.T = I + U1.T @ Z2 @ Z2.T @ U1
|
|
223
|
+
# Xi21 @ Xi11.T = Z2 @ Z2.T @ U1
|
|
224
|
+
#
|
|
225
|
+
# 3. Expand D_inv using tmp_1 = Xi11^{-1} @ U1.T:
|
|
226
|
+
# D_inv = I - tmp_1.T @ Xi21.T
|
|
227
|
+
# = I - U1 @ Xi11^{-T} @ Xi21.T
|
|
228
|
+
#
|
|
229
|
+
# 4. Apply Sylvester's determinant theorem:
|
|
230
|
+
# det(D_inv) = det(I - Xi11^{-T} @ Xi21.T @ U1)
|
|
231
|
+
#
|
|
232
|
+
# 5. Multiply interior by Xi11^{-1} @ Xi11 and substitute block identities:
|
|
233
|
+
# Let P = U1.T @ Z2 @ Z2.T @ U1
|
|
234
|
+
# det(D_inv) = det(I - (Xi11 @ Xi11.T)^{-1} @ (Xi21 @ Xi11.T).T @ U1)
|
|
235
|
+
# = det(I - (I + P)^{-1} @ P)
|
|
236
|
+
# = det((I + P)^{-1})
|
|
237
|
+
# = 1 / det(Xi11 @ Xi11.T)
|
|
238
|
+
# = det(Xi11)^{-2}
|
|
239
|
+
#
|
|
240
|
+
# 6. Simplify the log-determinant term in the log-likelihood:
|
|
241
|
+
# 0.5 * log(det(D_inv)) = -log(|det(Xi11)|)
|
|
242
|
+
#
|
|
243
|
+
# Since Xi11 is lower triangular, the log-determinant is the sum of the logs of its diagonal.
|
|
244
|
+
# Replace `0.5 * jnp.linalg.slogdet(D_inv)[1]` with:
|
|
245
|
+
# -jnp.sum(jnp.log(jnp.abs(jnp.diag(Xi11))))
|
|
246
|
+
# ell = ell1 + ell2 - 0.5 * t1 + 0.5 * jnp.linalg.slogdet(D_inv)[1]
|
|
247
|
+
ell = ell1 + ell2 - 0.5 * t1 - jnp.sum(jnp.log(jnp.abs(jnp.diag(Xi11))))
|
|
212
248
|
|
|
213
249
|
return FilterScanElement(A, b, U, eta, Z, ell)
|
|
@@ -0,0 +1,97 @@
|
|
|
1
|
+
"""Implements triangularization operator a matrix via QR decomposition."""
|
|
2
|
+
|
|
3
|
+
import jax
|
|
4
|
+
import jax.numpy as jnp
|
|
5
|
+
|
|
6
|
+
from cuthbertlib.types import Array
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def _adj(x: Array) -> Array:
|
|
10
|
+
"""Conjugate transpose for batched matrices."""
|
|
11
|
+
return jnp.swapaxes(x.conj(), -1, -2)
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@jax.custom_jvp
|
|
15
|
+
def tria(A: Array) -> Array:
|
|
16
|
+
"""A triangularization operator using QR decomposition.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
A: The matrix to triangularize.
|
|
20
|
+
|
|
21
|
+
Returns:
|
|
22
|
+
A lower triangular matrix R such that R @ R.T = A @ A.T.
|
|
23
|
+
|
|
24
|
+
References:
|
|
25
|
+
Paper: Arasaratnam and Haykin (2008): Square-Root Quadrature Kalman Filtering
|
|
26
|
+
https://ieeexplore.ieee.org/document/4524036
|
|
27
|
+
"""
|
|
28
|
+
_, R_qr = jnp.linalg.qr(_adj(A), mode="reduced")
|
|
29
|
+
return _adj(R_qr)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
@tria.defjvp
|
|
33
|
+
def _tria_jvp(primals, tangents):
|
|
34
|
+
# Derivation of the exact analytical JVP for the lower-triangularization operator:
|
|
35
|
+
#
|
|
36
|
+
# 1. The operation computes a lower triangular R such that:
|
|
37
|
+
# R @ R.T = A @ A.T
|
|
38
|
+
#
|
|
39
|
+
# 2. Taking the differential of both sides yields the Gramian differential identity:
|
|
40
|
+
# dR @ R.T + R @ dR.T = dA @ A.T + A @ dA.T
|
|
41
|
+
#
|
|
42
|
+
# 3. For rank-deficient A, the exact lower-triangular tangent dR decomposes into
|
|
43
|
+
# column-space and null-space components: dR = dR_col + dR_null.
|
|
44
|
+
#
|
|
45
|
+
# 4. To find dR_col, multiply the identity from the left by R^\dagger (pseudoinverse)
|
|
46
|
+
# and from the right by R^{\dagger T}:
|
|
47
|
+
# R^\dagger @ dR_col + dR_col.T @ R^{\dagger T} = R^\dagger @ dA @ A.T @ R^{\dagger T} + R^\dagger @ A @ dA.T @ R^{\dagger T}
|
|
48
|
+
#
|
|
49
|
+
# 5. Let Q be the active subspace orthogonal factor such that A = R @ Q.T.
|
|
50
|
+
# Substituting A.T @ R^{\dagger T} = Q and R^\dagger @ A = Q.T:
|
|
51
|
+
# R^\dagger @ dR_col + dR_col.T @ R^{\dagger T} = R^\dagger @ dA @ Q + Q.T @ dA.T @ R^{\dagger T}
|
|
52
|
+
#
|
|
53
|
+
# 6. Define K = R^\dagger @ dA @ Q. The right hand side becomes K + K.T:
|
|
54
|
+
# R^\dagger @ dR_col + (R^\dagger @ dR_col).T = K + K.T
|
|
55
|
+
#
|
|
56
|
+
# 7. Define dM_col = R^\dagger @ dR_col. Solving for the lower-triangular dM_col:
|
|
57
|
+
# dM_col = tril(K + K.T) - diag(K)
|
|
58
|
+
#
|
|
59
|
+
# 8. Recover the column-space differential dR_col by left-multiplying by R:
|
|
60
|
+
# dR_col = R @ dM_col
|
|
61
|
+
#
|
|
62
|
+
# 9. To satisfy the orthogonal cross-terms for arbitrary perturbations outside
|
|
63
|
+
# the column space of A, add the null-space component projected via (I - R @ R^\dagger):
|
|
64
|
+
# dR_null = (I - R @ R^\dagger) @ dA @ Q
|
|
65
|
+
#
|
|
66
|
+
# 10. The complete, exact JVP is the sum of both components:
|
|
67
|
+
# dR = dR_col + dR_null
|
|
68
|
+
|
|
69
|
+
(A,) = primals
|
|
70
|
+
(dA,) = tangents
|
|
71
|
+
|
|
72
|
+
A_T = jnp.swapaxes(A, -1, -2)
|
|
73
|
+
Q, R_qr = jnp.linalg.qr(A_T, mode="reduced")
|
|
74
|
+
|
|
75
|
+
R = jnp.swapaxes(R_qr, -1, -2)
|
|
76
|
+
|
|
77
|
+
# Q has shape (..., M, N). A is (..., N, M).
|
|
78
|
+
# A^T = Q R_qr => A = R_qr^T Q^T = R Q^T
|
|
79
|
+
|
|
80
|
+
R_pinv = jnp.linalg.pinv(R)
|
|
81
|
+
|
|
82
|
+
# K = R^{-1} dA Q
|
|
83
|
+
# R can be degenerate so we use the pseudoinverse to ensure the JVP is well-defined everywhere,
|
|
84
|
+
K = R_pinv @ dA @ Q
|
|
85
|
+
K_T = jnp.swapaxes(K, -1, -2)
|
|
86
|
+
|
|
87
|
+
# Solve for lower triangular perturbation dM + dM^T = K + K^T
|
|
88
|
+
I = jnp.eye(K.shape[-1], dtype=K.dtype)
|
|
89
|
+
dM = jnp.tril(K + K_T) - K * I
|
|
90
|
+
|
|
91
|
+
# Compute the null-space part
|
|
92
|
+
dR_null = (jnp.eye(R.shape[-2], dtype=R.dtype) - R @ R_pinv) @ dA @ Q
|
|
93
|
+
|
|
94
|
+
# Apply to get the tangent at R
|
|
95
|
+
dR = R @ dM + dR_null
|
|
96
|
+
|
|
97
|
+
return R, dR
|
|
@@ -63,6 +63,7 @@ def linearize_log_density(
|
|
|
63
63
|
|
|
64
64
|
Args:
|
|
65
65
|
log_density: A conditional log density of y given x. Returns a scalar.
|
|
66
|
+
x must be the first argument and y the second.
|
|
66
67
|
x: The input points.
|
|
67
68
|
y: The output points.
|
|
68
69
|
has_aux: Whether `log_density` returns an auxiliary value.
|
|
@@ -137,6 +138,7 @@ def linearize_log_density_given_chol_cov(
|
|
|
137
138
|
|
|
138
139
|
Args:
|
|
139
140
|
log_density: A conditional log density of y given x. Returns a scalar.
|
|
141
|
+
x must be the first argument and y the second.
|
|
140
142
|
x: The input points.
|
|
141
143
|
y: The output points.
|
|
142
144
|
chol_cov: The Cholesky factor of the covariance matrix of the Gaussian.
|
|
@@ -0,0 +1,53 @@
|
|
|
1
|
+
# Resampling
|
|
2
|
+
|
|
3
|
+
This sub-repository provides a unified interface for a variety of resampling
|
|
4
|
+
methods, which convert a set of weighted samples into an unweighted one which
|
|
5
|
+
likely contains duplicates.
|
|
6
|
+
|
|
7
|
+
A typical call to the library would be:
|
|
8
|
+
|
|
9
|
+
```python
|
|
10
|
+
sampling_key, resampling_key = jax.random.split(jax.random.key(0))
|
|
11
|
+
particles = jax.random.normal(sampling_key, (100, 2))
|
|
12
|
+
logits = jax.vmap(lambda x: jnp.where(jnp.all(x > 0), 0, -jnp.inf))(particles)
|
|
13
|
+
|
|
14
|
+
resampled_indices, _, resampled_particles = resampling.multinomial.resampling(resampling_key, logits, particles, 100)
|
|
15
|
+
```
|
|
16
|
+
|
|
17
|
+
Or for conditional resampling:
|
|
18
|
+
|
|
19
|
+
```python
|
|
20
|
+
# Here we resample but keep particle at index 0 fixed
|
|
21
|
+
conditional_resampled_indices, _, conditional_resampled_particles = resampling.multinomial.conditional_resampling(
|
|
22
|
+
resampling_key, logits, particles, 100, pivot_in=0, pivot_out=0
|
|
23
|
+
)
|
|
24
|
+
```
|
|
25
|
+
|
|
26
|
+
Adaptive resampling (i.e. resampling only when the effective sample size is below a
|
|
27
|
+
threshold) is also supported via a decorator:
|
|
28
|
+
|
|
29
|
+
```python
|
|
30
|
+
adaptive_resampling = resampling.adaptive.ess_decorator(
|
|
31
|
+
resampling.multinomial.resampling,
|
|
32
|
+
threshold=0.5,
|
|
33
|
+
)
|
|
34
|
+
adaptive_resampled_indices, _, adaptive_resampled_particles = adaptive_resampling(
|
|
35
|
+
resampling_key, logits, particles, 100
|
|
36
|
+
)
|
|
37
|
+
```
|
|
38
|
+
|
|
39
|
+
For consistent gradient estimates with respect to model parameters, the [stop-gradient particle filter](https://arxiv.org/abs/2106.10314) is also implemented as a decorator.
|
|
40
|
+
|
|
41
|
+
```python
|
|
42
|
+
differentiable_resampling = resampling.stop_gradient.stop_gradient_decorator(
|
|
43
|
+
resampling.multinomial.resampling
|
|
44
|
+
)
|
|
45
|
+
# can be combined with adaptive resampling
|
|
46
|
+
adaptive_and_differentiable_resampling = resampling.adaptive.ess_decorator(
|
|
47
|
+
differentiable_resampling,
|
|
48
|
+
threshold=0.5,
|
|
49
|
+
)
|
|
50
|
+
resampled_indices, _, resampled_particles = adaptive_and_differentiable_resampling(
|
|
51
|
+
resampling_key, logits, particles, 100
|
|
52
|
+
)
|
|
53
|
+
```
|
|
@@ -0,0 +1,11 @@
|
|
|
1
|
+
from cuthbertlib.resampling import (
|
|
2
|
+
adaptive,
|
|
3
|
+
autodiff,
|
|
4
|
+
killing,
|
|
5
|
+
multinomial,
|
|
6
|
+
systematic,
|
|
7
|
+
)
|
|
8
|
+
from cuthbertlib.resampling.adaptive import ess_decorator
|
|
9
|
+
from cuthbertlib.resampling.autodiff import stop_gradient_decorator
|
|
10
|
+
from cuthbertlib.resampling.protocols import ConditionalResampling, Resampling
|
|
11
|
+
from cuthbertlib.resampling.utils import inverse_cdf
|
|
@@ -0,0 +1,72 @@
|
|
|
1
|
+
"""Adaptive resampling decorator.
|
|
2
|
+
|
|
3
|
+
Provides a decorator to turn any Resampling function into an adaptive resampling
|
|
4
|
+
function which performs resampling only when the effective sample size (ESS)
|
|
5
|
+
falls below a threshold.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from functools import wraps
|
|
9
|
+
|
|
10
|
+
import jax
|
|
11
|
+
import jax.numpy as jnp
|
|
12
|
+
|
|
13
|
+
from cuthbertlib.resampling.protocols import Resampling
|
|
14
|
+
from cuthbertlib.smc.ess import log_ess
|
|
15
|
+
from cuthbertlib.types import Array, ArrayLike, ArrayTree, ArrayTreeLike
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def ess_decorator(func: Resampling, threshold: float) -> Resampling:
|
|
19
|
+
"""Wrap a Resampling function so that it only resamples when ESS < threshold.
|
|
20
|
+
|
|
21
|
+
The returned function is jitted and has `n` as a static argument. The
|
|
22
|
+
original resampler's docstring is appended to this wrapper's docstring so
|
|
23
|
+
IDEs and users can see the underlying algorithm documentation.
|
|
24
|
+
|
|
25
|
+
Args:
|
|
26
|
+
func: A resampling function with signature
|
|
27
|
+
(key, logits, positions, n) -> (indices, logits_out, positions_out).
|
|
28
|
+
threshold: Fraction of particle count specifying when to resample.
|
|
29
|
+
Resampling is triggered when ESS < ess_threshold * n.
|
|
30
|
+
|
|
31
|
+
Returns:
|
|
32
|
+
A Resampling function implementing adaptive resampling.
|
|
33
|
+
"""
|
|
34
|
+
# Build a descriptive docstring that includes the wrapped function doc
|
|
35
|
+
wrapped_doc = func.__doc__ or ""
|
|
36
|
+
doc = f"""
|
|
37
|
+
Adaptive resampling decorator (threshold={threshold}).
|
|
38
|
+
|
|
39
|
+
This wrapper will call the provided resampling function only when the
|
|
40
|
+
effective sample size (ESS) is below `ess_threshold * n`.
|
|
41
|
+
|
|
42
|
+
Wrapped resampler documentation:
|
|
43
|
+
{wrapped_doc}
|
|
44
|
+
"""
|
|
45
|
+
|
|
46
|
+
@wraps(func)
|
|
47
|
+
def _wrapped(
|
|
48
|
+
key: Array, logits: ArrayLike, positions: ArrayTreeLike, n: int
|
|
49
|
+
) -> tuple[Array, Array, ArrayTree]:
|
|
50
|
+
logits_arr = jnp.asarray(logits)
|
|
51
|
+
N = logits_arr.shape[0]
|
|
52
|
+
if n != N:
|
|
53
|
+
raise AssertionError(
|
|
54
|
+
"The number of sampled indices must be equal to the number of "
|
|
55
|
+
f"particles for `adaptive` resampling. Got {n} instead of {N}."
|
|
56
|
+
)
|
|
57
|
+
|
|
58
|
+
def _do_resample():
|
|
59
|
+
return func(key, logits_arr, positions, n)
|
|
60
|
+
|
|
61
|
+
def _no_resample():
|
|
62
|
+
return jnp.arange(n), logits_arr, positions
|
|
63
|
+
|
|
64
|
+
return jax.lax.cond(
|
|
65
|
+
log_ess(logits_arr) < jnp.log(threshold * n),
|
|
66
|
+
_do_resample,
|
|
67
|
+
_no_resample,
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
# Attach the composed docstring and return a jitted version
|
|
71
|
+
_wrapped.__doc__ = doc
|
|
72
|
+
return jax.jit(_wrapped, static_argnames=("n",))
|
|
@@ -0,0 +1,64 @@
|
|
|
1
|
+
"""Implements decorators for automatic differentiation of resampling schemes.
|
|
2
|
+
|
|
3
|
+
Current supported is the stop_gradient resampling scheme, which provides the
|
|
4
|
+
classical Fisher estimates for the score function via automatic differentiation.
|
|
5
|
+
This can be wrapped around a resampling scheme such as multinomial or systematic
|
|
6
|
+
resampling.
|
|
7
|
+
|
|
8
|
+
See [Scibior and Wood (2021)](https://arxiv.org/abs/2106.10314) for more details.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
from functools import wraps
|
|
12
|
+
|
|
13
|
+
import jax
|
|
14
|
+
import jax.numpy as jnp
|
|
15
|
+
|
|
16
|
+
from cuthbertlib.resampling.protocols import Resampling
|
|
17
|
+
from cuthbertlib.resampling.utils import apply_resampling_indices
|
|
18
|
+
from cuthbertlib.smc.ess import log_ess
|
|
19
|
+
from cuthbertlib.types import Array, ArrayLike, ArrayTree, ArrayTreeLike
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def stop_gradient_decorator(func: Resampling) -> Resampling:
|
|
23
|
+
"""Wrap a Resampling function to use stop gradient resampling.
|
|
24
|
+
|
|
25
|
+
Args:
|
|
26
|
+
func: A resampling function with signature
|
|
27
|
+
(key, logits, positions, n) -> (indices, logits_out, positions_out).
|
|
28
|
+
|
|
29
|
+
Returns:
|
|
30
|
+
A Resampling function implementing stop gradient resampling.
|
|
31
|
+
"""
|
|
32
|
+
# Build a descriptive docstring that includes the wrapped function doc
|
|
33
|
+
wrapped_doc = func.__doc__ or ""
|
|
34
|
+
doc = f"""
|
|
35
|
+
Stop gradient resampling decorator.
|
|
36
|
+
|
|
37
|
+
This wrapper will call the provided resampling function, and then apply
|
|
38
|
+
the stop gradient trick of [Scibior and Wood (2021)](https://arxiv.org/abs/2106.10314).
|
|
39
|
+
Resulting estimates of the score function (i.e., the gradient of the
|
|
40
|
+
log-likelihood with respect to model parameters) are unbiased,
|
|
41
|
+
corresponding to the classical Fisher estimate.
|
|
42
|
+
|
|
43
|
+
Wrapped resampler documentation:
|
|
44
|
+
{wrapped_doc}
|
|
45
|
+
"""
|
|
46
|
+
|
|
47
|
+
@wraps(func)
|
|
48
|
+
def _wrapped(
|
|
49
|
+
key: Array, logits: ArrayLike, positions: ArrayTreeLike, n: int
|
|
50
|
+
) -> tuple[Array, Array, ArrayTree]:
|
|
51
|
+
idx_base, logits_base, positions_base = func(
|
|
52
|
+
key, jax.lax.stop_gradient(logits), positions, n
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
logits = jnp.asarray(
|
|
56
|
+
logits_base
|
|
57
|
+
+ apply_resampling_indices(logits, idx_base)
|
|
58
|
+
- jax.lax.stop_gradient(apply_resampling_indices(logits, idx_base))
|
|
59
|
+
)
|
|
60
|
+
return idx_base, logits, positions_base
|
|
61
|
+
|
|
62
|
+
# Attach the composed docstring and return a jitted version
|
|
63
|
+
_wrapped.__doc__ = doc
|
|
64
|
+
return jax.jit(_wrapped, static_argnames=("n",))
|
|
@@ -11,7 +11,14 @@ from cuthbertlib.resampling.protocols import (
|
|
|
11
11
|
conditional_resampling_decorator,
|
|
12
12
|
resampling_decorator,
|
|
13
13
|
)
|
|
14
|
-
from cuthbertlib.
|
|
14
|
+
from cuthbertlib.resampling.utils import apply_resampling_indices
|
|
15
|
+
from cuthbertlib.types import (
|
|
16
|
+
Array,
|
|
17
|
+
ArrayLike,
|
|
18
|
+
ArrayTree,
|
|
19
|
+
ArrayTreeLike,
|
|
20
|
+
ScalarArrayLike,
|
|
21
|
+
)
|
|
15
22
|
|
|
16
23
|
_DESCRIPTION = """
|
|
17
24
|
The Killing resampling is a simple resampling mechanism that checks if
|
|
@@ -28,7 +35,9 @@ number of particles `logits.shape[0]`.
|
|
|
28
35
|
|
|
29
36
|
|
|
30
37
|
@partial(resampling_decorator, name="Killing", desc=_DESCRIPTION)
|
|
31
|
-
def resampling(
|
|
38
|
+
def resampling(
|
|
39
|
+
key: Array, logits: ArrayLike, positions: ArrayTreeLike, n: int
|
|
40
|
+
) -> tuple[Array, Array, ArrayTree]:
|
|
32
41
|
logits = jnp.asarray(logits)
|
|
33
42
|
key_1, key_2 = random.split(key)
|
|
34
43
|
N = logits.shape[0]
|
|
@@ -43,27 +52,30 @@ def resampling(key: Array, logits: ArrayLike, n: int) -> Array:
|
|
|
43
52
|
|
|
44
53
|
survived = log_uniforms <= logits - max_logit
|
|
45
54
|
if_survived = jnp.arange(N) # If the particle survives, it keeps its index
|
|
46
|
-
|
|
47
|
-
key_2, logits, N
|
|
55
|
+
otherwise_idx, _, _ = multinomial.resampling(
|
|
56
|
+
key_2, logits, positions, N
|
|
48
57
|
) # otherwise, it is replaced by another particle
|
|
49
|
-
idx = jnp.where(survived, if_survived,
|
|
50
|
-
|
|
58
|
+
idx = jnp.where(survived, if_survived, otherwise_idx)
|
|
59
|
+
# After resampling, all particles have equal weight
|
|
60
|
+
logits_out = jnp.zeros_like(logits)
|
|
61
|
+
return idx, logits_out, apply_resampling_indices(positions, idx)
|
|
51
62
|
|
|
52
63
|
|
|
53
64
|
@partial(conditional_resampling_decorator, name="Killing", desc=_DESCRIPTION)
|
|
54
65
|
def conditional_resampling(
|
|
55
66
|
key: Array,
|
|
56
67
|
logits: ArrayLike,
|
|
68
|
+
positions: ArrayTreeLike,
|
|
57
69
|
n: int,
|
|
58
70
|
pivot_in: ScalarArrayLike,
|
|
59
71
|
pivot_out: ScalarArrayLike,
|
|
60
|
-
) -> Array:
|
|
72
|
+
) -> tuple[Array, Array, ArrayTree]:
|
|
61
73
|
pivot_in = jnp.asarray(pivot_in)
|
|
62
74
|
pivot_out = jnp.asarray(pivot_out)
|
|
63
75
|
|
|
64
76
|
# Unconditional resampling
|
|
65
77
|
key_resample, key_shuffle = random.split(key)
|
|
66
|
-
|
|
78
|
+
idx_uncond, _, _ = resampling(key_resample, logits, positions, n)
|
|
67
79
|
|
|
68
80
|
# Conditional rolling pivot
|
|
69
81
|
max_logit = jnp.max(logits)
|
|
@@ -76,9 +88,11 @@ def conditional_resampling(
|
|
|
76
88
|
|
|
77
89
|
pivot_weights = jnp.exp(pivot_logits - logsumexp(pivot_logits))
|
|
78
90
|
pivot = random.choice(key_shuffle, n, p=pivot_weights)
|
|
79
|
-
idx = jnp.roll(
|
|
91
|
+
idx = jnp.roll(idx_uncond, pivot_in - pivot)
|
|
80
92
|
idx = idx.at[pivot_in].set(pivot_out)
|
|
81
|
-
|
|
93
|
+
# After resampling, all particles have equal weight
|
|
94
|
+
logits_out = jnp.zeros_like(logits)
|
|
95
|
+
return idx, logits_out, apply_resampling_indices(positions, idx)
|
|
82
96
|
|
|
83
97
|
|
|
84
98
|
def _log1mexp(x: ArrayLike) -> Array:
|
|
@@ -10,8 +10,14 @@ from cuthbertlib.resampling.protocols import (
|
|
|
10
10
|
conditional_resampling_decorator,
|
|
11
11
|
resampling_decorator,
|
|
12
12
|
)
|
|
13
|
-
from cuthbertlib.resampling.utils import inverse_cdf
|
|
14
|
-
from cuthbertlib.types import
|
|
13
|
+
from cuthbertlib.resampling.utils import apply_resampling_indices, inverse_cdf
|
|
14
|
+
from cuthbertlib.types import (
|
|
15
|
+
Array,
|
|
16
|
+
ArrayLike,
|
|
17
|
+
ArrayTree,
|
|
18
|
+
ArrayTreeLike,
|
|
19
|
+
ScalarArrayLike,
|
|
20
|
+
)
|
|
15
21
|
|
|
16
22
|
_DESCRIPTION = """
|
|
17
23
|
This has higher variance than other resampling schemes as it samples from
|
|
@@ -21,7 +27,9 @@ As a rule of thumb, you often don't."""
|
|
|
21
27
|
|
|
22
28
|
|
|
23
29
|
@partial(resampling_decorator, name="Multinomial", desc=_DESCRIPTION)
|
|
24
|
-
def resampling(
|
|
30
|
+
def resampling(
|
|
31
|
+
key: Array, logits: ArrayLike, positions: ArrayTreeLike, n: int
|
|
32
|
+
) -> tuple[Array, Array, ArrayTree]:
|
|
25
33
|
# In practice we don't have to sort the generated uniforms, but searchsorted
|
|
26
34
|
# works faster and is more stable if both inputs are sorted, so we use the
|
|
27
35
|
# _sorted_uniforms from N. Chopin, but still use searchsorted instead of his
|
|
@@ -32,23 +40,26 @@ def resampling(key: Array, logits: ArrayLike, n: int) -> Array:
|
|
|
32
40
|
key_uniforms, key_shuffle = random.split(key)
|
|
33
41
|
sorted_uniforms = _sorted_uniforms(key_uniforms, n)
|
|
34
42
|
idx = inverse_cdf(sorted_uniforms, logits)
|
|
35
|
-
|
|
43
|
+
idx = random.permutation(key_shuffle, idx)
|
|
44
|
+
logits_out = jnp.zeros_like(sorted_uniforms)
|
|
45
|
+
return idx, logits_out, apply_resampling_indices(positions, idx)
|
|
36
46
|
|
|
37
47
|
|
|
38
48
|
@partial(conditional_resampling_decorator, name="Multinomial", desc=_DESCRIPTION)
|
|
39
49
|
def conditional_resampling(
|
|
40
50
|
key: Array,
|
|
41
51
|
logits: ArrayLike,
|
|
52
|
+
positions: ArrayTreeLike,
|
|
42
53
|
n: int,
|
|
43
54
|
pivot_in: ScalarArrayLike,
|
|
44
55
|
pivot_out: ScalarArrayLike,
|
|
45
|
-
) -> Array:
|
|
56
|
+
) -> tuple[Array, Array, ArrayTree]:
|
|
46
57
|
pivot_in = jnp.asarray(pivot_in)
|
|
47
58
|
pivot_out = jnp.asarray(pivot_out)
|
|
48
59
|
|
|
49
|
-
idx = resampling(key, logits, n)
|
|
60
|
+
idx, logits_out, _ = resampling(key, logits, positions, n)
|
|
50
61
|
idx = idx.at[pivot_in].set(pivot_out)
|
|
51
|
-
return idx
|
|
62
|
+
return idx, logits_out, apply_resampling_indices(positions, idx)
|
|
52
63
|
|
|
53
64
|
|
|
54
65
|
@partial(jax.jit, static_argnames=("n",))
|
|
@@ -0,0 +1,98 @@
|
|
|
1
|
+
"""Shared protocols for resampling algorithms."""
|
|
2
|
+
|
|
3
|
+
from typing import Protocol, runtime_checkable
|
|
4
|
+
|
|
5
|
+
import jax
|
|
6
|
+
|
|
7
|
+
from cuthbertlib.types import (
|
|
8
|
+
Array,
|
|
9
|
+
ArrayLike,
|
|
10
|
+
ArrayTree,
|
|
11
|
+
ArrayTreeLike,
|
|
12
|
+
KeyArray,
|
|
13
|
+
ScalarArrayLike,
|
|
14
|
+
)
|
|
15
|
+
|
|
16
|
+
_RESAMPLING_DOC = """
|
|
17
|
+
Args:
|
|
18
|
+
key: JAX PRNG key.
|
|
19
|
+
logits: Logits.
|
|
20
|
+
positions: ArrayTreeLike
|
|
21
|
+
n: Number of indices to sample.
|
|
22
|
+
|
|
23
|
+
Returns:
|
|
24
|
+
ancestors: Array of resampling indices.
|
|
25
|
+
logits: Array of log-weights after resampling.
|
|
26
|
+
positions: ArrayTreeLike of resampled positions.
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
_CONDITIONAL_RESAMPLING_DOC = """
|
|
30
|
+
Args:
|
|
31
|
+
key: JAX PRNG key.
|
|
32
|
+
logits: Log-weights, possibly unnormalized.
|
|
33
|
+
positions: ArrayTreeLike
|
|
34
|
+
n: Number of indices to sample.
|
|
35
|
+
pivot_in: Index of the particle to keep.
|
|
36
|
+
pivot_out: Value of the output at index `pivot_in`.
|
|
37
|
+
|
|
38
|
+
Returns:
|
|
39
|
+
ancestors: Array of size n with indices to use for resampling.
|
|
40
|
+
logits: Array of log-weights after resampling.
|
|
41
|
+
positions: ArrayTreeLike of resampled positions.
|
|
42
|
+
"""
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
@runtime_checkable
|
|
46
|
+
class Resampling(Protocol):
|
|
47
|
+
"""Protocol for resampling operations."""
|
|
48
|
+
|
|
49
|
+
def __call__(
|
|
50
|
+
self, key: KeyArray, logits: ArrayLike, positions: ArrayTreeLike, n: int
|
|
51
|
+
) -> tuple[Array, Array, ArrayTree]:
|
|
52
|
+
f"""Computes resampling indices according to given logits.
|
|
53
|
+
{_RESAMPLING_DOC}
|
|
54
|
+
"""
|
|
55
|
+
...
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
@runtime_checkable
|
|
59
|
+
class ConditionalResampling(Protocol):
|
|
60
|
+
"""Protocol for conditional resampling operations."""
|
|
61
|
+
|
|
62
|
+
def __call__(
|
|
63
|
+
self,
|
|
64
|
+
key: KeyArray,
|
|
65
|
+
logits: ArrayLike,
|
|
66
|
+
positions: ArrayTreeLike,
|
|
67
|
+
n: int,
|
|
68
|
+
pivot_in: ScalarArrayLike,
|
|
69
|
+
pivot_out: ScalarArrayLike,
|
|
70
|
+
) -> tuple[Array, Array, ArrayTree]:
|
|
71
|
+
f"""Conditional resampling.
|
|
72
|
+
{_CONDITIONAL_RESAMPLING_DOC}
|
|
73
|
+
"""
|
|
74
|
+
...
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def resampling_decorator(func: Resampling, name: str, desc: str = "") -> Resampling:
|
|
78
|
+
"""Decorate Resampling function with unified docstring."""
|
|
79
|
+
doc = f"""
|
|
80
|
+
{name} resampling. {desc}
|
|
81
|
+
{_RESAMPLING_DOC}
|
|
82
|
+
"""
|
|
83
|
+
|
|
84
|
+
func.__doc__ = doc
|
|
85
|
+
return jax.jit(func, static_argnames=("n",))
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def conditional_resampling_decorator(
|
|
89
|
+
func: ConditionalResampling, name: str, desc: str = ""
|
|
90
|
+
) -> ConditionalResampling:
|
|
91
|
+
"""Decorate ConditionalResampling function with unified docstring."""
|
|
92
|
+
doc = f"""
|
|
93
|
+
{name} conditional resampling. {desc}
|
|
94
|
+
{_CONDITIONAL_RESAMPLING_DOC}
|
|
95
|
+
"""
|
|
96
|
+
|
|
97
|
+
func.__doc__ = doc
|
|
98
|
+
return jax.jit(func, static_argnames=("n",))
|
|
@@ -11,8 +11,14 @@ from cuthbertlib.resampling.protocols import (
|
|
|
11
11
|
conditional_resampling_decorator,
|
|
12
12
|
resampling_decorator,
|
|
13
13
|
)
|
|
14
|
-
from cuthbertlib.resampling.utils import inverse_cdf
|
|
15
|
-
from cuthbertlib.types import
|
|
14
|
+
from cuthbertlib.resampling.utils import apply_resampling_indices, inverse_cdf
|
|
15
|
+
from cuthbertlib.types import (
|
|
16
|
+
Array,
|
|
17
|
+
ArrayLike,
|
|
18
|
+
ArrayTree,
|
|
19
|
+
ArrayTreeLike,
|
|
20
|
+
ScalarArrayLike,
|
|
21
|
+
)
|
|
16
22
|
|
|
17
23
|
_DESCRIPTION = """
|
|
18
24
|
The Systematic resampling is a variance reduction which places marginally
|
|
@@ -21,19 +27,24 @@ uniform samples into the [0, 1] interval but only requires one uniform random.
|
|
|
21
27
|
|
|
22
28
|
|
|
23
29
|
@partial(resampling_decorator, name="Systematic", desc=_DESCRIPTION)
|
|
24
|
-
def resampling(
|
|
30
|
+
def resampling(
|
|
31
|
+
key: Array, logits: ArrayLike, positions: ArrayTreeLike, n: int
|
|
32
|
+
) -> tuple[Array, Array, ArrayTree]:
|
|
25
33
|
us = (random.uniform(key, ()) + jnp.arange(n)) / n
|
|
26
|
-
|
|
34
|
+
idx = inverse_cdf(us, logits)
|
|
35
|
+
logits_out = jnp.zeros_like(us)
|
|
36
|
+
return idx, logits_out, apply_resampling_indices(positions, idx)
|
|
27
37
|
|
|
28
38
|
|
|
29
39
|
@partial(conditional_resampling_decorator, name="Systematic", desc=_DESCRIPTION)
|
|
30
40
|
def conditional_resampling(
|
|
31
41
|
key: Array,
|
|
32
42
|
logits: ArrayLike,
|
|
43
|
+
positions: ArrayTreeLike,
|
|
33
44
|
n: int,
|
|
34
45
|
pivot_in: ScalarArrayLike,
|
|
35
46
|
pivot_out: ScalarArrayLike,
|
|
36
|
-
) -> Array:
|
|
47
|
+
) -> tuple[Array, Array, ArrayTree]:
|
|
37
48
|
logits = jnp.asarray(logits)
|
|
38
49
|
pivot_in = jnp.asarray(pivot_in)
|
|
39
50
|
pivot_out = jnp.asarray(pivot_out)
|
|
@@ -46,17 +57,17 @@ def conditional_resampling(
|
|
|
46
57
|
logits = jnp.roll(logits, -pivot_out)
|
|
47
58
|
arange = jnp.roll(arange, -pivot_out)
|
|
48
59
|
|
|
49
|
-
idx = conditional_resampling_0_to_0(key, logits, n)
|
|
60
|
+
idx, logits_out = conditional_resampling_0_to_0(key, logits, n)
|
|
50
61
|
idx = arange[idx]
|
|
51
62
|
idx = jnp.roll(idx, pivot_in)
|
|
52
|
-
return idx
|
|
63
|
+
return idx, logits_out, apply_resampling_indices(positions, idx)
|
|
53
64
|
|
|
54
65
|
|
|
55
66
|
def conditional_resampling_0_to_0(
|
|
56
67
|
key: Array,
|
|
57
68
|
logits: ArrayLike,
|
|
58
69
|
n: int,
|
|
59
|
-
) -> Array:
|
|
70
|
+
) -> tuple[Array, Array]:
|
|
60
71
|
logits = jnp.asarray(logits)
|
|
61
72
|
|
|
62
73
|
N = logits.shape[0]
|
|
@@ -81,4 +92,4 @@ def conditional_resampling_0_to_0(
|
|
|
81
92
|
roll_idx = jnp.floor(n_zero * W).astype(int)
|
|
82
93
|
|
|
83
94
|
idx = select(n_zero == 1, idx, jnp.roll(idx, -zero_loc[roll_idx]))
|
|
84
|
-
return jnp.clip(idx, 0, N - 1)
|
|
95
|
+
return jnp.clip(idx, 0, N - 1), jnp.zeros_like(linspace)
|
|
@@ -4,10 +4,11 @@ import jax
|
|
|
4
4
|
import jax.numpy as jnp
|
|
5
5
|
import numba as nb
|
|
6
6
|
import numpy as np
|
|
7
|
-
from jax.lax import platform_dependent
|
|
7
|
+
from jax.lax import platform_dependent, stop_gradient
|
|
8
8
|
from jax.scipy.special import logsumexp
|
|
9
|
+
from jax.tree_util import tree_map
|
|
9
10
|
|
|
10
|
-
from cuthbertlib.types import Array, ArrayLike
|
|
11
|
+
from cuthbertlib.types import Array, ArrayLike, ArrayTree, ArrayTreeLike
|
|
11
12
|
|
|
12
13
|
|
|
13
14
|
@jax.jit
|
|
@@ -31,7 +32,10 @@ def inverse_cdf(sorted_uniforms: ArrayLike, logits: ArrayLike) -> Array:
|
|
|
31
32
|
"""
|
|
32
33
|
weights = jnp.exp(logits - logsumexp(logits))
|
|
33
34
|
return platform_dependent(
|
|
34
|
-
sorted_uniforms,
|
|
35
|
+
sorted_uniforms,
|
|
36
|
+
stop_gradient(weights),
|
|
37
|
+
cpu=_inverse_cdf_cpu,
|
|
38
|
+
default=_inverse_cdf_default,
|
|
35
39
|
)
|
|
36
40
|
|
|
37
41
|
|
|
@@ -80,3 +84,8 @@ def _inverse_cdf_numba(su, ws, idx):
|
|
|
80
84
|
j += 1
|
|
81
85
|
s += ws[j]
|
|
82
86
|
idx[n] = j
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def apply_resampling_indices(positions: ArrayTreeLike, idx: Array) -> ArrayTree:
|
|
90
|
+
"""Apply resampling indices to positions."""
|
|
91
|
+
return tree_map(lambda x: x[idx], positions)
|
|
@@ -28,6 +28,7 @@ def log_weights_single(
|
|
|
28
28
|
x1: The current state.
|
|
29
29
|
log_weight_x0: The log weights of the previous state.
|
|
30
30
|
log_density: The log density function of x1 given x0.
|
|
31
|
+
x0 must be the first argument and x1 the second.
|
|
31
32
|
|
|
32
33
|
Returns:
|
|
33
34
|
The smoothing weight for sample x0 given a single sample x1.
|
|
@@ -97,6 +98,7 @@ def simulate(
|
|
|
97
98
|
x1_all: A collection of current states $x_1$.
|
|
98
99
|
log_weight_x0_all: The log weights of $x_0$.
|
|
99
100
|
log_density: The log density function of $x_1$ given $x_0$.
|
|
101
|
+
$x_0$ must be the first argument and $x_1$ the second.
|
|
100
102
|
x1_ancestor_indices: The ancestor indices of $x_1$. Not used.
|
|
101
103
|
|
|
102
104
|
Returns:
|
|
@@ -33,6 +33,7 @@ def simulate(
|
|
|
33
33
|
x1_all: A collection of current states $x_1$.
|
|
34
34
|
log_weight_x0_all: The log weights of $x_0$.
|
|
35
35
|
log_density: The log density function of $x_1$ given $x_0$.
|
|
36
|
+
$x_0$ must be the first argument and $x_1$ the second.
|
|
36
37
|
x1_ancestor_indices: The ancestor indices of $x_1$.
|
|
37
38
|
n_steps: Number of MCMC steps to perform.
|
|
38
39
|
|
|
@@ -55,8 +56,9 @@ def simulate(
|
|
|
55
56
|
idx, x0_res, idx_log_p = carry
|
|
56
57
|
key_prop, key_acc = keys_t
|
|
57
58
|
|
|
58
|
-
prop_idx = multinomial.resampling(
|
|
59
|
-
|
|
59
|
+
prop_idx, _, x0_prop = multinomial.resampling(
|
|
60
|
+
key_prop, log_weight_x0_all, x0_all, n_samples
|
|
61
|
+
)
|
|
60
62
|
prop_log_p = jax.vmap(log_density)(x0_prop, x1_all)
|
|
61
63
|
|
|
62
64
|
log_alpha = prop_log_p - idx_log_p
|
|
@@ -70,7 +72,7 @@ def simulate(
|
|
|
70
72
|
return (idx, x0_res, idx_log_p), None
|
|
71
73
|
|
|
72
74
|
x0_init = jax.tree.map(lambda z: z[x1_ancestor_indices], x0_all)
|
|
73
|
-
init_log_p = jax.vmap(log_density)(
|
|
75
|
+
init_log_p = jax.vmap(log_density)(x0_init, x1_all)
|
|
74
76
|
init = (x1_ancestor_indices, x0_init, init_log_p)
|
|
75
77
|
(out_index, out_samples, _), _ = jax.lax.scan(body, init, keys)
|
|
76
78
|
return out_samples, out_index
|
|
@@ -36,6 +36,7 @@ class BackwardSampling(Protocol):
|
|
|
36
36
|
x1_all: A collection of current states $x_1$.
|
|
37
37
|
log_weight_x0_all: The log weights of $x_0$.
|
|
38
38
|
log_density: The log density function of $x_1$ given $x_0$.
|
|
39
|
+
$x_0$ must be the first argument and $x_1$ the second.
|
|
39
40
|
x1_ancestor_indices: The ancestor indices of $x_1$.
|
|
40
41
|
|
|
41
42
|
Returns:
|
|
@@ -14,7 +14,9 @@ ScalarArray: TypeAlias = (
|
|
|
14
14
|
ScalarArrayLike: TypeAlias = ArrayLike # Object that will be cast to a ScalarArray
|
|
15
15
|
|
|
16
16
|
LogDensity: TypeAlias = Callable[[ArrayTreeLike], ScalarArray]
|
|
17
|
-
LogConditionalDensity: TypeAlias = Callable[
|
|
17
|
+
LogConditionalDensity: TypeAlias = Callable[
|
|
18
|
+
[ArrayTreeLike, ArrayTreeLike], ScalarArray
|
|
19
|
+
] # p(x_1 | x_0), where x_0 is the first argument and x_1 the second.
|
|
18
20
|
LogConditionalDensityAux: TypeAlias = Callable[
|
|
19
21
|
[ArrayTreeLike, ArrayTreeLike], tuple[ScalarArray, ArrayTree]
|
|
20
22
|
]
|
|
@@ -1,21 +0,0 @@
|
|
|
1
|
-
"""Implements triangularization operator a matrix via QR decomposition."""
|
|
2
|
-
|
|
3
|
-
import jax
|
|
4
|
-
|
|
5
|
-
from cuthbertlib.types import Array
|
|
6
|
-
|
|
7
|
-
|
|
8
|
-
def tria(A: Array) -> Array:
|
|
9
|
-
r"""A triangularization operator using QR decomposition.
|
|
10
|
-
|
|
11
|
-
Args:
|
|
12
|
-
A: The matrix to triangularize.
|
|
13
|
-
|
|
14
|
-
Returns:
|
|
15
|
-
A lower triangular matrix $R$ such that $R R^\top = A A^\top$.
|
|
16
|
-
|
|
17
|
-
Reference:
|
|
18
|
-
[Arasaratnam and Haykin (2008)](https://ieeexplore.ieee.org/document/4524036): Square-Root Quadrature Kalman Filtering
|
|
19
|
-
"""
|
|
20
|
-
_, R = jax.scipy.linalg.qr(A.T, mode="economic")
|
|
21
|
-
return R.T
|
|
@@ -1,26 +0,0 @@
|
|
|
1
|
-
# Resampling
|
|
2
|
-
|
|
3
|
-
This sub-repository provides a unified interface for a variety of resampling
|
|
4
|
-
methods, which convert a set of weighted samples into an unweighted one which
|
|
5
|
-
likely contains duplicates.
|
|
6
|
-
|
|
7
|
-
A typical call to the library would be:
|
|
8
|
-
|
|
9
|
-
```python
|
|
10
|
-
sampling_key, resampling_key = jax.random.split(jax.random.key(0))
|
|
11
|
-
particles = jax.random.normal(sampling_key, (100, 2))
|
|
12
|
-
logits = jax.vmap(lambda x: jnp.where(jnp.all(x > 0), 0, -jnp.inf))(particles)
|
|
13
|
-
|
|
14
|
-
resampled_indices = resampling.multinomial.resampling(resampling_key, logits, 100)
|
|
15
|
-
resampled_particles = particles[resampled_indices]
|
|
16
|
-
```
|
|
17
|
-
|
|
18
|
-
Or for conditional resampling:
|
|
19
|
-
|
|
20
|
-
```python
|
|
21
|
-
# Here we resample but keep particle at index 0 fixed
|
|
22
|
-
conditional_resampled_indices = resampling.multinomial.conditional_resampling(
|
|
23
|
-
resampling_key, logits, 100, pivot_in=0, pivot_out=0
|
|
24
|
-
)
|
|
25
|
-
conditional_resampled_particles = particles[conditional_resampled_indices]
|
|
26
|
-
```
|
|
@@ -1,92 +0,0 @@
|
|
|
1
|
-
"""Shared protocols for resampling algorithms."""
|
|
2
|
-
|
|
3
|
-
from typing import Protocol, runtime_checkable
|
|
4
|
-
|
|
5
|
-
import jax
|
|
6
|
-
|
|
7
|
-
from cuthbertlib.types import Array, ArrayLike, KeyArray, ScalarArrayLike
|
|
8
|
-
|
|
9
|
-
|
|
10
|
-
@runtime_checkable
|
|
11
|
-
class Resampling(Protocol):
|
|
12
|
-
"""Protocol for resampling operations."""
|
|
13
|
-
|
|
14
|
-
def __call__(self, key: KeyArray, logits: ArrayLike, n: int) -> Array:
|
|
15
|
-
"""Computes resampling indices according to given logits.
|
|
16
|
-
|
|
17
|
-
Args:
|
|
18
|
-
key: JAX PRNG key.
|
|
19
|
-
logits: Logits.
|
|
20
|
-
n: Number of indices to sample.
|
|
21
|
-
|
|
22
|
-
Returns:
|
|
23
|
-
Array of resampling indices.
|
|
24
|
-
"""
|
|
25
|
-
...
|
|
26
|
-
|
|
27
|
-
|
|
28
|
-
@runtime_checkable
|
|
29
|
-
class ConditionalResampling(Protocol):
|
|
30
|
-
"""Protocol for conditional resampling operations."""
|
|
31
|
-
|
|
32
|
-
def __call__(
|
|
33
|
-
self,
|
|
34
|
-
key: KeyArray,
|
|
35
|
-
logits: ArrayLike,
|
|
36
|
-
n: int,
|
|
37
|
-
pivot_in: ScalarArrayLike,
|
|
38
|
-
pivot_out: ScalarArrayLike,
|
|
39
|
-
) -> Array:
|
|
40
|
-
"""Conditional resampling.
|
|
41
|
-
|
|
42
|
-
Args:
|
|
43
|
-
key: JAX PRNG key.
|
|
44
|
-
logits: Log-weights, possibly unnormalized.
|
|
45
|
-
n: Number of indices to sample.
|
|
46
|
-
pivot_in: Index of the particle to keep.
|
|
47
|
-
pivot_out: Value of the output at index `pivot_in`.
|
|
48
|
-
|
|
49
|
-
Returns:
|
|
50
|
-
Array of size n with indices to use for resampling.
|
|
51
|
-
"""
|
|
52
|
-
...
|
|
53
|
-
|
|
54
|
-
|
|
55
|
-
def resampling_decorator(func: Resampling, name: str, desc: str = "") -> Resampling:
|
|
56
|
-
"""Decorate Resampling function with unified docstring."""
|
|
57
|
-
doc = f"""
|
|
58
|
-
{name} resampling. {desc}
|
|
59
|
-
|
|
60
|
-
Args:
|
|
61
|
-
key: PRNGKey to use in resampling
|
|
62
|
-
logits: Log-weights, possibly unnormalized.
|
|
63
|
-
n: Number of indices to sample.
|
|
64
|
-
|
|
65
|
-
Returns:
|
|
66
|
-
Array of size n with indices to use for resampling.
|
|
67
|
-
"""
|
|
68
|
-
|
|
69
|
-
func.__doc__ = doc
|
|
70
|
-
return jax.jit(func, static_argnames=("n",))
|
|
71
|
-
|
|
72
|
-
|
|
73
|
-
def conditional_resampling_decorator(
|
|
74
|
-
func: ConditionalResampling, name: str, desc: str = ""
|
|
75
|
-
) -> ConditionalResampling:
|
|
76
|
-
"""Decorate ConditionalResampling function with unified docstring."""
|
|
77
|
-
doc = f"""
|
|
78
|
-
{name} conditional resampling. {desc}
|
|
79
|
-
|
|
80
|
-
Args:
|
|
81
|
-
key: PRNGKey to use in resampling
|
|
82
|
-
logits: Log-weights, possibly unnormalized.
|
|
83
|
-
n: Number of indices to sample
|
|
84
|
-
pivot_in: Index of the particle to keep
|
|
85
|
-
pivot_out: Value of the output at index `pivot_in`
|
|
86
|
-
|
|
87
|
-
Returns:
|
|
88
|
-
Array of size n with indices to use for resampling.
|
|
89
|
-
"""
|
|
90
|
-
|
|
91
|
-
func.__doc__ = doc
|
|
92
|
-
return jax.jit(func, static_argnames=("n",))
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|