cuthbertlib 0.0.15__tar.gz → 0.1.0__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.15 → cuthbertlib-0.1.0}/PKG-INFO +1 -1
  2. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/ensemble_kalman/README.md +11 -0
  3. cuthbertlib-0.1.0/cuthbertlib/ensemble_kalman/filtering.py +463 -0
  4. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/pyproject.toml +1 -1
  5. cuthbertlib-0.0.15/cuthbertlib/ensemble_kalman/filtering.py +0 -168
  6. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/.gitignore +0 -0
  7. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/README.md +0 -0
  8. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/README.md +0 -0
  9. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/__init__.py +0 -0
  10. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/discrete/__init__.py +0 -0
  11. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/discrete/filtering.py +0 -0
  12. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/discrete/smoothing.py +0 -0
  13. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/ensemble_kalman/__init__.py +0 -0
  14. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/ensemble_kalman/localization.py +0 -0
  15. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/ensemble_kalman/smoothing.py +0 -0
  16. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/kalman/README.md +0 -0
  17. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/kalman/__init__.py +0 -0
  18. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/kalman/filtering.py +0 -0
  19. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/kalman/generate.py +0 -0
  20. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/kalman/sampling.py +0 -0
  21. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/kalman/smoothing.py +0 -0
  22. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linalg/README.md +0 -0
  23. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linalg/__init__.py +0 -0
  24. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linalg/collect_nans_chol.py +0 -0
  25. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linalg/marginal_sqrt_cov.py +0 -0
  26. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linalg/symmetric_inv_sqrt.py +0 -0
  27. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linalg/tria.py +0 -0
  28. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linearize/README.md +0 -0
  29. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linearize/__init__.py +0 -0
  30. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linearize/log_density.py +0 -0
  31. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linearize/moments.py +0 -0
  32. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linearize/taylor.py +0 -0
  33. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/README.md +0 -0
  34. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/__init__.py +0 -0
  35. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/common.py +0 -0
  36. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/cubature.py +0 -0
  37. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/gauss_hermite.py +0 -0
  38. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/linearize.py +0 -0
  39. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/unscented.py +0 -0
  40. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/utils.py +0 -0
  41. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/README.md +0 -0
  42. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/__init__.py +0 -0
  43. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/adaptive.py +0 -0
  44. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/autodiff.py +0 -0
  45. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/killing.py +0 -0
  46. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/multinomial.py +0 -0
  47. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/no_resampling.py +0 -0
  48. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/protocols.py +0 -0
  49. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/systematic.py +0 -0
  50. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/utils.py +0 -0
  51. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/README.md +0 -0
  52. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/__init__.py +0 -0
  53. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/ess.py +0 -0
  54. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/smoothing/__init__.py +0 -0
  55. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/smoothing/exact_sampling.py +0 -0
  56. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/smoothing/mcmc.py +0 -0
  57. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/smoothing/protocols.py +0 -0
  58. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/smoothing/tracing.py +0 -0
  59. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/stats/README.md +0 -0
  60. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/stats/__init__.py +0 -0
  61. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/stats/multivariate_normal.py +0 -0
  62. {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/types.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: cuthbertlib
3
- Version: 0.0.15
3
+ Version: 0.1.0
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
@@ -24,6 +24,17 @@ The core functions are:
24
24
  Together, `predict` and `filter_update` can be used to perform an online EnKF filtering step.
25
25
 
26
26
  The EnKF uses an ensemble of particles with a Kalman-style measurement update based on empirical covariances. Unlike the EKF, it does not require Jacobians, while naturally handling nonlinear dynamics.
27
+
28
+ ### Large observation dimensions
29
+
30
+ By default, `filter_update` forms the empirical cross-covariance $C_{xy}$ and a generalized Cholesky factor of the innovation covariance $S = C_{yy} + R$, costing $\mathcal{O}(y_{\rm dim}^3 + N y_{\rm dim} x_{\rm dim})$ and storing arrays of size $y_{\rm dim}^2$ and $x_{\rm dim}y_{\rm dim}$.
31
+
32
+ Passing `ensemble_subspace=True` instead carries out the analysis in the $N$-dimensional subspace spanned by the ensemble, using the Woodbury identity. The update becomes $X C^{-1} Y^\intercal R^{-1}\delta$ with $C = I_N + Y^\intercal R^{-1} Y$, so the only factorization is $N \times N$ and neither $C_{xy}$, $S$, nor the Kalman gain is ever formed. The cost is $\mathcal{O}(N^2 x_{\rm dim} + N^2 y_{\rm dim} + N^3)$ plus the cost of applying $R^{-1}$. This is algebraically exact and is preferable whenever $N \ll y_{\rm dim}$; for $y_{\rm dim} \lesssim N$ the default path is cheaper.
33
+
34
+ When $R^{-1}$ is applied by a dense Cholesky factor, this incurs a cost of $\mathcal{O}(Nd_y^2)$. One can reduce this to $\mathcal{O}(Nd_y)$ by passing a structured `chol_R`: a scalar for $\sigma^2 I$, or a 1D array of length $y_{\rm dim}$ for a diagonal factor. Note that the default path always requires a 2D `chol_R`.
35
+
36
+ Both localization hooks below are rejected with `ensemble_subspace=True`, as the Woodbury identity is inapplicable with tapering.
37
+
27
38
  <!-- --8<-- [end:filtering] -->
28
39
 
29
40
  ## Ensemble Rauch-Tung-Striebel smoothing
@@ -0,0 +1,463 @@
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, solve_triangular
13
+
14
+ from cuthbertlib.linalg import collect_nans_chol, tria
15
+ from cuthbertlib.stats import multivariate_normal
16
+ from cuthbertlib.types import Array, KeyArray, ScalarArray
17
+
18
+ ObservationFn = Callable[[Array], Array]
19
+ DynamicsFn = Callable[[Array, KeyArray], Array]
20
+ CrossCovarianceModifier = Callable[[Array], Array]
21
+ ConstructCholInnovationCovariance = Callable[[Array, Array], Array]
22
+
23
+
24
+ def no_covariance_modifier(covariance: Array) -> Array:
25
+ """Return an empirical covariance unchanged.
26
+
27
+ The identity covariance modifier, used as the default when no modification
28
+ (e.g. localization) is requested.
29
+
30
+ Args:
31
+ covariance: Empirical covariance matrix.
32
+
33
+ Returns:
34
+ The covariance matrix, unchanged.
35
+ """
36
+ return covariance
37
+
38
+
39
+ def _whiten(chol_R: Array, V: Array) -> Array:
40
+ """Apply the inverse observation-noise factor, ``chol_R^{-1} @ V``.
41
+
42
+ Args:
43
+ chol_R: Generalized Cholesky factor of the observation noise covariance.
44
+ A 2D array is applied by triangular solve, at O(y_dim ** 2) per column
45
+ of ``V``. A scalar or a 1D array is a multiple of the identity or a
46
+ diagonal factor respectively, and is applied by division at O(y_dim).
47
+ V: Array to whiten, shape (y_dim,) or (y_dim, m).
48
+
49
+ Returns:
50
+ Array with the shape of ``V``.
51
+
52
+ Raises:
53
+ ValueError: If ``chol_R`` has more than two dimensions.
54
+ """
55
+ if jnp.ndim(chol_R) == 2:
56
+ return solve_triangular(chol_R, V, lower=True)
57
+
58
+ if jnp.ndim(chol_R) <= 1:
59
+ # Broadcast the diagonal down the rows of V, which is (y_dim,) or (y_dim, m).
60
+ scale = chol_R if jnp.ndim(V) == 1 else jnp.reshape(chol_R, (-1, 1))
61
+ return V / scale
62
+
63
+ raise ValueError(
64
+ "chol_R must be a scalar, a 1D diagonal factor, or a 2D factor, but has "
65
+ f"{jnp.ndim(chol_R)} dimensions."
66
+ )
67
+
68
+
69
+ def _apply_chol(chol_R: Array, V: Array) -> Array:
70
+ """Apply the observation-noise factor, ``chol_R @ V``.
71
+
72
+ Args:
73
+ chol_R: Generalized Cholesky factor, as in [_whiten][cuthbertlib.ensemble_kalman.filtering._whiten].
74
+ V: Array to scale, shape (y_dim,) or (y_dim, m).
75
+
76
+ Returns:
77
+ Array with the shape of ``V``.
78
+
79
+ Raises:
80
+ ValueError: If ``chol_R`` has more than two dimensions.
81
+ """
82
+ if jnp.ndim(chol_R) == 2:
83
+ return chol_R @ V
84
+
85
+ if jnp.ndim(chol_R) <= 1:
86
+ # Broadcast the diagonal down the rows of V, which is (y_dim,) or (y_dim, m).
87
+ scale = chol_R if jnp.ndim(V) == 1 else jnp.reshape(chol_R, (-1, 1))
88
+ return scale * V
89
+
90
+ raise ValueError(
91
+ "chol_R must be a scalar, a 1D diagonal factor, or a 2D factor, but has "
92
+ f"{jnp.ndim(chol_R)} dimensions."
93
+ )
94
+
95
+
96
+ def _log_det_from_chol(chol: Array) -> ScalarArray:
97
+ """Log-determinant of ``chol @ chol.T`` from a 1D diagonal or 2D factor.
98
+
99
+ Uses the absolute diagonal, since a generalized Cholesky factor produced by
100
+ [tria][cuthbertlib.linalg.tria] (a QR) may carry negative diagonal entries.
101
+
102
+ A scalar factor is rejected rather than accepted: it carries no dimension, so
103
+ the determinant of the covariance it stands for is undefined without ``y_dim``.
104
+ Callers reach this only after
105
+ [collect_nans_chol][cuthbertlib.linalg.collect_nans_chol], which expands a
106
+ scalar factor to a 1D factor of length ``y_dim``.
107
+
108
+ Args:
109
+ chol: Generalized Cholesky factor, 1D (diagonal) or 2D.
110
+
111
+ Returns:
112
+ Scalar log-determinant.
113
+
114
+ Raises:
115
+ ValueError: If ``chol`` is not 1D or 2D.
116
+ """
117
+ if jnp.ndim(chol) == 2:
118
+ diagonal = jnp.diag(chol)
119
+ elif jnp.ndim(chol) == 1:
120
+ diagonal = chol
121
+ else:
122
+ raise ValueError(
123
+ "chol must be a 1D diagonal factor or a 2D factor, but has "
124
+ f"{jnp.ndim(chol)} dimensions."
125
+ )
126
+
127
+ return 2 * jnp.sum(jnp.log(jnp.abs(diagonal)))
128
+
129
+
130
+ def _quadratic_form_residual(A: Array, z: Array) -> ScalarArray:
131
+ r"""Evaluates a quadratic form as a least-squares residual.
132
+
133
+ $$z^\top z - z^\top AC^{-1}A^\top z$$
134
+
135
+ with $C = A^\top A + I_N$. Implements
136
+
137
+ $$\min_{v}\ \left\| \begin{bmatrix} z \\ 0 \end{bmatrix}
138
+ - \begin{bmatrix} A \\ I_N \end{bmatrix} v \right\|^2$$
139
+
140
+ which equals the form above, since the objective is
141
+ $z^\top z - 2v^\top g + v^\top Cv$ with $g = A^\top z$, minimized at $v = C^{-1}g$.
142
+ It is evaluated as $\|b - QQ^\top b\|^2$ for $b = [z;\,0]$ and a thin QR
143
+ factorization $[A;\,I_N] = QW$, the residual of the orthogonal projection of $b$
144
+ onto the column space.
145
+
146
+ Unlike evaluating the difference directly, nothing of comparable size is subtracted
147
+ and the minimizer is never formed, so relative accuracy does not degrade. The
148
+ identity block also makes $[A;\,I_N]$ full column rank whatever the rank of $A$.
149
+
150
+ Args:
151
+ A: matrix, shape (y_dim, N).
152
+ z: vector, shape (y_dim,).
153
+
154
+ Returns:
155
+ Scalar quadratic form.
156
+ """
157
+ n_particles = A.shape[1]
158
+ dtype = A.dtype
159
+
160
+ stacked = jnp.concatenate([A, jnp.eye(n_particles, dtype=dtype)], axis=0)
161
+ target = jnp.concatenate([z, jnp.zeros(n_particles, dtype=dtype)])
162
+
163
+ basis, _ = jnp.linalg.qr(stacked)
164
+ residual = target - basis @ (basis.T @ target)
165
+ return residual @ residual
166
+
167
+
168
+ def predict(
169
+ key: KeyArray,
170
+ ensemble: Array,
171
+ dynamics_fn: DynamicsFn,
172
+ inflation: float = 0.0,
173
+ ) -> Array:
174
+ """Propagate ensemble members through an arbitrary simulator p(x_{t+1} | x_t).
175
+
176
+ Args:
177
+ key: JAX PRNG key.
178
+ ensemble: Ensemble of state vectors, shape (N, x_dim).
179
+ dynamics_fn: Dynamics function mapping (state, key) -> state.
180
+ inflation: Multiplicative inflation factor applied to ensemble deviations.
181
+
182
+ Returns:
183
+ Predicted ensemble, shape (N, x_dim).
184
+ """
185
+ N, x_dim = ensemble.shape
186
+
187
+ # Propagate each member through the dynamics
188
+ keys = random.split(key, N)
189
+ propagated = jax.vmap(dynamics_fn, (0, 0))(ensemble, keys)
190
+
191
+ # Apply multiplicative inflation
192
+ mean = jnp.mean(propagated, axis=0)
193
+ propagated = mean + (1 + inflation) * (propagated - mean)
194
+
195
+ return propagated
196
+
197
+
198
+ def update(
199
+ key: KeyArray,
200
+ predicted_ensemble: Array,
201
+ observation_fn: ObservationFn,
202
+ chol_R: Array,
203
+ y: Array,
204
+ perturbed_obs: bool = True,
205
+ cross_covariance_modifier: CrossCovarianceModifier = no_covariance_modifier,
206
+ construct_chol_innovation_covariance: ConstructCholInnovationCovariance
207
+ | None = None,
208
+ ensemble_subspace: bool = False,
209
+ ) -> tuple[Array, ScalarArray]:
210
+ """Update ensemble members with an observation using the EnKF update.
211
+
212
+ NaNs in ``y`` are treated as missing dimensions and are excluded from the
213
+ update. When ``y`` is entirely NaN, the update is a no-op: the predicted
214
+ ensemble is returned unchanged with zero log-likelihood contribution.
215
+
216
+ Args:
217
+ key: JAX PRNG key.
218
+ predicted_ensemble: Predicted ensemble, shape (N, x_dim).
219
+ observation_fn: Observation function mapping state -> obs.
220
+ chol_R: Generalized Cholesky factor of the observation noise covariance,
221
+ shape (y_dim, y_dim). Square roots that are not generalized Cholesky
222
+ factors, such as a symmetric R ** 0.5, are not supported.
223
+ When ``ensemble_subspace`` is True this may instead be a scalar (a multiple
224
+ of the identity) or a 1D array of shape (y_dim,) (a diagonal factor).
225
+ Prefer those forms when the structure allows: a 2D factor must be applied
226
+ by triangular solve, at O(N * y_dim ** 2) instead of O(N * y_dim), and
227
+ stored densely in y_dim ** 2 entries. On either path, a 2D factor is also
228
+ refactored in O(y_dim ** 3) at any step where ``y`` has missing values,
229
+ whereas scalar and 1D factors handle missingness in O(y_dim). Steps with
230
+ nothing missing skip the refactor.
231
+ y: Observation vector, shape (y_dim,). NaNs indicate missing dimensions.
232
+ perturbed_obs: If True, use perturbed observations (stochastic EnKF).
233
+ If False, use deterministic update.
234
+ cross_covariance_modifier: Function that modifies the empirical
235
+ state-observation cross-covariance, shape (x_dim, y_dim), and returns
236
+ an array with the same shape. Defaults to the identity.
237
+ construct_chol_innovation_covariance: Optional function that
238
+ receives normalized observation deviations with shape (y_dim, N) and
239
+ ``chol_R`` with shape (y_dim, y_dim). It must return a generalized
240
+ Cholesky factor of the complete innovation covariance with shape
241
+ (y_dim, y_dim). The deviations have already been divided by
242
+ ``sqrt(N - 1)``. Both inputs use the original observation order.
243
+ ``None`` uses the standard, unlocalized square-root construction.
244
+ ensemble_subspace: If True, perform the analysis in the N-dimensional
245
+ ensemble subspace. This is algebraically exact and costs
246
+ O(N ** 2 * x_dim) in the state dimension rather than
247
+ O(N * x_dim * y_dim), so it is preferable when ``N << y_dim``. It is
248
+ incompatible with both localization arguments above, which it rejects.
249
+ Defaults to False; the choice is never made automatically.
250
+
251
+ Returns:
252
+ Tuple of (updated_ensemble, log_likelihood).
253
+
254
+ Raises:
255
+ ValueError: If ``ensemble_subspace`` is combined with either localization
256
+ argument, or if a non-2D ``chol_R`` is given without it.
257
+ """
258
+ if ensemble_subspace:
259
+ if cross_covariance_modifier is not no_covariance_modifier:
260
+ raise ValueError(
261
+ "ensemble_subspace=True is incompatible with cross_covariance_modifier: "
262
+ "the ensemble-subspace update never forms the state-observation "
263
+ "cross-covariance, so there is nothing to modify."
264
+ )
265
+ if construct_chol_innovation_covariance is not None:
266
+ raise ValueError(
267
+ "ensemble_subspace=True is incompatible with "
268
+ "construct_chol_innovation_covariance: tapering the innovation "
269
+ "covariance destroys the rank-N structure the update relies on."
270
+ )
271
+ return _update_ensemble_subspace(
272
+ key,
273
+ predicted_ensemble,
274
+ observation_fn,
275
+ chol_R,
276
+ y,
277
+ perturbed_obs,
278
+ )
279
+
280
+ if jnp.ndim(chol_R) != 2:
281
+ raise ValueError(
282
+ "chol_R must be 2D, of shape (y_dim, y_dim). Scalar and diagonal factors "
283
+ "are only supported with ensemble_subspace=True, because this path "
284
+ "factorizes the dense y_dim x y_dim innovation covariance."
285
+ )
286
+
287
+ N, x_dim = predicted_ensemble.shape
288
+
289
+ # Map ensemble to observation space
290
+ y_pred = jax.vmap(observation_fn, (0,))(predicted_ensemble)
291
+ x_mean = jnp.mean(predicted_ensemble, axis=0)
292
+ x_dev = predicted_ensemble - x_mean
293
+
294
+ missing = jnp.isnan(y)
295
+
296
+ # Modify or construct covariances before reordering due to NaNs.
297
+ argsort = jnp.argsort(missing, stable=True)
298
+ original_y_dev = y_pred - jnp.mean(y_pred, axis=0)
299
+ normalized_original_y_dev = original_y_dev.T / jnp.sqrt(N - 1)
300
+
301
+ C_xy = x_dev.T @ original_y_dev / (N - 1)
302
+ C_xy = cross_covariance_modifier(C_xy)
303
+
304
+ if construct_chol_innovation_covariance is not None:
305
+ original_chol_S = construct_chol_innovation_covariance(
306
+ normalized_original_y_dev, chol_R
307
+ )
308
+
309
+ # Handle partially-missing observations by reordering and zeroing missing dims.
310
+ # Use y_pred.T because y_pred is (N, y_dim) and we want to reorder along axis 0.
311
+ # Refactoring chol_R is O(y_dim ** 3); skip it when nothing is missing, in which
312
+ # case the reordering is the identity and the inputs are returned unchanged.
313
+ flag, chol_R, y, y_pred = jax.lax.cond(
314
+ jnp.any(missing),
315
+ lambda args: collect_nans_chol(missing, *args[1:]),
316
+ lambda args: args,
317
+ (missing, chol_R, y, y_pred.T),
318
+ )
319
+ y_pred = y_pred.T
320
+ y_dim = y.shape[0]
321
+
322
+ y_mean = jnp.mean(y_pred, axis=0)
323
+ y_dev = y_pred - y_mean
324
+ C_xy = C_xy[:, argsort]
325
+ C_xy = jnp.where(flag[None, :], 0.0, C_xy)
326
+
327
+ if construct_chol_innovation_covariance is None:
328
+ chol_S = tria(jnp.concatenate([y_dev.T / jnp.sqrt(N - 1), chol_R], axis=1))
329
+ else:
330
+ # The constructor sees the original indexing. Only collect and refactor its
331
+ # result when dimensions are missing; otherwise preserve its returned factor.
332
+ chol_S = jax.lax.cond(
333
+ jnp.any(missing),
334
+ lambda chol: collect_nans_chol(missing, chol)[1],
335
+ lambda chol: chol,
336
+ original_chol_S,
337
+ )
338
+
339
+ # Innovation per member
340
+ if perturbed_obs:
341
+ y_n = y[None, :] + (chol_R @ random.normal(key, (y_dim, N))).T
342
+ else:
343
+ y_n = jnp.broadcast_to(y[None, :], (N, y_dim))
344
+
345
+ innovations = y_n - y_pred
346
+
347
+ # Doing K = C_xy @ S^{-1}\delta right to left has cost O(Nd_y^2 + Nd_yd_x), left to right O(d_xd_y^2 + Nd_yd_x).
348
+ # If N < d_x, then right to left is cheaper; otherwise left to right is cheaper.
349
+ if N < x_dim:
350
+ increment = cho_solve((chol_S, True), innovations.T).T @ C_xy.T
351
+ else:
352
+ increment = innovations @ cho_solve((chol_S, True), C_xy.T)
353
+
354
+ # Update ensemble
355
+ updated = predicted_ensemble + increment
356
+
357
+ # Log-likelihood
358
+ ll = multivariate_normal.logpdf(y, y_mean, chol_S, nan_support=False)
359
+
360
+ return updated, jnp.asarray(ll)
361
+
362
+
363
+ def _update_ensemble_subspace(
364
+ key: KeyArray,
365
+ predicted_ensemble: Array,
366
+ observation_fn: ObservationFn,
367
+ chol_R: Array,
368
+ y: Array,
369
+ perturbed_obs: bool,
370
+ ) -> tuple[Array, ScalarArray]:
371
+ r"""EnKF update carried out in the N-dimensional ensemble subspace.
372
+
373
+ Algebraically identical to the dense update, more efficient when $N << d_y$. Writing
374
+ $X$ and $Y$ for the state and observation deviations scaled by
375
+ $1/\sqrt{N - 1}$, $\delta$ for the per-member innovations and
376
+ $C = I_N + Y^\top R^{-1} Y$, the update is
377
+
378
+ $$X_\text{new} = X_\text{old} + X\,C^{-1}Y^\top R^{-1}\delta$$
379
+
380
+ Working in whitened coordinates $A = \mathrm{chol}_R^{-1}Y$ makes
381
+ $C = I_N + A^\top A$, whose factor is obtained by
382
+ [tria][cuthbertlib.linalg.tria] without forming $C$. The log-likelihood follows
383
+ from $\log\det S = \log\det R + \log\det C$, with the quadratic form evaluated by
384
+ [_quadratic_form_residual][cuthbertlib.ensemble_kalman.filtering._quadratic_form_residual].
385
+
386
+ Args:
387
+ key: JAX PRNG key.
388
+ predicted_ensemble: Predicted ensemble, shape (N, x_dim).
389
+ observation_fn: Observation function mapping state -> obs.
390
+ chol_R: Generalized Cholesky factor of the observation noise covariance,
391
+ as a scalar, a 1D array of shape (y_dim,), or a 2D array.
392
+ y: Observation vector, shape (y_dim,). NaNs indicate missing dimensions.
393
+ perturbed_obs: If True, use perturbed observations (stochastic EnKF).
394
+
395
+ Returns:
396
+ Tuple of (updated_ensemble, log_likelihood).
397
+ """
398
+ N = predicted_ensemble.shape[0]
399
+
400
+ # Map ensemble to observation space
401
+ y_pred = jax.vmap(observation_fn, (0,))(predicted_ensemble)
402
+ x_dev = predicted_ensemble - jnp.mean(predicted_ensemble, axis=0)
403
+
404
+ # Handle partially-missing observations by reordering and zeroing missing dims.
405
+ # Use y_pred.T because y_pred is (N, y_dim) and we want to reorder along axis 0.
406
+ missing = jnp.isnan(y)
407
+ if jnp.ndim(chol_R) == 2:
408
+ # Refactoring a dense factor is O(y_dim ** 3); skip it when nothing is missing.
409
+ # Scalar and 1D factors are handled in O(y_dim), and a scalar is promoted to 1D,
410
+ # so an identity branch would not match shapes there.
411
+ chol_R, y, y_pred = jax.lax.cond(
412
+ jnp.any(missing),
413
+ lambda args: collect_nans_chol(missing, *args)[1:],
414
+ lambda args: args,
415
+ (chol_R, y, y_pred.T),
416
+ )
417
+ else:
418
+ _, chol_R, y, y_pred = collect_nans_chol(missing, chol_R, y, y_pred.T)
419
+ y_pred = y_pred.T
420
+ y_dim = y.shape[0]
421
+
422
+ y_mean = jnp.mean(y_pred, axis=0)
423
+ y_dev = y_pred - y_mean
424
+
425
+ scale = jnp.sqrt(jnp.asarray(N - 1, dtype=predicted_ensemble.dtype))
426
+
427
+ # Whitened observation deviations, chol_R^{-1} @ Y, shape (y_dim, N)
428
+ whitened_anomalies = _whiten(chol_R, y_dev.T / scale)
429
+
430
+ # chol_C @ chol_C.T = I_N + whitened_anomalies.T @ whitened_anomalies, without forming it
431
+ chol_C = tria(
432
+ jnp.concatenate(
433
+ [whitened_anomalies.T, jnp.eye(N, dtype=whitened_anomalies.dtype)],
434
+ axis=1,
435
+ )
436
+ )
437
+
438
+ # Innovation per member
439
+ if perturbed_obs:
440
+ noise = random.normal(key, (y_dim, N))
441
+ y_n = y[None, :] + _apply_chol(chol_R, noise).T
442
+ else:
443
+ y_n = jnp.broadcast_to(y[None, :], (N, y_dim))
444
+
445
+ whitened_member_innovations = _whiten(chol_R, (y_n - y_pred).T)
446
+
447
+ # Ensemble-subspace coefficients C^{-1} Y.T R^{-1} delta, shape (N, N)
448
+ coefficients = cho_solve(
449
+ (chol_C, True), whitened_anomalies.T @ whitened_member_innovations
450
+ )
451
+ updated = predicted_ensemble + coefficients.T @ (x_dev / scale)
452
+
453
+ # The log-likelihood uses the unperturbed innovation about the ensemble mean, which
454
+ # is a different object from the per-member innovations driving the update above.
455
+ whitened_mean_innovation = _whiten(chol_R, y - y_mean)
456
+ quadratic_form = _quadratic_form_residual(
457
+ whitened_anomalies, whitened_mean_innovation
458
+ )
459
+
460
+ log_det = _log_det_from_chol(chol_R) + _log_det_from_chol(chol_C)
461
+ ll = -0.5 * (y_dim * jnp.log(2 * jnp.pi) + log_det + quadratic_form)
462
+
463
+ return updated, jnp.asarray(ll)
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
4
4
 
5
5
  [project]
6
6
  name = "cuthbertlib"
7
- version = "0.0.15"
7
+ version = "0.1.0"
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,168 +0,0 @@
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, KeyArray, ScalarArray
17
-
18
- ObservationFn = Callable[[Array], Array]
19
- DynamicsFn = Callable[[Array, KeyArray], Array]
20
- CrossCovarianceModifier = Callable[[Array], Array]
21
- ConstructCholInnovationCovariance = Callable[[Array, Array], Array]
22
-
23
-
24
- def no_covariance_modifier(covariance: Array) -> Array:
25
- """Return an empirical covariance unchanged.
26
-
27
- The identity covariance modifier, used as the default when no modification
28
- (e.g. localization) is requested.
29
-
30
- Args:
31
- covariance: Empirical covariance matrix.
32
-
33
- Returns:
34
- The covariance matrix, unchanged.
35
- """
36
- return covariance
37
-
38
-
39
- def predict(
40
- key: KeyArray,
41
- ensemble: Array,
42
- dynamics_fn: DynamicsFn,
43
- inflation: float = 0.0,
44
- ) -> Array:
45
- """Propagate ensemble members through an arbitrary simulator p(x_{t+1} | x_t).
46
-
47
- Args:
48
- key: JAX PRNG key.
49
- ensemble: Ensemble of state vectors, shape (N, x_dim).
50
- dynamics_fn: Dynamics function mapping (state, key) -> state.
51
- inflation: Multiplicative inflation factor applied to ensemble deviations.
52
-
53
- Returns:
54
- Predicted ensemble, shape (N, x_dim).
55
- """
56
- N, x_dim = ensemble.shape
57
-
58
- # Propagate each member through the dynamics
59
- keys = random.split(key, N)
60
- propagated = jax.vmap(dynamics_fn, (0, 0))(ensemble, keys)
61
-
62
- # Apply multiplicative inflation
63
- mean = jnp.mean(propagated, axis=0)
64
- propagated = mean + (1 + inflation) * (propagated - mean)
65
-
66
- return propagated
67
-
68
-
69
- def update(
70
- key: KeyArray,
71
- predicted_ensemble: Array,
72
- observation_fn: ObservationFn,
73
- chol_R: Array,
74
- y: Array,
75
- perturbed_obs: bool = True,
76
- cross_covariance_modifier: CrossCovarianceModifier = no_covariance_modifier,
77
- construct_chol_innovation_covariance: ConstructCholInnovationCovariance
78
- | None = None,
79
- ) -> tuple[Array, ScalarArray]:
80
- """Update ensemble members with an observation using the EnKF update.
81
-
82
- NaNs in ``y`` are treated as missing dimensions and are excluded from the
83
- update. When ``y`` is entirely NaN, the update is a no-op: the predicted
84
- ensemble is returned unchanged with zero log-likelihood contribution.
85
-
86
- Args:
87
- key: JAX PRNG key.
88
- predicted_ensemble: Predicted ensemble, shape (N, x_dim).
89
- observation_fn: Observation function mapping state -> obs.
90
- chol_R: Cholesky factor of the observation noise covariance, shape (y_dim, y_dim).
91
- y: Observation vector, shape (y_dim,). NaNs indicate missing dimensions.
92
- perturbed_obs: If True, use perturbed observations (stochastic EnKF).
93
- If False, use deterministic update.
94
- cross_covariance_modifier: Function that modifies the empirical
95
- state-observation cross-covariance, shape (x_dim, y_dim), and returns
96
- an array with the same shape. Defaults to the identity.
97
- construct_chol_innovation_covariance: Optional function that
98
- receives normalized observation deviations with shape (y_dim, N) and
99
- ``chol_R`` with shape (y_dim, y_dim). It must return a generalized
100
- Cholesky factor of the complete innovation covariance with shape
101
- (y_dim, y_dim). The deviations have already been divided by
102
- ``sqrt(N - 1)``. Both inputs use the original observation order.
103
- ``None`` uses the standard, unlocalized square-root construction.
104
-
105
- Returns:
106
- Tuple of (updated_ensemble, log_likelihood).
107
- """
108
- N, x_dim = predicted_ensemble.shape
109
-
110
- # Map ensemble to observation space
111
- y_pred = jax.vmap(observation_fn, (0,))(predicted_ensemble)
112
- x_mean = jnp.mean(predicted_ensemble, axis=0)
113
- x_dev = predicted_ensemble - x_mean
114
-
115
- missing = jnp.isnan(y)
116
-
117
- # Modify or construct covariances before reordering due to NaNs.
118
- argsort = jnp.argsort(missing, stable=True)
119
- original_y_dev = y_pred - jnp.mean(y_pred, axis=0)
120
- normalized_original_y_dev = original_y_dev.T / jnp.sqrt(N - 1)
121
-
122
- C_xy = x_dev.T @ original_y_dev / (N - 1)
123
- C_xy = cross_covariance_modifier(C_xy)
124
-
125
- if construct_chol_innovation_covariance is not None:
126
- original_chol_S = construct_chol_innovation_covariance(
127
- normalized_original_y_dev, chol_R
128
- )
129
-
130
- # Handle partially-missing observations by reordering and zeroing missing dims.
131
- # Use y_pred.T because y_pred is (N, y_dim) and we want to reorder along axis 0.
132
- flag, chol_R, y, y_pred = collect_nans_chol(missing, chol_R, y, y_pred.T)
133
- y_pred = y_pred.T
134
- y_dim = y.shape[0]
135
-
136
- y_mean = jnp.mean(y_pred, axis=0)
137
- y_dev = y_pred - y_mean
138
- C_xy = C_xy[:, argsort]
139
- C_xy = jnp.where(flag[None, :], 0.0, C_xy)
140
-
141
- if construct_chol_innovation_covariance is None:
142
- chol_S = tria(jnp.concatenate([y_dev.T / jnp.sqrt(N - 1), chol_R], axis=1))
143
- else:
144
- # The constructor sees the original indexing. Only collect and refactor its
145
- # result when dimensions are missing; otherwise preserve its returned factor.
146
- chol_S = jax.lax.cond(
147
- jnp.any(missing),
148
- lambda chol: collect_nans_chol(missing, chol)[1],
149
- lambda chol: chol,
150
- original_chol_S,
151
- )
152
-
153
- # Kalman gain: K = C_xy @ S^{-1} = C_xy @ cho_solve(chol_S, I)
154
- K = cho_solve((chol_S, True), C_xy.T).T
155
-
156
- # Innovation per member
157
- if perturbed_obs:
158
- y_n = y[None, :] + (chol_R @ random.normal(key, (y_dim, N))).T
159
- else:
160
- y_n = jnp.broadcast_to(y[None, :], (N, y_dim))
161
-
162
- # Update ensemble
163
- updated = predicted_ensemble + (y_n - y_pred) @ K.T
164
-
165
- # Log-likelihood
166
- ll = multivariate_normal.logpdf(y, y_mean, chol_S, nan_support=False)
167
-
168
- return updated, jnp.asarray(ll)
File without changes
File without changes