pulsatrix
Loading...
Searching...
No Matches
linear_module.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <optional>
8#include <initializer_list>
9#include <vector>
10
11#include "pulsatrix/module.hpp"
12
13namespace pulsatrix {
14
37class LinearModule : public Module {
38public:
50 LinearModule(int64_t in_features, int64_t out_features, DeviceBackend* backend, DeviceType device);
51
58 LinearModule(int64_t in_features, int64_t out_features, DeviceBackend* backend);
59
74 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
75
77 [[nodiscard]] OpType op_type() const override { return OpType::Linear; }
78
80 void set_weight(std::initializer_list<float> values);
81
83 void set_bias(std::initializer_list<float> values);
84
86 void set_weight(const std::vector<float>& values);
87
89 void set_bias(const std::vector<float>& values);
90
91 [[nodiscard]] const Tensor& weight() const { return weight_; }
92 [[nodiscard]] const Tensor& bias() const { return bias_; }
93 [[nodiscard]] const Tensor& weight_grad() const { return weight_grad_; }
94 [[nodiscard]] const Tensor& bias_grad() const { return bias_grad_; }
95
118 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
119
121 [[nodiscard]] bool supports_lrp_rule(LRPRule) const override { return true; }
122
123 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override {
124 return {{"weight", {&weight_, &weight_grad_}}, {"bias", {&bias_, &bias_grad_}}};
125 }
126
127
129 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return weight_.device(); }
130
131protected:
132 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
133
134private:
135 int64_t in_features_;
136 int64_t out_features_;
137 DeviceBackend* backend_;
138 Tensor weight_;
139 Tensor bias_;
140 Tensor weight_grad_;
141 Tensor bias_grad_;
142 Tensor last_input_;
143 Tensor last_pre_bias_output_;
144 bool has_forwarded_ = false;
145};
146
147} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
y = x @ W + b, batched (x is (N, in_features), y is (N, out_features)) – migrated from the original u...
Definition linear_module.hpp:37
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
Epsilon-rule LRP relevance propagation (Bach et al. 2015), applied independently per example in the b...
Tensor backward(const Tensor &grad_output) override
Computes the gradient w.r.t. this module's input, and accumulates the weight/bias gradients internall...
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition linear_module.hpp:129
const Tensor & weight() const
Definition linear_module.hpp:91
LinearModule(int64_t in_features, int64_t out_features, DeviceBackend *backend, DeviceType device)
Constructs a linear layer with zero-initialized weight/bias.
bool supports_lrp_rule(LRPRule) const override
Implements every LRPRule.
Definition linear_module.hpp:121
std::vector< NamedParamRef > named_parameters() override
This module's trainable parameters, each with its hierarchical name – the one place a module declares...
Definition linear_module.hpp:123
void set_bias(const std::vector< float > &values)
Vector overload for runtime-sized sources – see Tensor's own vector ctor.
void set_weight(const std::vector< float > &values)
Vector overload for runtime-sized sources – see Tensor's own vector ctor.
void set_bias(std::initializer_list< float > values)
Overwrites the bias buffer – test/initialization use only.
OpType op_type() const override
Linear per charter's closed OpType set.
Definition linear_module.hpp:77
LinearModule(int64_t in_features, int64_t out_features, DeviceBackend *backend)
As above, on backend's own device (backend->device()).
Tensor forward_impl(const Tensor &input) override
The actual forward computation. Called by forward() after precondition checks.
const Tensor & bias_grad() const
Definition linear_module.hpp:94
const Tensor & weight_grad() const
Definition linear_module.hpp:93
const Tensor & bias() const
Definition linear_module.hpp:92
void set_weight(std::initializer_list< float > values)
Overwrites the weight buffer – test/initialization use only.
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
DeviceType device() const
Which device this tensor's buffer conceptually resides on.
Definition tensor.hpp:122
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