pulsatrix
Loading...
Searching...
No Matches
explainer_context.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <cmath>
8#include <cstddef>
9#include <cstdlib>
10#include <optional>
11#include <stdexcept>
12#include <string>
13#include <typeinfo>
14#include <unordered_map>
15#include <utility>
16#include <vector>
17
19#include "pulsatrix/assert.hpp"
24#include "pulsatrix/module.hpp"
26#include "pulsatrix/tensor.hpp"
27
28#if defined(__GNUG__)
29#include <cxxabi.h>
30#endif
31
32namespace pulsatrix {
33
34namespace detail {
36inline std::string module_type_name(const Module& module) {
37 const char* raw = typeid(module).name();
38#if defined(__GNUG__)
39 int status = 0;
40 char* demangled = abi::__cxa_demangle(raw, nullptr, nullptr, &status);
41 std::string name = (status == 0 && demangled != nullptr) ? demangled : raw;
42 std::free(demangled);
43 return name;
44#else
45 return raw;
46#endif
47}
48} // namespace detail
49
68public:
78 explicit ExplainerContext(std::vector<Module*> modules) : modules_(std::move(modules)) {
79 if (modules_.empty()) {
80 throw std::invalid_argument("ExplainerContext: modules must not be empty");
81 }
82 for (Module* m : modules_) {
83 if (m == nullptr) {
84 throw std::invalid_argument("ExplainerContext: modules must not contain a null pointer");
85 }
86 }
87 }
88
95 Tensor forward_pass(const Tensor& input) {
96 last_forward_was_patched_ = false;
97 graph_ = ComputationGraph{};
98 autograd_ = Autograd{};
99 activations_.clear();
100
101 input_node_ = graph_.add_node(OpType::Elementwise, input.shape(), "input");
102 activations_.emplace(input_node_, Tensor(input));
103
104 Tensor current = 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;
111 }
112 output_node_ = current_node;
113 return current;
114 }
115
153 Tensor forward_pass_with_patch(const Tensor& input, NodeId patch_node_id, const Tensor& patch_value) {
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");
157 }
158
160 Autograd autograd;
161 std::unordered_map<NodeId, Tensor> activations;
162
163 NodeId input_node = graph.add_node(OpType::Elementwise, input.shape(), "input");
164
165 Tensor current = input;
166 NodeId current_node = input_node;
167 if (patch_node_id == input_node) {
168 current = check_patch_shape(current, patch_value);
169 }
170 activations.emplace(input_node, Tensor(current));
171
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);
176 }
177 activations.emplace(node_id, Tensor(output));
178 current = std::move(output);
179 current_node = node_id;
180 }
181
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;
187 return current;
188 }
189
207 Tensor backward_pass(const Tensor& output_grad) {
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().");
214 }
215 autograd_.backward(graph_, output_node_, output_grad);
216 return Tensor(autograd_.gradient(input_node_));
217 }
218
234 Tensor relevance_pass(const Tensor& output_relevance, const LRPRuleConfig& config) {
235 return relevance_pass(output_relevance, std::vector<LRPRuleConfig>(modules_.size(), config));
236 }
237
244 Tensor relevance_pass(const Tensor& output_relevance, const std::vector<LRPRuleConfig>& configs) {
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.");
250 }
251 if (configs.size() != modules_.size()) {
252 throw std::invalid_argument("ExplainerContext::relevance_pass: need exactly one LRPRuleConfig per module");
253 }
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) + " (" +
257 detail::module_type_name(*modules_[i]) + ") does not implement the " +
258 lrp_rule_name(configs[i].rule) + " LRP rule");
259 }
260 }
261 Tensor relevance = output_relevance;
262 for (size_t i = modules_.size(); i-- > 0;) {
263 relevance = modules_[i]->propagate_relevance(relevance, configs[i]);
264 }
265 return relevance;
266 }
267
269 [[nodiscard]] const std::vector<Module*>& modules() const { return modules_; }
270
272 [[nodiscard]] const ComputationGraph& graph() const { return graph_; }
273
278 [[nodiscard]] const Tensor& activation(NodeId id) const {
279 auto it = activations_.find(id);
280 PULSATRIX_ASSERT(it != activations_.end());
281 return it->second;
282 }
283
294 [[nodiscard]] const Tensor& gradient(NodeId id) const { return autograd_.gradient(id); }
295
314 std::vector<NodeId> node_ids = graph_.topological_order();
315
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);
320 PULSATRIX_ASSERT(it != activations_.end());
321 activations.emplace(id, Tensor(it->second));
322 const Node& node = graph_.node(id);
323 metadata.emplace(id, ActivationSnapshot::NodeMetadata{node.op_type(), node.label()});
324 }
325
326 return ActivationSnapshot(std::move(activations), std::move(node_ids), std::move(metadata));
327 }
328
368 [[nodiscard]] Tensor logit_lens(NodeId node_id) const {
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 "
373 "forward pass");
374 }
375
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 "
381 "shape is unknown");
382 }
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");
387 }
388
389 Module* head = modules_.back();
390 return head->forward(it->second);
391 }
392
424 [[nodiscard]] Tensor attention_weights(NodeId node_id) const {
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)");
429 }
430 if (graph_.node(node_id).op_type() != OpType::Attention) {
431 throw std::invalid_argument(
432 "ExplainerContext::attention_weights: node_id does not name an attention layer");
433 }
434
435 auto* mha = dynamic_cast<MultiHeadAttentionModule*>(modules_[node_id - 1]);
436 PULSATRIX_ASSERT(mha != nullptr);
437 return Tensor(mha->last_attention_weights());
438 }
439
474 [[nodiscard]] CircuitGraph build_circuit_graph(const Tensor& input) {
475 const Tensor baseline = forward_pass(input);
476 const NodeId output_node = output_node_;
477 // Captured before any patching: every patched pass replaces graph_/activations_
478 // wholesale, so the baseline's shapes and per-node metadata are read from this
479 // self-contained copy rather than from state the loop itself overwrites.
481
482 std::vector<CircuitNode> nodes;
483 nodes.reserve(clean.node_ids().size());
484 for (NodeId id : clean.node_ids()) {
485 float ablation_effect = 0.0f;
486 if (id != output_node) {
487 Tensor zero_patch(clean.activation(id));
488 zero_patch.fill(0.0f);
489 const Tensor patched = forward_pass_with_patch(input, id, zero_patch);
490 ablation_effect = l2_distance(baseline, patched);
491 }
492 nodes.push_back(CircuitNode{id, clean.op_type(id), clean.label(id), ablation_effect});
493 }
494
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});
500 }
501 }
502
503 (void)forward_pass(input);
504 return CircuitGraph(std::move(nodes), std::move(edges));
505 }
506
508 [[nodiscard]] std::optional<std::string> layer_label(NodeId id) const { return graph_.node(id).label(); }
509
510private:
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");
525 }
526 return Tensor(patch_value);
527 }
528
539 [[nodiscard]] static float l2_distance(const Tensor& a, const Tensor& b) {
540 PULSATRIX_ASSERT(a.numel() == b.numel());
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;
545 }
546 return std::sqrt(sum_of_squares);
547 }
548
549 std::vector<Module*> modules_;
550 ComputationGraph graph_;
551 Autograd autograd_;
552 std::unordered_map<NodeId, Tensor> activations_;
553 NodeId input_node_ = 0;
554 NodeId output_node_ = 0;
568 bool last_forward_was_patched_ = false;
569};
570
571} // namespace pulsatrix
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.
Per-node metadata captured alongside the activation value.
Definition activation_snapshot.hpp:42
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).