pulsatrix
Loading...
Searching...
No Matches
pulsatrix::ExplainerContext Class Reference

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>

Public Member Functions

 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.
 

Detailed Description

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.

Constructor & Destructor Documentation

◆ ExplainerContext()

pulsatrix::ExplainerContext::ExplainerContext ( std::vector< Module * >  modules)
inlineexplicit

Constructs a context over an ordered module chain.

Parameters
modulesThe network, in forward-pass order. Not owned; each must outlive this ExplainerContext.
Exceptions
std::invalid_argumentif 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.

Member Function Documentation

◆ activation()

const Tensor & pulsatrix::ExplainerContext::activation ( NodeId  id) const
inline

The cached activation value at a node, from the most recent forward_pass().

Parameters
idNode id. Must have been produced by the most recent forward_pass() call.

◆ activation_snapshot()

ActivationSnapshot pulsatrix::ExplainerContext::activation_snapshot ( ) const
inline

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_idNode 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_argumentif 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_gradGradient 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_errorif 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()

CircuitGraph pulsatrix::ExplainerContext::build_circuit_graph ( const Tensor &  input)
inline

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
inputInput 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
inputInput 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
inputInput to the first module in the chain.
patch_node_idNode 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_valueValue 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_argumentif 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
idNode 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()

const ComputationGraph & pulsatrix::ExplainerContext::graph ( ) const
inline

The current graph (from the most recent forward_pass() call).

◆ 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_idNode 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_argumentif 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]

Tensor pulsatrix::ExplainerContext::relevance_pass ( const Tensor &  output_relevance,
const LRPRuleConfig &  config 
)
inline

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_relevanceRelevance at the network output, same shape as the most recent forward_pass()'s output.
configLRP rule configuration handed to every module.
Returns
Relevance at the input, same shape as the input to that forward_pass().
Exceptions
std::logic_errorif 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_argumentif 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_argumentif 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: