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.
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/PKG-INFO +1 -1
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/ensemble_kalman/README.md +11 -0
- cuthbertlib-0.1.0/cuthbertlib/ensemble_kalman/filtering.py +463 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/pyproject.toml +1 -1
- cuthbertlib-0.0.15/cuthbertlib/ensemble_kalman/filtering.py +0 -168
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/.gitignore +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/README.md +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/README.md +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/__init__.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/discrete/__init__.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/discrete/filtering.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/discrete/smoothing.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/ensemble_kalman/__init__.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/ensemble_kalman/localization.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/ensemble_kalman/smoothing.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/kalman/README.md +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/kalman/__init__.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/kalman/filtering.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/kalman/generate.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/kalman/sampling.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/kalman/smoothing.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linalg/README.md +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linalg/__init__.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linalg/collect_nans_chol.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linalg/marginal_sqrt_cov.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linalg/symmetric_inv_sqrt.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linalg/tria.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linearize/README.md +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linearize/__init__.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linearize/log_density.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linearize/moments.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/linearize/taylor.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/README.md +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/__init__.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/common.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/cubature.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/gauss_hermite.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/linearize.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/unscented.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/quadrature/utils.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/README.md +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/__init__.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/adaptive.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/autodiff.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/killing.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/multinomial.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/no_resampling.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/protocols.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/systematic.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/resampling/utils.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/README.md +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/__init__.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/ess.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/smoothing/__init__.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/smoothing/exact_sampling.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/smoothing/mcmc.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/smoothing/protocols.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/smc/smoothing/tracing.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/stats/README.md +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/stats/__init__.py +0 -0
- {cuthbertlib-0.0.15 → cuthbertlib-0.1.0}/cuthbertlib/stats/multivariate_normal.py +0 -0
- {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
|
|
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)
|
|
@@ -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
|
|
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
|
|
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
|