14#include <unordered_map>
37 const char* raw =
typeid(module).name();
40 char* demangled = abi::__cxa_demangle(raw,
nullptr,
nullptr, &status);
41 std::string name = (status == 0 && demangled !=
nullptr) ? demangled : raw;
79 if (modules_.empty()) {
80 throw std::invalid_argument(
"ExplainerContext: modules must not be empty");
82 for (
Module* m : modules_) {
84 throw std::invalid_argument(
"ExplainerContext: modules must not contain a null pointer");
96 last_forward_was_patched_ =
false;
102 activations_.emplace(input_node_,
Tensor(input));
105 NodeId current_node = input_node_;
106 for (
Module* module : modules_) {
107 auto [output, node_id] =
module->forward_traced(current, current_node, graph_, autograd_);
108 activations_.emplace(node_id,
Tensor(output));
109 current = std::move(output);
110 current_node = node_id;
112 output_node_ = current_node;
154 last_forward_was_patched_ =
true;
155 if (patch_node_id > modules_.size()) {
156 throw std::invalid_argument(
"ExplainerContext::forward_pass_with_patch: patch_node_id out of range");
161 std::unordered_map<NodeId, Tensor> activations;
166 NodeId current_node = input_node;
167 if (patch_node_id == input_node) {
168 current = check_patch_shape(current, patch_value);
170 activations.emplace(input_node,
Tensor(current));
172 for (
Module* module : modules_) {
173 auto [output, node_id] =
module->forward_traced(current, current_node, graph, autograd);
174 if (node_id == patch_node_id) {
175 output = check_patch_shape(output, patch_value);
177 activations.emplace(node_id,
Tensor(output));
178 current = std::move(output);
179 current_node = node_id;
182 graph_ = std::move(
graph);
183 autograd_ = std::move(autograd);
184 activations_ = std::move(activations);
185 input_node_ = input_node;
186 output_node_ = current_node;
208 if (last_forward_was_patched_) {
209 throw std::logic_error(
210 "ExplainerContext::backward_pass: the most recent forward pass was "
211 "forward_pass_with_patch(); a patched activation is a constant substitution, not a "
212 "differentiable function of the input, so gradients through it would be silently wrong. "
213 "Run an unpatched forward_pass() before backward_pass().");
215 autograd_.
backward(graph_, output_node_, output_grad);
235 return relevance_pass(output_relevance, std::vector<LRPRuleConfig>(modules_.size(), config));
245 if (last_forward_was_patched_) {
246 throw std::logic_error(
247 "ExplainerContext::relevance_pass: the most recent forward pass was "
248 "forward_pass_with_patch(); relevance through a patched activation would describe a "
249 "computation the input did not produce. Run an unpatched forward_pass() first.");
251 if (configs.size() != modules_.size()) {
252 throw std::invalid_argument(
"ExplainerContext::relevance_pass: need exactly one LRPRuleConfig per module");
254 for (
size_t i = 0; i < modules_.size(); ++i) {
255 if (!modules_[i]->supports_lrp_rule(configs[i].rule)) {
256 throw std::invalid_argument(
"ExplainerContext::relevance_pass: module " + std::to_string(i) +
" (" +
261 Tensor relevance = output_relevance;
262 for (
size_t i = modules_.size(); i-- > 0;) {
263 relevance = modules_[i]->propagate_relevance(relevance, configs[i]);
269 [[nodiscard]]
const std::vector<Module*>&
modules()
const {
return modules_; }
279 auto it = activations_.find(
id);
316 std::unordered_map<NodeId, Tensor> activations;
317 std::unordered_map<NodeId, ActivationSnapshot::NodeMetadata> metadata;
318 for (
NodeId id : node_ids) {
319 auto it = activations_.find(
id);
321 activations.emplace(
id,
Tensor(it->second));
326 return ActivationSnapshot(std::move(activations), std::move(node_ids), std::move(metadata));
369 auto it = activations_.find(node_id);
370 if (it == activations_.end()) {
371 throw std::invalid_argument(
372 "ExplainerContext::logit_lens: node_id has no cached activation from the most recent "
376 const NodeId head_input_node =
static_cast<NodeId>(modules_.size() - 1);
377 auto head_input = activations_.find(head_input_node);
378 if (head_input == activations_.end()) {
379 throw std::invalid_argument(
380 "ExplainerContext::logit_lens: no forward pass has run, so the final module's input "
383 if (!(it->second.shape() == head_input->second.shape())) {
384 throw std::invalid_argument(
385 "ExplainerContext::logit_lens: the node's cached activation shape does not match the "
386 "final module's input shape from the most recent forward pass");
389 Module* head = modules_.back();
390 return head->
forward(it->second);
425 if (node_id >= graph_.
node_count() || node_id > modules_.size()) {
426 throw std::invalid_argument(
427 "ExplainerContext::attention_weights: node_id is outside the range of nodes the most "
428 "recent forward pass produced (or no forward pass has run yet)");
431 throw std::invalid_argument(
432 "ExplainerContext::attention_weights: node_id does not name an attention layer");
437 return Tensor(mha->last_attention_weights());
476 const NodeId output_node = output_node_;
482 std::vector<CircuitNode> nodes;
483 nodes.reserve(clean.
node_ids().size());
485 float ablation_effect = 0.0f;
486 if (
id != output_node) {
488 zero_patch.
fill(0.0f);
490 ablation_effect = l2_distance(baseline, patched);
495 std::vector<CircuitEdge> edges;
496 if (!nodes.empty()) {
497 edges.reserve(nodes.size() - 1);
498 for (
size_t i = 0; i + 1 < nodes.size(); ++i) {
499 edges.push_back(
CircuitEdge{nodes[i].id, nodes[i + 1].id, nodes[i].ablation_effect});
504 return CircuitGraph(std::move(nodes), std::move(edges));
520 [[nodiscard]]
static Tensor check_patch_shape(
const Tensor& natural,
const Tensor& patch_value) {
521 if (!(patch_value.
shape() == natural.
shape())) {
522 throw std::invalid_argument(
523 "ExplainerContext::forward_pass_with_patch: patch_value shape must match the patched node's "
524 "natural output shape");
526 return Tensor(patch_value);
539 [[nodiscard]]
static float l2_distance(
const Tensor& a,
const Tensor& b) {
541 float sum_of_squares = 0.0f;
542 for (int64_t i = 0; i < a.numel(); ++i) {
543 const float diff = a.data()[i] - b.data()[i];
544 sum_of_squares += diff * diff;
546 return std::sqrt(sum_of_squares);
549 std::vector<Module*> modules_;
550 ComputationGraph graph_;
552 std::unordered_map<NodeId, Tensor> activations_;
568 bool last_forward_was_patched_ =
false;
Self-contained, enumerable copy of one forward pass's cached activations.
PULSATRIX_ASSERT – debug-only invariant check for programmer errors, distinct from throw (used for ca...
#define PULSATRIX_ASSERT(cond)
Aborts with a diagnostic message if cond is false. Debug-only – use for conditions that indicate a bu...
Definition assert.hpp:22
Reverse-mode autodiff – walks a ComputationGraph backward, accumulating gradients via per-node backwa...
Self-contained circuit-graph artifact – scored nodes and weighted edges.
A copyable, self-contained record of every activation cached during one forward pass,...
Definition activation_snapshot.hpp:39
OpType op_type(NodeId id) const
The op-type tag captured at a node.
Definition activation_snapshot.hpp:90
const std::vector< NodeId > & node_ids() const
Every captured node id, in topological order as of capture time.
Definition activation_snapshot.hpp:84
const Tensor & activation(NodeId id) const
The activation value captured at a node.
Definition activation_snapshot.hpp:73
std::optional< std::string > label(NodeId id) const
The optional label captured at a node.
Definition activation_snapshot.hpp:101
Computes gradients by walking a ComputationGraph in reverse topological order.
Definition autograd.hpp:36
void backward(const ComputationGraph &graph, NodeId root, const Tensor &grad_output)
Runs backward from a single root, seeding its gradient with grad_output.
const Tensor & gradient(NodeId id) const
Retrieves the accumulated gradient for a node after backward() has run.
A copyable, self-contained circuit graph: every node of one forward pass with an ablation importance ...
Definition circuit_graph.hpp:81
Owns every Node in a computation graph and exposes read access for graph-walking code (autograd's bac...
Definition computation_graph.hpp:27
size_t node_count() const
Number of nodes currently in the graph.
Definition computation_graph.hpp:50
NodeId add_node(OpType op_type, Shape shape, std::optional< std::string > label=std::nullopt, std::vector< NodeId > parent_ids={})
Adds a node to the graph and wires it to its parents.
std::vector< NodeId > topological_order() const
Returns every node id in a valid topological order (every node after all its parents).
const Node & node(NodeId id) const
Looks up a node by id.
Wraps an ordered chain of Modules, running them via Module::forward_traced to build a real Computatio...
Definition explainer_context.hpp:67
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 nod...
Definition explainer_context.hpp:368
const Tensor & gradient(NodeId id) const
The cached gradient at a node, from the most recent backward_pass() call.
Definition explainer_context.hpp:294
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 tha...
Definition explainer_context.hpp:153
ActivationSnapshot activation_snapshot() const
Captures the current activation cache into a self-contained ActivationSnapshot.
Definition explainer_context.hpp:313
const ComputationGraph & graph() const
The current graph (from the most recent forward_pass() call).
Definition explainer_context.hpp:272
Tensor attention_weights(NodeId node_id) const
The "attention lens": the per-head attention pattern produced by the attention layer at node_id durin...
Definition explainer_context.hpp:424
Tensor backward_pass(const Tensor &output_grad)
Runs Autograd::backward from the most recent forward_pass()'s output node.
Definition explainer_context.hpp:207
Tensor forward_pass(const Tensor &input)
Runs the full module chain forward, building a fresh graph and caching every node's activation value ...
Definition explainer_context.hpp:95
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 propaga...
Definition explainer_context.hpp:234
ExplainerContext(std::vector< Module * > modules)
Constructs a context over an ordered module chain.
Definition explainer_context.hpp:78
const Tensor & activation(NodeId id) const
The cached activation value at a node, from the most recent forward_pass().
Definition explainer_context.hpp:278
const std::vector< Module * > & modules() const
The module chain, in forward order (not owned).
Definition explainer_context.hpp:269
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 cha...
Definition explainer_context.hpp:474
std::optional< std::string > layer_label(NodeId id) const
Passthrough to Node::label() – e.g. a layer name, for debugging/display.
Definition explainer_context.hpp:508
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 orde...
Definition explainer_context.hpp:244
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
Tensor forward(const Tensor &input)
Runs this module's forward computation.
Definition module.hpp:73
softmax(Q @ K^T / sqrt(head_dim)) @ V, multi-head, with optional RoPE and optional QK-Norm....
Definition multihead_attention_module.hpp:52
A single computation graph node. Owned exclusively by its ComputationGraph (see computation_graph....
Definition node.hpp:29
OpType op_type() const
Definition node.hpp:42
const std::optional< std::string > & label() const
Definition node.hpp:44
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Tensor & fill(float value)
Sets every element to value. Safe no-op on a zero-element tensor.
const Shape & shape() const
This tensor's shape.
Definition tensor.hpp:113
Owns and exposes graph structure – the interpretability substrate every explainer (Phase 2+) walks.
Selects which LRP rule variant a Module::propagate_relevance() call uses.
Abstract base every layer subclasses – NVI forward(), pure-virtual LRP contract.
Multi-head scaled dot-product attention – this codebase's first Module composed out of other real Mod...
std::string module_type_name(const Module &module)
The dynamic type's readable name (e.g. "pulsatrix::SoftmaxModule"), for error messages.
Definition explainer_context.hpp:36
Definition acquisition_functions.hpp:16
size_t NodeId
Stable identifier for a Node within its owning ComputationGraph.
Definition node.hpp:18
std::string lrp_rule_name(LRPRule rule)
Lower-case rule name ("epsilon", "gamma", "alpha_beta", "zbox") for messages/metadata.
Definition lrp_rule_config.hpp:35
@ Attention
Multi-head (scaled dot-product) attention – Phase 3's MultiHeadAttentionModule.
One directed edge of a CircuitGraph, carrying a scalar weight.
Definition circuit_graph.hpp:43
One node of a CircuitGraph: a computation-graph node plus its importance score.
Definition circuit_graph.hpp:20
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).