From c6d6982211aba0a2c4d38ea904e0817b8a99a3bb Mon Sep 17 00:00:00 2001 From: Kristian Fossum Date: Wed, 30 Sep 2026 12:40:22 +0200 Subject: [PATCH] fix: carry EnIF-MDA information across assimilation passes --- docs/tutorials/enif.md | 46 +++++-- src/pipt/update_schemes/analysis/enif.py | 111 ++++++++++++---- src/pipt/update_schemes/esmda.py | 4 +- tests/assimilation/test_enif.py | 160 ++++++++++++++++++++++- 4 files changed, 271 insertions(+), 50 deletions(-) diff --git a/docs/tutorials/enif.md b/docs/tutorials/enif.md index 8b8d05d..eeacb15 100644 --- a/docs/tutorials/enif.md +++ b/docs/tutorials/enif.md @@ -41,11 +41,15 @@ The equivalent Python entry point is `ESMDA(keys_da, keys_en, sim, analysis="enif")`, and `("esmda", "enif")` resolves through the scheme registry like any other combination. -EnIF-MDA reruns the simulator and refits the regression and state precision -after each update. MDA requires positive, finite inflation factors satisfying -`sum(1 / alpha) = 1`. If you omit `inflation_param`, PET uses -`tot_assim_steps` for each factor. A scalar factor repeats across the schedule. -The schedule retains its original indexing on restart. +EnIF-MDA reruns the simulator and refits the response regression after each +update. It estimates the graph-based prior precision on the first pass, then +passes the accumulated posterior precision to the next pass, converting it to +that pass's standardized state coordinates. MDA requires positive, finite +inflation factors satisfying `sum(1 / alpha) = 1`. If you omit +`inflation_param`, PET uses `tot_assim_steps` for each factor. A scalar factor +repeats across the schedule. +Checkpoints retain both the schedule position and the accumulated precision; +restarting a multi-pass EnIF run requires a checkpoint with that information. ## Parameter graphs @@ -61,16 +65,20 @@ The regular-grid ordering matches PET's layered prior generator: with a different ordering, reduced active-cell arrays or irregular geometry, provide a graph whose node numbers match the imported parameter rows. -You can configure graphs and neighbourhood sizes under `enif`: +You can configure graphs and the initial precision-fitting neighbourhood under +`enif`: ```yaml enif: parameter_graphs: perm: perm_graph.npz neighbourhood_expansion: 2 - neighbor_propagation_order: 15 ``` +`neighbor_propagation_order` is accepted in existing configurations but is +ignored by EnIF-MDA: every retained state row is solved for at each pass so +the carried information is reflected in the ensemble. + Write a graph file as a symmetric sparse adjacency array with `scipy.sparse.save_npz`. For example, for five parameters arranged in a chain: @@ -96,9 +104,14 @@ increment, then applies PET's configured state limits. EnIF uses PET's perturbed observations, random-number stream, state clipping, forecast loop and misfit reporting. It scales observation covariance by the -current MDA factor once. It also estimates the unexplained response variance, -as in ERT. For a correlated observation covariance, it whitens the observations, -forecasts and perturbations before fitting the response map. +current MDA factor once. It estimates the unexplained response variance and +inflates the **total** noisy-residual variance: for step factor `alpha`, this +is `alpha * (observation variance + unexplained variance)`. PET's existing +observation perturbations already include the inflated measurement error; an +independent draw supplies the remaining `(alpha - 1) * unexplained variance`. +For a correlated observation covariance, PET whitens the observations, +forecasts and perturbations before fitting the response map and drawing this +additional noise in whitened coordinates. The analysis lives in `pipt/update_schemes/analysis/enif.py` and binds to the ES-MDA scheme like the `approx`, `full` and `subspace` flavours; it returns an @@ -106,10 +119,15 @@ additive state-space step and uses the direct sparse solver, matching ERT's non-iterative transport setting. After an update, the bound analysis object (`scheme.analysis`) exposes the -fitted `H`, `Prec_u`, `Prec_eps` and `Prec_posterior`. These matrices use -standardized, retained state rows; `enif_active_rows` maps them back to the -full state. With correlated observation errors, `H` and `Prec_eps` use -whitened observation coordinates. +fitted `H`, `Prec_u`, `Prec_eps` and `Prec_posterior`. `Prec_u` is the initial +graph-based fit on the first pass and the carried, rescaled posterior precision +on later passes. These matrices use standardized, retained state rows; +`enif_active_rows` maps them back to the full state. The scheme stores the +information needed by the next pass in `scheme.enif_information`, including +the posterior precision, scaling and +retained rows. Changes to the retained rows across passes are rejected rather +than silently discarding accumulated information. With correlated observation +errors, `H` and `Prec_eps` use whitened observation coordinates. This analysis requires at least two ensemble members and positive observation variances. It does not support PET's covariance localization, local analysis, diff --git a/src/pipt/update_schemes/analysis/enif.py b/src/pipt/update_schemes/analysis/enif.py index 26eb72e..afcb7bc 100644 --- a/src/pipt/update_schemes/analysis/enif.py +++ b/src/pipt/update_schemes/analysis/enif.py @@ -16,6 +16,7 @@ one-step schedule, ``mda={tot_assim_steps: 1}``. """ +from dataclasses import dataclass from os import PathLike import networkx as nx @@ -32,6 +33,15 @@ __all__ = ["enif_update"] +@dataclass(frozen=True) +class _EnIFInformation: + precision: sparse.sparray + scales: np.ndarray + active_rows: np.ndarray + groups: tuple + iteration: int + + class enif_update(AnalysisBase): """Graph-informed information-space update, as an ES-MDA analysis flavour. @@ -53,7 +63,8 @@ class enif_update(AnalysisBase): ``prior_`` block gets nearest-neighbour connectivity; a group without grid metadata is treated as independent. - ``neighbourhood_expansion``: precision fitting graph hops (default 2). - - ``neighbor_propagation_order``: update propagation hops (default 15). + - ``neighbor_propagation_order``: accepted for compatibility; MDA updates + all retained state rows to preserve the accumulated information. Covariance localization, local analysis, multilevel ensembles and ``emp_cov`` cannot be combined with this flavour; spatial dependence is @@ -81,8 +92,8 @@ def update(self, enX, enY, enE, **kwargs): Predicted data ensemble matrix, shape ``(nd, ne)``. enE : np.ndarray Perturbed observations with covariance - ``alpha * cov_data`` and the same shape as ``enY``. These are - used without adding more noise. + ``alpha * cov_data`` and the same shape as ``enY``. Additional + noise is drawn when the fitted response has unexplained variance. Returns ------- @@ -93,8 +104,9 @@ def update(self, enX, enY, enE, **kwargs): ----- Each parameter group has its own precision block. Parameters containing non-finite values, and parameters with no ensemble - spread, are held fixed. The regression and prior precision are - refitted at every MDA step. + spread, are held fixed. The regression is refitted at every MDA step; + the posterior precision is carried forward in the new standardized + state coordinates. """ scheme = self.scheme options = scheme.keys_da.get('enif', {}) @@ -115,7 +127,26 @@ def update(self, enX, enY, enE, **kwargs): active[finite] = np.ptp(enX[finite], axis=1) > 0 step = np.zeros(enX.shape, dtype=float) self.enif_active_rows = np.flatnonzero(active) + information = getattr(scheme, 'enif_information', None) + groups = tuple(sorted(scheme.idX.items(), key=lambda item: item[1][0])) + if scheme.iteration and information is None: + raise ValueError('EnIF-MDA requires the preceding posterior information to resume.') + if information is not None and ( + information.iteration != scheme.iteration + or information.groups != groups + or not np.array_equal(information.active_rows, self.enif_active_rows) + or information.precision.shape != (len(self.enif_active_rows),) * 2 + or information.scales.shape != (len(self.enif_active_rows),) + ): + raise ValueError('EnIF-MDA information does not match the current state rows or iteration.') if not active.any(): + scheme.enif_information = _EnIFInformation( + precision=sparse.csc_array((0, 0)), + scales=np.empty(0), + active_rows=self.enif_active_rows.copy(), + groups=groups, + iteration=scheme.iteration + 1, + ) return AnalysisResult(step=step) scaler = StandardScaler() @@ -123,36 +154,49 @@ def update(self, enX, enY, enE, **kwargs): Y, E, d, self.Prec_eps = self._observation_precision(enY, enE) self.H = linear_boost_ic_regression(U=U, Y=Y.T) - # Keep precision blocks in the same row order as the augmented state. - blocks = [] - for name, (start, stop) in sorted(scheme.idX.items(), key=lambda item: item[1][0]): - local_active = active[start:stop] - if not local_active.any(): - continue - graph = self._parameter_graph(name, stop - start) - graph = graph.subgraph(np.flatnonzero(local_active)) - graph = nx.convert_node_labels_to_integers(graph, ordering='sorted') - local_scaler = StandardScaler() - local_U = local_scaler.fit_transform(enX[start:stop][local_active].T) - blocks.append(fit_precision_cholesky_approximate( - local_U, - graph, - neighbourhood_expansion=options.get('neighbourhood_expansion', 2), - use_tqdm=self._use_tqdm(scheme), - )) - self.Prec_u = sparse.csc_array(sparse.block_diag(blocks, format='csc')) + if information is None: + # Keep precision blocks in the same row order as the augmented state. + blocks = [] + for name, (start, stop) in groups: + local_active = active[start:stop] + if not local_active.any(): + continue + graph = self._parameter_graph(name, stop - start) + graph = graph.subgraph(np.flatnonzero(local_active)) + graph = nx.convert_node_labels_to_integers(graph, ordering='sorted') + local_scaler = StandardScaler() + local_U = local_scaler.fit_transform(enX[start:stop][local_active].T) + blocks.append(fit_precision_cholesky_approximate( + local_U, + graph, + neighbourhood_expansion=options.get('neighbourhood_expansion', 2), + use_tqdm=self._use_tqdm(scheme), + )) + self.Prec_u = sparse.csc_array(sparse.block_diag(blocks, format='csc')) + else: + change_of_scale = sparse.diags_array(scaler.scale_ / information.scales, format='csc') + self.Prec_u = (change_of_scale @ information.precision @ change_of_scale).tocsc() gtmap = EnIF(Prec_u=self.Prec_u, Prec_eps=self.Prec_eps, H=self.H) - self.update_indices = gtmap.get_update_indices( - neighbor_propagation_order=options.get('neighbor_propagation_order', 15), - ) + self.update_indices = None canonical = gtmap.pushforward_to_canonical(U) residuals = gtmap.response_residual(U, Y.T) - # ERT transport draws noise internally. Use PET's existing perturbations - # instead: d - (residuals + d - E) == E - residuals. + alpha = scheme.alpha[scheme.iteration] + extra_variance = (alpha - 1) * gtmap.unexplained_variance + if alpha > 1: + self.Prec_eps = sparse.diags_array( + 1 / (1 / self.Prec_eps.diagonal() + extra_variance), format='csc', + ) + gtmap.Prec_eps = self.Prec_eps + extra_noise = scheme.ensemble.rng.standard_normal(residuals.shape) * np.sqrt(extra_variance) + else: + extra_noise = 0 + # PET already drew the measurement noise in E. The extra independent + # draw inflates the full noisy-residual variance, including the fitted + # unexplained response variance. canonical = gtmap.update_canonical( canonical=canonical, - residual_noisy=residuals + d - E.T, + residual_noisy=residuals + d - E.T + extra_noise, d=d, ) updated = gtmap.pullback_from_canonical( @@ -163,6 +207,13 @@ def update(self, enX, enY, enE, **kwargs): ) self.Prec_posterior = gtmap.Prec_u step[active] = scaler.inverse_transform(updated).T - enX[active] + scheme.enif_information = _EnIFInformation( + precision=self.Prec_posterior, + scales=scaler.scale_.copy(), + active_rows=self.enif_active_rows.copy(), + groups=groups, + iteration=scheme.iteration + 1, + ) return AnalysisResult(step=step) # ------------------------------------------------------------------ @@ -235,6 +286,8 @@ def _observation_precision(self, enY, enE): covariance = np.asarray(scheme.cov_data, dtype=float) alpha = scheme.alpha[scheme.iteration] nd = enY.shape[0] + if not np.isfinite(alpha) or alpha < 1: + raise ValueError('EnIF-MDA inflation must be finite and at least one.') if not np.all(np.isfinite(covariance)): raise ValueError('EnIF observation covariance must be finite.') if covariance.ndim == 2: diff --git a/src/pipt/update_schemes/esmda.py b/src/pipt/update_schemes/esmda.py index ad4f666..6006869 100644 --- a/src/pipt/update_schemes/esmda.py +++ b/src/pipt/update_schemes/esmda.py @@ -116,7 +116,7 @@ class ESMDA(AssimilationScheme): # The perturbed observations are redrawn every step (from the ensemble's # stream, whose state travels with the ensemble); the misfit is scored # against the un-inflated draw taken at construction (`enObs_conv`). - RESTART_ATTRIBUTES = ("enObs", "enObs_conv", "scale_data") + RESTART_ATTRIBUTES = ("enObs", "enObs_conv", "scale_data", "enif_information") def __init__(self, keys_da, keys_en, sim, analysis=None, ensemble=None): """Build the ensemble from the config (or take the one given) and bind the analysis. @@ -135,6 +135,8 @@ def __init__(self, keys_da, keys_en, sim, analysis=None, ensemble=None): # The analysis flavour is a parameter of the algorithm, not a different # algorithm, so it selects an analysis object rather than a class. self.bind_analysis(self.resolve_analysis(analysis, ensemble.keys_da)) + if self.analysis_name == 'enif': + self.enif_information = None self.prev_data_misfit_mean = None diff --git a/tests/assimilation/test_enif.py b/tests/assimilation/test_enif.py index 1b91dc5..6c15e81 100644 --- a/tests/assimilation/test_enif.py +++ b/tests/assimilation/test_enif.py @@ -12,6 +12,9 @@ posterior is known analytically. """ +import importlib +from types import SimpleNamespace + import numpy as np import pandas as pd import pytest @@ -42,6 +45,7 @@ def __init__(self): self.iteration = 0 self.vecObs = np.array([1.2, -0.3]) self.cov_data = np.array([0.2, 0.5]) + self.ensemble = SimpleNamespace(rng=np.random.RandomState(19)) @pytest.fixture(autouse=True) @@ -56,9 +60,8 @@ def scheme_double(): return FakeScheme() -@pytest.mark.parametrize('alpha', [1.0, 4.0]) -def test_matches_ert_transport(scheme_double, alpha): - """Compare the PET step with ERT's fit-and-transport recipe, member by member.""" +def test_matches_single_step_ert_transport(scheme_double): + """The one-pass update retains the EnIF transport equations.""" rng = np.random.default_rng(13) X = rng.normal(size=(6, 80)) * np.arange(1, 7)[:, None] + 5 Y = np.vstack((X[0] + 0.3 * X[1] ** 2, X[4] - X[5])) @@ -70,18 +73,17 @@ def test_matches_ert_transport(scheme_double, alpha): precision = fit_precision_cholesky_approximate(U, graph, use_tqdm=False) reference = EnIF( Prec_u=precision, - Prec_eps=sparse.diags_array(1 / (alpha * scheme_double.cov_data), format='csc'), + Prec_eps=sparse.diags_array(1 / scheme_double.cov_data, format='csc'), H=H, ) noise = reference.generate_observation_noise(X.shape[1], seed=19) expected = reference.transport( U, Y.T, scheme_double.vecObs, - update_indices=reference.get_update_indices(neighbor_propagation_order=15), + update_indices=None, iterative=False, seed=19, ) expected = scaler.inverse_transform(expected).T - scheme_double.alpha = [alpha] E = scheme_double.vecObs[:, None] - noise.T analysis = enif_update(scheme_double) result = analysis.update(X, Y, E) @@ -90,6 +92,107 @@ def test_matches_ert_transport(scheme_double, alpha): np.testing.assert_allclose(analysis.Prec_posterior.toarray(), reference.Prec_u.toarray()) +def test_mda_inflates_total_residual_and_posterior_spread(scheme_double, monkeypatch): + """An imperfect response map must dilute information and its stochastic update.""" + enif_module = importlib.import_module('pipt.update_schemes.analysis.enif') + monkeypatch.setattr(enif_module, 'fit_precision_cholesky_approximate', + lambda *args, **kwargs: sparse.eye_array(1, format='csc')) + monkeypatch.setattr(enif_module, 'linear_boost_ic_regression', + lambda **kwargs: sparse.csc_array([[1.0]])) + rng = np.random.default_rng(30) + X = rng.normal(size=(1, 3000)) + X = (X - X.mean()) / X.std() + Y = X + rng.normal(scale=0.9, size=X.shape) + scheme_double.idX = {'field': (0, 1)} + scheme_double.prior_info = {'field': {}} + scheme_double.vecObs = np.array([1.2]) + scheme_double.cov_data = np.array([0.25]) + scheme_double.alpha = [5.0] + E = scheme_double.vecObs[:, None] + rng.normal(scale=np.sqrt(5 * 0.25), size=Y.shape) + + analysis = enif_update(scheme_double) + updated = X + analysis.update(X, Y, E).step + expected_variance = 1 / (1 + 1 / (5 * (0.25 + 0.9**2))) + + np.testing.assert_allclose(analysis.Prec_eps.diagonal(), + 1 / (5 * 0.25 + 4 * np.var(Y - X)), rtol=0.01) + np.testing.assert_allclose(analysis.Prec_posterior.toarray(), + [[1 + 1 / (5 * (0.25 + 0.9**2))]], atol=0.01) + np.testing.assert_allclose(updated.var(), expected_variance, atol=0.035) + + +def test_mda_carries_precision_in_new_coordinates(scheme_double, monkeypatch): + enif_module = importlib.import_module('pipt.update_schemes.analysis.enif') + calls = [] + + def fit_precision(*args, **kwargs): + calls.append(1) + return sparse.eye_array(1, format='csc') + + def fit_response(U, Y): + return sparse.csc_array([[float(U[:, 0] @ Y[:, 0] / (U[:, 0] @ U[:, 0]))]]) + + monkeypatch.setattr(enif_module, 'fit_precision_cholesky_approximate', fit_precision) + monkeypatch.setattr(enif_module, 'linear_boost_ic_regression', fit_response) + scheme_double.idX = {'field': (0, 1)} + scheme_double.prior_info = {'field': {}} + scheme_double.vecObs = np.array([1.2]) + scheme_double.cov_data = np.array([0.5]) + scheme_double.alpha = [2.0, 2.0] + X = np.linspace(-2, 2, 100)[None, :] + initial_precision = 1 / X.var() + analysis = enif_update(scheme_double) + + for iteration in range(2): + scheme_double.iteration = iteration + E = np.broadcast_to(scheme_double.vecObs[:, None], X.shape) + X = X + analysis.update(X, X, E).step + information = scheme_double.enif_information + np.testing.assert_allclose(information.precision.toarray()[0, 0] / information.scales[0]**2, + initial_precision + (iteration + 1) / (2 * 0.5), atol=1e-12) + assert information.iteration == iteration + 1 + assert analysis.update_indices is None + + assert len(calls) == 1 + scheme_double.iteration = 2 + with pytest.raises(ValueError, match='does not match'): + analysis.update(np.ones_like(X), X, E) + + +def test_carried_precision_preserves_cross_parameter_coupling(scheme_double, monkeypatch): + enif_module = importlib.import_module('pipt.update_schemes.analysis.enif') + fitted = sparse.csc_array([[2.0, 0.6], [0.6, 3.0]]) + monkeypatch.setattr(enif_module, 'fit_precision_cholesky_approximate', + lambda *args, **kwargs: fitted) + monkeypatch.setattr(enif_module, 'linear_boost_ic_regression', + lambda U, Y: sparse.csc_array([[float(U[:, 0] @ Y[:, 0] / (U[:, 0] @ U[:, 0])), 0.0]])) + scheme_double.idX = {'field': (0, 2)} + scheme_double.prior_info = {'field': {}} + scheme_double.vecObs = np.array([1.2]) + scheme_double.cov_data = np.array([0.5]) + scheme_double.alpha = [2.0, 2.0] + X = np.vstack((np.linspace(-2, 2, 80), np.linspace(1, 3, 80))) + analysis = enif_update(scheme_double) + E = np.broadcast_to(scheme_double.vecObs[:, None], (1, X.shape[1])) + X = X + analysis.update(X, X[:1], E).step + previous = scheme_double.enif_information + physical_precision = previous.precision.toarray() / np.outer(previous.scales, previous.scales) + + scheme_double.iteration = 1 + analysis.update(X, X[:1], E) + scales = X.std(axis=1) + np.testing.assert_allclose(analysis.Prec_u.toarray() / np.outer(scales, scales), physical_precision) + assert analysis.Prec_u[0, 1] != 0 + + +def test_mda_rejects_missing_preceding_information(scheme_double): + scheme_double.iteration = 1 + scheme_double.alpha = [2.0, 2.0] + X = np.ones((6, 20)) + with pytest.raises(ValueError, match='preceding posterior information'): + enif_update(scheme_double).update(X, X[:2], X[:2]) + + def test_parameter_grid_order_and_custom_graphs(scheme_double, tmp_path): analysis = enif_update(scheme_double) scheme_double.prior_info['field']['nz'] = 2 @@ -176,6 +279,17 @@ def test_constant_and_nonfinite_parameters(scheme_double): analysis.update(X, Y * np.nan, Y) +def test_constant_parameters_remain_constant_over_multiple_passes(scheme_double): + scheme_double.alpha = [2.0, 2.0] + X = np.ones((6, 20)) + Y = np.ones((2, 20)) + analysis = enif_update(scheme_double) + for iteration in range(2): + scheme_double.iteration = iteration + np.testing.assert_array_equal(analysis.update(X, Y, Y).step, 0) + assert scheme_double.enif_information.iteration == 2 + + # ---------------------------------------------------------------------- # Configuration validation, at binding time # ---------------------------------------------------------------------- @@ -286,3 +400,37 @@ def test_state_limits(pet_inputs): scheme.run_assimilation() assert np.min(scheme.enX) >= -0.1 assert np.max(scheme.enX) <= 0.1 + + +def test_enif_restart_preserves_accumulated_information(pet_inputs, monkeypatch, tmp_path): + keys_da, keys_en, sim = pet_inputs + keys_da['mda'] = {'tot_assim_steps': 3} + checkpoint = tmp_path / 'enif_restart.pkl' + np.random.seed(14) + reference = ESMDA(keys_da, keys_en, sim) + reference.run_assimilation() + + keys_da.update(restartsave=True, restart_file=str(checkpoint)) + np.random.seed(14) + partial = ESMDA(keys_da, keys_en, sim) + update_step = partial.update_step + + def interrupt_after_checkpoint(): + if partial.iteration == 1: + raise InterruptedError + return update_step() + + monkeypatch.setattr(partial, 'update_step', interrupt_after_checkpoint) + with pytest.raises(InterruptedError): + partial.run_assimilation() + assert checkpoint.exists() + + keys_da.update(restart=True, restartsave=False) + np.random.seed(12345) + resumed = ESMDA(keys_da, keys_en, sim) + resumed.run_assimilation() + + np.testing.assert_array_equal(resumed.enX, reference.enX) + np.testing.assert_array_equal(resumed.analysis.Prec_posterior.toarray(), + reference.analysis.Prec_posterior.toarray()) + assert resumed.enif_information.iteration == 3