diff --git a/src/underworld3/checkpoint/snapshot.py b/src/underworld3/checkpoint/snapshot.py index eff45baca..878280caa 100644 --- a/src/underworld3/checkpoint/snapshot.py +++ b/src/underworld3/checkpoint/snapshot.py @@ -134,6 +134,11 @@ class Snapshot: # state-bearers. List preserves capture order — informational only, # since lookup is by key. state_bearers: list = field(default_factory=list) + # Expression captures: list of (stable_key, _sym, _wrapped). A parameter is + # not a mesh, a swarm or a state-bearer, so before this it was not captured + # at all and a restore returned the fields of one step with the parameters of + # another, silently. See _capture_expressions. + expressions: list = field(default_factory=list) metadata: dict[str, Any] = field(default_factory=dict) @@ -194,6 +199,7 @@ def snapshot(model, *, path: Optional[str] = None) -> Snapshot: _capture_swarm(snap, swarm) for obj in list(model._state_bearers): _capture_state_bearer(snap, obj) + _capture_expressions(snap) return snap @@ -243,6 +249,69 @@ def _capture_state_bearer(snap: Snapshot, obj) -> None: snap.state_bearers.append((_state_bearer_key(obj), copy.deepcopy(state))) +def _expression_key(expr) -> str: + """Stable per-process key for an expression, matching the state-bearer + convention: the class name and ``uw_object.instance_number``.""" + return f"{type(expr).__name__}_{expr.instance_number}" + + +def _capture_expressions(snap: Snapshot) -> None: + """Capture every live expression's contents. + + A run that ramps a parameter (``kappa.sym = 7.0`` between steps) used to be + restorable in its FIELDS only: the mesh variables came back and the + parameters stayed wherever the run had left them, so a replayed step solved + a different problem from the one the transcript recorded — with no error, in + the one place a wrong answer is hardest to notice. + + Contents are stored by REFERENCE, not deep-copied. ``_sym`` is a sympy + object and sympy objects are immutable, so the reference cannot go stale + under a later ``sym =`` assignment: that assignment REPLACES the object + rather than mutating it. Deep-copying instead would be actively wrong here — + an expression's ``_sym`` can carry mesh-variable symbols, and cloning that + graph would detach the restored parameter from the live mesh. + """ + from underworld3.function.expressions import live_expressions + + for expr in live_expressions(): + try: + snap.expressions.append( + (_expression_key(expr), expr._sym, expr._wrapped) + ) + except AttributeError: + # Not every expression subclass carries both slots; one that does + # not is one whose value is derived, and deriving it again is right. + continue + + +def _restore_expressions(snap: Snapshot) -> None: + """Put captured expression contents back. + + Matched by key against the expressions alive NOW. Two cases are deliberately + quiet rather than fatal, because neither means the restore is wrong: + + * captured but no longer alive — the expression was dropped since; there is + nothing to write to; + * alive but not captured — it was created after the snapshot, so the + snapshot has no opinion about what it should hold. + + Both differ from the state-bearer rule above, which raises: a missing + state-bearer means the snapshot came from a different Model, whereas + expressions are created and dropped freely throughout a normal run. + """ + if not snap.expressions: + return + from underworld3.function.expressions import live_expressions + + live = {_expression_key(e): e for e in live_expressions()} + for key, captured_sym, captured_wrapped in snap.expressions: + expr = live.get(key) + if expr is None: + continue + expr._sym = captured_sym + expr._wrapped = captured_wrapped + + def _capture_swarm(snap: Snapshot, swarm) -> None: payload = swarm.snapshot_payload() name = payload["name"] @@ -352,6 +421,8 @@ def restore(model, snap: Snapshot) -> None: ) obj.state = copy.deepcopy(captured_state) + _restore_expressions(snap) + def _build_mesh_payload(snap: Snapshot, mesh_name: str) -> dict: return { diff --git a/src/underworld3/function/expressions.py b/src/underworld3/function/expressions.py index e34c094b1..6ca0c29c1 100644 --- a/src/underworld3/function/expressions.py +++ b/src/underworld3/function/expressions.py @@ -14,6 +14,7 @@ - UWexpression.to() simply calls uw.convert_units(self, target) """ +import weakref import sympy import numpy as np from sympy import Symbol, simplify, Number @@ -602,6 +603,29 @@ def substitute_expr(fn, sub_expr, keep_constants=True, return_self=True): # UWexpression Class - Simplified (no UWQuantity inheritance) # ============================================================================ +def live_expressions(): + """The persistent expression containers, in a stable order. + + Reads ``UWexpression._expr_names``, which is the registry that already + defines what a UW expression IS: a container with identity by name, looked + up rather than rebuilt, so the same ``uw.expression(r"\\eta", ...)`` reaches + the same object and a formula written against it keeps seeing later edits to + its contents. That is exactly the set whose CONTENTS a snapshot has to + capture — the parameters of the run. + + ``_ephemeral_expr_names`` is deliberately NOT included. Those are the + ``_unique_name_generation=True`` expressions made for derivative lowering + and template substitution: their contents are derived, they are held weakly, + and they are rebuilt from the persistent ones. Restoring them would write + over a value the machinery is about to recompute. + + A list, not the live dict: capture must not iterate a mapping that other + construction can mutate underneath it. Sorted by name so two captures of the + same state read the same way. + """ + return [UWexpression._expr_names[k] for k in sorted(UWexpression._expr_names)] + + class UWexpression(MathematicalMixin, uw_object, Symbol): """ A SymPy Symbol that wraps a value for lazy evaluation. diff --git a/tests/test_0007_snapshot_inmemory.py b/tests/test_0007_snapshot_inmemory.py index d52fee33a..68af29ed4 100644 --- a/tests/test_0007_snapshot_inmemory.py +++ b/tests/test_0007_snapshot_inmemory.py @@ -767,3 +767,115 @@ def test_continuation_bit_identical_across_stash_and_recover(): _assert_bit_identical(ctrl, stash, "stash-and-recover") + + +# --------------------------------------------------------------------------- # +# Expressions: parameters must rewind with the fields +# --------------------------------------------------------------------------- # +# A snapshot captured meshes, swarms and registered state-bearers. A parameter +# is none of those, so it was not captured at all: restore returned the FIELDS +# of the captured step and left every `uw.expression` wherever the run had since +# moved it. A replayed step then solved a different problem from the one the +# transcript recorded, with no error raised — and ramping a parameter between +# steps (`kappa.sym = ...`) is the ordinary way to write such a run. + + +def test_expression_contents_are_restored_with_the_fields(): + """The contract. Scribble on a field AND a parameter; both come back. + + The field half already held; it is asserted alongside so a regression that + breaks both cannot pass by breaking them symmetrically. + """ + import sympy + + uw, model, mesh = _fresh_model_and_mesh() + T = uw.discretisation.MeshVariable("T_expr_snap", mesh, 1, degree=2) + kappa = uw.expression(r"\kappa", 1.0, "diffusivity") + + T.array[...] = 3.0 + snap = model.save_state() + + kappa.sym = sympy.Float(7.0) + T.array[...] = -42.0 + + model.load_state(snap) + + assert float(kappa.sym) == pytest.approx(1.0), ( + f"the parameter did not rewind with the field; kappa is {kappa.sym}") + assert np.allclose(np.asarray(T.array[...]), 3.0) + + +def test_a_symbolic_expression_value_is_restored_not_just_a_number(): + """A parameter's contents can be an expression, not only a float. + + Restoring must put the whole symbolic value back. Stored by reference + rather than deep-copied, because `_sym` can carry mesh-variable symbols and + cloning that graph would detach the restored parameter from the live mesh — + so this also checks the reference has not gone stale. + """ + import sympy + + uw, model, mesh = _fresh_model_and_mesh() + x, y = mesh.X + eta = uw.expression(r"\eta", sympy.sympify(1), "viscosity") + + eta.sym = 2 + x * y + captured = eta.sym + snap = model.save_state() + + eta.sym = sympy.Float(99.0) + model.load_state(snap) + + assert eta.sym == captured + assert eta.sym.free_symbols == captured.free_symbols, ( + "the restored value lost its coordinate symbols — it was cloned rather " + "than referenced") + + +def test_an_expression_created_after_the_snapshot_is_left_alone(): + """A snapshot has no opinion about something that did not exist when it was + taken. Leaving it is right; raising would break the ordinary case of a run + that builds a new parameter after a restore point.""" + uw, model, mesh = _fresh_model_and_mesh() + snap = model.save_state() + + latecomer = uw.expression(r"\alpha", 5.0, "made after the snapshot") + model.load_state(snap) + + assert float(latecomer.sym) == pytest.approx(5.0) + + +def test_capture_reads_the_persistent_container_registry(): + """Capture is driven by the registry that defines what an expression IS — + ``_expr_names``, the by-name container store — so anything a user can reach + later by name is captured, and the ephemeral derivative/template expressions + are not.""" + from underworld3.function.expressions import UWexpression, live_expressions + + uw, model, mesh = _fresh_model_and_mesh() + mine = uw.expression(r"\beta", 2.0, "made here") + + assert any(e is mine for e in live_expressions()) + + snap = model.save_state() + captured = {key for key, _sym, _wrapped in snap.expressions} + assert f"{type(mine).__name__}_{mine.instance_number}" in captured + + # ephemerals are held in their own weak registry and deliberately skipped + ephemeral_names = {k[0] for k in list(UWexpression._ephemeral_expr_names)} + persistent_names = set(UWexpression._expr_names) + assert not (ephemeral_names & persistent_names & {r"\beta"}) + + +def test_a_reused_name_is_one_container_and_is_captured_once(): + """Identity is the NAME: asking for the same name returns the same object, + which is what lets a formula written early keep seeing later edits. Capture + must therefore record it once, not once per construction site.""" + uw, model, mesh = _fresh_model_and_mesh() + first = uw.expression(r"\gamma_{shared}", 1.0, "first use") + second = uw.expression(r"\gamma_{shared}", 2.0, "second use") + assert first is second + + snap = model.save_state() + key = f"{type(first).__name__}_{first.instance_number}" + assert [k for k, _s, _w in snap.expressions].count(key) == 1