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.
Files changed (62) hide show
  1. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/PKG-INFO +1 -1
  2. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/__init__.py +1 -0
  3. cuthbertlib-0.0.10/cuthbertlib/enkf/README.md +14 -0
  4. cuthbertlib-0.0.10/cuthbertlib/enkf/__init__.py +1 -0
  5. cuthbertlib-0.0.10/cuthbertlib/enkf/filtering.py +119 -0
  6. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/kalman/filtering.py +37 -1
  7. cuthbertlib-0.0.10/cuthbertlib/linalg/tria.py +97 -0
  8. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linearize/log_density.py +2 -0
  9. cuthbertlib-0.0.10/cuthbertlib/resampling/README.md +53 -0
  10. cuthbertlib-0.0.10/cuthbertlib/resampling/__init__.py +11 -0
  11. cuthbertlib-0.0.10/cuthbertlib/resampling/adaptive.py +72 -0
  12. cuthbertlib-0.0.10/cuthbertlib/resampling/autodiff.py +64 -0
  13. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/resampling/killing.py +24 -10
  14. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/resampling/multinomial.py +18 -7
  15. cuthbertlib-0.0.10/cuthbertlib/resampling/protocols.py +98 -0
  16. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/resampling/systematic.py +20 -9
  17. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/resampling/utils.py +12 -3
  18. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/smoothing/exact_sampling.py +2 -0
  19. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/smoothing/mcmc.py +5 -3
  20. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/smoothing/protocols.py +1 -0
  21. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/types.py +3 -1
  22. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/pyproject.toml +1 -1
  23. cuthbertlib-0.0.8/cuthbertlib/linalg/tria.py +0 -21
  24. cuthbertlib-0.0.8/cuthbertlib/resampling/README.md +0 -26
  25. cuthbertlib-0.0.8/cuthbertlib/resampling/__init__.py +0 -3
  26. cuthbertlib-0.0.8/cuthbertlib/resampling/protocols.py +0 -92
  27. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/.gitignore +0 -0
  28. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/README.md +0 -0
  29. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/README.md +0 -0
  30. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/discrete/__init__.py +0 -0
  31. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/discrete/filtering.py +0 -0
  32. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/discrete/smoothing.py +0 -0
  33. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/kalman/README.md +0 -0
  34. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/kalman/__init__.py +0 -0
  35. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/kalman/generate.py +0 -0
  36. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/kalman/sampling.py +0 -0
  37. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/kalman/smoothing.py +0 -0
  38. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linalg/README.md +0 -0
  39. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linalg/__init__.py +0 -0
  40. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linalg/collect_nans_chol.py +0 -0
  41. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linalg/marginal_sqrt_cov.py +0 -0
  42. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linalg/symmetric_inv_sqrt.py +0 -0
  43. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linearize/README.md +0 -0
  44. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linearize/__init__.py +0 -0
  45. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linearize/moments.py +0 -0
  46. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/linearize/taylor.py +0 -0
  47. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/README.md +0 -0
  48. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/__init__.py +0 -0
  49. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/common.py +0 -0
  50. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/cubature.py +0 -0
  51. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/gauss_hermite.py +0 -0
  52. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/linearize.py +0 -0
  53. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/unscented.py +0 -0
  54. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/quadrature/utils.py +0 -0
  55. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/README.md +0 -0
  56. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/__init__.py +0 -0
  57. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/ess.py +0 -0
  58. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/smoothing/__init__.py +0 -0
  59. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/smc/smoothing/tracing.py +0 -0
  60. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/stats/README.md +0 -0
  61. {cuthbertlib-0.0.8 → cuthbertlib-0.0.10}/cuthbertlib/stats/__init__.py +0 -0
  62. {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.8
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
@@ -5,6 +5,7 @@ del version
5
5
 
6
6
  from cuthbertlib import (
7
7
  discrete,
8
+ enkf,
8
9
  kalman,
9
10
  linalg,
10
11
  linearize,
@@ -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
- ell = ell1 + ell2 - 0.5 * t1 + 0.5 * jnp.linalg.slogdet(D_inv)[1]
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.types import Array, ArrayLike, ScalarArrayLike
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(key: Array, logits: ArrayLike, n: int) -> Array:
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
- otherwise = multinomial.resampling(
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, otherwise)
50
- return idx
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
- idx = resampling(key_resample, logits, n)
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(idx, pivot_in - pivot)
91
+ idx = jnp.roll(idx_uncond, pivot_in - pivot)
80
92
  idx = idx.at[pivot_in].set(pivot_out)
81
- return idx
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 Array, ArrayLike, ScalarArrayLike
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(key: Array, logits: ArrayLike, n: int) -> Array:
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
- return random.permutation(key_shuffle, idx)
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 Array, ArrayLike, ScalarArrayLike
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(key: Array, logits: ArrayLike, n: int) -> Array:
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
- return inverse_cdf(us, logits)
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, weights, cpu=_inverse_cdf_cpu, default=_inverse_cdf_default
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(key_prop, log_weight_x0_all, n_samples)
59
- x0_prop = jax.tree.map(lambda z: z[prop_idx], x0_all)
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)(x1_all, x0_init)
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[[ArrayTreeLike, ArrayTreeLike], ScalarArray]
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
  ]
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
4
4
 
5
5
  [project]
6
6
  name = "cuthbertlib"
7
- version = "0.0.8"
7
+ version = "0.0.10"
8
8
  description = "Atomic building blocks for state-space model inference with JAX"
9
9
  requires-python = ">=3.10"
10
10
  readme = "README.md"
@@ -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,3 +0,0 @@
1
- from cuthbertlib.resampling import killing, multinomial, systematic
2
- from cuthbertlib.resampling.protocols import ConditionalResampling, Resampling
3
- from cuthbertlib.resampling.utils import inverse_cdf
@@ -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