From 1b0efd29d07a7b7e0e0a3e85b760f1d1135bed4b Mon Sep 17 00:00:00 2001 From: Martyn Garcia Date: Sat, 26 Sep 2026 09:55:59 -0600 Subject: [PATCH 1/3] Validate each input signature once instead of every port on every call GraphModule._execute() checked the shape, dtype and device of every port of every node, and replayed the registered-state walk, on every forward call: a fixed ~10 us per node that made a five-node MLP at batch 32 run ~1.7x hand-written PyTorch. The first call for a given input signature (every external input's shape, dtype and device, plus training mode and autocast state) now runs the full checked program and the state check, as before. The signature is then remembered, and later calls with it compare the signature and run the modules back to back from a slot-indexed program. What was validated is invalidated on the events that can break it: - .to()/.cuda()/.double()/... (_apply) drops every signature and forces a full state re-check; - registering a parameter, buffer or submodule on any module in the graph (register_* or attribute assignment) advances a watch generation through torch's global registration hooks, filtered to modules of built graphs, so the next call re-checks state and revalidates every port; a registration during a fast call is checked at the end of that call, as before. Errors are readable: port mismatches name the node, operation and source line and show the contract in HNDL notation ("[B=32, 64]:float32 on cpu") beside the tensor that arrived; torch errors raised inside a node are wrapped as E_RUNTIME with the node's inputs and any state change that explains them; OOM errors pass through unwrapped. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/hndl/torch.py | 563 ++++++++++++++++++++++++++++++++++++++------ tests/test_torch.py | 274 ++++++++++++++++++++- 2 files changed, 758 insertions(+), 79 deletions(-) diff --git a/src/hndl/torch.py b/src/hndl/torch.py index bebe9cb..2aab59e 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,95 @@ #: 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 + +#: Errors that pass out of a node untouched: callers catch these by type to +#: react (shrink the batch, free memory), so wrapping them would break that. +_UNWRAPPED = (torch.cuda.OutOfMemoryError, MemoryError) + +_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 +180,24 @@ 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. A torch error raised inside a node is reported as + ``E_RUNTIME`` naming the node, its operation and its source line, with the + tensors it received, and with any change to registered state that explains + it. + """ 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,94 @@ 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 _node_failure(self, node_id, error, bound, batch): + """The ``E_RUNTIME`` to raise for an exception out of one node's module. + + Returns ``None`` when the original exception should propagate as is: + out-of-memory errors, which callers catch by type, and HNDL errors that + already name their node. Everything here runs only after a failure. + """ + if isinstance(error, _UNWRAPPED): + return None + name = self._node_name(node_id) + raised = f"{name} raised {type(error).__name__}: {error}" + if not self._state_matches(): + # State removed between calls breaks forward rather than reporting + # itself; say so, because that is the cause the torch error hides. + return self._state_error("build", then=raised.replace("\n", "\n ")) + if isinstance(error, HNDLError): + if error.node is not None: + return None + return self._error(f"{name}: {error.message}", node_id, code=error.code) + lines = [raised.replace("\n", "\n ")] + 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}") + return self._error("\n".join(lines), node_id) + + def _fast_failure(self, node, error, values): + """Diagnose an exception from the unchecked loop, after the fact. + + The call's input signature was validated earlier, so the node's inputs + are re-checked against its contract first: a mismatch there means an + upstream module changed what it produces since, which is the useful + thing to report. + """ + if isinstance(error, _UNWRAPPED): + return None + ins, _, _, _, node_id = node + bound = [values[slot] for slot in ins] + leading = values[0] + batch = leading.shape[0] if isinstance(leading, torch.Tensor) and leading.ndim else None + 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) + if self._state_matches(): + for (_, expected, dtype, port), value in zip(entry[1], bound): + try: + self._check(value, expected, dtype, port, batch) + except HNDLError as mismatch: + return self._error( + f"{mismatch.message}\n (an earlier call with the same input signature passed " + f"every check, so an upstream module now produces a different tensor)\n" + f" {self._node_name(node_id)} then raised {type(error).__name__}: {error}", + node_id) + return self._node_failure(node_id, error, bound, batch) + + # -- registered state --------------------------------------------------- + def _state_snapshot(self): """Record what ``named_parameters``/``named_buffers`` see, per module. @@ -200,18 +418,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 +445,95 @@ 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_error(self, when, then=None): + """``E_RUNTIME`` naming the node whose registered state changed, and how.""" + 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) + subject = self._node_name(node_id) if node_id is not None else "the network" + listed = "; ".join(text for _, text in changes[:6]) + if len(changes) > 6: + listed += f"; and {len(changes) - 6} more" + phrase = "changed during forward" if when == "forward" else "no longer matches the build" + lines = [f"registered state of {subject} {phrase}: {listed}"] + if then is not None: + lines.append(f" {then}") + 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 +572,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 +677,113 @@ 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: + failure = self._node_failure(node_id, error, bound, batch) + if failure is None: + raise + raise failure from error + 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: + failure = self._fast_failure(node, error, values) + if failure is None: + raise + raise failure from error + 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/test_torch.py b/tests/test_torch.py index 85b29da..56c00e9 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,256 @@ 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 + + +@pytest.mark.parametrize("fail_from_call", [1, 2]) +def test_a_torch_error_inside_a_node_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(HNDLError) as caught: + model(torch.randn(2, 4)) + error = caught.value + assert (error.code, error.node, error.line, error.column) == ("E_RUNTIME", "odd", 2, 1) + assert str(error).splitlines() == [ + "E_RUNTIME (node odd; line 2, column 1): flaky 'odd' raised RuntimeError: " + "mat1 and mat2 shapes cannot be multiplied (2x4 and 3x3)", + " input 'x': got [2, 4]:float32 on cpu; contract [B=2, 4]:float32 on cpu", + ] + assert isinstance(error.__cause__, RuntimeError) + + +def test_errors_callers_catch_by_type_leave_nodes_unwrapped(): + 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"): + model(torch.randn(2, 4)) + + +def test_an_hndl_error_inside_a_node_gains_the_node(): + 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 (node split_here; line 1, column 1): " + "flaky 'split_here': extent 5 does not divide by 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(HNDLError) as caught: + model(torch.randn(2, 4)) + lines = str(caught.value).splitlines() + assert lines[0] == ("E_RUNTIME (node head; line 2, column 1): linear 'head' input 'x' (from node:d/out): " + "expected shape [B=2, 4], got [2, 2]") + assert lines[-1].startswith(" linear 'head' then raised RuntimeError: mat1 and mat2") + + +def _replaces_the_head_weight(model): + model["head"].weight = nn.Parameter(torch.randn(2, 7)) + + +@pytest.mark.parametrize("mutate, message", [ + (lambda model: setattr(model["hidden"], "extra", nn.Linear(2, 2)), + "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))), + "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), + "registered state of batch_norm 'norm' no longer matches the build: " + "buffer 'nodes.n_norm.running_mean' was set to None"), + (lambda model: delattr(model["head"], "weight"), + "registered state of linear 'head' no longer matches the build: " + "parameter 'nodes.n_head.weight' was removed"), + (_replaces_the_head_weight, + "linear 'head' raised 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, 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(HNDLError, match="^E_RUNTIME") as caught: + model(x) + assert message in str(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) From 532b681d6f269090a9d4078697dfbba1c5a10595 Mon Sep 17 00:00:00 2001 From: Martyn Garcia Date: Sat, 26 Sep 2026 10:05:19 -0600 Subject: [PATCH 2/3] Document validation once per input signature CHANGELOG, README, IMPLEMENTATION and SPEC describe when contracts are checked, what invalidates a validated signature, the new error wording, and the registered-state edits PyTorch runs no hook for. The parity benchmark's docstring now describes the cheap path it actually times. Co-Authored-By: Claude Opus 5.5 (1M context) --- CHANGELOG.md | 21 +++++++++++++++ IMPLEMENTATION.md | 17 ++++++++++++ README.md | 2 +- SPEC.md | 8 +++--- .../test_parity_vs_handwritten_pytorch.py | 26 +++++++++++-------- 5 files changed, 58 insertions(+), 16 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 1910908..c644451 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -27,6 +27,27 @@ `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 is now an `E_RUNTIME` naming the node, with its inputs and + any registered-state change that explains it, and the original exception as + `__cause__` (out-of-memory errors still propagate unchanged). 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..7f1494e 100644 --- a/IMPLEMENTATION.md +++ b/IMPLEMENTATION.md @@ -67,6 +67,23 @@ 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 is wrapped the same way, with the original +kept as `__cause__`. 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..cea2f6a 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, or an error raised inside a layer, is an `E_RUNTIME` error naming the layer, its operation and its source line, with the expected and actual shapes. 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..e5a2312 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 is reported as `E_RUNTIME` naming the node, its operation, and its source position, with the tensors it received, except out-of-memory errors, which propagate unchanged. - 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; a node's module raises during forward; 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/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. """ From 5de1dea1068a2419fb36f43f5307dfe504c7929e Mon Sep 17 00:00:00 2001 From: Martyn Garcia Date: Sat, 26 Sep 2026 10:13:13 -0600 Subject: [PATCH 3/3] Keep exceptions raised inside a node as their own type, with a note Wrapping a torch error raised inside a node in HNDLError (a ValueError) broke callers that catch RuntimeError around a forward call. The original exception now propagates unchanged --- type, message and traceback --- and gains one add_note() naming the node, its operation and source line, its inputs against their contract, any registered-state change that explains the failure, the upstream port that drifted after validation, and a hint when the error came from torch.compile. An exception passing out through nested graphs keeps only the innermost node's note. HNDLError (E_RUNTIME) is kept for what HNDL's own checks find: first-call validation and state mismatches. Co-Authored-By: Claude Opus 5.5 (1M context) --- CHANGELOG.md | 9 ++- IMPLEMENTATION.md | 6 +- README.md | 2 +- SPEC.md | 4 +- src/hndl/torch.py | 134 ++++++++++++++++++++++---------------------- tests/test_torch.py | 97 +++++++++++++++++++++----------- 6 files changed, 143 insertions(+), 109 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index c644451..924d46e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -40,9 +40,12 @@ 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 is now an `E_RUNTIME` naming the node, with its inputs and - any registered-state change that explains it, and the original exception as - `__cause__` (out-of-memory errors still propagate unchanged). Registered + 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 diff --git a/IMPLEMENTATION.md b/IMPLEMENTATION.md index 7f1494e..3ac7b66 100644 --- a/IMPLEMENTATION.md +++ b/IMPLEMENTATION.md @@ -77,8 +77,10 @@ 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 is wrapped the same way, with the original -kept as `__cause__`. PyTorch runs no registration hook for `del` of a +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 diff --git a/README.md b/README.md index cea2f6a..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. 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, or an error raised inside a layer, is an `E_RUNTIME` error naming the layer, its operation and its source line, with the expected and actual shapes. 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 e5a2312..5ed1f6a 100644 --- a/SPEC.md +++ b/SPEC.md @@ -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. -- 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 is reported as `E_RUNTIME` naming the node, its operation, and its source position, with the tensors it received, except out-of-memory errors, which propagate unchanged. +- 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; a node's module raises during forward; registered state no longer matches the build; 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 2aab59e..caba5fe 100644 --- a/src/hndl/torch.py +++ b/src/hndl/torch.py @@ -36,9 +36,8 @@ #: forgets them all and starts again; bounds the cache under varying batch sizes. _MAX_SIGNATURES = 64 -#: Errors that pass out of a node untouched: callers catch these by type to -#: react (shrink the batch, free memory), so wrapping them would break that. -_UNWRAPPED = (torch.cuda.OutOfMemoryError, MemoryError) +#: 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 @@ -193,10 +192,11 @@ class GraphModule(nn.Module): 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. A torch error raised inside a node is reported as - ``E_RUNTIME`` naming the node, its operation and its source line, with the - tensors it received, and with any change to registered state that explains - it. + 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): @@ -351,63 +351,61 @@ def _port_error(self, port, value, dtype, batch, problem): 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 _node_failure(self, node_id, error, bound, batch): - """The ``E_RUNTIME`` to raise for an exception out of one node's module. + 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. - Returns ``None`` when the original exception should propagate as is: - out-of-memory errors, which callers catch by type, and HNDL errors that - already name their node. Everything here runs only after a failure. + 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, _UNWRAPPED): - return None - name = self._node_name(node_id) - raised = f"{name} raised {type(error).__name__}: {error}" + 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 torch error hides. - return self._state_error("build", then=raised.replace("\n", "\n ")) - if isinstance(error, HNDLError): - if error.node is not None: - return None - return self._error(f"{name}: {error.message}", node_id, code=error.code) - lines = [raised.replace("\n", "\n ")] + # 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}") - return self._error("\n".join(lines), node_id) - - def _fast_failure(self, node, error, values): - """Diagnose an exception from the unchecked loop, after the fact. - - The call's input signature was validated earlier, so the node's inputs - are re-checked against its contract first: a mismatch there means an - upstream module changed what it produces since, which is the useful - thing to report. - """ - if isinstance(error, _UNWRAPPED): - return None - ins, _, _, _, node_id = node - bound = [values[slot] for slot in ins] - leading = values[0] - batch = leading.shape[0] if isinstance(leading, torch.Tensor) and leading.ndim else None - 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) - if self._state_matches(): - for (_, expected, dtype, port), value in zip(entry[1], bound): - try: - self._check(value, expected, dtype, port, batch) - except HNDLError as mismatch: - return self._error( - f"{mismatch.message}\n (an earlier call with the same input signature passed " - f"every check, so an upstream module now produces a different tensor)\n" - f" {self._node_name(node_id)} then raised {type(error).__name__}: {error}", - node_id) - return self._node_failure(node_id, error, bound, batch) + 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 --------------------------------------------------- @@ -511,18 +509,21 @@ def _state_changes(self): changes.append((node_id, f"the submodules of '{path}' were re-registered in a different order")) return changes - def _state_error(self, when, then=None): - """``E_RUNTIME`` naming the node whose registered state changed, and how.""" + 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) - subject = self._node_name(node_id) if node_id is not None else "the network" 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}"] - if then is not None: - lines.append(f" {then}") 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") @@ -695,10 +696,8 @@ def _run_checked(self, inputs, compiling): try: result = module(*bound) except Exception as error: - failure = self._node_failure(node_id, error, bound, batch) - if failure is None: - raise - raise failure from 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) @@ -764,10 +763,11 @@ def _execute(self, inputs): for slot, value in zip(outs, self._split_result(result, out_ports, out_set, node_id)): values[slot] = value except Exception as error: - failure = self._fast_failure(node, error, values) - if failure is None: - raise - raise failure from 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 diff --git a/tests/test_torch.py b/tests/test_torch.py index 56c00e9..587888b 100644 --- a/tests/test_torch.py +++ b/tests/test_torch.py @@ -701,41 +701,68 @@ def forward(self, 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_names_the_node(fail_from_call): +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(HNDLError) as caught: + with pytest.raises(RuntimeError, match=r"^mat1 and mat2 shapes cannot be multiplied \(2x4 and 3x3\)\n") as caught: model(torch.randn(2, 4)) - error = caught.value - assert (error.code, error.node, error.line, error.column) == ("E_RUNTIME", "odd", 2, 1) - assert str(error).splitlines() == [ - "E_RUNTIME (node odd; line 2, column 1): flaky 'odd' raised RuntimeError: " - "mat1 and mat2 shapes cannot be multiplied (2x4 and 3x3)", - " input 'x': got [2, 4]:float32 on cpu; contract [B=2, 4]:float32 on cpu", - ] - assert isinstance(error.__cause__, RuntimeError) + 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_leave_nodes_unwrapped(): +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"): + 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_gains_the_node(): +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 (node split_here; line 1, column 1): " - "flaky 'split_here': extent 5 does not divide by 2") + 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(): @@ -755,43 +782,45 @@ def forward(self, x): 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(HNDLError) as caught: + with pytest.raises(RuntimeError, match="^mat1 and mat2") as caught: model(torch.randn(2, 4)) - lines = str(caught.value).splitlines() - assert lines[0] == ("E_RUNTIME (node head; line 2, column 1): linear 'head' input 'x' (from node:d/out): " - "expected shape [B=2, 4], got [2, 2]") - assert lines[-1].startswith(" linear 'head' then raised RuntimeError: mat1 and mat2") + 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, message", [ - (lambda model: setattr(model["hidden"], "extra", nn.Linear(2, 2)), - "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))), +@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), + (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"), - (lambda model: delattr(model["head"], "weight"), - "registered state of linear 'head' no longer matches the build: " - "parameter 'nodes.n_head.weight' was removed"), - (_replaces_the_head_weight, - "linear 'head' raised RuntimeError: mat1 and mat2 shapes cannot be multiplied (3x5 and 7x2)"), + # 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, message): +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(HNDLError, match="^E_RUNTIME") as caught: + with pytest.raises(raised) as caught: model(x) - assert message in str(caught.value) + 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():