"""Layer-output diagnostics; reductions are explicit local measurement choices.""" from contextlib import contextmanager import torch @contextmanager def capture_layer_outputs(layers, loop_count): states, handles = [], [] width = len(layers) if width == 0 or loop_count <= 0: raise ValueError('positive layer and loop counts required') def capture(index, output): if index != len(states) % width or len(states) >= width * loop_count: raise ValueError('unexpected recurrent layer invocation order') value = output[0] if isinstance(output, tuple) else output if not isinstance(value, torch.Tensor) or value.ndim != 3: raise ValueError('expected batch/token/hidden layer output') states.append(value.detach().to(device='cpu', dtype=torch.float32, copy=True)) try: for index, layer in enumerate(layers): handles.append(layer.register_forward_hook( lambda module, inputs, output, index=index: capture(index, output))) yield states if len(states) != width * loop_count: raise ValueError('incomplete recurrent layer trajectory') finally: for handle in handles: handle.remove() def compare_layer_outputs(reference, observed, *, layer_count, loop_count): if layer_count <= 0 or loop_count <= 0: raise ValueError('positive layer and loop counts required') if len(reference) != layer_count * loop_count or len(observed) != len(reference): raise ValueError('incomplete recurrent layer trajectory') rows = [] for index, (ref, actual) in enumerate(zip(reference, observed)): if ref.shape != actual.shape or ref.ndim != 3 or ref.numel() == 0: raise ValueError('layer output shapes differ or are empty') ref, actual = ref.double(), actual.double() if not torch.isfinite(ref).all() or not torch.isfinite(actual).all(): raise ValueError('nonfinite layer output') delta = actual - ref norm = ref.norm().item() error = delta.norm().item() rows.append(dict(loop=index // layer_count, layer=index % layer_count, trajectory_index=index, tokens=ref.shape[0] * ref.shape[1], l2_frobenius=error, mean_token_l2=delta.norm(dim=-1).mean().item(), relative_frobenius=error / norm if norm else None, reference_frobenius=norm, max_abs_error=delta.abs().max().item())) return rows