Topology: hooks, runtime dataflow and torch.fx¶
Why hooks alone are not enough¶
Forward hooks give module inputs/outputs and execution order. They cannot see functional
operations in a parent's forward: torch.cat([up, skip]), x + residual, x[:, 0]. Those are
exactly the operations that create skip connections, fusion and branching.
Runtime dataflow tracing (default, topology="auto")¶
During the metadata pass a torch.overrides.TorchFunctionMode (capture.DataflowTracer) sees
every tensor operation. Each produced tensor is tagged with an op-node id (a Python attribute on
the tensor object). Each op records its parent op ids and the module call active on the hook
stack at the time. Model inputs are passed as detached views tagged as input nodes, so user
tensors are never modified. Nothing is retained except small integer records; the tags die with
the tensors.
Stage edges come from walking backwards from each selected stage's input tensors through the op graph. The walk stops at the first op owned by another selected stage (or a model input). Ownership uses the module call tree: an op belongs to the nearest selected ancestor call. Model-output edges are found the same way. This recovers:
- skip / residual connections (
enc1 → dec1), classified as skip when the source is also an ancestor of the primary path; - multimodal merge edges (
clinical_mlp → fusion); - branching heads (one stage feeding several heads);
- reused modules (each call is its own stage).
It also works with data-dependent control flow, because it observes the actual execution. The primary incoming edge of a stage is the deepest pathway (most ancestors), which keeps the trunk straight in the layout.
If tracing raises (for example, exotic tensor subclasses), the pass is re-run without the tracer, with a warning, and topology falls back to execution order.
torch.fx investigation (topology="fx")¶
torch.fx.symbolic_trace can supply topology while hooks supply runtime tensors. With a custom
Tracer that treats the selected stage modules as leaves, FX yields a static graph whose
call_module nodes are the stages. Propagating "producing stage" sets through the other nodes
gives stage edges. On traceable models (e.g. the test U-Net) this reproduces the runtime edges
exactly.
Findings:
* FX fails on data-dependent Python control flow (if x.mean() > 0:), many HuggingFace models, and
code that inspects tensor values. That is why it is not the default.
* FX gives no activations. It complements hooks and cannot replace them.
* On failure neural_flow warns "model could not be symbolically traced … Using runtime execution
order instead." and uses the sequential fallback. Activation visualisation is never blocked.
Sequential fallback (topology="sequential")¶
Stages are chained in execution order from the first input. Any stage that ended up without a predecessor in the other modes is also attached to the previous stage, so the figure is always connected.
Capture pass¶
The second pass registers hooks only on modules that own a selected stage call, counting calls
per module to hit the right invocation of reused modules. Summaries are computed inside the hook
(see tensor_rendering.md), so large activations are reduced before anything moves to
capture_device. An AttentionTracer mode records attention weights from
F.multi_head_attention_forward / F.scaled_dot_product_attention and attributes them to the
enclosing selected stage. All hooks and modes are removed in finally blocks, and the model's
training flag is restored.