From cf803e5f7958afdc79839cd139a29f91cabd9336 Mon Sep 17 00:00:00 2001 From: Kristian Fossum Date: Mon, 28 Sep 2026 10:41:03 +0200 Subject: [PATCH] feat!: add EnIF analysis flavour for ES-MDA Use ERT's graphite-maps estimators as a bound analysis flavour of ES-MDA, following the restructured schemes/analyses split: the flavour selects through COMPATIBLE_ANALYSES, and the original single-update EnIF is the one-step MDA schedule. Include graphite-maps in the standard installation, document configuration, and test numerical parity and assimilation. BREAKING CHANGE: PET requires Python 3.12 through 3.14 to support the required graphite-maps dependency. --- .github/workflows/tests.yml | 2 +- README.md | 3 + docs/architecture.md | 7 +- docs/configuration.md | 15 +- docs/tutorials/README.md | 1 + docs/tutorials/enif.md | 117 ++++++++ pyproject.toml | 9 +- src/input_output/config.py | 4 +- src/pipt/update_schemes/analysis/__init__.py | 11 +- src/pipt/update_schemes/analysis/enif.py | 292 +++++++++++++++++++ src/pipt/update_schemes/esmda.py | 8 +- tests/assimilation/test_enif.py | 288 ++++++++++++++++++ tests/assimilation/test_scheme_factory.py | 8 +- 13 files changed, 744 insertions(+), 21 deletions(-) create mode 100644 docs/tutorials/enif.md create mode 100644 src/pipt/update_schemes/analysis/enif.py create mode 100644 tests/assimilation/test_enif.py diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index b5bca57..52a8f2b 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -37,7 +37,7 @@ jobs: strategy: fail-fast: false matrix: - python-version: ["3.10", 3.11, 3.12] + python-version: ["3.12", "3.13", "3.14"] steps: - uses: actions/checkout@v4 diff --git a/README.md b/README.md index ee11cf9..591bd45 100644 --- a/README.md +++ b/README.md @@ -13,6 +13,9 @@ at NORCE Norwegian Research Centre AS. ## Installation +PET requires Python 3.12 through 3.14. The standard installation includes +EnIF and EnIF-MDA with their dependencies. + Before installing ensure you have python3 pre-requisites. On a Debian system run: ``` diff --git a/docs/architecture.md b/docs/architecture.md index 9a63084..1ba4b5e 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -56,9 +56,10 @@ An analysis (`pipt.update_schemes.analysis`) is a class with `update(enX, enY, enE, **kwargs) -> AnalysisResult`, returning exactly one of a state-space `step`, a weight-space `w_step` or a `W_step`. The scheme turns it into a proposal with `propose_state(result, step_scale)`. The flavours are -`approx`, `full`, `subspace`, `margis` and the multilevel `hybrid`; -`register_analysis` adds one. Analyses read what they need from the scheme: -`state_scaling`, `scale_data`, `proj`, `cov_data`, `trunc_energy`, `lam`. +`approx`, `full`, `subspace`, `margis`, the multilevel `hybrid` and `enif` +(ES-MDA); `register_analysis` adds one. Analyses read what they need from the +scheme: `state_scaling`, `scale_data`, `proj`, `cov_data`, `trunc_energy`, +`lam`. ## Data on the analysis path diff --git a/docs/configuration.md b/docs/configuration.md index 4a4fc11..654f952 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -20,7 +20,7 @@ Flags accept `true`/`false`, `yes`/`no` and the Python booleans. A key marked | Key | Meaning | Default | | --- | --- | --- | | `scheme` | Algorithm: `esmda`, `es`, `enkf`, `lmenrml`, `gnenrml`. `pipt.available_schemes()` lists every `(scheme, analysis)` pair. | required | -| `analysis` | Analysis flavour the scheme runs: `approx`, `full`, `subspace` (all schemes); `subspace2` (ES-MDA, LM-EnRML, GN-EnRML); `margis` (GN-EnRML). `subspace2` solves for the ensemble transform directly and uses the analytic data covariance, so it reads neither `energy` nor `iteration.energy`. | `approx` | +| `analysis` | Analysis flavour the scheme runs: `approx`, `full`, `subspace` (all schemes); `subspace2` (ES-MDA, LM-EnRML, GN-EnRML); `margis` (GN-EnRML); `enif` (ES-MDA, see the `[dataassim.enif]` block below). `subspace2` solves for the ensemble transform directly and uses the analytic data covariance, so it reads neither `energy` nor `iteration.energy`. | `approx` | | `energy` | Truncation energy of the SVD in ES-MDA, ES and EnKF; a fraction, or a percentage when greater than 1. The iterative schemes read `iteration.energy`. | `0.98` | | `emp_cov` | The variance file holds an ensemble of observation errors; the analyses use that empirical covariance. Flag. | off | @@ -59,6 +59,19 @@ Flags accept `true`/`false`, `yes`/`no` and the Python booleans. A key marked | `tot_assim_steps` | Number of inflated assimilation steps; one update each. | required | | `inflation_param` | Inflation factor per step (a list) or one factor for all. The inverses must sum to 1. | `tot_assim_steps` for every step | +### EnIF: the `[dataassim.enif]` block + +Settings for the `enif` analysis flavour of ES-MDA (see +[the EnIF tutorial](tutorials/enif.md)). Spatial dependence is specified by +parameter graphs rather than localization: the flavour rejects `localization`, +`localanalysis`, `multilevel` and `emp_cov`. + +| Key | Meaning | Default | +| --- | --- | --- | +| `parameter_graphs` | Maps state names to NetworkX graphs, SciPy sparse adjacency arrays, or `.npz` files written with `scipy.sparse.save_npz`. Without one, a state with `grid` metadata in `prior_` gets nearest-neighbour connectivity, and a state without it independent nodes. | none | +| `neighbourhood_expansion` | Graph hops used when fitting the prior precision. | `2` | +| `neighbor_propagation_order` | Graph hops the update propagates through. | `15` | + ### Localization: the `[dataassim.localization]` block `name` selects the strategy; `pipt.localization.available_localizations()` diff --git a/docs/tutorials/README.md b/docs/tutorials/README.md index 427fb06..14ea895 100644 --- a/docs/tutorials/README.md +++ b/docs/tutorials/README.md @@ -21,3 +21,4 @@ Here are some tutorials. - [`adding_an_analysis.ipynb`](pipt/extending/adding_an_analysis): Write a new analysis flavour and bind it to a scheme - [`adding_a_scheme.ipynb`](pipt/extending/adding_a_scheme): Write a new scheme and register it for config-driven use +- [EnIF, an out-of-tree analysis brought in-tree](enif.md): The graph-informed information-filter flavour for ES-MDA -- installation, settings and parameter graphs diff --git a/docs/tutorials/enif.md b/docs/tutorials/enif.md new file mode 100644 index 0000000..8b8d05d --- /dev/null +++ b/docs/tutorials/enif.md @@ -0,0 +1,117 @@ +# Ensemble information filter (EnIF) + +PET offers the EnIF analysis as a flavour of ES-MDA. It uses the sparse +regression and precision estimation from ERT's +[`_enif_update.py`](https://github.com/equinor/ert/blob/main/src/ert/analysis/_enif_update.py) +through `graphite-maps`. + +## Installation + +Install PET from your checkout, including EnIF and its dependencies: + +```sh +python -m pip install -e . +``` + +PET requires Python 3.12 through 3.14, matching its `graphite-maps` dependency. + +## Select the analysis + +Keep your existing ensemble, observation and simulator settings. EnIF is an +analysis flavour of ES-MDA, so select it in the `dataassim` section: + +```yaml +scheme: esmda +analysis: enif +mda: + tot_assim_steps: 3 + inflation_param: [2, 4, 4] +``` + +The original, single-update EnIF is the one-step schedule: + +```yaml +scheme: esmda +analysis: enif +mda: + tot_assim_steps: 1 +``` + +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. + +## Parameter graphs + +EnIF estimates a separate prior precision block for each state in `idX`. +By default: + +- A state with `grid` metadata in `prior_` uses nearest-neighbour + connectivity. PET's prior parser converts `grid` to `nx`, `ny` and `nz`. +- A state without grid metadata uses independent graph nodes. + +The regular-grid ordering matches PET's layered prior generator: +`row = z * nx * ny + x * ny + y` (y varies fastest). For imported ensembles +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`: + +```yaml +enif: + parameter_graphs: + perm: perm_graph.npz + neighbourhood_expansion: 2 + neighbor_propagation_order: 15 +``` + +Write a graph file as a symmetric sparse adjacency array with +`scipy.sparse.save_npz`. For example, for five parameters arranged in a chain: + +```python +import networkx as nx +from scipy import sparse + +graph = nx.path_graph(5) +sparse.save_npz('perm_graph.npz', nx.to_scipy_sparse_array(graph, format='csc')) +``` + +Python configurations can also supply NetworkX graphs or SciPy sparse +adjacency arrays directly in `parameter_graphs`. Use local node numbers +`0` through `number_of_parameter_rows - 1` for each state. Graph weights do +not affect the fit; EnIF uses connectivity. + +EnIF excludes rows containing non-finite values and rows with zero ensemble +spread from estimation. It removes their graph nodes without connecting +neighbours across the resulting gaps. The analysis gives those rows a zero +increment, then applies PET's configured state limits. + +## Update and diagnostics + +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. + +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 +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. + +This analysis requires at least two ensemble members and positive observation +variances. It does not support PET's covariance localization, local analysis, +multilevel ensembles or `emp_cov` sample input. Use parameter graphs to specify +spatial dependence. diff --git a/pyproject.toml b/pyproject.toml index 2ae54bc..f6121bd 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -15,15 +15,15 @@ maintainers = [ ] license = { file = "LICENSE" } readme = "README.md" -requires-python = ">=3.10" +requires-python = ">=3.12,<3.15" classifiers = [ "Development Status :: 4 - Beta", "Intended Audience :: Science/Research", "License :: OSI Approved :: GNU General Public License v3 (GPLv3)", "Programming Language :: Python :: 3", - "Programming Language :: Python :: 3.10", - "Programming Language :: Python :: 3.11", "Programming Language :: Python :: 3.12", + "Programming Language :: Python :: 3.13", + "Programming Language :: Python :: 3.14", "Topic :: Scientific/Engineering", ] dependencies = [ @@ -34,6 +34,7 @@ dependencies = [ "tqdm", "PyWavelets", "geostat @ git+https://github.com/Python-Ensemble-Toolbox/Geostatistics@3f9f0c876815db140fae3404d322892190cb6728", + "graphite-maps>=0.0.11,<0.1", "pandas", "p_tqdm", "tomli", @@ -48,7 +49,7 @@ pet = "pet_cli.__main__:main" [project.optional-dependencies] minires = [ # Pure-Python two-phase TPFA simulator, ref src/simulator/minires.py. - # Optional: it needs Python >= 3.12, whereas PET supports 3.10. + # Optional: keeps the toy simulator out of the core dependency set. "minires>=0.3.2", ] dev = [ diff --git a/src/input_output/config.py b/src/input_output/config.py index c09f9b9..e523bbc 100644 --- a/src/input_output/config.py +++ b/src/input_output/config.py @@ -67,7 +67,7 @@ def pairs_to_dict(entries) -> dict: ALIASES_DATAASSIM = {"truedata": "data", "var": "datavar", "save_folder": "savefolder", "restartfile": "restart_file"} FLAGS_DATAASSIM = ("emp_cov", "restart", "restartsave", "obsvarsave", "screendata", "post_process_forecast", "scale_data", "logit") -BLOCKS_DATAASSIM = ("iteration", "mda", "compress", "localization", "localanalysis") +BLOCKS_DATAASSIM = ("iteration", "mda", "compress", "localization", "localanalysis", "enif") ALIASES_ENSEMBLE = {"importstaticvar": "importstate", "save_folder": "savefolder"} FLAGS_ENSEMBLE = ("save_prior", "disable_tqdm", "natural_gradient") @@ -86,7 +86,7 @@ def pairs_to_dict(entries) -> dict: KNOWN_DATAASSIM = frozenset({ "scheme", "analysis", "data", "datavar", "obsname", "datatype", "truedataindex", "assimindex", "energy", "emp_cov", "iteration", "mda", "compress", "localization", "localanalysis", "actnum", "scale_data", "scale", - "screendata", "post_process_forecast", "remove_outliers", + "screendata", "post_process_forecast", "remove_outliers", "enif", "savefolder", "nosave", "savedata", "analysisdebug", "iterinfo", "obsvarsave", "qa", "qc", "restart", "restartsave", "restart_file", "logit", "logger_name", # legacy text files keep the ensemble's keys in DATAASSIM diff --git a/src/pipt/update_schemes/analysis/__init__.py b/src/pipt/update_schemes/analysis/__init__.py index 93f87a6..fbf2248 100644 --- a/src/pipt/update_schemes/analysis/__init__.py +++ b/src/pipt/update_schemes/analysis/__init__.py @@ -15,10 +15,11 @@ :class:`AnalysisBase` -- the shared contract and helpers. ``approx``, ``full``, ``subspace``, ``subspace2`` The four registered flavours. -``hybrid``, ``margis`` - Flavours consumed as mixins rather than through the registry: ``hybrid`` - belongs to the multilevel scheme and ``margis`` is backed by a private - package when installed. +``hybrid``, ``margis``, ``enif`` + Flavours that live outside the registry: ``hybrid`` belongs to the + multilevel scheme, ``margis`` is backed by a private package when + installed, and ``enif`` (ES-MDA only) uses the graphite-maps + estimators -- see :mod:`pipt.update_schemes.analysis.enif`. ``registry`` Name-to-class lookup, plus :func:`register_analysis` for out-of-tree flavours. @@ -35,6 +36,7 @@ from .hybrid import hybrid_update from .subspace import subspace_update from .subspace2 import subspace2_update +from .enif import enif_update from .registry import ( ANALYSES, available_analyses, @@ -50,6 +52,7 @@ "subspace_update", "subspace2_update", "hybrid_update", + "enif_update", "ANALYSES", "available_analyses", "get_analysis", diff --git a/src/pipt/update_schemes/analysis/enif.py b/src/pipt/update_schemes/analysis/enif.py new file mode 100644 index 0000000..26eb72e --- /dev/null +++ b/src/pipt/update_schemes/analysis/enif.py @@ -0,0 +1,292 @@ +"""Graph-informed ensemble information-filter update (EnIF). + +The ensemble information filter replaces the ensemble covariance of the +smoother updates with two sparse graph-informed estimates: a per-parameter +precision matrix fitted on a graph of the parameter connectivity, and a +boosted linear regression of the responses on the (standardised) state. The +estimators are ERT's, through the ``graphite-maps`` dependency -- which is +also why PET requires Python 3.12 through 3.14; PET supplies the MDA +lifecycle around them -- perturbed observations, inflation schedule, +forecasting, state limits and scoring. + +The flavour is ES-MDA-specific, the way ``hybrid`` belongs to the multilevel +scheme and ``margis`` to GN-EnRML: it is wired into +``ESMDA.COMPATIBLE_ANALYSES`` rather than the global analysis registry, and +selectable as ``analysis='enif'``. The original single-update EnIF is the +one-step schedule, ``mda={tot_assim_steps: 1}``. +""" + +from os import PathLike + +import networkx as nx +import numpy as np +from graphite_maps.enif import EnIF +from graphite_maps.linear_regression import linear_boost_ic_regression +from graphite_maps.precision_estimation import fit_precision_cholesky_approximate +from scipy import linalg, sparse +from sklearn.preprocessing import StandardScaler + +from pipt.update_schemes.analysis.base import AnalysisBase, AnalysisResult +import pipt.misc_tools.extract_tools as extract + +__all__ = ["enif_update"] + + +class enif_update(AnalysisBase): + """Graph-informed information-space update, as an ES-MDA analysis flavour. + + Parameters + ---------- + scheme : object, optional + The ES-MDA scheme this analysis computes updates for. ``None`` leaves + it unbound; the configuration is validated when it is bound, so an + incompatible ``dataassim``/``ensemble`` combination or a bad ``enif`` + block fails at construction rather than mid-run. + + Notes + ----- + The ``enif`` block of the ``dataassim`` section accepts: + + - ``parameter_graphs``: maps state names to NetworkX graphs, sparse + adjacency arrays, or files written with ``scipy.sparse.save_npz``. + Without one, a group with ``nx``/``ny`` (``nz``) grid metadata in its + ``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). + + Covariance localization, local analysis, multilevel ensembles and + ``emp_cov`` cannot be combined with this flavour; spatial dependence is + specified by the parameter graphs. + + Diagnostics of the last update -- the fitted regression ``H``, the prior + and posterior precisions ``Prec_u``/``Prec_posterior``, the observation + precision ``Prec_eps``, the ``update_indices`` and the active rows + ``enif_active_rows`` -- are kept on the analysis object, not the scheme. + """ + + def __init__(self, scheme=None): + super().__init__(scheme) + if scheme is not None: + self._validate_configuration(scheme) + + def update(self, enX, enY, enE, **kwargs): + """Compute the graph-informed update step. + + Parameters + ---------- + enX : np.ndarray + State ensemble matrix, shape ``(nx, ne)``. + enY : np.ndarray + 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. + + Returns + ------- + AnalysisResult + The additive state-space ``step``. + + Notes + ----- + 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. + """ + scheme = self.scheme + options = scheme.keys_da.get('enif', {}) + + if enX.ndim != 2 or enX.shape[1] < 2: + raise ValueError('EnIF requires at least two ensemble members.') + if enY.ndim != 2 or enY.shape[1] != enX.shape[1] or enE.shape != enY.shape: + raise ValueError('EnIF state, forecast and observation ensembles have incompatible shapes.') + if enY.shape[0] == 0 or scheme.vecObs.shape != (enY.shape[0],): + raise ValueError('EnIF requires observations matching the forecast rows.') + if not all(np.all(np.isfinite(value)) for value in (enY, enE, scheme.vecObs)): + raise ValueError('EnIF observations and forecasts must be finite.') + + finite = np.all(np.isfinite(enX), axis=1) + if not finite.any(): + raise ValueError('No finite parameter rows available for EnIF.') + active = finite.copy() + active[finite] = np.ptp(enX[finite], axis=1) > 0 + step = np.zeros(enX.shape, dtype=float) + self.enif_active_rows = np.flatnonzero(active) + if not active.any(): + return AnalysisResult(step=step) + + scaler = StandardScaler() + U = scaler.fit_transform(enX[active].T) + 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')) + + 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), + ) + 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. + canonical = gtmap.update_canonical( + canonical=canonical, + residual_noisy=residuals + d - E.T, + d=d, + ) + updated = gtmap.pullback_from_canonical( + updated_canonical=canonical, + update_indices=self.update_indices, + U_prior=U, + iterative=False, + ) + self.Prec_posterior = gtmap.Prec_u + step[active] = scaler.inverse_transform(updated).T - enX[active] + return AnalysisResult(step=step) + + # ------------------------------------------------------------------ + # Helpers + # ------------------------------------------------------------------ + def _parameter_graph(self, name, size): + """Load a group graph or build nearest-neighbour connectivity from its grid. + + Graph nodes are local parameter rows, numbered ``0 .. size-1``. + Regular grids use y-fastest ordering, then x, then z, matching PET's + layered prior ensembles. Without grid metadata, parameters are + independent. + + Parameters + ---------- + name : str + State-variable name whose graph to build. + size : int + Number of parameter rows in the group. + + Returns + ------- + networkx.Graph + """ + scheme = self.scheme + graph = scheme.keys_da.get('enif', {}).get('parameter_graphs', {}).get(name) + if graph is not None: + if isinstance(graph, (str, PathLike)): + graph = sparse.load_npz(graph) + if sparse.issparse(graph): + if graph.shape != (size, size) or (graph != graph.T).nnz: + raise ValueError(f'EnIF graph for {name} must be a symmetric ({size}, {size}) adjacency.') + graph = nx.from_scipy_sparse_array(graph) + if not isinstance(graph, nx.Graph) or graph.is_directed() or graph.is_multigraph(): + raise ValueError(f'EnIF graph for {name} must be an undirected simple graph.') + if set(graph.nodes) != set(range(size)): + raise ValueError(f'EnIF graph for {name} must have nodes 0 through {size - 1}.') + return graph.copy() + + info = scheme.prior_info[name] + if not all(key in info for key in ('nx', 'ny')): + return nx.empty_graph(size) + shape = (int(info.get('nz', 1)), int(info['nx']), int(info['ny'])) + if min(shape) < 1 or np.prod(shape) != size: + raise ValueError( + f'EnIF grid for {name} has {np.prod(shape)} cells but {size} parameter rows. ' + 'Provide parameter_graphs for a reduced or irregular grid.' + ) + cells = np.arange(size).reshape(shape) + graph = nx.empty_graph(size) + for axis in range(3): + left = [slice(None)] * 3 + right = [slice(None)] * 3 + left[axis] = slice(None, -1) + right[axis] = slice(1, None) + graph.add_edges_from(zip(cells[tuple(left)].ravel(), cells[tuple(right)].ravel())) + return graph + + def _observation_precision(self, enY, enE): + """Inflate observation covariance once; whiten correlated observation errors. + + Returns + ------- + Y, E, d, Prec_eps + The (possibly whitened) forecast and perturbation ensembles and + observation vector, with the observation precision of the + inflated covariance. + """ + scheme = self.scheme + covariance = np.asarray(scheme.cov_data, dtype=float) + alpha = scheme.alpha[scheme.iteration] + nd = enY.shape[0] + if not np.all(np.isfinite(covariance)): + raise ValueError('EnIF observation covariance must be finite.') + if covariance.ndim == 2: + if covariance.shape != (nd, nd) or not np.allclose(covariance, covariance.T): + raise ValueError('EnIF observation covariance must be square and symmetric.') + if np.count_nonzero(covariance - np.diag(covariance.diagonal())): + chol = linalg.cholesky(covariance, lower=True) + Y = linalg.solve_triangular(chol, enY, lower=True) + E = linalg.solve_triangular(chol, enE, lower=True) + d = linalg.solve_triangular(chol, scheme.vecObs, lower=True) + precision = sparse.diags_array(np.full(nd, 1.0 / alpha), format='csc') + return Y, E, d, precision + covariance = covariance.diagonal() + if covariance.shape != (nd,) or np.any(covariance <= 0): + raise ValueError('EnIF requires one strictly positive observation variance per forecast row.') + precision = sparse.diags_array(1.0 / (alpha * covariance), format='csc') + return enY, enE, scheme.vecObs, precision + + def _validate_configuration(self, scheme): + """Reject keys EnIF cannot honour, and validate the ``enif`` block.""" + keys_da = scheme.keys_da + keys_en = self._scheme_keys_en(scheme) + for key in ('localization', 'localanalysis', 'multilevel'): + if key in keys_da or key in keys_en: + raise ValueError(f'EnIF does not support {key}.') + if extract.is_enabled(keys_da.get('emp_cov', False)): + raise ValueError('EnIF requires observation variances, not emp_cov samples.') + + options = keys_da.get('enif', {}) + if not isinstance(options, dict): + raise ValueError('ENIF settings must be a dictionary.') + for key, minimum in (('neighbourhood_expansion', 1), ('neighbor_propagation_order', 0)): + value = options.get(key, minimum) + if not isinstance(value, (int, np.integer)) or isinstance(value, bool) or value < minimum: + raise ValueError(f'EnIF {key} must be an integer >= {minimum}.') + graphs = options.get('parameter_graphs', {}) + if not isinstance(graphs, dict) or set(graphs) - set(scheme.idX): + raise ValueError('EnIF parameter_graphs must map known state names to graphs.') + + @staticmethod + def _scheme_keys_en(scheme): + """The ensemble section of the config, wherever the scheme keeps it. + + A real scheme exposes it as ``scheme.ensemble.keys_en``; a flat test + double may carry ``keys_en`` directly; otherwise there is none. + """ + keys_en = getattr(scheme, 'keys_en', None) + if keys_en is None: + keys_en = getattr(getattr(scheme, 'ensemble', None), 'keys_en', None) + return keys_en or {} + + @classmethod + def _use_tqdm(cls, scheme): + """Show graphite-maps' progress bars unless the config disabled them.""" + return not bool(cls._scheme_keys_en(scheme).get('disable_tqdm', False)) diff --git a/src/pipt/update_schemes/esmda.py b/src/pipt/update_schemes/esmda.py index 10daf2c..ad4f666 100644 --- a/src/pipt/update_schemes/esmda.py +++ b/src/pipt/update_schemes/esmda.py @@ -13,6 +13,7 @@ from pipt.update_schemes.analysis.full import full_update from pipt.update_schemes.analysis.subspace import subspace_update from pipt.update_schemes.analysis.subspace2 import subspace2_update +from pipt.update_schemes.analysis.enif import enif_update import pipt.misc_tools.analysis_tools as at __all__ = ['ESMDA'] @@ -47,11 +48,13 @@ class ESMDA(AssimilationScheme): variable names, and the ``prior_`` blocks describing each. sim : object Forward simulator instance, e.g. ``simulator.opm.flow``. - analysis : {'approx', 'full', 'subspace'}, optional + analysis : {'approx', 'full', 'subspace', 'subspace2', 'enif'}, optional Analysis flavour, i.e. how the ensemble-approximated sensitivity is inverted. Defaults to the ``analysis`` key in ``keys_da``, falling back to ``'approx'``. The flavours differ in cost and in how they handle a - rank-deficient ensemble; they solve the same update equation. + rank-deficient ensemble; they solve the same update equation. The + ``'enif'`` flavour is the graph-informed information filter; see its + module for the ``enif`` settings block. Attributes ---------- @@ -107,6 +110,7 @@ class ESMDA(AssimilationScheme): "full": full_update, "subspace": subspace_update, "subspace2": subspace2_update, + "enif": enif_update, } # The perturbed observations are redrawn every step (from the ensemble's diff --git a/tests/assimilation/test_enif.py b/tests/assimilation/test_enif.py new file mode 100644 index 0000000..1b91dc5 --- /dev/null +++ b/tests/assimilation/test_enif.py @@ -0,0 +1,288 @@ +"""Numerical and PET lifecycle tests for the EnIF analysis flavour. + +The first block binds ``enif_update`` to a flat scheme double, following +``test_analysis_binding``: a real scheme exposes the ensemble-owned context +(``idX``, ``prior_info``, ``cov_data``) as properties of its own, so a double +just needs those names present. These tests pin the numerics against ERT's +own fit-and-transport recipe, the graph construction, and the input +validation. + +The second block runs ES-MDA with the flavour bound, through the ordinary +config path, on a single-parameter identity model where the Gaussian +posterior is known analytically. +""" + +import numpy as np +import pandas as pd +import pytest + +from graphite_maps.enif import EnIF +from graphite_maps.linear_regression import linear_boost_ic_regression +from graphite_maps.precision_estimation import fit_precision_cholesky_approximate +from scipy import sparse +from sklearn.preprocessing import StandardScaler +import networkx as nx + +from misc.structures import PETDataFrame +from pipt import ESMDA +from pipt.update_schemes import registry +from pipt.update_schemes.analysis.enif import enif_update +from simulator.simple_models import lin_1d + + +class FakeScheme: + """The context the EnIF analysis reads, and nothing else (see module docstring).""" + + def __init__(self): + self.keys_da = {} + self.keys_en = {'disable_tqdm': True} + self.idX = {'field': (0, 6)} + self.prior_info = {'field': {'nx': 3, 'ny': 2, 'nz': 1}} + self.alpha = [1.0] + self.iteration = 0 + self.vecObs = np.array([1.2, -0.3]) + self.cov_data = np.array([0.2, 0.5]) + + +@pytest.fixture(autouse=True) +def preserve_random_state(): + state = np.random.get_state() + yield + np.random.set_state(state) + + +@pytest.fixture +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.""" + 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])) + graph = nx.grid_2d_graph(3, 2) + graph = nx.convert_node_labels_to_integers(graph) + scaler = StandardScaler() + U = scaler.fit_transform(X.T) + H = linear_boost_ic_regression(U=U, Y=Y.T) + 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'), + 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), + 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) + + np.testing.assert_allclose(X + result.step, expected, rtol=1e-11, atol=1e-11) + np.testing.assert_allclose(analysis.Prec_posterior.toarray(), reference.Prec_u.toarray()) + + +def test_parameter_grid_order_and_custom_graphs(scheme_double, tmp_path): + analysis = enif_update(scheme_double) + scheme_double.prior_info['field']['nz'] = 2 + graph = analysis._parameter_graph('field', 12) + assert set(graph.neighbors(0)) == {1, 2, 6} + assert set(graph.neighbors(5)) == {3, 4, 11} + assert graph.number_of_edges() == 20 + + custom = nx.path_graph(6) + filename = tmp_path / 'graph.npz' + sparse.save_npz(filename, nx.to_scipy_sparse_array(custom)) + scheme_double.keys_da['enif'] = {'parameter_graphs': {'field': filename}} + assert set(analysis._parameter_graph('field', 6).edges) == set(custom.edges) + + scheme_double.keys_da['enif']['parameter_graphs']['field'] = nx.empty_graph(6) + assert analysis._parameter_graph('field', 6).number_of_edges() == 0 + scheme_double.keys_da['enif']['parameter_graphs']['field'] = nx.path_graph(5) + with pytest.raises(ValueError, match='nodes 0 through 5'): + analysis._parameter_graph('field', 6) + scheme_double.keys_da['enif'] = {} + with pytest.raises(ValueError, match='12 cells but 6 parameter rows'): + analysis._parameter_graph('field', 6) + + +def test_masks_and_group_precision(scheme_double): + rng = np.random.default_rng(5) + X = rng.normal(size=(6, 60)) + X[1] = np.nan + X[3, 0] = np.nan + X[4] = 2.0 + Y = np.vstack((X[0], X[5])) + scheme_double.idX = {'other': (5, 6), 'field': (0, 5)} + scheme_double.prior_info = {'field': {'nx': 5, 'ny': 1}, 'other': {}} + analysis = enif_update(scheme_double) + result = analysis.update(X, Y, np.tile(scheme_double.vecObs[:, None], (1, X.shape[1]))) + + np.testing.assert_array_equal(analysis.enif_active_rows, [0, 2, 5]) + np.testing.assert_array_equal(result.step[[1, 3, 4]], 0) + assert np.isfinite(result.step).all() + assert np.linalg.norm(result.step[[0, 5]]) > 0 + # Removing inactive nodes must not bridge across the hole in the field. + assert analysis.Prec_u[0, 1] == 0 + assert analysis.Prec_u[:2, 2:].nnz == 0 + + +def test_correlated_observations(scheme_double): + """Whitened EnIF agrees with a Gaussian update using the full covariance.""" + rng = np.random.default_rng(17) + X = rng.normal(size=(6, 400)) + X -= X.mean(axis=1, keepdims=True) + X /= X.std(axis=1, keepdims=True) + Y = np.vstack((X[0], X[1])) + covariance = np.array([[0.4, 0.2], [0.2, 0.6]]) + scheme_double.cov_data = covariance + scheme_double.prior_info = {'field': {}} + E = np.tile(scheme_double.vecObs[:, None], (1, X.shape[1])) + analysis = enif_update(scheme_double) + result = analysis.update(X, Y, E) + + expected_mean = np.linalg.solve(np.eye(2) + covariance, scheme_double.vecObs) + np.testing.assert_allclose((X + result.step).mean(axis=1)[:2], expected_mean, atol=0.025) + + +@pytest.mark.parametrize('covariance', [np.array([0.0, 1.0]), np.array([-1.0, 1.0]), + np.array([np.nan, 1.0]), np.ones(3), + np.array([[1.0, 0.2], [0.0, 1.0]])]) +def test_invalid_observation_covariance(scheme_double, covariance): + scheme_double.cov_data = covariance + analysis = enif_update(scheme_double) + with pytest.raises(ValueError): + analysis._observation_precision(np.ones((2, 20)), np.ones((2, 20))) + + +def test_constant_and_nonfinite_parameters(scheme_double): + X = np.ones((6, 20)) + Y = np.ones((2, 20)) + analysis = enif_update(scheme_double) + np.testing.assert_array_equal(analysis.update(X, Y, Y).step, 0) + with pytest.raises(ValueError, match='No finite parameter rows'): + analysis.update(X * np.nan, Y, Y) + with pytest.raises(ValueError, match='at least two ensemble members'): + analysis.update(X[:, :1], Y[:, :1], Y[:, :1]) + with pytest.raises(ValueError, match='must be finite'): + analysis.update(X, Y * np.nan, Y) + + +# ---------------------------------------------------------------------- +# Configuration validation, at binding time +# ---------------------------------------------------------------------- +@pytest.mark.parametrize('section, key', [ + ('keys_da', 'localization'), ('keys_da', 'localanalysis'), + ('keys_da', 'multilevel'), ('keys_en', 'multilevel'), + ('keys_en', 'localization'), +]) +def test_rejects_unsupported_options(scheme_double, section, key): + getattr(scheme_double, section)[key] = {} + with pytest.raises(ValueError, match=f'EnIF does not support {key}'): + enif_update(scheme_double) + + +def test_rejects_emp_cov(scheme_double): + scheme_double.keys_da['emp_cov'] = True + with pytest.raises(ValueError, match='emp_cov'): + enif_update(scheme_double) + + +@pytest.mark.parametrize('options', [ + 'not-a-dictionary', + {'neighbourhood_expansion': 0}, + {'neighbourhood_expansion': 1.5}, + {'neighbourhood_expansion': True}, + {'neighbor_propagation_order': -1}, +]) +def test_rejects_invalid_enif_options(scheme_double, options): + scheme_double.keys_da['enif'] = options + with pytest.raises(ValueError): + enif_update(scheme_double) + + +def test_rejects_unknown_parameter_graph_names(scheme_double): + scheme_double.keys_da['enif'] = {'parameter_graphs': {'unknown': nx.path_graph(6)}} + with pytest.raises(ValueError, match='known state names'): + enif_update(scheme_double) + + +# ---------------------------------------------------------------------- +# ES-MDA with the EnIF flavour, through the ordinary config path +# ---------------------------------------------------------------------- +@pytest.fixture +def pet_inputs(tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + rng = np.random.default_rng(42) + np.savez('prior.npz', field=rng.normal(size=(1, 400))) + + data = PETDataFrame(pd.DataFrame({'value': [1.0]}, index=pd.Index([0], name='index'))) + var = PETDataFrame(pd.DataFrame({'value': [['abs', 0.25]]}, index=pd.Index([0], name='index'))) + data.to_pickle('true_data.pkl') + var.to_pickle('var.pkl') + + keys_da = { + 'scheme': 'esmda', 'analysis': 'enif', + 'obsname': 'index', 'truedataindex': [0], 'assimindex': [0], + 'datatype': ['value'], 'data': 'true_data.pkl', 'datavar': 'var.pkl', + } + keys_en = { + 'ne': 400, 'state': ['field'], 'prior_field': {'mean': 0.0, 'var': 1.0}, + 'importstate': 'prior.npz', 'disable_tqdm': True, + } + sim = lin_1d({'reporttype': 'index', 'reportpoint': [0], 'datatype': ['value']}) + return keys_da, keys_en, sim + + +@pytest.mark.parametrize('mda, steps', [ + ({'tot_assim_steps': 1}, 1), + ({'tot_assim_steps': 3}, 3), + ({'tot_assim_steps': 3, 'inflation_param': [2, 4, 4]}, 3), +]) +def test_esmda_enif_assimilation_loop(pet_inputs, mda, steps): + keys_da, keys_en, sim = pet_inputs + keys_da['mda'] = mda + np.random.seed(21) + scheme = ESMDA(keys_da, keys_en, sim) + prior = scheme.prior_enX.copy() + result = scheme.run_assimilation() + + assert isinstance(scheme.analysis, enif_update) + assert scheme.analysis_name == 'enif' + assert scheme.iteration == steps + assert result.data_misfit < result.prior_data_misfit + np.testing.assert_array_equal(scheme.prior_enX, prior) + # N(0, 1) prior observed at 1 with variance 0.25 has N(0.8, 0.2) posterior. + np.testing.assert_allclose(scheme.enX.mean(), 0.8, atol=0.07) + np.testing.assert_allclose(scheme.enX.var(), 0.2, atol=0.05) + posterior = np.load('Results/posterior_state_estimate.npz')['field'] + np.testing.assert_array_equal(posterior, scheme.enX) + np.testing.assert_allclose(scheme.pred_data.matrix, scheme.enX) + + +def test_registry_offers_the_flavour_on_esmda_only(): + assert ('esmda', 'enif') in registry.available_schemes() + ctor = registry.get_scheme('esmda', 'enif') + assert ctor.func is ESMDA + assert ctor.keywords == {'analysis': 'enif'} + from pipt.update_schemes.analysis.registry import available_analyses + assert 'enif' not in available_analyses() + + +def test_state_limits(pet_inputs): + keys_da, keys_en, sim = pet_inputs + keys_da['mda'] = {'tot_assim_steps': 1} + keys_en['prior_field']['limits'] = [-0.1, 0.1] + np.random.seed(23) + scheme = ESMDA(keys_da, keys_en, sim) + scheme.run_assimilation() + assert np.min(scheme.enX) >= -0.1 + assert np.max(scheme.enX) <= 0.1 diff --git a/tests/assimilation/test_scheme_factory.py b/tests/assimilation/test_scheme_factory.py index 80a5415..44773d3 100644 --- a/tests/assimilation/test_scheme_factory.py +++ b/tests/assimilation/test_scheme_factory.py @@ -50,10 +50,10 @@ def test_five_algorithms_cover_every_registered_combination(): def test_registry_size_matches_the_algorithms_specials_and_historical_names(): """Down from eighteen hand-written classes: 5 algorithms x 3 flavours, ``subspace2`` - on the three schemes that can apply an ensemble transform, the two combinations - backed by a distinct implementation, and the two historical names (co_lm_enrml, - gn_enrml) that each pin a single flavour.""" - assert len(registry.available_schemes()) == 5 * 3 + 3 + 2 + 2 + on the three schemes that can apply an ensemble transform, ``enif`` on ES-MDA, + the two combinations backed by a distinct implementation, and the two historical + names (co_lm_enrml, gn_enrml) that each pin a single flavour.""" + assert len(registry.available_schemes()) == 5 * 3 + 3 + 1 + 2 + 2 assert len(ALGORITHMS) == 5 # the public constructors above assert len(registry.ALGORITHMS) == 5 + 2 # plus the two historical names