pulsatrix
Loading...
Searching...
No Matches
residual_module.hpp
Go to the documentation of this file.
1
6#pragma once
7
8#include <optional>
10
11namespace pulsatrix {
12
38class ResidualModule : public Module {
39public:
54
66 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
67
69 [[nodiscard]] OpType op_type() const override { return OpType::Elementwise; }
70
81 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
82
84 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override;
85
87 void set_training(bool training) override;
88
89 [[nodiscard]] Module& inner() { return *inner_; }
90
91
93 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
94
95protected:
101 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
102
103private:
104 Module* inner_;
105 DeviceBackend* backend_;
106
107 Tensor last_x_;
108 Tensor last_f_x_;
109 bool has_forwarded_ = false;
110};
111
112} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
virtual DeviceType device() const noexcept=0
Which device this backend's buffers reside on.
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
y = x + inner->forward(x) for an arbitrary already-built Module. The classic ResNet shortcut connecti...
Definition residual_module.hpp:38
std::vector< NamedParamRef > named_parameters() override
inner_'s own named_parameters(), prefixed inner. โ€“ this module owns none of its own.
void set_training(bool training) override
Cascades to inner_, the same way SequentialModule/MultiHeadAttentionModule do.
OpType op_type() const override
Elementwise per this module's own op_type() note above.
Definition residual_module.hpp:69
Tensor backward(const Tensor &grad_output) override
Gradient w.r.t. this module's input: both paths receive grad_output unchanged (real gradient of a pla...
Module & inner()
Definition residual_module.hpp:89
Tensor forward_impl(const Tensor &input) override
y = x + inner_->forward(x).
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition residual_module.hpp:93
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
LRP relevance propagation: the residual epsilon/z-rule split between x and inner_->forward(x),...
ResidualModule(Module *inner, DeviceBackend *backend)
Constructs a residual wrapper around an existing Module.
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
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57