Efficient EnKF analysis step (including N << d_obs) - #279
Conversation
The current implementation does
$$
C_{XY}S^{-1}\delta = (\delta^\intercal S^{-1}C_{XY}^{\intercal})^{\intercal}
$$
From right to left, this incurs a cost of $d_xd_y^2 + Nd_yd_x$ (assuming the inversion cost has already been paid). Using the original formulation (the one on the left), the complexity is $Nd_y^2 + Nd_yd_x$.
`ensemble_subspace=True` performs the analysis in the N-dimensional
ensemble subspace via the Woodbury identity:
X_new = X + X C^{-1} Y^T R^{-1} delta, C = I_N + Y^T R^{-1} Y
In whitened coordinates A = chol_R^{-1} Y this is C = I_N + A^T A, whose
factor comes from tria([A^T, I_N]) without forming C. Only an N x N system
is factorized; the gain, C_xy and S are never built. The state-dimension
cost drops from O(N d_x d_y) to O(N^2 d_x), and the O(d_y^3) factorization
disappears.
The log-likelihood uses log det S = log det R + log det C. Its quadratic
form is evaluated as the least-squares residual
min_v || [z; 0] - [A; I_N] v ||^2
which equals z^T z - z^T A C^{-1} A^T z but avoids the cancellation, severe
once the ensemble spread is large relative to R
(`cancellation_free_log_likelihood=False` selects the difference form).
Since R is never combined with a dense C_yy, chol_R may be scalar, diagonal
or dense on this path. Both localization hooks are rejected: tapering
raises the rank the update relies on being low. The path is never selected
automatically.
Dealing with missingness results in a $d_y^3$ cost even with no missing data. added a check to make sure it is required and skipped otherwise
|
|
||
| Returns: | ||
| Tuple with observation function, Cholesky factor of the observation noise covariance, and observation vector. | ||
| Tuple with observation function, lower-triangular Cholesky factor of the observation noise covariance, and observation vector. |
There was a problem hiding this comment.
Why does chol_R have to be lower-triangular rather than generalised Cholesky factor for this?
There was a problem hiding this comment.
Actually, I believe the "generalised Cholesky factor" is lower triangular, just not necessarily the unique output from jnp.linalg.cholesky. See
cuthbert/cuthbertlib/linalg/tria.py
Line 15 in 689dd02
There was a problem hiding this comment.
Yes, "generalized Cholesky factor" is fine and I should've worded it that way.
Claude was worried that people could pass general square root factors (i.e. not necessarily lower triangular), hence the slightly strange wording.
Another thing: the factors cannot be singular (i.e.
There was a problem hiding this comment.
Cool cool yeah I think let's use "generalized Cholesky factor" here as it matches our terminology. Although we do need to state clearly what that means in the docs xD - this is issue #59
Fixed to pass test
SamDuffield
left a comment
There was a problem hiding this comment.
This looks good, great stuff! And as always we can come back and refactor the interface in future if we see good reason
Addresses issue #277.
What was implemented
0. Efficiency gain when
N < d_xThe current implementation computes the Kalman gain matrix
Assuming$S$ is already factored, computing this from left to right incurs a cost of $\mathcal{O}(d_xd_y^2 + Nd_yd_x)$ , whereas from right to left the cost is $\mathcal{O}(Nd_y^2 + Nd_yd_x)$ . I added a simple dimension check to use the more efficient path.
1. Efficient implementation when
N << d_obsFor the math, see #277. I provide a new implementation ("new branch") as an alternative to the previous implementation ("previous branch"); the new branch is more efficient when
N << d_obs. The goal is to preserve the previous branch's behavior exactly.Choices that were made:
ensemble_subspace(defaults toFalse) passed tobuild_filterand toupdate. The default preserves the previous branch, but I believe the new branch could be the default in some cases.ensemble_subspace=Truetogether withmodify_cross_covariance != no_covariance_modifierorconstruct_chol_innovation_covariance != Noneraises an error (both at filter construction and in the filter update).Rthat is either a scalar or a 1D vector (implicitly assuming the Cholesky factor is diagonal, with either that scalar or that vector on the diagonal), as well as the full Cholesky factor. The reason for this choice is that computingN << d_obs#277) is implemented as a least-squares problem, which avoids subtracting two large values. Claude claims this is more stable at a small extra cost. See below for the mathematical details. A test in the test suite ensures that both formulations match.collect_nans_cholstill costsNaNs. In both branches, I added a check that skips this step when noNaNs are present. This function also converted square roots into lower-triangular factors; is this dangerous?2. Tests
tests/cuthbertlib/ensemble_kalman/test_filtering.pytest_update_ensemble_subspace_matches_dense: basic consistency test. The previous branch and the new branch give the same answer (they differ only in computational complexity).test_update_ensemble_subspace_chol_R_forms: scalar, 1D and 2D factors of the sameRreproduce the previous branch.test_quadratic_form_residual_matches_difference:_quadratic_form_residual(the least-squares implementation) matches the naive difference formula.test_update_ensemble_subspace_missing_observations: with partially missing observations, the new branch matches the previous branch, for both a dense and a 1Dchol_R.test_update_ensemble_subspace_rejects_invalid_arguments:updaterejectsensemble_subspace=Truecombined withcross_covariance_modifierorconstruct_chol_innovation_covariance.tests/cuthbert/ensemble_kalman/test_enkf.pytest_build_filter_ensemble_subspace_matches_default: checks that a 10-step filter run withbuild_filter(..., ensemble_subspace=True)matchesbuild_filter(..., ensemble_subspace=False)in ensemble and log normalizing constant, testing end-to-end behavior.test_build_filter_ensemble_subspace_rejects_localization:build_filterraises at construction time whenensemble_subspace=Trueis combined with either localization callback.3. Sharp edges and limitations
Raschol_R(instead of a lower-triangular Cholesky factor) can silently produce wrong results. This is reflected in the documentation and, as far as I can tell,chol_Rwas never intended to be anything other than a lower-triangular factor, but it might be worth double-checking that this is acceptable. In the previous branch this is not an issue: results remain correct, and only the realization of the perturbed observations changes.4. Mathematical details of the least-squares formulation
where$w^\ast = C^{-1}g$ and $g := Y^\top R^{-1}d$ .
Proof:
Taking the gradient and setting it to zero yields the optimum$w^\ast = C^{-1}g$ . We now compute