Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 32 additions & 14 deletions docs/tutorials/enif.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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:

Expand All @@ -96,20 +104,30 @@ 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
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,
Expand Down
111 changes: 82 additions & 29 deletions src/pipt/update_schemes/analysis/enif.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
one-step schedule, ``mda={tot_assim_steps: 1}``.
"""

from dataclasses import dataclass
from os import PathLike

import networkx as nx
Expand All @@ -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.

Expand All @@ -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
Expand Down Expand Up @@ -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
-------
Expand All @@ -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', {})
Expand All @@ -115,44 +127,76 @@ 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()
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'))
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(
Expand All @@ -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)

# ------------------------------------------------------------------
Expand Down Expand Up @@ -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:
Expand Down
4 changes: 3 additions & 1 deletion src/pipt/update_schemes/esmda.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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

Expand Down
Loading
Loading