pulsatrix
Loading...
Searching...
No Matches
relu_module.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <optional>
9
10namespace pulsatrix {
11
15class ReluModule : public Module {
16public:
27
29 explicit ReluModule(DeviceBackend* backend);
30
42 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
43
45 [[nodiscard]] OpType op_type() const override { return OpType::Activation; }
46
56 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
57
59 [[nodiscard]] bool supports_lrp_rule(LRPRule) const override { return true; }
60
61
63 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return device_; }
64
65protected:
66 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
67
68private:
69 DeviceBackend* backend_;
70 DeviceType device_;
71 Tensor last_input_;
72};
73
74} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
y = max(x, 0), elementwise. No parameters, no parameter gradients.
Definition relu_module.hpp:15
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition relu_module.hpp:63
Tensor backward(const Tensor &grad_output) override
Computes the gradient w.r.t. this module's input.
bool supports_lrp_rule(LRPRule) const override
Pass-through relevance is the same under every rule: supports all of them.
Definition relu_module.hpp:59
ReluModule(DeviceBackend *backend)
As above, on backend's own device (backend->device()).
OpType op_type() const override
Activation per charter's closed OpType set.
Definition relu_module.hpp:45
Tensor forward_impl(const Tensor &input) override
The actual forward computation. Called by forward() after precondition checks.
ReluModule(DeviceBackend *backend, DeviceType device)
Constructs a ReLU module.
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
Pass-through LRP relevance propagation.
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
DeviceType
Which physical device a Tensor's buffer resides on.
Definition device_backend.hpp:17
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