pulsatrix
Loading...
Searching...
No Matches
sequential_module.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <vector>
8
10
11namespace pulsatrix {
12
27class SequentialModule : public Module {
28public:
36 explicit SequentialModule(std::vector<Module*> layers);
37
46 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
47
50 [[nodiscard]] OpType op_type() const override { return OpType::Composite; }
51
61 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
62
64 [[nodiscard]] bool supports_lrp_rule(LRPRule rule) const override {
65 for (const Module* layer : layers_) {
66 if (!layer->supports_lrp_rule(rule)) {
67 return false;
68 }
69 }
70 return true;
71 }
72
74 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override;
75
80 void set_training(bool training) override;
81
83 [[nodiscard]] const std::vector<Module*>& layers() const { return layers_; }
84
85protected:
91 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
92
93private:
94 std::vector<Module*> layers_;
95 bool has_forwarded_ = false;
96};
97
98} // namespace pulsatrix
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
Composes layers_[0..n-1] in forward() order; backward()/propagate_relevance() chain layers_[n-1....
Definition sequential_module.hpp:27
SequentialModule(std::vector< Module * > layers)
Constructs a container over an ordered sequence of layers.
bool supports_lrp_rule(LRPRule rule) const override
A rule is supported iff every contained layer supports it (the config is forwarded to all).
Definition sequential_module.hpp:64
Tensor forward_impl(const Tensor &input) override
Chains forward() across layers_ in order.
OpType op_type() const override
Composite per charter's closed OpType set โ€“ a container wrapping several ops is genuinely not any sin...
Definition sequential_module.hpp:50
std::vector< NamedParamRef > named_parameters() override
Every contained layer's named_parameters(), prefixed with its index (0.weight).
const std::vector< Module * > & layers() const
The contained layers, in forward-execution order.
Definition sequential_module.hpp:83
void set_training(bool training) override
Sets this container's own training flag and cascades to every contained layer.
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
Chains propagate_relevance() across layers_ in reverse order.
Tensor backward(const Tensor &grad_output) override
Chains backward() across layers_ in reverse order.
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) โ€“ th...
Definition tensor.hpp:29
Abstract base every layer subclasses โ€“ NVI forward(), pure-virtual LRP contract.
Definition acquisition_functions.hpp:16
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
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57