diff --git a/CHANGELOG.md b/CHANGELOG.md index 1910908..924d46e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -27,6 +27,30 @@ `ResolvedPlan.from_json` drops its 16 MiB input cap and duplicate `max_nodes` check; revalidation still applies `max_nodes`. +- **Contracts are checked once per input signature, not on every call.** The + first forward whose inputs have a given shape, dtype and device (in a given + training mode and autocast state) still checks every node's ports and that + no module created or removed registered state; later calls with the same + signature compare only the inputs and run the layers back to back. Moving + or casting the model, or registering a parameter, buffer or submodule on + any of its modules, makes the next call check everything again. The fixed + per-call cost of a five-node MLP drops from ~30-40 us to ~6-9 us: at batch + 32 it runs ~1.2x hand-written PyTorch instead of ~1.8x. + Errors say more: `E_RUNTIME` names the node, its operation and source line, + and prints the contract in HNDL notation beside the tensor that arrived + (`expected [B=32, 64]:float32 on cpu` / `got [32, 63]:float32 on cpu`); + dtypes are spelled `float32`, not `torch.float32`. An exception raised + inside a layer keeps its type, message and traceback (a `RuntimeError` or + out-of-memory error is still caught as one) and gains a note, via + `add_note`, naming the node, its operation and source line, its inputs + against their contract, and any registered-state change that explains it; + an exception passing out through nested graphs carries only the innermost + node's note. Registered + state edits PyTorch runs no hook for --- `del` of a registered name, + `module.param = None`, direct writes to `_parameters`/`_buffers`/`_modules` + --- are reported at the next full check or when a layer then fails, not on + the very next call. + - Add `broadcast_mul`, `coordinate_grid`, `fourier_features`, and `grid_sample` for style-conditioned coordinate renderers composed in HNDL. Fourier tables and coordinate grids are persistent buffers; sampling follows PyTorch semantics. diff --git a/IMPLEMENTATION.md b/IMPLEMENTATION.md index 490bbdb..3ac7b66 100644 --- a/IMPLEMENTATION.md +++ b/IMPLEMENTATION.md @@ -67,6 +67,25 @@ graph input is integer. Edge dtypes are checked at resolution (`E_DTYPE`), so an integer tensor cannot reach a floating-point port. Reduced precision is qualified on CUDA. +A built model checks its contracts once per input signature, not on every +call. The first call whose inputs have a given shape, dtype and device (in a +given training mode and autocast state) checks every node's ports and that no +module created or removed registered state; later calls with the same +signature compare only the inputs and run the layers back to back. Moving or +casting the model, or registering a parameter, buffer or submodule on any of +its modules, makes the next call check everything again. Failures are +`E_RUNTIME` errors that name the node, its operation and its source line and +print the contract beside the tensor that arrived, for example +`linear 'head' input 'x' (from node:hidden/out): expected shape [B=32, 64], got [32, 63]`; +an exception raised inside a layer propagates unchanged (a `RuntimeError` +stays a `RuntimeError`) with a note, printed under its message in the +traceback, that names the node, its operation and source line, the tensors it +received, and any registered-state change that explains it. PyTorch runs no registration hook for `del` of a +registered name, for assigning `None` over a registered parameter, or for +writing `_parameters`/`_buffers`/`_modules` directly, so those edits are +reported at the next full check or when a layer then fails, not on the very +next call. + Built-in unary operations take an optional leading tensor or `x=`. Custom unary operations use their declared input-port keyword. Every operator's arguments, defaults, bounds, shape relation, and examples are listed in diff --git a/README.md b/README.md index d370497..b38f3a3 100644 --- a/README.md +++ b/README.md @@ -73,7 +73,7 @@ last_layer = model[-1] features = model[:2] # nn.Sequential sharing these layers ``` -Unnamed operations receive IDs such as `n0`; `model["n0"]` and `model[0]` return the same module. Optional names appear in the branching example below. Slices reuse their parameters, so training a slice also updates the original model. The complete model retains its resolved shape contract; a slice is a regular PyTorch sequence. Standard `state_dict()`, `train()`, and `eval()` remain available. Moving the model with `.cpu()` or `.cuda()` and casting it with `.double()`, `.half()`, `.bfloat16()`, `.float()`, or `.to(dtype=...)` retarget the runtime checks too, so a cast model takes tensors of its new compute dtype; ports declared as integers, such as token ids, keep the dtype the plan declared. And `copy.deepcopy(model)` returns an independent model --- its own parameters and buffers, the same trainability flags and training mode, and no draw on the random state --- which is what a moving-average copy of a model needs. The resolved plan is immutable, so the copy shares it. To store a model, save `model.plan.to_json()` next to `torch.save(model.state_dict())` and rebuild it; pickling the module itself is not supported. +Unnamed operations receive IDs such as `n0`; `model["n0"]` and `model[0]` return the same module. Optional names appear in the branching example below. Slices reuse their parameters, so training a slice also updates the original model. The complete model retains its resolved shape contract; a slice is a regular PyTorch sequence. Standard `state_dict()`, `train()`, and `eval()` remain available. Moving the model with `.cpu()` or `.cuda()` and casting it with `.double()`, `.half()`, `.bfloat16()`, `.float()`, or `.to(dtype=...)` retarget the runtime checks too, so a cast model takes tensors of its new compute dtype; ports declared as integers, such as token ids, keep the dtype the plan declared. Those checks run once: the first call with a given input shape, dtype and device checks every layer's inputs and outputs against the plan, and later calls with the same inputs run the layers directly, at close to hand-written speed. A mismatch is an `E_RUNTIME` error naming the layer, its operation and its source line, with the expected and actual shapes; an error raised inside a layer keeps its own type and gains a note with the same details. And `copy.deepcopy(model)` returns an independent model --- its own parameters and buffers, the same trainability flags and training mode, and no draw on the random state --- which is what a moving-average copy of a model needs. The resolved plan is immutable, so the copy shares it. To store a model, save `model.plan.to_json()` next to `torch.save(model.state_dict())` and rebuild it; pickling the module itself is not supported. Initial weights use PyTorch’s normal random state. For repeatable initialization in the same environment, call [`torch.manual_seed(7)`](https://docs.pytorch.org/docs/stable/notes/randomness.html#pytorch-random-number-generator) before constructing the network. diff --git a/SPEC.md b/SPEC.md index a4d68f9..5ed1f6a 100644 --- a/SPEC.md +++ b/SPEC.md @@ -50,7 +50,7 @@ A tensor contract is an ordered tuple of dimensions including batch, interpreted Batch passes through every operator unchanged except the two that move tensors across axis 0. Joining along that axis stacks examples, so the result holds several batches at once: the entry is then written `"k*B"`, meaning `k` times the plan's batch. `"B"` is one batch and has no other spelling — `"1*B"`, `"0*B"`, `"B*2"` and leading zeros are rejected with `E_SCHEMA` — so a plan that never touches the batch axis carries exactly the entries, and therefore the digests, it carried before multiples existed. Multiples run from 2 to 1024; exceeding that fails with `E_RESOURCE`. A `k*B` port requires `k` times the runtime batch of the call, checked like any other extent. External contracts — `input_shape` and `output_shape`, named or not — stay one plan batch: a graph must chunk a joined tensor back before publishing it, and a declared `"2*B"` contract fails with `E_SCHEMA`. Internal port contracts in a saved plan may carry multiples, and a malformed entry at axis 0 fails `E_SCHEMA` on load. -A plan carries a compute `dtype` of `float32` (the default), `float16`, or `bfloat16`, and an `input_dtype` that defaults to the compute dtype. The graph input may instead be an integer contract — `int64`, `int32`, or `bool` — for token ids and masks; a floating-point `input_dtype` must equal the compute dtype. Operator ports may declare their own dtype, including `any` for a port that accepts whatever its producer carries. Every edge's dtype is checked once during resolution and a mismatch fails with `E_DTYPE`, so an integer tensor cannot reach a floating-point port. The backend constructs parameters in the compute dtype and checks each port's dtype, shape, and device at runtime. +A plan carries a compute `dtype` of `float32` (the default), `float16`, or `bfloat16`, and an `input_dtype` that defaults to the compute dtype. The graph input may instead be an integer contract — `int64`, `int32`, or `bool` — for token ids and masks; a floating-point `input_dtype` must equal the compute dtype. Operator ports may declare their own dtype, including `any` for a port that accepts whatever its producer carries. Every edge's dtype is checked once during resolution and a mismatch fails with `E_DTYPE`, so an integer tensor cannot reach a floating-point port. The backend constructs parameters in the compute dtype and checks each port's dtype, shape, and device at runtime, once per input signature (§11). Both frontend APIs accept `input_shape` and `output_shape`, both including batch, plus optional `dtype` and `input_dtype` keywords. Structured graph contracts use the same tuples. @@ -631,7 +631,7 @@ linear() Declarations are trusted code, for custom operators as much as for built-ins: relation functions and module constructors run in the host's process. A registered relation is allowed to be arbitrary Python, and custom operators may therefore express any rule a built-in can. What remains untrusted is data: configuration source and saved plans can name only an alias or identity the host already registered, can never import a class, a module path, or a callback, and are bounded by the same parser and resolver limits. A plan stores `example.silu@1` plus concrete arguments and port shapes, not executable Python, and building or restoring it without that registration fails with `E_STATE_VERSION`. -The constructor receives every declared argument by keyword, plus any shape symbol it names as a keyword-only parameter (`*, D`) and, when requested by name, `input_shapes` / `output_shapes` mappings of resolved port shapes. It runs under `torch.device(device)`, so tensors may be created normally; parameters are then cast to the plan's compute dtype. `forward` receives tensors positionally in declared input-port order and returns one tensor, or a tuple/list in declared order or a dictionary with exactly the declared names for multiple outputs. HNDL validates each result's shape, dtype, and device. Within a built model each node owns an independent module instance and registered state; reusing a module or registered storage across nodes is rejected. Sharing a producer tensor across branches, or obtaining a slice of an already built chain, remains supported. A genuine unary chain supports indexing and shared sequential slices even when its ports use names other than `x` and `out`. +The constructor receives every declared argument by keyword, plus any shape symbol it names as a keyword-only parameter (`*, D`) and, when requested by name, `input_shapes` / `output_shapes` mappings of resolved port shapes. It runs under `torch.device(device)`, so tensors may be created normally; parameters are then cast to the plan's compute dtype. `forward` receives tensors positionally in declared input-port order and returns one tensor, or a tuple/list in declared order or a dictionary with exactly the declared names for multiple outputs. HNDL validates each result's shape, dtype, and device on the first call for each input signature (§11). Within a built model each node owns an independent module instance and registered state; reusing a module or registered storage across nodes is rejected. Sharing a producer tensor across branches, or obtaining a slice of an already built chain, remains supported. A genuine unary chain supports indexing and shared sequential slices even when its ports use names other than `x` and `out`. Saved plans contain concrete shapes and arguments plus exact operator identities and versions. Declarations remain in the explicitly supplied registry; artifacts cannot import them. Restoring a custom plan revalidates its concrete equations against that registry, and any mismatch fails with `E_INTEGRITY` rather than being re-inferred (§12 lists the two kinds of missing value restoration fills in). An operator must raise its semantic version when it changes a shape relation, an argument schema, its numerical meaning, or its parameter layout. @@ -727,7 +727,7 @@ The **artifact digest** covers the complete saved plan except its own digest fie - Construction revalidates the plan, then runs an allocation-free `meta` pass over every node to bound registered storage before allocating anything (§7). The build receipt records `torch_version`, `device`, `dtype`, `initialization_seed`, `seed_mode`, and `state_bytes`. - Parameters and buffers are constructed in the plan's compute dtype. A node is constructed under the target device so constructors may allocate normally. - A `GraphModule` validates declared external input contracts in `forward(**inputs)`, which accepts exactly the declared input names, and returns a dictionary keyed by public output names in declaration order, including for a single output. A missing, repeated, or undeclared runtime input fails with `E_BINDING` before any node runs. Both network facades instead accept the declared inputs positionally in declaration order, by keyword, or both, and return the selected output tensor for a single public output or the same dictionary for several, preserving the contract checks even for a branched graph. -- Every port is checked for shape, dtype, and device before and after applying its node, and a module that creates or removes registered state during forward fails with `E_RUNTIME`. +- Contracts are checked once per input signature: the first call whose external inputs have a given shape, dtype, and device — in a given training mode and autocast state — checks every port for shape, dtype, and device before and after applying its node, and then checks that no module created or removed registered state; a module that did fails with `E_RUNTIME`. Later calls with a validated signature compare only that signature and then run the nodes. Moving or casting the module, and registering a parameter, buffer, or submodule on any module of the graph, invalidate what was validated: the next call re-checks registered state against the build and every port again. An exception raised inside a node's module propagates unchanged, keeping its type, message, and traceback, with one note (`add_note`) naming the node, its operation, and its source position, the tensors it received, and any change to registered state that explains it; an exception passing out through nested graphs keeps only the innermost node's note. `E_RUNTIME` reports what the backend's own checks find. - Modules register once under `nodes.n_`; the prefix avoids collisions with module attribute names. Stateless nodes keep execution/diagnostic identities without state entries. Two nodes sharing a module instance, tensor, or storage fail with `E_REGISTRY`. - Nodes run in stable topological order, breaking ties by declaration order. The order is saved in the plan. - Parameters/buffers are fully materialized before optimizer or distributed setup. No first-forward parameter creation is allowed. @@ -837,7 +837,7 @@ Failures must expose a stable code, source/node/field location, affected constra | `E_INTEGRITY` | A saved plan's digests or concrete equations do not match its contents | | `E_RESOURCE` | Source size, nesting depth, graph, state-size, or resolution budget is exceeded | | `E_STATE_VERSION` | Saved state requires an unavailable compatible operator implementation or plan version | -| `E_RUNTIME` | A runtime tensor violates its declared shape, dtype, or device, or the device is unavailable | +| `E_RUNTIME` | A runtime tensor violates its declared shape, dtype, or device; registered state no longer matches the build; or the device is unavailable | Configuration failures report original source line/column locations through the dedent map. Native capture locations are best effort and may use call-site information, but must always identify the affected node/field without requiring function-source inspection. `print(model)` and `repr(model)` must include every layer, its ID/operator, and complete input/output shapes without executing the network. `print(plan)` exposes the equivalent named-port table without a backend. `plan.describe()` must additionally expose provenance sufficiently to explain why a field changed between separately resolved specifications. diff --git a/src/hndl/torch.py b/src/hndl/torch.py index bebe9cb..caba5fe 100644 --- a/src/hndl/torch.py +++ b/src/hndl/torch.py @@ -4,10 +4,13 @@ from collections.abc import Mapping from contextlib import contextmanager import copy +import itertools from types import MappingProxyType +import weakref import torch from torch import nn +from torch.nn.modules import module as _torch_module from .errors import HNDLError from .settings import MATRIX_SCHEMES @@ -16,7 +19,8 @@ # Build metadata a copied network shares with its original: immutable records # describing the resolved architecture, never the parameters that train. SHARED_METADATA = frozenset({"plan", "build_receipt", "_port_orders", "_port_dtypes", - "_input_names", "_input_dtypes", "_output_dtypes", "_build_dtype", + "_input_names", "_input_set", "_input_dtypes", "_output_dtypes", + "_build_dtype", "_node_labels", "_spec_inputs", "_spec_outputs", "_spec_in", "_spec_out"}) DTYPES = {"float32": torch.float32, "float16": torch.float16, "bfloat16": torch.bfloat16, @@ -28,6 +32,94 @@ #: The resolved-shape cache before any call: batch, dtype, then three programs. _UNCOMPILED = (_UNSET, None, (), (), ()) +#: How many distinct validated input signatures a graph remembers before it +#: forgets them all and starts again; bounds the cache under varying batch sizes. +_MAX_SIGNATURES = 64 + +#: How the note HNDL adds to an exception raised inside a node begins. +_NOTE_PREFIX = "HNDL:" + +_is_compiling = torch.compiler.is_compiling +# One C call answering "is any autocast region active"; the public query needs +# a device type per call and costs three times as much. +_autocast_enabled = getattr(torch._C, "_is_any_autocast_enabled", None) or ( + lambda: torch.is_autocast_enabled("cpu") or torch.is_autocast_enabled("cuda")) + +# Every module inside a built graph is watched. Registering a parameter, +# buffer or submodule on one --- ``register_*`` or plain attribute assignment, +# including assigning ``None`` over a registered buffer or submodule --- +# advances ``_watch_generation``, which is how a graph learns, for the price +# of one integer comparison per call, that its registered state may have +# changed. ``itertools.count`` hands out unique values even under concurrent +# bumps, so a generation observed once is never observed again after a change. +# +# PyTorch runs no hook for three edits: ``del`` of a registered name, assigning +# ``None`` over a registered *parameter*, and writing ``_parameters``, +# ``_buffers`` or ``_modules`` directly. The first call for every input +# signature, every ``.to()``/cast, and the failure path of any node that then +# raises still compare the whole registered state against the build, so those +# edits are reported there rather than on the next call. +# +# Keyed by ``id`` rather than held in a WeakSet so that a module defining +# ``__eq__`` without ``__hash__`` can still be watched; an entry disappears +# with its module, before that ``id`` can be reused. +_watched = weakref.WeakValueDictionary() +_ticks = itertools.count(1) +_watch_generation = 0 + +_Tensor = torch.Tensor + + +def _registration_hook(module, name, value): + global _watch_generation + if _watched.get(id(module)) is module: + _watch_generation = next(_ticks) + # Returning None leaves the registration exactly as the caller asked. + + +def _watch(state_program): + for row in state_program: + _watched[id(row[0])] = row[0] + + +for _register in (_torch_module.register_module_parameter_registration_hook, + _torch_module.register_module_buffer_registration_hook, + _torch_module.register_module_module_registration_hook): + _register(_registration_hook) +del _register + + +def _dtype_text(dtype): + return "any" if dtype is None else str(dtype).removeprefix("torch.") + + +def _contract_text(spec, batch): + """A compiled contract shape in HNDL notation, batch entries bound: ``[B=32, 64]``.""" + parts = [] + for dimension in spec: + if type(dimension) is tuple: + multiple, text = dimension + parts.append(text if multiple is None or batch is None else f"{text}={multiple * batch}") + else: + parts.append(str(dimension)) + return "[" + ", ".join(parts) + "]" + + +def _tensor_text(value): + if not isinstance(value, torch.Tensor): + return type(value).__name__ + return f"{_shape_text(tuple(value.shape))}:{_dtype_text(value.dtype)} on {value.device}" + + +def _module_kind(module): + return "None" if module is None else type(module).__name__ + + +def _source_position(node): + """The ``(line, column)`` a declarative frontend recorded for a node, if any.""" + source = node.source if isinstance(node.source, Mapping) else {} + return source.get("line"), source.get("column") + def _compile_shape(shape): """A contract shape with its batch entries pre-parsed, once, at build time. @@ -87,7 +179,25 @@ def _is_chain(plan): class GraphModule(nn.Module): - """Execute a frozen graph once per node and return named output tensors.""" + """Execute a frozen graph once per node and return named output tensors. + + Validation happens once per *input signature*, not once per call. The + first call whose external inputs have a given shape, dtype and device --- + in a given training mode and autocast state --- checks every port of every + node against the plan and then checks that no module created or removed + registered state. Later calls with the same signature compare that + signature and then run the modules back to back. + + What was validated is forgotten when it may no longer hold: moving or + casting the module (``.to()``, ``.cuda()``, ``.double()``, ...) and + registering a parameter, buffer or submodule on any module in the graph + both re-check the registered state before the next call and revalidate + every port on it. An exception raised inside a node's module propagates + as itself --- same type, message and traceback --- with one note naming the + node, its operation and source line, the tensors it received, and any + change to registered state that explains it. ``E_RUNTIME`` is kept for + what HNDL's own checks find. + """ def __init__(self, plan, modules, device, receipt, port_orders): super().__init__() @@ -102,6 +212,7 @@ def __init__(self, plan, modules, device, receipt, port_orders): self._chain = _is_chain(plan) self._port_orders = port_orders self._input_names = tuple(plan.inputs) + self._input_set = frozenset(self._input_names) self._input_dtypes = MappingProxyType( {name: DTYPES[entry["dtype"]] for name, entry in plan.inputs.items()}) self._port_dtypes = {} @@ -114,6 +225,9 @@ def __init__(self, plan, modules, device, receipt, port_orders): produced[f"node:{node.id}/{port}"] = plan.dtype if declared[port] == "any" else declared[port] self._output_dtypes = MappingProxyType( {name: DTYPES[produced[entry["ref"]]] for name, entry in plan.outputs.items()}) + # What an error needs to name a node: its operation and source position. + self._node_labels = MappingProxyType( + {node.id: (self._alias(node.op), *_source_position(node)) for node in plan.nodes}) self._state_program = self._state_snapshot() # Contract shapes with their batch entries pre-parsed; _resolve_shapes # turns these into concrete expectations once per distinct batch size. @@ -124,6 +238,13 @@ def __init__(self, plan, modules, device, receipt, port_orders): self._spec_out = {node.id: {port: _compile_shape(shape) for port, shape in node.output_shapes.items()} for node in plan.nodes} self._compiled = _UNCOMPILED + self._fast = self._build_fast_program() + # Input signatures whose calls passed every check, and the watch + # generation at which the registered state last matched the build + # (``None`` forces a full re-check before the next call). + self._validated = {} + _watch(self._state_program) + self._state_seen = _watch_generation self.build_receipt = MappingProxyType(receipt) def __setattr__(self, name, value): @@ -153,9 +274,13 @@ def __deepcopy__(self, memo): # instance dictionary directly, exactly as unpickling would. object.__setattr__(result, name, value if name in SHARED_METADATA else copy.deepcopy(value, memo)) - # Cheap insurance: the clone rebuilds its baked program against its own - # modules rather than trusting a structure copied mid-flight. + # Cheap insurance: the clone rebuilds its baked programs against its own + # modules rather than trusting a structure copied mid-flight, validates + # its first call afresh, and watches its own modules for registrations. object.__setattr__(result, "_compiled", _UNCOMPILED) + object.__setattr__(result, "_fast", result._build_fast_program()) + object.__setattr__(result, "_validated", {}) + _watch(result._state_program) return result def __copy__(self): @@ -180,6 +305,11 @@ def _apply(self, fn, recurse=True): # only; a probe that came back integral means no compute-dtype change. if probe.is_floating_point() or probe.is_complex(): self._runtime_dtype = probe.dtype + # Every validated signature named the old device and dtype. A move is + # also a natural point to re-check registered state in full, which + # catches the edits no registration hook reports (see _state_seen). + self._validated = {} + self._state_seen = None return self def _effective_dtype(self, dtype): @@ -193,6 +323,92 @@ def _effective_dtype(self, dtype): """ return self._runtime_dtype if dtype is not None and dtype == self._build_dtype else dtype + # -- errors ------------------------------------------------------------- + + def _node_name(self, node_id): + """``linear 'head'``: the operation and the node, as an error names them.""" + return f"{self._node_labels[node_id][0]} {node_id!r}" + + def _error(self, message, node_id=None, code="E_RUNTIME"): + if node_id is None: + return HNDLError(code, message) + _, line, column = self._node_labels[node_id] + return HNDLError(code, message, node=node_id, line=line, column=column) + + def _port_error(self, port, value, dtype, batch, problem): + """An ``E_RUNTIME`` for one port, with the contract beside the tensor.""" + where, spec, node_id = port + lines = [f"{where}: {problem}"] + if batch is not None and batch <= 0: + batch = None # an empty batch binds no B worth printing + if isinstance(value, torch.Tensor): + lines.append(f" expected {_contract_text(spec, batch)}:{_dtype_text(dtype)} on {self._runtime_device}") + lines.append(f" got {_tensor_text(value)}") + first = self._input_names[0] + if (batch is not None and where != f"input {first!r}" and value.ndim == len(spec) + and any(type(d) is tuple and d[0] is not None and size != d[0] * batch + for d, size in zip(spec, value.shape))): + lines.append(f" B={batch} is this call's batch size, read from input {first!r}") + return self._error("\n".join(lines), node_id) + + def _annotate_failure(self, node_id, error, bound, batch, after_validation): + """Add HNDL's context to an exception out of one node, as a note. + + The exception itself propagates unchanged --- same type, same message, + same traceback --- so callers that catch ``RuntimeError`` (or an + out-of-memory error, to shrink the batch) keep working; the note is + what the traceback prints after the message. Everything here runs only + after a failure, and an exception passing out through several graphs + (a graph used as a node of another) keeps the innermost node's note. + """ + if isinstance(error, HNDLError) and error.node is not None: + return # an HNDL check already named its node + if any(isinstance(note, str) and note.startswith(_NOTE_PREFIX) + for note in getattr(error, "__notes__", ())): + return + try: + error.add_note(self._failure_note(node_id, error, bound, batch, after_validation)) + except Exception: # noqa: BLE001 - a note must never replace the real error + pass + + def _failure_note(self, node_id, error, bound, batch, after_validation): + label, line, column = self._node_labels[node_id] + position = "".join(f", {key} {value}" for key, value in (("line", line), ("column", column)) + if value is not None) + lines = [f"{_NOTE_PREFIX} raised inside node {node_id!r} ({label}{position})"] + if not self._state_matches(): + # State removed between calls breaks forward rather than reporting + # itself; say so, because that is the cause the error hides. + _, listed = self._state_summary() + lines.append(f" registered state no longer matches the build: {listed}") + elif after_validation: + # The signature was validated earlier, so a node input that now + # breaks its contract means an upstream module changed its output. + compiled = self._compiled + if batch != compiled[0] or self._runtime_dtype is not compiled[1]: + compiled = self._resolve_shapes(batch) + entry = next(entry for entry in compiled[3] if entry[5] == node_id) + for (_, expected, dtype, port), value in zip(entry[1], bound): + try: + self._check(value, expected, dtype, port, batch) + except HNDLError as mismatch: + lines.extend(" " + text for text in mismatch.message.splitlines()) + lines.append(" (an earlier call with the same input signature passed every check, " + "so an upstream module now produces a different tensor)") + break + effective = self._effective_dtype + for port, value in zip(self._port_orders[node_id], bound): + spec = self._spec_in[node_id][port] + dtype = effective(self._port_dtypes[node_id][port]) + lines.append(f" input {port!r}: got {_tensor_text(value)}; contract " + f"{_contract_text(spec, batch)}:{_dtype_text(dtype)} on {self._runtime_device}") + if type(error).__module__.startswith(("torch._dynamo", "torch._inductor")): + lines.append(" the error came from torch.compile compiling or running this node's module; " + "run it without torch.compile to see the eager error") + return "\n".join(lines) + + # -- registered state --------------------------------------------------- + def _state_snapshot(self): """Record what ``named_parameters``/``named_buffers`` see, per module. @@ -200,18 +416,23 @@ def _state_snapshot(self): things: the order ``modules()`` visits, each module's path from the root, and the keys each module contributes. Freezing the visited modules in one flat tuple --- alongside the child mapping that fixes - every path --- lets :meth:`_execute` re-derive the same answer without - recursing through ``named_modules`` or building a single string. - - Each row is ``(module, parameter_keys, buffer_keys, children)``. - ``children`` is the module's ``_modules`` mapping as key/value pairs, - so a submodule swapped out under an unchanged key is still a change. - The key tuples are filtered exactly as PyTorch filters them: ``None`` - slots are skipped, and a tensor already seen earlier in the walk is - skipped, with parameters and buffers de-duplicated independently. + every path --- lets :meth:`_state_matches` re-derive the same answer + without recursing through ``named_modules`` or building a single string. + + Each row is ``(module, parameter_keys, buffer_keys, children, path, + node_id)``. ``children`` is the module's ``_modules`` mapping as + key/value pairs, so a submodule swapped out under an unchanged key is + still a change. The key tuples are filtered exactly as PyTorch filters + them: ``None`` slots are skipped, and a tensor already seen earlier in + the walk is skipped, with parameters and buffers de-duplicated + independently. ``path`` and ``node_id`` are only read to word errors. """ + owners = {} + for key, layer in self.nodes._modules.items(): + for module in layer.modules(): + owners.setdefault(id(module), key.removeprefix("n_")) program, seen_parameters, seen_buffers = [], set(), set() - for module in self.modules(): + for path, module in self.named_modules(): keys = [] for store, seen in ((module._parameters, seen_parameters), (module._buffers, seen_buffers)): @@ -222,9 +443,98 @@ def _state_snapshot(self): seen.add(id(value)) contributed.append(key) keys.append(tuple(contributed)) - program.append((module, keys[0], keys[1], tuple(module._modules.items()))) + program.append((module, keys[0], keys[1], tuple(module._modules.items()), + path, owners.get(id(module)))) return tuple(program) + def _state_matches(self): + """Whether every module still registers exactly what it did at build. + + Replays the build-time walk against the flat program instead of + recursing through the module tree again. Pinning every module's child + mapping keeps the two trees identical --- nothing can be grafted in + unseen --- so comparing each module's contributed keys answers exactly + what comparing the dotted name tuples used to answer. ``set.add`` + returns None, so its clause always passes and only records the tensor. + """ + seen_parameters, seen_buffers = set(), set() + for module, parameters, buffers, children, _, _ in self._state_program: + if (tuple(module._modules.items()) != children + or tuple([key for key, value in module._parameters.items() + if value is not None and id(value) not in seen_parameters + and not seen_parameters.add(id(value))]) != parameters + or tuple([key for key, value in module._buffers.items() + if value is not None and id(value) not in seen_buffers + and not seen_buffers.add(id(value))]) != buffers): + return False + return True + + def _state_changes(self): + """What differs from the build-time walk, as ``(node_id, text)`` pairs.""" + changes = [] + seen = {"parameter": set(), "buffer": set()} + for module, parameters, buffers, children, path, node_id in self._state_program: + prefix = f"{path}." if path else "" + for kind, store, recorded in (("parameter", module._parameters, parameters), + ("buffer", module._buffers, buffers)): + ids, current = seen[kind], [] + for key, value in store.items(): + if value is not None and id(value) not in ids: + ids.add(id(value)) + current.append(key) + if tuple(current) == recorded: + continue + for key in recorded: + if key not in current: + how = ("removed" if key not in store else "set to None" if store[key] is None + else "re-registered as an alias of an earlier tensor") + changes.append((node_id, f"{kind} '{prefix}{key}' was {how}")) + changes.extend((node_id, f"{kind} '{prefix}{key}' was added") + for key in current if key not in recorded) + if set(current) == set(recorded): + changes.append((node_id, f"the {kind}s of '{path}' were re-registered in a different order")) + now, before = module._modules, dict(children) + if tuple(now.items()) == children: + continue + for key, child in before.items(): + if key not in now: + changes.append((node_id, f"submodule '{prefix}{key}' was removed")) + elif now[key] is not child: + how = (f"filled in with {_module_kind(now[key])} (it was None at build)" if child is None + else f"replaced ({_module_kind(child)} -> {_module_kind(now[key])})") + changes.append((node_id, f"submodule '{prefix}{key}' was {how}")) + changes.extend((node_id, f"submodule '{prefix}{key}' was added ({_module_kind(child)})") + for key, child in now.items() if key not in before) + if set(now) == set(before) and all(now[key] is child for key, child in before.items()): + changes.append((node_id, f"the submodules of '{path}' were re-registered in a different order")) + return changes + + def _state_summary(self): + """The first node whose registered state changed, and every change, as text.""" + changes = self._state_changes() or [(None, "the module tree differs from the one recorded at build")] + node_id = next((owner for owner, _ in changes if owner is not None), None) + listed = "; ".join(text for _, text in changes[:6]) + if len(changes) > 6: + listed += f"; and {len(changes) - 6} more" + return node_id, listed + + def _state_error(self, when): + """``E_RUNTIME`` naming the node whose registered state changed, and how.""" + node_id, listed = self._state_summary() + subject = self._node_name(node_id) if node_id is not None else "the network" + phrase = "changed during forward" if when == "forward" else "no longer matches the build" + lines = [f"registered state of {subject} {phrase}: {listed}"] + lines.append(" HNDL fixes every parameter, buffer and submodule when it builds a plan, so state " + "added or removed later is not trained, seeded or saved as the plan describes; " + "create it in __init__, or resolve and build a new plan to change the architecture") + return self._error("\n".join(lines), node_id) + + def _verify_state(self, when): + if not self._state_matches(): + raise self._state_error(when) + + # -- the checked program ------------------------------------------------ + @staticmethod def _resolve(spec, batch): """A compiled shape as concrete sizes, or the entry that is not a batch axis. @@ -263,41 +573,100 @@ def _resolve_shapes(self, batch): compiled = ( batch, self._runtime_dtype, tuple((name, f"input:{name}", resolve(self._spec_inputs[name], batch), - effective(self._input_dtypes[name])) + effective(self._input_dtypes[name]), + (f"input {name!r}", self._spec_inputs[name], None)) for name in self._input_names), self._build_program(expected_in, expected_out, effdt), tuple((name, entry["ref"], resolve(self._spec_outputs[name], batch), - effective(self._output_dtypes[name])) + effective(self._output_dtypes[name]), + (f"output {name!r}", self._spec_outputs[name], None)) for name, entry in self.plan.outputs.items())) self._compiled = compiled return compiled def _build_program(self, expected_in, expected_out, effdt): """Bake the per-call constants --- labels, module handles, port sets --- into tuples.""" - return tuple( - (self.nodes[f"n_{node.id}"], - tuple((node.inputs[port], expected_in[node.id][port], f"{node.id}/{port}", - effdt[node.id][port]) - for port in self._port_orders[node.id]), - tuple((f"node:{node.id}/{port}", expected_out[node.id][port], f"{node.id}/{port}", - effdt[node.id][port]) - for port in node.outputs), - node.outputs, frozenset(node.outputs), node.id) - for node in self.plan.nodes) - - def _check(self, value, expected, location, dtype): + program = [] + for node in self.plan.nodes: + name = self._node_name(node.id) + program.append(( + self.nodes[f"n_{node.id}"], + tuple((node.inputs[port], expected_in[node.id][port], effdt[node.id][port], + (f"{name} input {port!r} (from {node.inputs[port]})", + self._spec_in[node.id][port], node.id)) + for port in self._port_orders[node.id]), + tuple((f"node:{node.id}/{port}", expected_out[node.id][port], effdt[node.id][port], + (f"{name} output {port!r}", self._spec_out[node.id][port], node.id)) + for port in node.outputs), + node.outputs, frozenset(node.outputs), node.id)) + return tuple(program) + + def _build_fast_program(self): + """The unchecked loop: every tensor lives in a numbered slot. + + External inputs take the first slots in declaration order and every + node output the next free one, so a node's arguments are list indices + rather than string-keyed lookups. Each row is ``(module, arguments, + single output slot or None, node)``: ``arguments`` is a bare slot for + a one-input node, the common case, and a tuple of slots otherwise; + ``node`` holds what only the multi-output and failure paths read. + """ + slots = {f"input:{name}": index for index, name in enumerate(self._input_names)} + program = [] + for node in self.plan.nodes: + ins = tuple(slots[node.inputs[port]] for port in self._port_orders[node.id]) + outs = [] + for port in node.outputs: + slots[f"node:{node.id}/{port}"] = len(slots) + outs.append(slots[f"node:{node.id}/{port}"]) + program.append((self.nodes[f"n_{node.id}"], ins[0] if len(ins) == 1 else ins, + outs[0] if len(outs) == 1 else None, + (ins, tuple(outs), node.outputs, frozenset(node.outputs), node.id))) + padding = (None,) * (len(slots) - len(self._input_names)) + outputs = tuple((name, slots[entry["ref"]]) for name, entry in self.plan.outputs.items()) + return padding, tuple(program), outputs + + def _check(self, value, expected, dtype, port, batch): if not isinstance(value, torch.Tensor): - raise HNDLError("E_RUNTIME", f"{location} must be a tensor") + raise self._port_error(port, value, dtype, batch, f"expected a tensor, got {type(value).__name__}") if type(expected) is str: - raise HNDLError("E_RUNTIME", f"{location}: contract entry {expected!r} is not a batch dimension") + raise self._port_error(port, value, dtype, batch, + f"contract entry {expected!r} is not a batch dimension") if value.shape != expected or value.ndim == 0 or value.shape[0] <= 0: - raise HNDLError("E_RUNTIME", f"{location}: expected shape {expected}, got {tuple(value.shape)}") + if value.ndim == 0: + problem = f"expected shape {_contract_text(port[1], batch)}, got a 0-dimensional tensor" + elif value.shape[0] <= 0 and value.ndim == len(expected): + problem = f"batch size must be positive, got shape {_shape_text(tuple(value.shape))}" + else: + problem = (f"expected shape {_contract_text(port[1], batch)}, " + f"got {_shape_text(tuple(value.shape))}") + raise self._port_error(port, value, dtype, batch, problem) if dtype is not None and value.dtype != dtype: - raise HNDLError("E_RUNTIME", f"{location}: expected dtype {dtype}, got {value.dtype}") + raise self._port_error(port, value, dtype, batch, + f"expected dtype {_dtype_text(dtype)}, got {_dtype_text(value.dtype)}") if value.device != self._runtime_device: - raise HNDLError("E_RUNTIME", f"{location}: expected device {self._runtime_device}, got {value.device}") - - def _execute(self, inputs): + raise self._port_error(port, value, dtype, batch, + f"expected device {self._runtime_device}, got {value.device}") + + def _split_result(self, result, out_ports, out_set, node_id): + """A module's return value as one value per declared output port.""" + if isinstance(result, dict) and set(result) == out_set: + return tuple(result[port] for port in out_ports) + if isinstance(result, (tuple, list)) and len(result) == len(out_ports): + return tuple(result) + if len(out_ports) == 1: + return (result,) + if isinstance(result, dict): + got = f"a dict with keys {tuple(result)}" + elif isinstance(result, (tuple, list)): + got = f"a {type(result).__name__} of {len(result)} values" + else: + got = f"one {type(result).__name__}" + raise self._error(f"{self._node_name(node_id)} returned {got}; expected output ports {out_ports}, " + "as a tuple or list in that order or a dict with exactly those keys", node_id) + + def _run_checked(self, inputs, compiling): + """One call with every port checked before and after its node.""" leading = inputs[self._input_names[0]] batch = leading.shape[0] if isinstance(leading, torch.Tensor) and leading.ndim else None # The resolved expectations depend on the batch and on the dtype casts @@ -309,54 +678,112 @@ def _execute(self, inputs): _, _, input_program, node_program, output_program = compiled check = self._check values = {} - for name, key, expected, dtype in input_program: + for name, key, expected, dtype, port in input_program: value = inputs[name] - check(value, expected, key, dtype) + check(value, expected, dtype, port, batch) values[key] = value for module, ins, outs, out_ports, out_set, node_id in node_program: bound = [] - for ref, expected, location, dtype in ins: + for ref, expected, dtype, port in ins: value = values[ref] - check(value, expected, location, dtype) + check(value, expected, dtype, port, batch) bound.append(value) - result = module(*bound) - if isinstance(result, dict) and set(result) == out_set: - results = tuple(result[port] for port in out_ports) - elif isinstance(result, (tuple, list)) and len(result) == len(out_ports): - results = tuple(result) - elif len(out_ports) == 1: - results = (result,) + if compiling: + # No handler for dynamo to trace; a failure while compiling + # surfaces as the compiler reports it. + result = module(*bound) else: - raise HNDLError("E_RUNTIME", f"{node_id}: expected output ports {out_ports}") - for (key, expected, location, dtype), value in zip(outs, results): - check(value, expected, location, dtype) + try: + result = module(*bound) + except Exception as error: + self._annotate_failure(node_id, error, bound, batch, after_validation=False) + raise + for (key, expected, dtype, port), value in zip( + outs, self._split_result(result, out_ports, out_set, node_id)): + check(value, expected, dtype, port, batch) values[key] = value outputs = {} - for name, ref, expected, dtype in output_program: + for name, ref, expected, dtype, port in output_program: value = values[ref] - check(value, expected, name, dtype) + check(value, expected, dtype, port, batch) outputs[name] = value - # Replay the build-time walk against the flat program instead of - # recursing through the module tree again. Pinning every module's child - # mapping keeps the two trees identical --- nothing can be grafted in - # unseen --- so comparing each module's contributed keys answers exactly - # what comparing the dotted name tuples used to answer. ``set.add`` - # returns None, so its clause always passes and only records the tensor. - seen_parameters, seen_buffers = set(), set() - for module, parameters, buffers, children in self._state_program: - if (tuple(module._modules.items()) != children - or tuple([key for key, value in module._parameters.items() - if value is not None and id(value) not in seen_parameters - and not seen_parameters.add(id(value))]) != parameters - or tuple([key for key, value in module._buffers.items() - if value is not None and id(value) not in seen_buffers - and not seen_buffers.add(id(value))]) != buffers): - raise HNDLError("E_RUNTIME", "A module created or removed registered state during forward") + return outputs + + def _signature(self, inputs): + """What a validated call is keyed on, or ``None`` for a non-tensor input.""" + key = [self.training, _autocast_enabled()] + for name in self._input_names: + value = inputs[name] + if not isinstance(value, torch.Tensor): + return None + key += (value.shape, value.dtype, value.device) + return tuple(key) + + def _validate(self, inputs, signature, generation): + """The first call for a signature: check everything, then remember it.""" + if generation != self._state_seen: + # A parameter, buffer or submodule was registered on a module of + # this graph since the last check: confirm the state still matches + # the build, and forget every validated signature, because a + # parameter replaced under the same name can change any shape. + self._verify_state("build") + self._validated = {} + self._state_seen = generation + outputs = self._run_checked(inputs, compiling=False) + generation = _watch_generation + self._verify_state("forward") + self._state_seen = generation + validated = self._validated + if len(validated) >= _MAX_SIGNATURES: + validated.clear() + validated[signature] = True + return outputs + + def _execute(self, inputs): + if _is_compiling(): + # Traced once per guard set, so the full checks cost nothing per + # call there and dynamo turns them into guards. + outputs = self._run_checked(inputs, compiling=True) + self._verify_state("forward") + return outputs + generation = _watch_generation + signature = self._signature(inputs) + if generation != self._state_seen or signature not in self._validated: + return self._validate(inputs, signature, generation) + padding, program, output_slots = self._fast + values = [*map(inputs.__getitem__, self._input_names), *padding] + try: + for module, ins, out, node in program: + result = (module(values[ins]) if type(ins) is int + else module(*[values[slot] for slot in ins])) + if out is not None and type(result) is _Tensor: + values[out] = result + else: + _, outs, out_ports, out_set, node_id = node + for slot, value in zip(outs, self._split_result(result, out_ports, out_set, node_id)): + values[slot] = value + except Exception as error: + ins, _, _, _, node_id = node + leading = values[0] + self._annotate_failure(node_id, error, [values[slot] for slot in ins], + leading.shape[0] if leading.ndim else None, after_validation=True) + raise + outputs = {name: values[slot] for name, slot in output_slots} + if _watch_generation != generation: + # Something registered state while this call ran --- perhaps one + # of its own modules. Check it now, as the first call would have. + generation = _watch_generation + self._verify_state("forward") + self._state_seen = generation return outputs def _bind(self, args, kwargs): """Bind runtime tensors to the declared external inputs, in order.""" names = self._input_names + if not kwargs and len(args) == len(names): + return dict(zip(names, args)) + if not args and kwargs.keys() == self._input_set: + return kwargs expected = "Expected exactly the external tensor inputs " + ", ".join(repr(n) for n in names) if len(args) > len(names): raise HNDLError("E_BINDING", expected) diff --git a/tests/benchmark/test_parity_vs_handwritten_pytorch.py b/tests/benchmark/test_parity_vs_handwritten_pytorch.py index effd02b..a6adff1 100644 --- a/tests/benchmark/test_parity_vs_handwritten_pytorch.py +++ b/tests/benchmark/test_parity_vs_handwritten_pytorch.py @@ -12,21 +12,25 @@ convolutional discriminator trunk, and a transformer block. The comparison is not free of overhead on purpose: ``GraphModule._execute`` -validates every node's shape, dtype and device on each forward call (see -``_check`` in ``src/hndl/torch.py``), which a hand-written module never does. +validates every node's shape, dtype and device on the first call for each +input signature (see ``_run_checked`` in ``src/hndl/torch.py``), which a +hand-written module never does, and on every later call compares that +signature before running the modules from a slot-indexed program. The timing +loop repeats one signature, so it measures the second, cheap path. :data:`TOLERANCE` is the generous multiple of hand-written time that overhead is allowed to cost. These assertions are meant to fail loudly rather than skip if that overhead ever grows unreasonable. -That overhead is a roughly **constant** cost per forward call — about 5 us per -node on this machine, so ~27 us for the five-node MLP, independent of batch -size — which means the ratio a case reports depends on how much arithmetic the -batch gives it to amortize against. The MLP is timed at batch 256 for that -reason; at batch 32 the same network measures about 1.6x, which is the fixed -overhead weighing on a 40 us forward pass rather than a per-element slowdown. -(Both figures were ~2.5x larger before the resolved-shape and baked-program -caches landed: the per-call cost used to be ~13 us per node and batch 32 -measured ~2.5x.) +That overhead is a roughly **constant** cost per forward call — about 2 us per +call plus well under 1 us per node on this machine, so ~6-9 us for the +five-node MLP, independent of batch size — which means the ratio a case +reports depends on how much arithmetic the batch gives it to amortize against. +The MLP is timed at batch 256, where the fixed cost was once large; at batch +32 the same network now measures about 1.15-1.2x. (Before validation moved to +once per signature, every call checked every port and replayed the +registered-state walk: ~5 us per node after the resolved-shape and +baked-program caches, ~13 us before them, and batch 32 measured ~1.7x and +~2.5x respectively.) Run with ``-s`` to see each case's two timings and their ratio. """ diff --git a/tests/test_torch.py b/tests/test_torch.py index 85b29da..587888b 100644 --- a/tests/test_torch.py +++ b/tests/test_torch.py @@ -388,11 +388,17 @@ def test_deepcopy_keeps_lookup_moves_and_the_state_consistency_check(): assert clone._runtime_device == torch.device("cpu") x = torch.randn(2, 4) clone.eval()(x) - module, _, buffers, children = clone._state_program[-1] - object.__setattr__(clone, "_state_program", - clone._state_program[:-1] + ((module, ("ghost",), buffers, children),)) - with pytest.raises(HNDLError, match="E_RUNTIME.*registered state"): + model.eval()(x) + # The clone records and watches its own modules: state registered on one + # of them after a validated call is caught on the clone's next call, and + # the original, whose modules are untouched, keeps running. + clone["head"].ghost = nn.Parameter(torch.zeros(1)) + with pytest.raises(HNDLError, match="E_RUNTIME.*registered state.*'nodes.n_head.ghost' was added"): clone(x) + assert model(x).shape == (2, 2) + model["norm"].ghost = nn.Parameter(torch.zeros(1)) + with pytest.raises(HNDLError, match="E_RUNTIME.*registered state"): + model(x) def test_state_registered_during_forward_is_rejected(): @@ -551,24 +557,24 @@ def test_floating_point_casts_move_the_runtime_dtype_checks(cast, dtype): assert all(p.dtype == dtype for p in model.parameters()) result = _forward_or_skip(model, torch.randn(3, 4, dtype=dtype)) assert result.dtype == dtype and result.shape == (3, 2) - with pytest.raises(HNDLError, match=f"E_RUNTIME.*expected dtype {dtype}, got torch.float32"): + with pytest.raises(HNDLError, match=f"E_RUNTIME.*expected dtype {str(dtype).removeprefix('torch.')}, got float32"): model(torch.randn(3, 4)) restored = model.float() assert restored is model assert model(torch.randn(3, 4)).dtype == torch.float32 - with pytest.raises(HNDLError, match="E_RUNTIME.*expected dtype torch.float32, got torch.float64"): + with pytest.raises(HNDLError, match="E_RUNTIME.*expected dtype float32, got float64"): model(torch.randn(3, 4, dtype=torch.float64)) def test_double_tracks_parameterless_graphs_and_device_only_moves(): model = network("relu()", input_shape=("B", 4), output_shape=("B", 4), device="cpu").double() assert model(torch.randn(2, 4, dtype=torch.float64)).dtype == torch.float64 - with pytest.raises(HNDLError, match="E_RUNTIME.*expected dtype torch.float64"): + with pytest.raises(HNDLError, match="E_RUNTIME.*expected dtype float64"): model(torch.randn(2, 4)) unmoved = _mlp(device="cpu").to("cpu") assert unmoved(torch.randn(2, 4)).dtype == torch.float32 - with pytest.raises(HNDLError, match="E_RUNTIME.*expected dtype torch.float32"): + with pytest.raises(HNDLError, match="E_RUNTIME.*expected dtype float32"): unmoved(torch.randn(2, 4, dtype=torch.float64)) @@ -577,9 +583,9 @@ def test_casts_leave_integer_and_declared_dtypes_alone(): output_shape=("B", 5, 3), input_dtype="int64", device="cpu").double() tokens = torch.randint(0, 20, (2, 5)) assert model(tokens).dtype == torch.float64 - with pytest.raises(HNDLError, match="E_RUNTIME.*expected dtype torch.int64, got torch.int32"): + with pytest.raises(HNDLError, match="E_RUNTIME.*expected dtype int64, got int32"): model(tokens.to(torch.int32)) - with pytest.raises(HNDLError, match="E_RUNTIME.*expected dtype torch.int64, got torch.float64"): + with pytest.raises(HNDLError, match="E_RUNTIME.*expected dtype int64, got float64"): model(tokens.to(torch.float64)) @@ -587,10 +593,285 @@ def test_deepcopy_carries_and_isolates_the_runtime_dtype(): model = _mlp(device="cpu") clone = copy.deepcopy(model.double()) assert clone(torch.randn(2, 4, dtype=torch.float64)).dtype == torch.float64 - with pytest.raises(HNDLError, match="E_RUNTIME.*expected dtype torch.float64"): + with pytest.raises(HNDLError, match="E_RUNTIME.*expected dtype float64"): clone(torch.randn(2, 4)) original = _mlp(device="cpu") copy.deepcopy(original).double() assert original(torch.randn(2, 4)).dtype == torch.float32 assert original._input_dtypes["x"] == torch.float32 and model._input_dtypes["x"] == torch.float32 + + +# -- Validation once per input signature --------------------------------------- + + +def _count_checks(model): + """Count port checks on ``model``: zero on a call means it took the fast path.""" + calls = [] + check = model._check + + def counting(*args): + calls.append(args) + return check(*args) + + model._check = counting + return calls + + +def test_ports_are_validated_once_per_input_signature(): + model = _stateful(device="cpu") + checks = _count_checks(model) + x = torch.randn(3, 4) + expected = model(x) + assert len(checks) == 1 + 2 * 4 + 1 # external input, every node's ports, output + checks.clear() + torch.testing.assert_close(model(x), expected) + assert model(torch.randn(3, 4)).shape == (3, 2) + assert checks == [] + model(torch.randn(5, 4)) # a new batch size is a new signature + assert checks + checks.clear() + model(torch.randn(3, 4)) # ... and the earlier one is still remembered + assert checks == [] + model.eval()(x) # so is a new training mode + assert checks + checks.clear() + model.eval()(x) + assert checks == [] + model.to("cpu")(x) # a move or cast forgets everything validated + assert checks + checks.clear() + nn.Linear(3, 3).register_parameter("elsewhere", nn.Parameter(torch.zeros(1))) + model(x) # modules outside the graph registering state do not matter + assert checks == [] + + +def test_a_broken_input_after_the_first_call_is_reported_readably(): + model = network('linear(8, name="hidden")\nrelu()\nlinear(2, name="head")', + input_shape=("B", 4), output_shape=("B", 2), device="cpu") + model(torch.randn(3, 4)) + with pytest.raises(HNDLError) as caught: + model(torch.randn(3, 5)) + assert str(caught.value) == ("E_RUNTIME: input 'x': expected shape [B=3, 4], got [3, 5]\n" + " expected [B=3, 4]:float32 on cpu\n" + " got [3, 5]:float32 on cpu") + with pytest.raises(HNDLError, match="^E_RUNTIME: input 'x': expected dtype float32, got float64\n"): + model(torch.randn(3, 4, dtype=torch.float64)) + with pytest.raises(HNDLError, match="^E_RUNTIME: input 'x': expected a tensor, got list$"): + model([[0.0] * 4] * 3) + with pytest.raises(HNDLError, match="^E_RUNTIME: input 'x': batch size must be positive, got shape " + r"\[0, 4\]\n expected \[B, 4\]:float32"): + model(torch.randn(0, 4)) + assert model(torch.randn(3, 4)).shape == (3, 2) + + +def test_a_second_input_names_the_batch_it_disagrees_with(): + model = network("a = linear(x, 4)\nconcat(a, y, axis=1)", input_shape={"x": ("B", 3), "y": ("B", 2)}, + output_shape=("B", 6), device="cpu") + model(torch.randn(3, 3), torch.randn(3, 2)) + with pytest.raises(HNDLError) as caught: + model(torch.randn(3, 3), torch.randn(2, 2)) + assert str(caught.value).splitlines() == [ + "E_RUNTIME: input 'y': expected shape [B=3, 2], got [2, 2]", + " expected [B=3, 2]:float32 on cpu", + " got [2, 2]:float32 on cpu", + " B=3 is this call's batch size, read from input 'x'", + ] + + +def _flaky_registry(fail_from_call, error=None): + """An operator that works until call ``fail_from_call`` and then raises.""" + registry = Registry.builtins() + + @registry.operator("flaky", identity="tests.flaky", summary="Fail after a few calls.", + shape="x[B, F] -> out[B, F]") + class Flaky(nn.Module): + def __init__(self): + super().__init__() + self.calls = 0 + + def forward(self, x): + self.calls += 1 + if self.calls >= fail_from_call: + if error is not None: + raise error + return x @ torch.ones(3, 3) + return x + + return registry + + +def _hndl_notes(error): + """The notes HNDL added to an exception raised inside a node.""" + return [note for note in getattr(error, "__notes__", ()) if note.startswith("HNDL:")] + + +@pytest.mark.parametrize("fail_from_call", [1, 2]) +def test_a_torch_error_inside_a_node_keeps_its_type_and_names_the_node(fail_from_call): + """Raised during the validating first call or on the unchecked path, it reads the same.""" + model = network('linear(4)\nflaky(name="odd")\nrelu()', input_shape=("B", 4), output_shape=("B", 4), + device="cpu", registry=_flaky_registry(fail_from_call)) + for _ in range(fail_from_call - 1): + model(torch.randn(2, 4)) + with pytest.raises(RuntimeError, match=r"^mat1 and mat2 shapes cannot be multiplied \(2x4 and 3x3\)\n") as caught: + model(torch.randn(2, 4)) + assert type(caught.value) is RuntimeError + assert _hndl_notes(caught.value) == [ + "HNDL: raised inside node 'odd' (flaky, line 2, column 1)\n" + " input 'x': got [2, 4]:float32 on cpu; contract [B=2, 4]:float32 on cpu"] + + +def test_errors_callers_catch_by_type_keep_their_type(): + registry = _flaky_registry(2, error=torch.cuda.OutOfMemoryError("CUDA out of memory")) + model = network("flaky()", input_shape=("B", 4), output_shape=("B", 4), device="cpu", registry=registry) + model(torch.randn(2, 4)) + with pytest.raises(torch.cuda.OutOfMemoryError, match="^CUDA out of memory\n") as caught: + model(torch.randn(2, 4)) + assert _hndl_notes(caught.value)[0].startswith("HNDL: raised inside node 'n0' (flaky, line 1, column 1)") + + +def test_an_hndl_error_inside_a_node_keeps_its_message_and_gains_a_note(): + registry = _flaky_registry(1, error=HNDLError("E_RUNTIME", "extent 5 does not divide by 2")) + model = network('flaky(name="split_here")', input_shape=("B", 4), output_shape=("B", 4), + device="cpu", registry=registry) + with pytest.raises(HNDLError) as caught: + model(torch.randn(2, 4)) + assert str(caught.value) == "E_RUNTIME: extent 5 does not divide by 2" + assert _hndl_notes(caught.value)[0].startswith("HNDL: raised inside node 'split_here' (flaky, line 1") + + +def test_an_error_passing_out_through_nested_graphs_keeps_one_note(): + """A built graph used as a node of another: only the innermost node is named.""" + registry = _flaky_registry(2) + inner = network('linear(4)\nflaky(name="deep")', input_shape=("B", 4), output_shape=("B", 4), + device="cpu", registry=registry) + + @registry.operator("wrapped", identity="tests.wrapped", summary="Run an inner graph.", + shape="x[B, F] -> out[B, F]") + class Wrapped(nn.Module): + def __init__(self): + super().__init__() + self.inner = copy.deepcopy(inner) + + def forward(self, x): + return self.inner(x) + + outer = network('wrapped(name="outer_node")', input_shape=("B", 4), output_shape=("B", 4), + device="cpu", registry=registry) + outer(torch.randn(2, 4)) + with pytest.raises(RuntimeError) as caught: + outer(torch.randn(2, 4)) + notes = _hndl_notes(caught.value) + assert len(notes) == 1 and notes[0].startswith("HNDL: raised inside node 'deep' (flaky, line 2") + + +def test_an_upstream_node_changing_its_output_after_validation_is_reported_at_the_port(): + registry = Registry.builtins() + + @registry.operator("drifts", identity="tests.drifts", summary="Narrow its output after one call.", + shape="x[B, F] -> out[B, F]") + class Drifts(nn.Module): + def __init__(self): + super().__init__() + self.calls = 0 + + def forward(self, x): + self.calls += 1 + return x if self.calls == 1 else x[:, :2] + + model = network('drifts(name="d")\nlinear(4, name="head")', input_shape=("B", 4), output_shape=("B", 4), + device="cpu", registry=registry) + model(torch.randn(2, 4)) + with pytest.raises(RuntimeError, match="^mat1 and mat2") as caught: + model(torch.randn(2, 4)) + assert _hndl_notes(caught.value)[0].splitlines()[:2] == [ + "HNDL: raised inside node 'head' (linear, line 2, column 1)", + " linear 'head' input 'x' (from node:d/out): expected shape [B=2, 4], got [2, 2]", + ] + + +def _replaces_the_head_weight(model): + model["head"].weight = nn.Parameter(torch.randn(2, 7)) + + +@pytest.mark.parametrize("mutate, raised, message", [ + # Registrations the hooks see are HNDL's own finding, before any node runs. + (lambda model: setattr(model["hidden"], "extra", nn.Linear(2, 2)), HNDLError, + "E_RUNTIME (node hidden; line 1, column 1): registered state of linear 'hidden' no longer matches " + "the build: submodule 'nodes.n_hidden.extra' was added (Linear)"), + (lambda model: model["hidden"].register_parameter("scale", nn.Parameter(torch.ones(1))), HNDLError, + "registered state of linear 'hidden' no longer matches the build: " + "parameter 'nodes.n_hidden.scale' was added"), + (lambda model: setattr(model["norm"], "running_mean", None), HNDLError, + "registered state of batch_norm 'norm' no longer matches the build: " + "buffer 'nodes.n_norm.running_mean' was set to None"), + # A deletion runs no hook: the node's own error surfaces, explained by a note. + (lambda model: delattr(model["head"], "weight"), AttributeError, + "HNDL: raised inside node 'head' (linear, line 1, column 60)\n" + " registered state no longer matches the build: parameter 'nodes.n_head.weight' was removed"), + (_replaces_the_head_weight, RuntimeError, "mat1 and mat2 shapes cannot be multiplied (3x5 and 7x2)"), +]) +def test_a_submodule_or_parameter_mutated_after_the_first_call_is_reported(mutate, raised, message): + model = _stateful(device="cpu").eval() + x = torch.randn(3, 4) + model(x) + mutate(model) + for _ in range(2): # reported on every call until the state is repaired + with pytest.raises(raised) as caught: + model(x) + assert type(caught.value) is raised + assert message in "\n".join([str(caught.value), *_hndl_notes(caught.value)]) + + +def test_a_parameter_set_to_none_is_reported_at_the_next_revalidation(): + """PyTorch runs no registration hook for ``module.param = None``. + + So the unchecked path keeps running --- here ``linear`` without its bias --- + until something revalidates the graph: a new input signature, a move or + cast, or any registration the hooks do see. + """ + model = _mlp(device="cpu") + x = torch.randn(3, 4) + model(x) + model["hidden"].bias = None + assert model(x).shape == (3, 2) + for revalidate in (lambda: model(torch.randn(5, 4)), lambda: model.to("cpu")(x)): + with pytest.raises(HNDLError, match="^E_RUNTIME .*parameter 'nodes.n_hidden.bias' was set to None"): + revalidate() + + +def test_state_registered_during_a_later_forward_is_rejected_after_that_call(): + registry = Registry.builtins() + + @registry.operator("grows_late", identity="tests.grows_late", summary="Grow state on call three.", + shape="data[B, F] -> value[B, F]") + class GrowsLate(nn.Module): + def __init__(self): + super().__init__() + self.calls = 0 + + def forward(self, data): + self.calls += 1 + if self.calls == 3: + self.extra = nn.Parameter(torch.zeros(1)) + return data + + model = network("grows_late()", input_shape=("B", 4), output_shape=("B", 4), device="cpu", registry=registry) + x = torch.randn(2, 4) + model(x) + model(x) + with pytest.raises(HNDLError, match="^E_RUNTIME .*registered state of grows_late 'n0' changed during " + "forward: parameter 'nodes.n_n0.extra' was added"): + model(x) + with pytest.raises(HNDLError, match="registered state .* no longer matches the build"): + model(x) + + +def test_torch_compile_traces_the_checked_program(): + model = _stateful(device="cpu").eval() + x = torch.randn(3, 4) + expected = model(x) + compiled = torch.compile(model, backend="eager", fullgraph=True) + torch.testing.assert_close(compiled(x), expected) + torch.testing.assert_close(compiled(x), expected) + torch.testing.assert_close(model(x), expected)