Wraps an ordered chain of Modules, running them via Module::forward_traced to build a real ComputationGraph and Autograd backward wiring – graph-native explainers (Missions 2-3) use graph()/backward_pass()/activation(); surrogate explainers (Phase 3) would use only forward_pass(), per the charter's stated interface segregation.
More...
#include <explainer_context.hpp>
|
| | ExplainerContext (std::vector< Module * > modules) |
| | Constructs a context over an ordered module chain.
|
| |
| Tensor | forward_pass (const Tensor &input) |
| | Runs the full module chain forward, building a fresh graph and caching every node's activation value as it goes.
|
| |
| Tensor | forward_pass_with_patch (const Tensor &input, NodeId patch_node_id, const Tensor &patch_value) |
| | Runs the full module chain forward, substituting patch_value for the natural output of the module that produces patch_node_id – the causal-intervention (activation patching) primitive.
|
| |
| Tensor | backward_pass (const Tensor &output_grad) |
| | Runs Autograd::backward from the most recent forward_pass()'s output node.
|
| |
| Tensor | relevance_pass (const Tensor &output_relevance, const LRPRuleConfig &config) |
| | Propagates LRP relevance from the network output back to the input through every module's own propagate_relevance(), in reverse order – the relevance counterpart of backward_pass(), and what LRP::explain() runs.
|
| |
| Tensor | relevance_pass (const Tensor &output_relevance, const std::vector< LRPRuleConfig > &configs) |
| | relevance_pass() with a per-module rule choice: configs[i] is handed to the i-th module (forward order) – e.g. LRP composites (lrp.hpp).
|
| |
| const std::vector< Module * > & | modules () const |
| | The module chain, in forward order (not owned).
|
| |
| const ComputationGraph & | graph () const |
| | The current graph (from the most recent forward_pass() call).
|
| |
| const Tensor & | activation (NodeId id) const |
| | The cached activation value at a node, from the most recent forward_pass().
|
| |
| const Tensor & | gradient (NodeId id) const |
| | The cached gradient at a node, from the most recent backward_pass() call.
|
| |
| ActivationSnapshot | activation_snapshot () const |
| | Captures the current activation cache into a self-contained ActivationSnapshot.
|
| |
| Tensor | logit_lens (NodeId node_id) const |
| | The "logit lens": runs the chain's final module – its read-out head – on the cached activation at node_id, answering "what would the network predict if this
layer's representation were already final?".
|
| |
| Tensor | attention_weights (NodeId node_id) const |
| | The "attention lens": the per-head attention pattern produced by the attention layer at node_id during the most recent forward pass – "what did each head
attend to", keyed by node id rather than by a direct module reference.
|
| |
| CircuitGraph | build_circuit_graph (const Tensor &input) |
| | Builds a CircuitGraph for this chain at the given input: every node scored by how much zeroing it changes the network's output, plus the chain's edges.
|
| |
| std::optional< std::string > | layer_label (NodeId id) const |
| | Passthrough to Node::label() – e.g. a layer name, for debugging/display.
|
| |
Wraps an ordered chain of Modules, running them via Module::forward_traced to build a real ComputationGraph and Autograd backward wiring – graph-native explainers (Missions 2-3) use graph()/backward_pass()/activation(); surrogate explainers (Phase 3) would use only forward_pass(), per the charter's stated interface segregation.
- Note
- Owns a std::vector<Module*>, not a new Sequential container – no generic module-sequence type exists yet in this codebase (XorNetwork hardcodes its own chain), and this mission doesn't need one beyond what ExplainerContext itself requires. Modules are not owned; every pointer must outlive this ExplainerContext.
-
forward_pass() replaces its ComputationGraph/Autograd members with fresh instances on every call, rather than mutating a shared graph – ComputationGraph deliberately has no clear/reset method (Phase 0's persistence-past-backward guarantee applies within one forward_pass()'s lifetime), but Integrated Gradients (Mission 2) needs multiple independent forward passes at different interpolated inputs, each with its own graph. See mission_explainer_context.md's Recon.
◆ ExplainerContext()
| pulsatrix::ExplainerContext::ExplainerContext |
( |
std::vector< Module * > |
modules | ) |
|
|
inlineexplicit |
Constructs a context over an ordered module chain.
- Parameters
-
| modules | The network, in forward-pass order. Not owned; each must outlive this ExplainerContext. |
- Exceptions
-
| std::invalid_argument | if modules is empty or contains a nullptr – external boundary (campaign_exai_dl_library_adversarial_hardening.md, Mission 2, finding 5): without this check, a nullptr element was an unconditional null-pointer dereference on the next forward_pass() call. |
◆ activation()
| const Tensor & pulsatrix::ExplainerContext::activation |
( |
NodeId |
id | ) |
const |
|
inline |
The cached activation value at a node, from the most recent forward_pass().
- Parameters
-
◆ activation_snapshot()
Captures the current activation cache into a self-contained ActivationSnapshot.
- Returns
- A snapshot of every activation from the most recent forward_pass(), the node ids in topological order, and each node's op_type/label – all copied, with no reference back to this context or its graph.
- Note
- Does not run (or require) a new forward pass; it reads what is already cached. Called before any forward_pass(), it returns an empty snapshot – there is nothing cached yet, which is a well-defined state, not an error.
-
This is the generalization of activation(): that method is a single-node lookup into a cache the next forward_pass() overwrites in place, so two runs' activations could never be held at once. A snapshot survives any number of subsequent forward passes, which is what activation patching (campaign campaign_exai_dl_library_mechanistic_interpretability, Phase 4) needs – a "clean" and a "corrupted" run alive simultaneously. Since modules_ is fixed for this context's lifetime, snapshots from different forward passes are NodeId-comparable.
◆ attention_weights()
| Tensor pulsatrix::ExplainerContext::attention_weights |
( |
NodeId |
node_id | ) |
const |
|
inline |
The "attention lens": the per-head attention pattern produced by the attention layer at node_id during the most recent forward pass – "what did each head
attend to", keyed by node id rather than by a direct module reference.
- Parameters
-
| node_id | Node whose attention pattern is returned. Must come from the most recent forward pass and must name a node whose op_type() is OpType::Attention. |
- Returns
- That layer's (N, num_heads, L, L) softmax output, copied. Raw data only – no head aggregation, no ranking, no plotting (campaign campaign_exai_dl_library_mechanistic_interpretability, Phase 5, is deliberately scoped data-only and takes no visualization dependency).
- Exceptions
-
| std::invalid_argument | if node_id lies outside the node range the most recent forward pass produced (including the case where no forward pass has run yet, so there are no nodes at all), or if it names a node that is not an attention layer – both external boundaries (caller-supplied node id), hence throw rather than PULSATRIX_ASSERT, matching logit_lens and forward_pass_with_patch rather than activation()'s older pattern, so the behavior is identical in Debug and Release. |
- Note
- Node id 0 is the input node, which is never an attention layer, so it falls into the op-type throw rather than needing a case of its own.
-
Uses the same node-to-module index correspondence forward_pass_with_patch and logit_lens already rely on: forward_pass() assigns node id i+1 to modules_[i]'s output, so modules_[node_id - 1] is the module that produced that node.
-
The dynamic_cast is PULSATRIX_ASSERTed, not thrown on: only MultiHeadAttentionModule reports OpType::Attention anywhere in this codebase, so a failure here would mean the op-type/module invariant itself broke – an internal-consistency violation, not a caller mistake. Checked rather than assumed, per the assert-vs-throw classification in context_tdd_adversarial_boundary_testing.md.
-
Reads the module's own cache, so it reflects that module's most recent forward() – which, for a module driven only through this context, is the most recent forward_pass()/forward_pass_with_patch() call. A patched pass's pattern is the real pattern that ran, which is exactly what a causal-intervention study wants.
◆ backward_pass()
| Tensor pulsatrix::ExplainerContext::backward_pass |
( |
const Tensor & |
output_grad | ) |
|
|
inline |
Runs Autograd::backward from the most recent forward_pass()'s output node.
- Parameters
-
| output_grad | Gradient w.r.t. the chain's output. |
- Returns
- Gradient w.r.t. the chain's input.
- Note
- Must be called after forward_pass().
- Exceptions
-
| std::logic_error | if the most recent forward pass was a forward_pass_with_patch() call. A patched activation is a constant substitution, not a differentiable function of the input, yet Module::forward_traced registers each module's ordinary backward closure regardless – so without this guard Autograd::backward() would return a perfectly plausible gradient taken through a link that does not exist in the computation it claims to differentiate. Rejected loudly rather than answered wrongly, the same principle as Phase 1.5's device guards; classified external boundary (a caller-sequencing mistake, like the constructor's checks) so it is a throw and behaves identically in Debug and Release. Re-arm with an ordinary forward_pass() – the guard is per-call history, not a permanent latch. |
◆ build_circuit_graph()
Builds a CircuitGraph for this chain at the given input: every node scored by how much zeroing it changes the network's output, plus the chain's edges.
- Parameters
-
| input | Input the circuit is built at. A circuit graph is input-conditional – ablation importance is "how much does this node matter *for this input*", not a property of the weights alone. |
- Returns
- A self-contained CircuitGraph: one CircuitNode per graph node (in topological order), and one CircuitEdge per adjacent pair. Raw data only – rendering is explicitly deferred (see circuit_graph.hpp's own note and campaign campaign_exai_dl_library_mechanistic_interpretability, Phase 5 Mission 3).
- Note
- Scoring method: for each non-output node, forward_pass_with_patch() substitutes a zero-Tensor of that node's own natural shape, and ablation_effect is the L2 distance between that patched output and the real, unpatched one – the standard "ablation importance" score, built entirely from Phase 4's patching primitive with no new causal-inference machinery.
-
The zero patch is a copy of the node's own cached activation, filled with 0.0f, so it matches that node's natural shape, backend and device by construction – Tensor exposes no backend() accessor to rebuild one from a Shape alone.
-
The output node's ablation_effect is 0.0f by convention, not by computation: forward_pass_with_patch() on the output node returns the patch value itself (nothing runs after it), so self-patching it would score the arbitrary magnitude of the real output rather than any causal quantity. It is still present in nodes(), for structural completeness.
-
Edges are the chain's inherent adjacency (i -> i+1), weighted by node i's own ablation_effect. This is exact only because every graph this codebase builds is a single-parent linear chain (forward_pass()'s loop); see CircuitEdge::weight.
-
Cost is one forward pass per node plus two unpatched passes – O(node_count) forward passes. Fine for the small chains this codebase builds; a large network would want a sampled or grouped variant, which is not built here.
-
Runs an ordinary forward_pass(input) last, so the context is handed back exactly as a plain forward pass would leave it: cached activations from the real run, and backward_pass()'s patched-pass guard disarmed. Without that, every caller would silently inherit the state of the final ablation run.
◆ forward_pass()
| Tensor pulsatrix::ExplainerContext::forward_pass |
( |
const Tensor & |
input | ) |
|
|
inline |
Runs the full module chain forward, building a fresh graph and caching every node's activation value as it goes.
- Parameters
-
| input | Input to the first module in the chain. |
- Returns
- The final module's output.
◆ forward_pass_with_patch()
| Tensor pulsatrix::ExplainerContext::forward_pass_with_patch |
( |
const Tensor & |
input, |
|
|
NodeId |
patch_node_id, |
|
|
const Tensor & |
patch_value |
|
) |
| |
|
inline |
Runs the full module chain forward, substituting patch_value for the natural output of the module that produces patch_node_id – the causal-intervention (activation patching) primitive.
- Parameters
-
| input | Input to the first module in the chain. |
| patch_node_id | Node whose activation is overridden. Node ids are assigned deterministically per forward pass: 0 is the input node, then one per module in chain order, so a node id captured from an earlier pass over this same context names the same logical layer here. |
| patch_value | Value substituted at that node. Must have the same shape as the node's natural (unpatched) output. |
- Returns
- The chain's output, computed downstream from the substituted value. Patching the output node returns patch_value itself – nothing runs after it.
- Exceptions
-
| std::invalid_argument | if patch_node_id exceeds the largest node id this forward pass produces, or if patch_value's shape differs from the target node's natural output shape – both external boundaries (caller-supplied), hence throw rather than PULSATRIX_ASSERT, consistent with the constructor. |
- Note
- Every module before the patch point is unaffected (a forward-only computation has no upstream influence); every module after it computes from patch_value. Patching the input node is a supported degenerate case, equivalent to forward_pass(patch_value).
-
Builds into local graph/autograd/activation state and only commits it on success, so a throw leaves this context exactly as the previous forward_pass() left it – a failed patch attempt must not corrupt a usable context.
-
Single-node patch per call by design (campaign campaign_exai_dl_library_mechanistic_interpretability, Phase 4 Mission 1); multi-node patching would be an explicit extension, not assumed here.
-
backward_pass()'s guard against this method covers Autograd-based gradients only. ExplainerContext exposes no LRP/propagate_relevance traversal of its own to guard – Module::propagate_relevance() is invoked directly by explainer code outside this class (module.hpp), not through ExplainerContext. If a future explainer manually chains propagate_relevance() calls across modules using state left behind by a patched forward pass, the same causal-inconsistency hazard backward_pass() guards against applies there too, unguarded – that explainer's own author is responsible for it, the same way any code bypassing ExplainerContext's own accessors already is.
◆ gradient()
| const Tensor & pulsatrix::ExplainerContext::gradient |
( |
NodeId |
id | ) |
const |
|
inline |
The cached gradient at a node, from the most recent backward_pass() call.
- Parameters
-
| id | Node id. Must have accumulated a gradient during the most recent backward_pass() call (i.e. lie on the path between the seeded output and the input). |
- Note
- Mirrors activation()'s shape – Autograd::backward() already populates a gradient at every node it walks through, not just the input node backward_pass() itself returns; this exposes that directly, needed for Grad-CAM's target-conv-layer gradient (Phase 2 Mission 3).
◆ graph()
◆ layer_label()
| std::optional< std::string > pulsatrix::ExplainerContext::layer_label |
( |
NodeId |
id | ) |
const |
|
inline |
Passthrough to Node::label() – e.g. a layer name, for debugging/display.
◆ logit_lens()
| Tensor pulsatrix::ExplainerContext::logit_lens |
( |
NodeId |
node_id | ) |
const |
|
inline |
The "logit lens": runs the chain's final module – its read-out head – on the cached activation at node_id, answering "what would the network predict if this
layer's representation were already final?".
- Parameters
-
| node_id | Node whose cached activation is projected through the head. Must come from the most recent forward pass, and its activation's shape must match what the head actually consumed in that pass. |
- Returns
- The head's output for that intermediate representation. Raw data only – no softmax, no ranking, no plotting (campaign campaign_exai_dl_library_mechanistic_interpretability, Phase 5 Mission 1, is deliberately scoped data-only and takes no visualization dependency).
- Exceptions
-
| std::invalid_argument | if node_id has no cached activation (an id no forward pass produced, or no forward pass has run yet), or if that activation's shape differs from the head's actual input shape in the most recent forward pass – both external boundaries (caller-supplied node id), hence throw rather than PULSATRIX_ASSERT, matching forward_pass_with_patch and the constructor, and so the behavior is identical in Debug and Release. |
- Note
- The shape precondition is checked here, before the head runs, so the message names the mismatch at the point the mistake was made rather than surfacing as a confusing failure inside the head's own gemm. It is a real precondition of the technique, not an assumption: only a chain whose hidden width stays constant up to the head (the analogue of a transformer's fixed-width residual stream) has compatible earlier nodes at all.
-
The head's expected input shape is taken from the activation the head actually received in the most recent forward pass (the node at index modules_.size() - 1), not from any declared per-module input-shape accessor – Module exposes no such accessor, and inventing one across every subclass is far outside this mission.
-
Applied to that same node, this returns the real forward pass's output exactly: it is literally the computation forward_pass() already performed, which is this method's zero-tolerance correctness oracle.
-
const with respect to this – no ExplainerContext state is read-modified. It does run the head module's own forward(), which overwrites that module's internal forward cache (so a subsequent head.backward() would refer to this call's input). Same caveat this campaign's SparseAutoencoder::reconstruct() mission already documented and accepted for read-only scoring paths.
-
The head is the chain's literal last module, not a separately designated "unembedding" – this codebase has no such distinct concept. A caller-specifiable head would be an explicit future extension, not assumed here.
◆ modules()
| const std::vector< Module * > & pulsatrix::ExplainerContext::modules |
( |
| ) |
const |
|
inline |
The module chain, in forward order (not owned).
◆ relevance_pass() [1/2]
Propagates LRP relevance from the network output back to the input through every module's own propagate_relevance(), in reverse order – the relevance counterpart of backward_pass(), and what LRP::explain() runs.
- Parameters
-
| output_relevance | Relevance at the network output, same shape as the most recent forward_pass()'s output. |
| config | LRP rule configuration handed to every module. |
- Returns
- Relevance at the input, same shape as the input to that forward_pass().
- Exceptions
-
| std::logic_error | if the most recent forward pass was forward_pass_with_patch(), for the same reason backward_pass() refuses: the modules' cached state would describe a computation the input did not produce. |
| std::invalid_argument | if a module does not implement config.rule (Module::supports_lrp_rule()) – checked for every module before any propagation; the rule is never silently replaced by epsilon. |
◆ relevance_pass() [2/2]
| Tensor pulsatrix::ExplainerContext::relevance_pass |
( |
const Tensor & |
output_relevance, |
|
|
const std::vector< LRPRuleConfig > & |
configs |
|
) |
| |
|
inline |
relevance_pass() with a per-module rule choice: configs[i] is handed to the i-th module (forward order) – e.g. LRP composites (lrp.hpp).
- Exceptions
-
| std::invalid_argument | if configs.size() != number of modules, or a module does not implement its rule; std::logic_error as for the uniform overload. |
The documentation for this class was generated from the following file: