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
71 changes: 71 additions & 0 deletions src/underworld3/checkpoint/snapshot.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)


Expand Down Expand Up @@ -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


Expand Down Expand Up @@ -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"]
Expand Down Expand Up @@ -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 {
Expand Down
24 changes: 24 additions & 0 deletions src/underworld3/function/expressions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.
Expand Down
112 changes: 112 additions & 0 deletions tests/test_0007_snapshot_inmemory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading