pulsatrix
Loading...
Searching...
No Matches
module.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <optional>
8#include <stdexcept>
9#include <string>
10#include <utility>
11#include <vector>
12
13#include "pulsatrix/assert.hpp"
17#include "pulsatrix/node.hpp"
18#include "pulsatrix/op_type.hpp"
19#include "pulsatrix/tensor.hpp"
20
21namespace pulsatrix {
22
28
38 std::string name;
40};
41
58class Module {
59public:
60 virtual ~Module() = default;
61
73 [[nodiscard]] Tensor forward(const Tensor& input) {
74 if (input.numel() <= 0) {
75 throw std::invalid_argument("Module::forward: input must not be empty");
76 }
77 if (const std::optional<DeviceType> device = compute_device()) {
78 require_device(input, *device, "Module::forward");
79 }
80 return forward_impl(input);
81 }
82
91 [[nodiscard]] virtual std::optional<DeviceType> compute_device() const { return std::nullopt; }
92
111 [[nodiscard]] std::pair<Tensor, NodeId> forward_traced(const Tensor& input, NodeId input_node,
112 ComputationGraph& graph, Autograd& autograd) {
113 Tensor output = forward(input);
114 NodeId node_id = graph.add_node(op_type(), output.shape(), std::nullopt, {input_node});
115 autograd.register_backward(node_id, [this](const Tensor& grad_output) { return this->backward(grad_output); });
116 return {std::move(output), node_id};
117 }
118
125 [[nodiscard]] virtual Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) = 0;
126
135 [[nodiscard]] virtual bool supports_lrp_rule(LRPRule rule) const { return rule == LRPRule::Epsilon; }
136
148 [[nodiscard]] virtual Tensor backward(const Tensor& grad_output) = 0;
149
156 [[nodiscard]] virtual OpType op_type() const = 0;
157
167 [[nodiscard]] virtual std::vector<NamedParamRef> named_parameters() { return {}; }
168
180 [[nodiscard]] virtual std::vector<ParamRef> parameters() {
181 std::vector<ParamRef> params;
182 for (const NamedParamRef& p : named_parameters()) {
183 params.push_back(p.ref);
184 }
185 return params;
186 }
187
199 void set_requires_grad(bool requires_grad, const std::string& prefix = "") {
200 if (prefix.empty()) {
201 // parameters(), not named_parameters(): also reaches a legacy module that has no names.
202 for (ParamRef p : parameters()) {
203 p.value->set_requires_grad(requires_grad);
204 }
205 return;
206 }
207 std::vector<Tensor*> selected;
208 for (const NamedParamRef& p : named_parameters()) {
209 if (p.name == prefix || p.name.rfind(prefix + ".", 0) == 0) {
210 selected.push_back(p.ref.value);
211 }
212 }
213 if (selected.empty()) {
214 throw std::invalid_argument("Module::set_requires_grad: no parameter named or under '" + prefix + "'");
215 }
216 for (Tensor* value : selected) {
217 value->set_requires_grad(requires_grad);
218 }
219 }
220
234 virtual void set_training(bool training) { training_ = training; }
235
237 [[nodiscard]] bool is_training() const { return training_; }
238
239protected:
241 [[nodiscard]] virtual Tensor forward_impl(const Tensor& input) = 0;
242
243private:
244 bool training_ = true;
245};
246
254inline void append_named_parameters(std::vector<NamedParamRef>& out, const std::string& prefix, Module& child) {
255 std::vector<NamedParamRef> named = child.named_parameters();
256 if (named.empty()) {
257 std::vector<ParamRef> plain = child.parameters();
258 for (size_t i = 0; i < plain.size(); ++i) {
259 named.push_back({std::to_string(i), plain[i]});
260 }
261 }
262 for (NamedParamRef& p : named) {
263 out.push_back({prefix + "." + p.name, p.ref});
264 }
265}
266
267} // namespace pulsatrix
PULSATRIX_ASSERT – debug-only invariant check for programmer errors, distinct from throw (used for ca...
Reverse-mode autodiff – walks a ComputationGraph backward, accumulating gradients via per-node backwa...
Computes gradients by walking a ComputationGraph in reverse topological order.
Definition autograd.hpp:36
void register_backward(NodeId id, BackwardFn fn)
Registers how to compute this node's input gradient from its output gradient.
Owns every Node in a computation graph and exposes read access for graph-walking code (autograd's bac...
Definition computation_graph.hpp:27
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.
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
virtual std::vector< ParamRef > parameters()
This module's trainable parameters and their gradients, for an optimizer to update uniformly across m...
Definition module.hpp:180
virtual Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config)=0
Computes this module's contribution to LRP relevance propagation.
virtual std::vector< NamedParamRef > named_parameters()
This module's trainable parameters, each with its hierarchical name – the one place a module declares...
Definition module.hpp:167
void set_requires_grad(bool requires_grad, const std::string &prefix="")
Freezes (false) or unfreezes (true) parameters by name (roadmap FND-2).
Definition module.hpp:199
virtual OpType op_type() const =0
This module's operation-category tag, for ComputationGraph node tagging.
virtual void set_training(bool training)
Sets this module's training/eval mode. Defaults to training (matches every mainstream framework's Mod...
Definition module.hpp:234
Tensor forward(const Tensor &input)
Runs this module's forward computation.
Definition module.hpp:73
virtual Tensor forward_impl(const Tensor &input)=0
The actual forward computation. Called by forward() after precondition checks.
bool is_training() const
Whether this module is currently in training mode.
Definition module.hpp:237
std::pair< Tensor, NodeId > forward_traced(const Tensor &input, NodeId input_node, ComputationGraph &graph, Autograd &autograd)
Runs forward() while also registering a ComputationGraph node (tagged with this module's op_type(),...
Definition module.hpp:111
virtual Tensor backward(const Tensor &grad_output)=0
Computes the gradient w.r.t. this module's input, given the gradient w.r.t. its output....
virtual bool supports_lrp_rule(LRPRule rule) const
Whether propagate_relevance() implements rule (no silent fallback: callers such as ExplainerContext::...
Definition module.hpp:135
virtual ~Module()=default
virtual std::optional< DeviceType > compute_device() const
The device this module computes on, so forward() can reject an input on another device before any ker...
Definition module.hpp:91
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
int64_t numel() const
Total element count – shape().numel().
Definition tensor.hpp:116
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.
Definition acquisition_functions.hpp:16
size_t NodeId
Stable identifier for a Node within its owning ComputationGraph.
Definition node.hpp:18
void require_device(const Tensor &t, DeviceType expected, const char *where)
Throws unless t lives on expected – the check every module and loss runs on the tensors handed to it ...
Definition tensor.hpp:285
void append_named_parameters(std::vector< NamedParamRef > &out, const std::string &prefix, Module &child)
Appends child's named parameters to out, each renamed to prefix.name – the one step every container's...
Definition module.hpp:254
OpType
The op-type tag a Node carries. Charter Part 2 §3: nodes are tagged by a small closed set of op types...
Definition op_type.hpp:19
LRPRule
The LRP rule family a module applies. Semantics follow Zennit 1.0.0 exactly (Anders et al....
Definition lrp_rule_config.hpp:19
Computation graph node – op type, shape, optional label, parent/child edges.
Closed set of operation categories every graph Node is tagged with.
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57
A parameter together with its hierarchical, dot-separated name relative to the module that reported i...
Definition module.hpp:37
std::string name
Definition module.hpp:38
ParamRef ref
Definition module.hpp:39
A trainable parameter and its accumulated gradient, as owned by some Module.
Definition module.hpp:24
Tensor * value
Definition module.hpp:25
Tensor * grad
Definition module.hpp:26
N-dimensional tensor – owns a buffer via DeviceBackend*, RAII (Rule of Five).