diff --git a/CHANGELOG.md b/CHANGELOG.md index d8e23234..5fa8d0cd 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -215,6 +215,7 @@ guidance; the itemized changes follow. (stateful machines), the params pytree, `sysid`, retracing and binary frames ### Changed +- **A step that changes the layout of the state raises `ValueError`; it used to be stored** (released behaviour, corrected: `MADD-ANO-220`): `GraphManager.step`, `run`, `run_adaptive` and `FmuSidecar.step` refuse a stepped state whose leaf has another shape or kind of dtype than the state it replaces, or whose fields differ, naming the node and the leaf; nothing is stored. A node's `update` must return the layout `initial_state()` built: build a state that grew at its first step at its final shape. - **`add_coupling_group` refuses a boolean option that is not a boolean, and `rtol=0` under the norms that divide by it**: `diagnostics`, `subcycling` and `strict_convergence` raise `TypeError` for anything but `True`/`False` (`diagnostics="off"` turned the diagnostics on); `rtol=0` under `convergence_norm="mixed"` or `"interface"` raises `ValueError` (the residual was `0/0`, the step ran to its cap and the report raised `ZeroDivisionError`). Action: pass a bool; use a positive `rtol`. - **A multi-rate graph whose schedule would not keep a node's clock no longer compiles** (MADD-ANO-227, in every release to 0.3.1): timesteps more than about 1e9 apart (`1.0` and `1e-10` got rate dividers of 1 and 0 and ran on different clocks, silently), or with no common step the float GCD finds, are a `ValueError` at `compile()` naming the two nodes, an error from `validate()`, and a 400 from `POST /graph/nodes` and `DELETE /graph/nodes/{name}`. The dividers of every graph that is kept are unchanged. Give the nodes timesteps that are whole multiples of one step. - **REST numbers are strict, and counts are held to their declared range** (MADD-ANO-228): `"timestep": true` (a node at 1.0 s), `"0.5"` and `" 0.25 "` are a 422, as are a boolean or text for any number of a request model; an integer query parameter is decimal digits only; `POST /sim/profile` answers 422 for `n_steps` outside 1..1000 or `n_warmup` outside 0..50 where it clamped them without saying so. Send JSON numbers, and counts inside the documented range. @@ -417,6 +418,7 @@ guidance; the itemized changes follow. The `[verify]` extra now only pulls `hypothesis`. ### Fixed +- **`GraphManager.step`, `run` and `run_adaptive` no longer store a step that reshapes a state leaf** (`MADD-ANO-220`, in 0.1.0 to 0.3.1): a list given for a scalar constant (`BallNode(initial_velocity=[1.0, 2.0])`), or an external input of another shape at a later step, broadcast a scalar leaf, and the state's checkpoint then did not load after `reset_state()`. The step is refused by name (compared once per trace, host-side: no compiled program changes), as `run_scan` always did; the FMU sidecar and bridge refuse it too. - **`coupling_diagnostics()`: three numbers that read wrong with their flag set, and one found beside them** (MADD-ANO-222, 225, 226, 229, never released): `gradient_relative_error_bound` at the float floor took the change along one stand-in direction (5.1e-6 for 3.0e-5); `rho_spectral` read settled past eight interface scalars (0.273 for 0.219), from sampled rounding on a float32 group far from normal (0.250 for 0.206), and lost a loop below a field's rounding on a sub-cycled group (1e-12 for 1.25e-4). The gradient bound's undirected distance now takes an operator norm; a Krylov space still growing at the cap is never settled (`spectral_usable=False` above seven interface scalars, eight where they are the whole state); a certificate over every perturbation of the measured size replaces the samples; a sub-cycled boundary value is differentiated as `(1 - alpha) a + alpha b` (values unchanged, gradients of sub-cycled groups move at rounding level). Action: none; expect `spectral_usable=False` on more float32 and 16-bit groups whose Jacobian is far from normal (usable fraction on the search's draws 0.979 to 0.976). @@ -424,7 +426,7 @@ guidance; the itemized changes follow. - **`run_pod.py --summarise` reads a goal file with no cells or no checks INVALID (exit 3) for every goal** (never released): an emptied `forward.json`, `gradient.json` or `exchange.json` read "no checks" and the summary exited 0. - **A constructor's refusal at `POST /graph/nodes` names the parameter** it can be told from, the value sent and the class's default (`n_cells: 8.0` was "'float' object cannot be interpreted as an integer"). - **An FMU parameter whose `ParamSpec` accepts no value advertises an empty `min` / `max`** (never released; the spec-level refusals are 0.4.0's own): a `logit` spec with a subnormal float32 bound, or a `log` / `logit` bound or width the leaf's dtype does not hold, is refused by `check_params`, the sidecar and REST for every value, while `build_model_description` advertised the neighbours of its bounds (`min = 1.18e-38`, `max` just under `1.42e-14` for `logit` on `(1.4e-45, 1.42e-14)`), so a bridge over a sidecar built without `param_specs` took every value between them. The variable now advertises `min = tiny > max = -tiny`, and the bridge refuses to serve the description, with or without specs. Action: none; narrow the bounds as the refusal says. -- **`POST /graph/nodes` accepts only a node the graph can step with its state as built** (the REST door of `MADD-ANO-220`, which stays open in process: `GraphManager.step` stores a step that changes a leaf's shape; use `run_scan`, which refuses it): one update is traced as the graph calls it (params pytree, with and without boundary inputs) and held to `initial_state()`'s fields, shapes and kinds of dtype; `venous_pressure: null` and a list for a scalar constant are a 400 naming the parameter. Also: lists and objects count against the 10^6 params bound; `POST /checkpoint/load` names why a graph cannot compile; `POST /surrogate/deactivate` (experimental) lists the edges it drops (`dropped_edges`). +- **`POST /graph/nodes` accepts only a node the graph can step with its state as built** (the REST door of `MADD-ANO-220`; in process the step itself refuses such a node, see the entry above): one update is traced as the graph calls it (params pytree, with and without boundary inputs) and held to `initial_state()`'s fields, shapes and kinds of dtype; `venous_pressure: null` and a list for a scalar constant are a 400 naming the parameter. Also: lists and objects count against the 10^6 params bound; `POST /checkpoint/load` names why a graph cannot compile; `POST /surrogate/deactivate` (experimental) lists the edges it drops (`dropped_edges`). - **`run_pod.py` refuses an option value no goal can use with exit 2, before the backend loads** (never released): `--cells 0`, `--warmup -1`, `--repeats 0`, `--steps 0`, a `--mesh` that cannot be read. They used to pass, or exit 5 as a crashed goal. - **`rho_spectral` and `spectral_usable` under `diagnostics=True`** (MADD-ANO-221, 223, 224, never released): the spectral estimate took a Krylov direction below 1e-5 of a product for rounding in every dtype, took the compressed radius by repeated squaring, and called a spectrum settled on the Arnoldi residual alone; `rho_spectral` read 0.379 for 0.5 (float64), 0.0512 for 0.0500 (float32) and 0.941 for 0.735 (twelve non-normal scalars), and `spectral_error_bound` 0.61x the distance, each with `spectral_usable=True`. The breakdown test is now eight units of the products' own rounding, the radius is `eigvals`, and one more Jacobian-vector product (nine per group per step, was eight) checks the estimate against itself: where it moves the radius by more than 5% of `1 - rho_spectral` the flag is `False`. diff --git a/docs/developer_guide/node_authoring.md b/docs/developer_guide/node_authoring.md index 723c57d9..8a0fc95f 100644 --- a/docs/developer_guide/node_authoring.md +++ b/docs/developer_guide/node_authoring.md @@ -9,6 +9,7 @@ A MADDENING {term}`node ` is a **{term}`pure function ` wra - `initial_state(params) -> dict` — returns the initial state arrays - `update(state, boundary_inputs, dt) -> new_state` — returns a new state dict - `update()` must be **{term}`JAX-traceable`**: no Python-level side effects, no data-dependent control flow, no print statements. Use `jnp.where` instead of `if/else`. +- `update()` returns the state layout it was given: the same fields, each with the shape and kind of dtype `initial_state()` built. A step that changes one (a list passed for a scalar constant broadcasts a scalar field) is refused by `GraphManager.step`, `run` and `run_adaptive` with a `ValueError` naming the node and the field, as `run_scan` refuses it. - State is **immutable** — return a new dict, don't mutate in place - Parameters live in `self.params`, not in state diff --git a/docs/developer_guide/sharding_topology.md b/docs/developer_guide/sharding_topology.md index 015df6c7..2f43924a 100644 --- a/docs/developer_guide/sharding_topology.md +++ b/docs/developer_guide/sharding_topology.md @@ -253,12 +253,14 @@ What 0.4.0 added: integral can leave out the padding of a short shard's block; * refusals of a node declaring a Cartesian `halo_width()` and of a node whose cell count is not the layout's (both silently wrong before); -* a domain integral carried in the state (a graph's second step, or an - initial value the node declares) is placed as the step returns it — - replicated once reduced, stacked along the mesh axis otherwise — rather - than partitioned, so such a node runs in a `GraphManager` (the stencil - wrapper had the same defect for a vector or per-shard integral, fixed the - same way); +* a domain integral carried in the state (an initial value the node + declares, or a step's output handed back to the wrapper) is placed as + the step returns it — replicated once reduced, stacked along the mesh + axis otherwise — rather than partitioned, so such a node runs in a + `GraphManager` (the stencil wrapper had the same defect for a vector or + per-shard integral, fixed the same way). In a graph the node gives the + integral an initial value in `initial_state()`: a step that adds a field + to the state is refused (`GraphManager.step`, as `run_scan`); * the session runner `benchmarks/multigpu/run_pod.py` (and its CPU `--dry-run`). diff --git a/docs/release_notes/v0.4.0.md b/docs/release_notes/v0.4.0.md index 98d75ebd..f6502d0b 100644 --- a/docs/release_notes/v0.4.0.md +++ b/docs/release_notes/v0.4.0.md @@ -7148,8 +7148,7 @@ fixed](#the-known-coupling-sysid-and-fmu-findings-fixed)". - **MADD-ANO-217** (**open**; carried by no release before this one): `fit_lm` from a stiffness started 18 to 80 times too high can stop short, `converged=False`, on the end of a wide damping range; start it again from the returned parameters. -`MADD-ANO-220` — *GraphManager.step, run and run_adaptive store a step that -changed the shape of a state leaf* — **open**. +- **MADD-ANO-220** (resolved; shipped in every tagged release): `step`, `run` and `run_adaptive` stored a step that changed the shape of a state leaf (a list given for a scalar constant), after which the state's checkpoint did not reload; such a step now raises `ValueError` and stores nothing. `MADD-ANO-230` — *A direction the Arnoldi breakdown test takes for rounding can carry the dominant mode of a Jacobian far from normal: float32 fan-out hubs diff --git a/docs/validation/known_anomalies.yaml b/docs/validation/known_anomalies.yaml index 7641eb68..2a57916b 100644 --- a/docs/validation/known_anomalies.yaml +++ b/docs/validation/known_anomalies.yaml @@ -4287,7 +4287,7 @@ anomalies: - "tests/cloud/multigpu/test_integral_listed_as_state_field.py::test_an_integral_listed_in_state_fields_is_summed_over_the_shards" - "tests/cloud/multigpu/test_unstructured_layout_contract.py::test_a_domain_integral_fed_back_as_state_steps_again" - "tests/cloud/multigpu/test_unstructured_layout_contract.py::test_a_node_declaring_its_integral_runs_in_a_graph_like_the_unsharded_node" - - "tests/cloud/multigpu/test_unstructured_layout_contract.py::test_a_node_without_an_initial_integral_steps_in_a_graph" + - "tests/cloud/multigpu/test_unstructured_layout_contract.py::test_a_node_without_an_initial_integral_is_handed_its_own_output_back" - "tests/cloud/multigpu/test_stencil_domain_integral_state.py::test_a_vector_domain_integral_in_the_state_runs_in_a_graph_like_the_unsharded_node" - "tests/cloud/multigpu/test_stencil_domain_integral_state.py::test_a_per_shard_domain_integral_fed_back_by_a_graph_steps_again" - "tests/cloud/multigpu/test_stencil_domain_integral_state.py::test_the_initial_value_of_an_integral_is_placed_as_the_step_returns_it" @@ -14929,12 +14929,23 @@ anomalies: cannot use (`HeartPumpNode` `venous_pressure: null`), after which every `POST /sim/step` answered 400, is refused there too. - In process the step entry points are unchanged. A check of the - stepped state's keys and shapes in `step` (once per compile, - host-side) was written and withdrawn: it refuses the fixture the - compile-count gate's own tests use to provoke a retrace - (`tests/core/test_compile_counts.py`, a state that grows on its - first step), so it waits for a decision on that fixture. + Resolved in 0.4.0 in process too. `step` and `run` compare the + stepped state with the state it replaces -- keys, shapes and kinds + of dtype (boolean, integer, float, complex) -- once per trace of the + compiled step, host-side, and `run_adaptive` compares the first + step it keeps; a difference is a `ValueError` naming the node and + the leaf, and nothing is stored (the graph holds the state it held, + no observer is told of a step, no callback is called). A step that + is traced again (a new compile; an external input of another shape, + which broadcast a scalar leaf at a later step) is compared again. + A leaf whose kind of dtype changed (an integer count returned as a + float) was stored the same way and is refused the same way. The + FMU sidecar and bridge, which run the compiled step on a state of + their own, make the same comparison (`FmuSidecar.step`; an error + reply from the bridge, nothing committed). Nothing is traced and + no node's `update` is called for it: the compiled programs are + unchanged. A dtype's width is not compared (with x64 enabled a + stock node returns float64 for a float32 state). severity: "major" safety_relevance: "context_dependent" safety_relevance_rationale: > @@ -14947,22 +14958,39 @@ anomalies: - "maddening.core.graph_manager.GraphManager.step" - "maddening.core.graph_manager.GraphManager.run" - "maddening.core.graph_manager.GraphManager.run_adaptive" - affected_versions: ">=0.1.0" + - "maddening.fmi.sidecar.FmuSidecar.step" + affected_versions: ">=0.1.0, <0.4.0" workaround: > - Give a node's constants the rank its state has (a scalar for a scalar - field), and advance a graph you have not checked with `run_scan`, - which refuses a step that changes the state's layout. Over REST, - 0.4.0 refuses such a node where it is added. - resolution_status: "open" - resolution_version: null - residual_risk: null + On 0.1.0 to 0.3.1 give a node's constants the rank its state has (a + scalar for a scalar field), and advance a graph you have not checked + with `run_scan`, which refuses a step that changes the state's + layout. + resolution_status: "resolved" + resolution_version: "0.4.0" + residual_risk: > + Changed behaviour: a step that changes the shape or the kind of + dtype of a state leaf, or the fields of a node's state, raises + `ValueError` where it was stored; a graph that relied on its state + growing at the first step must build the state at its final layout + in `initial_state()`. That includes a node whose `update` returns a + field `initial_state()` does not build (a domain integral it only + emits): `step` stored the first step of such a node, sharded or + not, and now refuses it; the node gives the field an initial value. + A dtype of another width of the same kind is still stored as + returned. verification: - # A strict xfail: a change that makes the step entry points keep or - # refuse the layout is noticed. - "tests/core/test_a_step_keeps_the_state_layout.py::test_the_step_entry_points_keep_the_shape_of_every_state_leaf" + - "tests/core/test_a_step_keeps_the_state_layout.py::test_a_step_that_changes_a_leafs_kind_of_dtype_or_the_fields_is_refused" + - "tests/core/test_a_step_keeps_the_state_layout.py::test_a_later_step_that_is_traced_again_is_compared_again" + - "tests/core/test_a_step_keeps_the_state_layout.py::test_the_step_of_a_graph_compiled_again_is_compared_again" + - "tests/core/test_a_step_keeps_the_state_layout.py::test_the_comparison_is_made_once_per_trace_and_calls_no_update" - "tests/core/test_a_step_keeps_the_state_layout.py::test_the_scan_entry_points_refuse_a_step_that_changes_a_leafs_shape" + - "tests/fmi/test_a_sidecar_step_keeps_the_state_layout.py::test_a_sidecar_step_that_reshapes_a_leaf_is_refused_and_nothing_is_committed" + - "tests/cloud/multigpu/test_a_sharded_step_keeps_the_state_layout.py::test_a_node_beside_a_sharded_one_that_reshapes_its_state_is_refused" - "tests/api/test_a_new_node_is_one_the_graph_can_step.py::test_a_value_of_a_rank_that_reshapes_the_state_at_the_first_step_is_a_400" + - "tests/api/test_a_new_node_is_one_the_graph_can_step.py::test_a_node_the_door_never_saw_is_refused_where_the_graph_steps" - "tests/property/test_rest_requests_generated_from_the_schema.py::test_a_new_node_of_any_kind_is_refused_or_steps_with_its_layout_and_reloads" + - anomaly_id: "MADD-ANO-221" title: "The Arnoldi breakdown test was 1e-5 of the product in every dtype: rho_spectral read 0.379 for 0.5 and spectral_error_bound 0.61x the true distance, spectral_usable=True" description: > diff --git a/docs/validation/rest_runpod_claims.yaml b/docs/validation/rest_runpod_claims.yaml index c7d6b729..120db959 100644 --- a/docs/validation/rest_runpod_claims.yaml +++ b/docs/validation/rest_runpod_claims.yaml @@ -2295,7 +2295,10 @@ claims: a configuration error found when the step is traced (an edge of mismatched shapes, constraints the step cannot use), and a run-time check in the step (equinox.error_if) raising at the first step, or in the sixth doubling - slice of a run after 32 steps were stored; a step of an edited graph + slice of a run after 32 steps were stored; a step that changes the layout + of the state (a node the graph was given in process whose scalar leaf + broadcasts, which the step stored before 0.4.0: MADD-ANO-220), refused by + the graph on /sim/step and /sim/run; a step of an edited graph whose compile refuses; a failure injected at every point of a step's body. Not claimed for simultaneous requests. @@ -2303,6 +2306,7 @@ claims: tests: - tests/api/test_params_writes_the_step_cannot_run_with.py::test_a_graph_configured_so_it_cannot_step_is_a_400_naming_why - tests/api/test_runner_pacing_and_status.py::test_an_edge_the_graph_cannot_validate_is_a_400_on_every_route_that_steps + - tests/api/test_a_new_node_is_one_the_graph_can_step.py::test_a_node_the_door_never_saw_is_refused_where_the_graph_steps - tests/api/test_routes_answer_what_they_did.py::test_a_single_step_that_raises_is_a_400_and_stores_nothing - tests/api/test_routes_answer_what_they_did.py::test_a_run_that_raises_at_its_first_step_is_a_400_and_moves_nothing - tests/api/test_routes_answer_what_they_did.py::test_a_run_that_raises_part_way_says_how_many_steps_it_took diff --git a/docs/validation/soup_package.md b/docs/validation/soup_package.md index 404aa7bc..225a57ce 100644 --- a/docs/validation/soup_package.md +++ b/docs/validation/soup_package.md @@ -282,7 +282,7 @@ stale copy fails CI rather than shipping. | MADD-ANO-217 | fit_lm from a stiffness started far too high can stop short, unconverged, on the end of a wide damping range where the loss still falls | `minor` | `not_safety_relevant` | `open` | >=0.4.0.dev0 | | MADD-ANO-218 | A value handed to step or a scan ran in the dtype it arrived in, not the external input's declared dtype: an x64 graph with a default-declared (float32) input ran 0.1 as a float64 where its own FMU ran float32(0.1) | `major` | `context_dependent` | `resolved` (in 0.4.0) | none | | MADD-ANO-219 | build_model_description(default_step_size=) wrote stepSize="np.float64(0.05)": a modelDescription.xml no FMI importer reads, with nothing said at export | `minor` | `not_safety_relevant` | `resolved` (in 0.4.0) | >=0.3.0, <0.4.0 | -| MADD-ANO-220 | GraphManager.step, run and run_adaptive store a step that changed the shape of a state leaf | `major` | `context_dependent` | `open` | >=0.1.0 | +| MADD-ANO-220 | GraphManager.step, run and run_adaptive store a step that changed the shape of a state leaf | `major` | `context_dependent` | `resolved` (in 0.4.0) | >=0.1.0, <0.4.0 | | MADD-ANO-221 | The Arnoldi breakdown test was 1e-5 of the product in every dtype: rho_spectral read 0.379 for 0.5 and spectral_error_bound 0.61x the true distance, spectral_usable=True | `major` | `context_dependent` | `resolved` (in 0.4.0) | none | | MADD-ANO-222 | A coupling loop that passes through a change of a field below that field's float rounding is not in rho_spectral: a float32 sub-cycled ring with a field 1e-6 of its driver reads 1e-12 for 1.25e-4, spectral_usable=True | `major` | `context_dependent` | `resolved` (in 0.4.0) | none | | MADD-ANO-223 | The spectral radius of the compressed Jacobian was taken by repeated squaring, which is not backward stable for a non-normal matrix: a float32 ring of five read rho_spectral 0.0512 for 0.0500, and graded matrices of five 0.9555 for 0.9152, spectral_usable=True | `major` | `context_dependent` | `resolved` (in 0.4.0) | none | @@ -294,7 +294,7 @@ stale copy fails CI rather than shipping. | MADD-ANO-227 | Nodes whose timesteps are more than about 1e9 apart, or have no common step the float GCD finds, ran on different clocks with no error | `major` | `context_dependent` | `resolved` (in 0.4.0) | >=0.1.0, <0.4.0 | | MADD-ANO-228 | REST read a boolean or numeric text as a timestep (`true` added a node stepping at 1.0 s), and clamped the counts of POST /sim/profile without saying so | `major` | `context_dependent` | `resolved` (in 0.4.0) | >=0.1.0, <0.4.0 | -*230 anomalies registered. 38 have a defect reachable in this version — every entry whose `resolution_status` is not `resolved` or `duplicate`, which is 29 `open` plus 9 `partially_resolved` whose residual risk is still live. The Affected Versions column is a PEP 440 specifier set read against this document's version; `none` marks a defect introduced and fixed within one development cycle, which no release carried. The convention, and the gate that holds every range to it, are in the header of `known_anomalies.yaml`. Rationale, workaround, affected components and verification evidence for each: `known_anomalies.yaml`.* +*230 anomalies registered. 37 have a defect reachable in this version — every entry whose `resolution_status` is not `resolved` or `duplicate`, which is 28 `open` plus 9 `partially_resolved` whose residual risk is still live. The Affected Versions column is a PEP 440 specifier set read against this document's version; `none` marks a defect introduced and fixed within one development cycle, which no release carried. The convention, and the gate that holds every range to it, are in the header of `known_anomalies.yaml`. Rationale, workaround, affected components and verification evidence for each: `known_anomalies.yaml`.* ## 4. Verification Evidence diff --git a/src/maddening/api/server.py b/src/maddening/api/server.py index 4ae393bb..e31a70d2 100644 --- a/src/maddening/api/server.py +++ b/src/maddening/api/server.py @@ -135,6 +135,7 @@ _hook_outputs, _leaf_values_equal, _node_update_layout_drift, + _prng_key_leaves, _node_with_params, _params_holders, ) @@ -1051,9 +1052,9 @@ def _dry_run_node(node, state: Any = None) -> None: pytree (``HeartPumpNode`` ``venous_pressure: null``) was accepted and every later step was a 400; and a list where the node's state is a scalar (``BallNode`` ``initial_velocity: [0, 0]``) was accepted, the - state changed shape at the first step (``GraphManager.step`` stores - such a state: MADD-ANO-220), and the graph's own checkpoint no longer - loaded after a reset. + state changed shape at the first step (which ``GraphManager.step`` + then stored, and now refuses: MADD-ANO-220), and the graph's own + checkpoint no longer loaded after a reset. *state* is the node's initial state when the caller has already built it (the size check does), so it is not allocated twice. @@ -1063,6 +1064,15 @@ def _dry_run_node(node, state: Any = None) -> None: spec = _NodeSpec(node=node, update_fn=node.update, timestep=node.delta_t, accepts_params=accepts) leaves = node.params_pytree() if accepts else None + # A PRNG key in the state has no JSON form: every reply that carries + # the state (``POST /sim/step``) would fail. Refused here by name (it + # used to be refused by accident, as an internal error of the layout + # comparison, which now compares a key as a kind of its own). + keyed = _prng_key_leaves(state) + if keyed: + raise ValueError( + f"its state holds a PRNG key ({', '.join(map(repr, keyed))}), which has no " + "JSON form, so no reply could carry the state") # With every declared input delivered (the node once its edges are # added), then with none, as the graph steps it until then: a constant # read only in the absence of an input is read by every step of the diff --git a/src/maddening/core/_param_probes.py b/src/maddening/core/_param_probes.py index 175b0c27..c44a6ee8 100644 --- a/src/maddening/core/_param_probes.py +++ b/src/maddening/core/_param_probes.py @@ -209,9 +209,20 @@ def _leaf_layout(leaf: Any) -> tuple[tuple, Any]: dtype = getattr(leaf, "dtype", None) if dtype is None: dtype = jnp.result_type(leaf) + if jnp.issubdtype(dtype, jax.dtypes.extended): + # A PRNG key (``key``): not a NumPy dtype, and a kind of its own. + return shape, dtype return shape, np.dtype(jax.dtypes.canonicalize_dtype(dtype)) +def _dtype_kind(dtype: Any) -> Any: + """The kind of a :func:`_leaf_layout` dtype: NumPy's one-letter kind + (boolean, integer, float, complex), and for an extended dtype -- a PRNG + key, which has none -- the dtype itself, so a key is a key of the same + implementation and nothing else.""" + return getattr(dtype, "kind", dtype) + + def _leaf_path(path: tuple) -> str: """``node/field`` for a pytree key path.""" parts = [] @@ -225,19 +236,30 @@ def _leaf_path(path: tuple) -> str: return "/".join(parts) +def _prng_key_leaves(state: Any) -> list[str]: + """The leaves of *state* that are PRNG keys (an extended dtype), by + path. The REST server refuses a new node that has one: a key has no + JSON form, so no reply could carry the node's state.""" + return [_leaf_path(path) + for path, leaf in jax.tree_util.tree_flatten_with_path(state)[0] + if jnp.issubdtype(_leaf_layout(leaf)[1], jax.dtypes.extended)] + + def _state_layout_drift(before: Any, after: Any, *, dtypes: bool = True) -> list[str]: """How the state layout *after* differs from *before*: one line per leaf that is missing, new, of another shape or (with *dtypes*) of - another kind of dtype (boolean, integer, float, complex); empty when + another kind of dtype (boolean, integer, float, complex, or a PRNG + key of one implementation); empty when the two trees have the same layout. A dtype's width is not compared: with x64 enabled every stock node's update returns ``float64`` for the ``float32`` its ``initial_state()`` builds, and the state reloads. The comparison behind the REST server's dry run of a new node - (``POST /graph/nodes``): a node whose ``update`` returns a leaf of - another shape than its ``initial_state()`` built broadcasts it at the - first step, and :meth:`GraphManager.step` stores the result - (MADD-ANO-220). Host-side: it reads shapes and dtypes, never values, + (``POST /graph/nodes``) and behind the check of a stepped state + (:meth:`GraphManager.step`, ``FmuSidecar.step``): a node whose + ``update`` returns a leaf of another shape than its + ``initial_state()`` built broadcasts it at the first step, which was + stored (MADD-ANO-220). Host-side: it reads shapes and dtypes, never values, so tracers and the abstract values of :func:`jax.eval_shape` compare like arrays. """ @@ -253,7 +275,7 @@ def _state_layout_drift(before: Any, after: Any, *, dtypes: bool = True) -> list if shape_before != shape_after: found.append(f"{name!r} has shape {shape_before} before the update " f"and {shape_after} after it") - elif dtypes and dtype_before.kind != dtype_after.kind: + elif dtypes and _dtype_kind(dtype_before) != _dtype_kind(dtype_after): found.append(f"{name!r} has dtype {dtype_before} before the update " f"and {dtype_after} after it") found.extend(f"{name!r} appears only after the update" for name in new if name not in old) diff --git a/src/maddening/core/graph_manager.py b/src/maddening/core/graph_manager.py index 33d19c54..bf83d2a1 100644 --- a/src/maddening/core/graph_manager.py +++ b/src/maddening/core/graph_manager.py @@ -252,6 +252,11 @@ def __init__(self) -> None: # already warned about, which are never warned about again. self._underflow_check_pending = False self._underflow_warned: set[str] = set() + # The state-layout check of a stepped state + # (``_store_stepped_state``): which trace of which compile the + # last compared step was (``None`` until one has been compared), + # so a new compile and a retrace each bring the next comparison. + self._layout_checked_trace: Optional[tuple[int, int]] = None # Escaped-tracer bookkeeping; see ``_recover_from_escaped_tracers``. self._state_traced = False self._state_before_trace: Optional[dict] = None @@ -3804,6 +3809,54 @@ def _resolve_external_inputs( # Escaped tracers # ------------------------------------------------------------------ + def _store_stepped_state(self, new_state: dict) -> None: + """:meth:`_store_state` for the result of one compiled step, which + must have the layout of the state it replaces. + + A node whose ``update`` returns a leaf of another shape -- a + constant given as a list where the node's state is a scalar, an + edge or an external input delivering a vector into a scalar field + -- broadcast the leaf at that step, and :meth:`step` and + :meth:`run` stored the result without a word: the state no longer + had the layout of ``initial_state()``, so a checkpoint saved after + the step did not load into the graph after :meth:`reset_state` or + into the graph rebuilt from :meth:`to_dict`. :meth:`run_scan` and + its siblings have always refused it, loudly, as a scan carry of + another type. Refused here by name, and nothing is stored. + + Compared once per trace of the step: a traced step is one + program, so the layout it returns is fixed by the layout it is + given, and the next comparison is due when the step is traced + again (a new compile; an external input of another shape). The + comparison is of keys, shapes and kinds of dtype, host-side: no + trace, no call of a node's ``update``, and nothing is added to + the step. Every other step costs one tuple comparison. + """ + trace = (self._compile_generation, self._n_traces) + if self._layout_checked_trace != trace: + self._refuse_layout_drift(new_state) + self._layout_checked_trace = trace + self._store_state(new_state) + + def _refuse_layout_drift(self, new_state: dict) -> None: + """Raise if *new_state* has not the keys, shapes and kinds of + dtype of the state the graph holds (see + :meth:`_store_stepped_state`). A dtype's width is not compared: + with x64 enabled a stock node's update returns ``float64`` for the + ``float32`` its ``initial_state()`` builds.""" + drift = _param_probes._state_layout_drift(self._state, new_state) + if drift: + raise ValueError( + "one step changes the layout of the graph's state, so the " + "stepped state could not be saved and reloaded, reset or " + "scanned: " + "; ".join(drift) + " (leaves are named " + "node/field). A node's update() must return its state " + "with the shapes and kinds of dtype initial_state() gave " + "it: check the node's constructor values and the shapes " + "its edges and external inputs deliver. Nothing was " + "stored." + ) + def _store_state(self, new_state: dict) -> None: """Write *new_state* back, remembering whether it is traced. @@ -4893,7 +4946,7 @@ def step( external_inputs = self._resolve_external_inputs(external_inputs) params = self._params_or_default(params) - self._store_state(self._call_surfacing_strict( + self._store_stepped_state(self._call_surfacing_strict( step_fn, self._state, external_inputs, params)) user_state = self._user_state(self._state) self._notify(EVENT_STEP, user_state) @@ -4954,7 +5007,7 @@ def run( params = self._params_or_default(params) for i in range(n_steps): - self._store_state(self._call_surfacing_strict( + self._store_stepped_state(self._call_surfacing_strict( step_fn, self._state, external_inputs, params)) user_state = self._user_state(self._state) self._notify(EVENT_STEP, user_state) @@ -5874,6 +5927,13 @@ def run_adaptive( # Use the more accurate (half-step) result -- unless a # solve the step keeps did not converge. _reports._raise_if_a_kept_solve_failed(strict_messages, (verdicts_1, verdicts_2)) + # The first step the run keeps is compared with the state + # the graph holds (``_refuse_layout_drift``), before any + # warning about it, and before a callback or an observer is + # handed it: every later one comes from the same program, + # given the layout this one has. + if n_steps == 0: + self._refuse_layout_drift(state_half) if forced: warnings.warn( f"Adaptive stepper hit dt_min={dt_min} at t={t:.6g} " diff --git a/src/maddening/fmi/sidecar.py b/src/maddening/fmi/sidecar.py index 53b83a4a..238f2644 100644 --- a/src/maddening/fmi/sidecar.py +++ b/src/maddening/fmi/sidecar.py @@ -70,6 +70,7 @@ from maddening.core.compliance.metadata import StabilityLevel from maddening.core.compliance.stability import stability from maddening.core._exact_integers import lost_as_integer +from maddening.core._param_probes import _state_layout_drift from maddening.core.params import check_bounds from maddening.fmi.directional_derivatives import ( DirectionalDerivativeKind, @@ -559,6 +560,9 @@ def __init__(self, config: SidecarConfig) -> None: #: spec, so the accepted values are the intersection of the two. self._advertised: dict[tuple[str, str], Any] = {} self._input_resolver = config.input_resolver + #: Which trace of the graph's compiled step the last layout + #: comparison was made on (:meth:`_refuse_layout_drift`). + self._layout_checked_trace: Optional[int] = None def _refuse_new_values_for(self, fixed: Mapping[str, str]) -> None: """Add ``{".params.": reason}`` to the parameters whose @@ -681,8 +685,53 @@ def _advanced(self, state: dict[str, dict[str, Any]], if self._input_resolver is not None: external_inputs = self._input_resolver(external_inputs) if self._params is None: - return self._config.step_fn(state, external_inputs) - return self._config.step_fn(state, external_inputs, self._params) + advanced = self._config.step_fn(state, external_inputs) + else: + advanced = self._config.step_fn(state, external_inputs, self._params) + self._refuse_layout_drift(state, advanced) + return advanced + + def _refuse_layout_drift(self, state: Any, advanced: Any) -> None: + """Raise if one step gave *advanced* other keys, shapes or kinds of + dtype than *state* has, as ``GraphManager.step`` does. + + The sidecar runs the graph's compiled step on a state of its own, + so the graph's check never sees it: a node that broadcasts a state + leaf at its first step (a list given for a scalar constant) left + the sidecar holding a state that is not the one its model + description declares. Nothing is committed: :meth:`step` and the + bridge commit only what this returns. + + A graph's compiled step is compared once per trace (the layout a + traced program returns is fixed by the layout it is given); any + other callable, and the step of a graph compiled again since, is + compared at every step. Host-side: shapes and dtypes are read, + nothing is traced or computed. + + Raises + ------ + ValueError + Naming each leaf (``node/field``) and how it differs. + """ + trace = None + owner = _step_compile(self._config.step_fn) + if owner is not None and owner[0] is not None: + graph, generation = owner + if getattr(graph, "_compile_generation", None) == generation: + trace = getattr(graph, "_n_traces", None) + if trace is not None and trace == self._layout_checked_trace: + return + drift = _state_layout_drift(state, advanced) + if drift: + raise ValueError( + "one step changes the layout of the state, which is then not the " + "state the model description declares: " + "; ".join(drift) + + " (leaves are named node/field). A node's update() must return " + "its state with the shapes and kinds of dtype initial_state() gave " + "it: check the node's constructor values and the shapes its edges " + "and inputs deliver. Nothing was advanced." + ) + self._layout_checked_trace = trace def get_params(self) -> dict[str, Any]: """``{".params.": value}`` for every FMI ``parameter`` diff --git a/tests/api/test_a_new_node_is_one_the_graph_can_step.py b/tests/api/test_a_new_node_is_one_the_graph_can_step.py index 3d248580..22c10c2e 100644 --- a/tests/api/test_a_new_node_is_one_the_graph_can_step.py +++ b/tests/api/test_a_new_node_is_one_the_graph_can_step.py @@ -15,6 +15,7 @@ os.environ.setdefault("JAX_PLATFORMS", "cpu") +import jax import jax.numpy as jnp import pytest @@ -80,7 +81,19 @@ def update(self, state, boundary_inputs, dt): return {"x": state["x"]} +class DrawsFromAKey(SimulationNode): + """Carries a PRNG key in its state.""" + + def initial_state(self): + return {"key": jax.random.key(0), "x": jnp.zeros((), jnp.float32)} + + def update(self, state, boundary_inputs, dt): + key, sub = jax.random.split(state["key"]) + return {"key": key, "x": state["x"] + dt * jax.random.normal(sub)} + + REGISTRY = {cls.__name__: cls for cls in ( + DrawsFromAKey, BallNode, HeartPumpNode, RigidBody2DNode, RigidBodyNode, SpringDamperNode, TableNode, NeedsItsInput, CountsInIntegers, DropsAField, ReadsThePytreeOnly)} @@ -207,3 +220,34 @@ def test_a_node_that_reads_its_constants_from_the_params_pytree_only_is_added(se # ... and the same class with a gain the update cannot multiply by. resp = _add(client, "ReadsThePytreeOnly", {"gain": [1.0, 2.0]}, name="other") assert resp.status_code == 400 and "params.gain" in resp.json()["detail"] + + +def test_a_node_the_door_never_saw_is_refused_where_the_graph_steps(tmp_path): + """The dry run guards ``POST /graph/nodes``. A graph the server was + handed, or one edited in process, does not pass that door: its step is + held to the state's layout by the graph itself (``GraphManager.step`` + and ``run``), which both step routes answer as a 400 with nothing + stored.""" + gm = GraphManager() + gm.add_node(SpringDamperNode("s", DT, stiffness=30.0, damping=2.0, initial_position=1.0)) + gm.add_node(BallNode("b", DT, initial_velocity=[1.0, 2.0])) + gm.compile() + server = SimulationServer(node_registry=REGISTRY, graph_manager=gm, + checkpoint_root=str(tmp_path)) + with TestClient(server.create_app(), raise_server_exceptions=False) as client: + held = gm._state + for resp in (client.post("/sim/step"), client.post("/sim/run", params={"n_steps": 3})): + assert resp.status_code == 400, resp.text + assert "'b/position' has shape () before the update and (2,) after it" in resp.text + assert gm._state is held + + +def test_a_node_whose_state_holds_a_prng_key_is_refused_by_name(served): + """A key has no JSON form, so no reply could carry the state. The + refusal used to be an accident of the layout comparison (an internal + error about an extended dtype); the comparison now takes a key, and + the door says what it refuses.""" + client, gm = served + resp = _add(client, "DrawsFromAKey", {}) + _assert_refused_whole(client, gm, resp, "PRNG key", "'key'", "no JSON form") + assert "canonicalize_dtype" not in resp.text diff --git a/tests/cloud/multigpu/test_a_sharded_step_keeps_the_state_layout.py b/tests/cloud/multigpu/test_a_sharded_step_keeps_the_state_layout.py new file mode 100644 index 00000000..08c5b4aa --- /dev/null +++ b/tests/cloud/multigpu/test_a_sharded_step_keeps_the_state_layout.py @@ -0,0 +1,62 @@ +"""A graph with a sharded node is held to its state's layout at a step, +like any other (``GraphManager._store_stepped_state``, MADD-ANO-220). + +The stepped state of a sharded node is compared as it is stored: leaves +placed over the mesh have their global shapes, which a step keeps. A +node beside the sharded one that broadcasts a leaf at its first step is +refused by name, and the sharded state is the one the graph held. +""" + +from __future__ import annotations + +import jax +import numpy as np +import pytest + +from maddening.cloud.multigpu.device_mesh import create_device_mesh +from maddening.cloud.multigpu.sharded_node import ShardedStencilNode +from maddening.core.graph_manager import GraphManager +from maddening.nodes import BallNode +from maddening.nodes.heat import HeatNode + +DT = 1e-4 + +pytestmark = pytest.mark.skipif(len(jax.devices()) < 4, reason="needs >=4 devices") + + +def _graph(velocity) -> GraphManager: + gm = GraphManager() + rod = HeatNode("rod", DT, n_cells=32, thermal_diffusivity=0.1, initial_temperature=300.0) + gm.add_node(ShardedStencilNode(rod, create_device_mesh(n_devices=4), + axis_map={"devices": 0})) + gm.add_node(BallNode("b", DT, initial_velocity=velocity)) + gm.compile() + return gm + + +def _layout(gm: GraphManager) -> dict: + return {f"{node}/{key}": (np.shape(leaf), np.asarray(leaf).dtype.kind) + for node, fields in gm._state.items() for key, leaf in fields.items()} + + +def test_a_healthy_sharded_graph_steps_and_keeps_its_layout(): + gm = _graph(1.0) + layout = _layout(gm) + gm.step() + gm.run(2) + gm.run_adaptive(3 * DT, dt_initial=DT, dt_max=DT) + # (The step of this graph is traced twice -- the plain node's leaves + # come back placed on the mesh -- so it is compared twice, and passes.) + assert _layout(gm) == layout + + +@pytest.mark.parametrize("entry", ["step", "run", "run_adaptive"]) +def test_a_node_beside_a_sharded_one_that_reshapes_its_state_is_refused(entry): + gm = _graph([1.0, 2.0]) + held = gm._state + call = {"step": gm.step, "run": lambda: gm.run(2), + "run_adaptive": lambda: gm.run_adaptive(3 * DT, dt_initial=DT, dt_max=DT)}[entry] + with pytest.raises(ValueError, match=r"'b/position' has shape \(\) before the update " + r"and \(2,\) after it"): + call() + assert gm._state is held diff --git a/tests/cloud/multigpu/test_unstructured_layout_contract.py b/tests/cloud/multigpu/test_unstructured_layout_contract.py index 51ab90ef..e790670d 100644 --- a/tests/cloud/multigpu/test_unstructured_layout_contract.py +++ b/tests/cloud/multigpu/test_unstructured_layout_contract.py @@ -353,12 +353,28 @@ def test_a_node_declaring_its_integral_runs_in_a_graph_like_the_unsharded_node(s assert np.all(np.isfinite(np.asarray(gm.get_node_state("src")["total"]))) -def test_a_node_without_an_initial_integral_steps_in_a_graph(): - """``gm.step()`` feeds step 1's output (now carrying the integral) back - in: three steps, the integral of the last.""" - gm = _graph(ShardedUnstructuredNode(_Source(8), create_device_mesh(shape=(2,)), - _chain_layout(8, 2))) +def test_a_node_without_an_initial_integral_is_handed_its_own_output_back(): + """The wrapper takes step 1's output (now carrying the integral) back + in: three steps, the integral of the last. The steps are the wrapper's + own, jitted as a graph jits them. + + A graph does not store that first step any more: its state would gain + a field ``initial_state()`` did not build, so a checkpoint of it would + not fit the graph after a reset (``GraphManager.step`` refuses a step + that changes the state's layout, MADD-ANO-220, as ``run_scan`` always + has). In a graph the node declares its integral, as the test above + does.""" + node = ShardedUnstructuredNode(_Source(8), create_device_mesh(shape=(2,)), + _chain_layout(8, 2)) + step = jax.jit(lambda state: node.update(state, {}, 0.1)) + state = node.initial_state() + assert "total" not in state for _ in range(3): + state = step(state) + assert float(np.asarray(state["total"])) == pytest.approx(36.0 + 8 * 0.3, rel=1e-6) + + gm = _graph(node) + held = gm._state + with pytest.raises(ValueError, match="'src/total' appears only after the update"): gm.step() - assert float(np.asarray(gm.get_node_state("src")["total"])) == pytest.approx( - 36.0 + 8 * 0.3, rel=1e-6) + assert gm._state is held diff --git a/tests/core/test_a_step_keeps_the_state_layout.py b/tests/core/test_a_step_keeps_the_state_layout.py index a8bc1ac6..9b7c859b 100644 --- a/tests/core/test_a_step_keeps_the_state_layout.py +++ b/tests/core/test_a_step_keeps_the_state_layout.py @@ -1,15 +1,20 @@ -"""A step should leave the graph's state with the layout it had. +"""A step leaves the graph's state with the layout it had. A node whose ``update`` returns a leaf of another shape than its ``initial_state()`` built -- a list given for a scalar constant, a vector delivered into a scalar field -- broadcasts the leaf at its first step. ``run_scan`` and its siblings refuse such a graph (a scan carry of another -type). ``GraphManager.step``, ``run`` and ``run_adaptive`` store the -result without a word, as they have since 0.1.0 (MADD-ANO-220, open): a -checkpoint saved after the step does not load after ``reset_state``. The -REST server refuses such a node where it is added (``POST /graph/nodes``), -with the comparison tested here (``_param_probes._state_layout_drift``). +type). ``GraphManager.step``, ``run`` and ``run_adaptive`` stored the +result without a word from 0.1.0 (MADD-ANO-220): a checkpoint saved after +the step did not load after ``reset_state``. They now refuse it, by the +name of the node and the leaf, and store nothing: the stepped state is +compared with the one it replaces once per trace of the step, host-side +(``GraphManager._store_stepped_state``). The REST server refuses such a +node where it is added (``POST /graph/nodes``), with the same comparison +(``_param_probes._state_layout_drift``). """ +import warnings + import jax import jax.numpy as jnp import numpy as np @@ -47,20 +52,264 @@ def test_the_scan_entry_points_refuse_a_step_that_changes_a_leafs_shape(entry): assert _shapes(gm) == shapes -@pytest.mark.xfail(strict=True, reason=( - "MADD-ANO-220 (open): step, run and run_adaptive store a stepped state whose leaf " - "has another shape than the state it replaces; the state no longer has the layout " - "of initial_state(), and its checkpoint does not load after a reset")) -@pytest.mark.parametrize("entry", ["step", "run", "run_adaptive"]) +class _Recorder: + """An observer and a run callback: what a refused step must not reach.""" + + def __init__(self, gm: GraphManager): + self.events, self.calls = [], 0 + gm.add_observer(lambda event, data: self.events.append(event)) + + def callback(self, *args): + self.calls += 1 + + +def _entry_points(gm: GraphManager, seen: _Recorder, **kw) -> dict: + return {"step": lambda: gm.step(**kw), + "run": lambda: gm.run(3, callback=seen.callback, **kw), + "run_adaptive": lambda: gm.run_adaptive(5 * DT, callback=seen.callback, **kw)} + + +ENTRIES = ["step", "run", "run_adaptive"] + + +@pytest.mark.parametrize("entry", ENTRIES) def test_the_step_entry_points_keep_the_shape_of_every_state_leaf(entry): + """MADD-ANO-220: the stepped state was stored. Refused, by the name of + the node and the leaf, and the graph is as it was: the very state + object (its clock and step count are in it), no observer told of a + step, no callback called.""" gm = _broadcasting_ball() - shapes = _shapes(gm) - try: - {"step": gm.step, "run": lambda: gm.run(3), - "run_adaptive": lambda: gm.run_adaptive(5 * DT)}[entry]() - except (ValueError, TypeError): - pass # a refusal that stores nothing would do - assert _shapes(gm) == shapes + gm.compile() + held, shapes = gm._state, _shapes(gm) + seen = _Recorder(gm) + with pytest.raises(ValueError, match=r"'b/position' has shape \(\) before the update " + r"and \(2,\) after it") as refusal: + _entry_points(gm, seen)[entry]() + assert "Nothing was stored" in str(refusal.value) + assert gm._state is held and _shapes(gm) == shapes + assert "step" not in seen.events and seen.calls == 0 + # ... and it is refused again: a refusal does not use the comparison up. + with pytest.raises(ValueError, match="changes the layout"): + _entry_points(gm, seen)[entry]() + assert gm._state is held and "step" not in seen.events and seen.calls == 0 + + +class _Halves(SimulationNode): + """An integer count its update returns as a float.""" + + def initial_state(self): + return {"n": jnp.zeros((), jnp.int32), "x": jnp.zeros((), jnp.float32)} + + def update(self, state, boundary_inputs, dt, *, params=None): + return {"n": state["n"] + 0.5, "x": state["x"] + dt} + + +class _Widens(SimulationNode): + """An ``int16`` count its update returns as ``int32``: another width of + the same kind of dtype.""" + + def initial_state(self): + return {"n": jnp.zeros((), jnp.int16)} + + def update(self, state, boundary_inputs, dt, *, params=None): + return {"n": state["n"].astype(jnp.int32) + 1} + + +class _DropsAField(SimulationNode): + def initial_state(self): + return {"x": jnp.zeros((), jnp.float32), "y": jnp.zeros((), jnp.float32)} + + def update(self, state, boundary_inputs, dt, *, params=None): + return {"x": state["x"] + dt} + + +@pytest.mark.parametrize("entry", ENTRIES) +@pytest.mark.parametrize("node, needle", [ + (_Halves, "'h/n' has dtype int32 before the update and float32 after it"), + (_DropsAField, "'h/y' is missing after the update"), +], ids=["kind-of-dtype", "missing-field"]) +def test_a_step_that_changes_a_leafs_kind_of_dtype_or_the_fields_is_refused(entry, node, needle): + """The same loss by another route: the scans refuse a carry of another + dtype, and a checkpoint of the float count does not load into the + integer state of the graph after a reset.""" + gm = GraphManager() + gm.add_node(node("h", DT)) + gm.compile() + held, seen = gm._state, _Recorder(gm) + with pytest.raises(ValueError) as refusal: + _entry_points(gm, seen)[entry]() + assert needle in str(refusal.value) + assert gm._state is held and "step" not in seen.events and seen.calls == 0 + + +def test_a_dtype_of_another_width_is_not_what_the_comparison_refuses(): + """The rule's edge, stated: only the kind of dtype is compared. With + x64 enabled a stock node returns ``float64`` for the ``float32`` its + ``initial_state()`` builds, and a comparison of widths would refuse + every one of them.""" + gm = GraphManager() + gm.add_node(_Widens("w", DT)) + gm.step() + gm.step() + assert gm._state["w"]["n"].dtype == jnp.int32 and int(gm._state["w"]["n"]) == 2 + + +class _Adds(SimulationNode): + """Adds whatever its ``drive`` input delivers to a scalar.""" + + def initial_state(self): + return {"total": jnp.zeros((), jnp.float32)} + + def update(self, state, boundary_inputs, dt, *, params=None): + return {"total": state["total"] + boundary_inputs["drive"]} + + def boundary_input_spec(self): + return {"drive": BoundaryInputSpec(shape=(), description="what it adds")} + + +def _adder() -> GraphManager: + gm = GraphManager() + gm.add_node(_Adds("a", DT)) + gm.add_external_input("a", "drive") + return gm + + +@pytest.mark.parametrize("entry", ENTRIES) +def test_a_later_step_that_is_traced_again_is_compared_again(entry): + """The comparison is once per trace of the step, not once per graph: + an external input of another shape retraces the step, and the leaf it + broadcasts is refused as it is at a first step. The graph is where + the last good step left it, and goes on from there to the bits.""" + one, vector = {"a": {"drive": 1.0}}, {"a": {"drive": jnp.ones(3, jnp.float32)}} + gm = _adder() + gm.step(external_inputs=one) + gm.step(external_inputs=one) + assert gm.trace_count == 1 + held, seen = gm._state, _Recorder(gm) + with pytest.raises(ValueError, match=r"'a/total' has shape \(\) before the update " + r"and \(3,\) after it"): + _entry_points(gm, seen, external_inputs=vector)[entry]() + assert gm._state is held and "step" not in seen.events and seen.calls == 0 + gm.step(external_inputs=one) + reference = _adder() + for _ in range(3): + reference.step(external_inputs=one) + assert np.asarray(gm._state["a"]["total"]).tobytes() == \ + np.asarray(reference._state["a"]["total"]).tobytes() + assert np.shape(gm._state["a"]["total"]) == () + + +def test_the_step_of_a_graph_compiled_again_is_compared_again(): + """Each compile is a new program. A graph that has stepped is given a + node that broadcasts: its next step is the first of the new compile.""" + gm = GraphManager() + gm.add_node(BallNode("ok", DT)) + for _ in range(3): + gm.step() + gm.add_node(BallNode("b", DT, initial_velocity=[1.0, 2.0])) + gm.compile() + assert gm.trace_count == 0 + held = gm._state + with pytest.raises(ValueError, match="'b/position' has shape"): + gm.step() + assert gm._state is held + # The user's way on: the node replaced by one of the right rank. + gm.remove_node("b") + gm.add_node(BallNode("b", DT, initial_velocity=1.0)) + gm.step() + assert np.shape(gm._state["b"]["position"]) == () + + +class _CountsItsTraces(SimulationNode): + traces = 0 + + def initial_state(self): + return {"x": jnp.zeros((2,), jnp.float32)} + + def update(self, state, boundary_inputs, dt, *, params=None): + type(self).traces += 1 + return {"x": state["x"] + dt} + + +def test_the_comparison_is_made_once_per_trace_and_calls_no_update(monkeypatch): + """It reads the shapes of the state the step returned, host-side: the + node's ``update`` runs as often as the step's one trace runs it, and + nine steps of one program are compared once. (The patch is of the + name the graph reads: the count below shows it is in effect.)""" + compared = [] + real = _param_probes._state_layout_drift + monkeypatch.setattr(_param_probes, "_state_layout_drift", + lambda *a, **kw: compared.append(1) or real(*a, **kw)) + gm = GraphManager() + gm.add_node(_CountsItsTraces("c", DT)) + gm.compile() + _CountsItsTraces.traces = 0 + for _ in range(5): + gm.step() + gm.run(4) + assert gm.trace_count == 1 and _CountsItsTraces.traces == 1 + assert len(compared) == 1 + gm.compile() + gm.step() + gm.step() + assert len(compared) == 2 + + +def test_a_graph_with_a_coupling_group_is_held_to_its_layout_too(): + """The group's own keys (its iteration count, its residual) are leaves + of the stepped state like any other, and a healthy group steps; a + node beside the group that broadcasts is refused, and the group's + bookkeeping is where it was.""" + def build(velocity): + gm = GraphManager() + gm.add_node(SpringDamperNode("p", DT, stiffness=30.0, damping=2.0, initial_position=1.0)) + gm.add_node(SpringDamperNode("q", DT, stiffness=20.0, damping=2.0, initial_position=0.5)) + gm.add_edge("p", "q", "position", "anchor_position") + gm.add_edge("q", "p", "position", "anchor_position") + gm.add_coupling_group(["p", "q"], max_iterations=8, tolerance=1e-6) + gm.add_node(BallNode("b", DT, initial_velocity=velocity)) + with warnings.catch_warnings(): + warnings.simplefilter("ignore") + gm.compile() + return gm + + healthy = build(1.0) + before = _param_probes._state_layout_drift(healthy._state, healthy._state) + layout = jax.tree.map(lambda leaf: (np.shape(leaf), np.asarray(leaf).dtype.kind), + healthy._state) + for _ in range(3): + healthy.step() + assert before == [] and layout == jax.tree.map( + lambda leaf: (np.shape(leaf), np.asarray(leaf).dtype.kind), healthy._state) + + gm = build([1.0, 2.0]) + held = gm._state + for call in (gm.step, lambda: gm.run(2), lambda: gm.run_adaptive(5 * DT)): + with pytest.raises(ValueError, match="'b/position' has shape"): + call() + assert gm._state is held + + +def test_a_step_inside_a_transform_is_compared_on_the_traced_state(): + """Under ``jax.grad`` the stepped state is a tree of tracers, whose + shapes are as real as an array's: a healthy graph differentiates, and + one that broadcasts is refused inside the trace as outside it.""" + def loss(gm): + def f(v): + params = gm.params + params["nodes"]["b"]["gravity"] = v + gm.step(params=params) + return jnp.sum(gm.step(params=params)["b"]["position"]) + return f + + healthy = GraphManager() + healthy.add_node(BallNode("b", DT, initial_position=5.0)) + healthy.compile() + assert np.isfinite(float(jax.grad(loss(healthy))(jnp.float32(-9.81)))) + with pytest.raises(ValueError, match="'b/position' has shape"): + broadcasting = _broadcasting_ball() + broadcasting.compile() + jax.grad(loss(broadcasting))(jnp.float32(-9.81)) def _drift(node: SimulationNode, boundary_inputs=None, **kw) -> list: @@ -124,3 +373,46 @@ def test_the_comparison_names_every_kind_of_difference(): # The abstract values of a trace compare like arrays. abstract = jax.eval_shape(lambda: before) assert _param_probes._state_layout_drift(before, abstract) == [] + + +class _Draws(SimulationNode): + """Carries a PRNG key in its state, as a stochastic node does.""" + + def initial_state(self): + return {"key": jax.random.key(0), "x": jnp.zeros((), jnp.float32)} + + def update(self, state, boundary_inputs, dt, *, params=None): + key, sub = jax.random.split(state["key"]) + return {"key": key, "x": state["x"] + dt * jax.random.normal(sub)} + + +class _SpendsItsKey(_Draws): + def update(self, state, boundary_inputs, dt, *, params=None): + return {"key": jax.random.key_data(state["key"]), "x": state["x"]} + + +def test_a_prng_key_in_the_state_is_a_leaf_of_a_kind_of_its_own(): + """A key's dtype (``key``) is not a NumPy dtype and has no kind + letter: it is compared as itself. The comparison used to raise on it + (``canonicalize_dtype called on extended dtype``), which refused every + step of a graph that carries a key.""" + gm = GraphManager() + gm.add_node(_Draws("d", DT)) + gm.step() + gm.run(2) + assert gm._state["d"]["key"].dtype == jax.random.key(0).dtype + assert _drift(_Draws("d", DT)) == [] + key, other = jax.random.key(0), jax.random.key(0, impl="rbg") + assert _param_probes._state_layout_drift({"k": key}, {"k": jax.random.split(key)[0]}) == [] + for after in (other, jnp.zeros((), jnp.uint32)): + (found,) = _param_probes._state_layout_drift({"k": key}, {"k": after}) + assert "'k' has dtype key before the update" in found + assert _param_probes._state_layout_drift({"k": key}, {"k": other}, dtypes=False) == [] + # ... and a node that returns the key's bits for the key is refused. + spent = GraphManager() + spent.add_node(_SpendsItsKey("d", DT)) + spent.compile() + held = spent._state + with pytest.raises(ValueError, match="'d/key' has shape"): + spent.step() + assert spent._state is held diff --git a/tests/core/test_compile_counts.py b/tests/core/test_compile_counts.py index 8b4fbd4d..67c1acbe 100644 --- a/tests/core/test_compile_counts.py +++ b/tests/core/test_compile_counts.py @@ -79,27 +79,31 @@ def _coupled() -> GraphManager: class AvalChangesOnce(SimulationNode): """A node that forces exactly one extra compile of the step. - The first step is traced against a ``(1,)`` leaf and returns a - ``(2,)`` one, so the second step's argument has a different aval and - JAX must trace and compile the whole step again. From the third - step on the aval is stable, so the retrace count settles at 2 rather - than growing without bound -- the mutation has to be as + Its ``x`` starts as a Python float, which the first step is traced + against as a weakly typed ``float32`` scalar; the step returns a + strongly typed one, so the second step's argument has a different + aval and JAX must trace and compile the whole step again. From the + third step on the aval is stable, so the retrace count settles at 2 + rather than growing without bound -- the mutation has to be as deterministic as the thing it is testing. This is the shape of the real bug the counts exist to catch: nothing here is wrong enough to raise, the results are all correct, and the - only symptom is a second XLA compile on every run. + only symptom is a second XLA compile on every run. The state keeps + its layout throughout -- the same fields, shapes and kinds of dtype -- + so the step is one ``GraphManager.step`` stores (a step that reshapes + a leaf is refused), and the state saves, reloads and scans. """ def initial_state(self): - return {"x": jnp.zeros(1), "n": jnp.array(0, jnp.int32)} + return {"x": 0.0, "n": jnp.array(0, jnp.int32)} def state_fields(self): return ["x", "n"] def update(self, state, boundary_inputs, dt, *, params=None): - x = state["x"] if state["x"].shape[0] == 2 else jnp.zeros(2) - return {"x": x + dt, "n": state["n"] + 1} + x = jnp.asarray(state["x"], jnp.float32) # strongly typed from here on + return {"x": x + jnp.float32(dt), "n": state["n"] + 1} # --------------------------------------------------------------------------- @@ -254,6 +258,11 @@ def test_a_changing_state_aval_is_reported_as_an_extra_compile(): gm.step() assert gm.trace_count == 2, "the mutation did not actually retrace" assert compile_counts(gm).retrace_count == 2 + # ... and it is a retrace of a step that keeps the state's layout: the + # weak type is the only thing that changed. + assert jnp.shape(gm._state["g"]["x"]) == () and gm._state["g"]["x"].dtype == jnp.float32 + assert not gm._state["g"]["x"].weak_type + assert int(gm._state["g"]["n"]) == 8 # four steps here, four of the measurement's # --------------------------------------------------------------------------- diff --git a/tests/fmi/test_a_sidecar_step_keeps_the_state_layout.py b/tests/fmi/test_a_sidecar_step_keeps_the_state_layout.py new file mode 100644 index 00000000..de1fe60d --- /dev/null +++ b/tests/fmi/test_a_sidecar_step_keeps_the_state_layout.py @@ -0,0 +1,126 @@ +"""The FMU sidecar and bridge hold a step to the layout of the state it +was given, as ``GraphManager.step`` does. + +The sidecar runs a graph's compiled step on a state of its own, so the +graph's comparison never sees it. A node that broadcasts a state leaf at +its first step (``BallNode(initial_velocity=[1.0, 2.0])``: the scalar +``position`` becomes ``(2,)``) left the sidecar holding a state that was +not the one its model description declares (MADD-ANO-220). The step is +refused by the name of the node and the leaf, and nothing is committed. +""" + +from __future__ import annotations + +import jax.numpy as jnp +import numpy as np +import pytest + +from maddening.core.graph_manager import GraphManager +from maddening.fmi import build_model_description +from maddening.fmi import sidecar as sidecar_module +from maddening.fmi.sidecar import FmuSidecar, SidecarConfig +from maddening.fmi.tcp_bridge import FmuTcpBridge +from maddening.nodes import BallNode, SpringDamperNode + +DT = 0.01 +NEEDLE = r"'b/position' has shape \(\) before the update and \(2,\) after it" + + +def _graph(velocity) -> GraphManager: + gm = GraphManager() + gm.add_node(SpringDamperNode("s", DT, stiffness=30.0, damping=2.0, initial_position=1.0)) + gm.add_node(BallNode("b", DT, initial_velocity=velocity)) + gm.compile() + return gm + + +def _sidecar(gm: GraphManager, md) -> FmuSidecar: + return FmuSidecar(SidecarConfig( + schema_token=md.instantiation_token, step_fn=gm._compiled_step, + initial_state=gm._state, params=gm.params, param_specs=gm.param_specs(), + input_resolver=gm._resolve_external_inputs)) + + +def test_a_sidecar_step_that_reshapes_a_leaf_is_refused_and_nothing_is_committed(): + gm = _graph([1.0, 2.0]) + sidecar = _sidecar(gm, build_model_description(gm, model_name="m")) + held = sidecar._state + for _ in range(2): # a refusal does not use the comparison up + with pytest.raises(ValueError, match=NEEDLE) as refusal: + sidecar.step(None) + assert "Nothing was advanced" in str(refusal.value) + assert sidecar._state is held + assert np.shape(sidecar.state["b"]["position"]) == () + + +def test_a_bridge_step_that_reshapes_a_leaf_is_an_error_reply_and_the_clock_stays(): + gm = _graph([1.0, 2.0]) + md = build_model_description(gm, model_name="m") + bridge = FmuTcpBridge(_sidecar(gm, md), md, master_dt=DT) + held = bridge._sidecar._state + reply = bridge.handle({"op": "step", "t": 0.0, "dt": md.default_step_size}) + assert reply["ok"] is False and "'b/position' has shape ()" in reply["error"], reply + assert bridge._sidecar._state is held and bridge._time == 0.0 + + +def test_a_graphs_step_is_compared_once_per_trace(monkeypatch): + """Not at every step: the layout a traced program returns is fixed by + the layout it is given. The patch is of the name the sidecar reads.""" + calls = [] + real = sidecar_module._state_layout_drift + monkeypatch.setattr(sidecar_module, "_state_layout_drift", + lambda *a, **kw: calls.append(1) or real(*a, **kw)) + gm = _graph(1.0) + sidecar = _sidecar(gm, build_model_description(gm, model_name="m")) + for _ in range(5): + sidecar.step(None) + assert len(calls) == 1 and gm.trace_count == 1 + # The sidecar's five steps are the graph's five, to the bits. + reference = _graph(1.0) + reference.run(5) + for node in ("s", "b"): + assert np.asarray(sidecar.state[node]["position"]).tobytes() == \ + np.asarray(reference._state[node]["position"]).tobytes() + + +def test_a_step_that_is_not_a_graphs_is_compared_at_every_step(monkeypatch): + """Nothing says when a wrapper or a double is a new program, so it is + never assumed to be the old one: a leaf that grows at the third step + is refused at the third step.""" + calls = [] + real = sidecar_module._state_layout_drift + monkeypatch.setattr(sidecar_module, "_state_layout_drift", + lambda *a, **kw: calls.append(1) or real(*a, **kw)) + taken = [] + + def step_fn(state, external_inputs): + taken.append(1) + x = state["n"]["x"] + 1.0 + return {"n": {"x": jnp.broadcast_to(x, (2,)) if len(taken) == 3 else x}} + + sidecar = FmuSidecar(SidecarConfig( + schema_token="t", step_fn=step_fn, + initial_state={"n": {"x": jnp.zeros((), jnp.float32)}})) + sidecar.step(None) + sidecar.step(None) + held = sidecar._state + with pytest.raises(ValueError, match=r"'n/x' has shape \(\) before the update"): + sidecar.step(None) + assert sidecar._state is held and float(held["n"]["x"]) == 2.0 and len(calls) == 3 + + +def test_the_step_of_a_graph_compiled_again_since_is_compared_at_every_step(monkeypatch): + """Its trace count is the new compile's, not this step's.""" + gm = _graph(1.0) + sidecar = _sidecar(gm, build_model_description(gm, model_name="m")) + sidecar.step(None) + calls = [] + real = sidecar_module._state_layout_drift + monkeypatch.setattr(sidecar_module, "_state_layout_drift", + lambda *a, **kw: calls.append(1) or real(*a, **kw)) + sidecar.step(None) + assert calls == [] + gm.compile() + sidecar._advanced(sidecar._state, None) + sidecar._advanced(sidecar._state, None) + assert len(calls) == 2