pulsatrix
Loading...
Searching...
No Matches
conv2d_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
31class Conv2DModule : public Module {
32public:
46 Conv2DModule(int64_t in_channels, int64_t out_channels, int64_t kernel_h, int64_t kernel_w,
47 DeviceBackend* backend, int64_t stride = 1, int64_t padding = 0);
48
50 [[nodiscard]] int64_t stride() const { return stride_; }
52 [[nodiscard]] int64_t padding() const { return padding_; }
53
64 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
65
67 [[nodiscard]] OpType op_type() const override { return OpType::Conv; }
68
70 void set_kernel(std::initializer_list<float> values);
71
73 void set_bias(std::initializer_list<float> values);
74
76 void set_kernel(const std::vector<float>& values);
77
79 void set_bias(const std::vector<float>& values);
80
81 [[nodiscard]] const Tensor& kernel() const { return kernel_; }
82 [[nodiscard]] const Tensor& bias() const { return bias_; }
83 [[nodiscard]] const Tensor& kernel_grad() const { return kernel_grad_; }
84 [[nodiscard]] const Tensor& bias_grad() const { return bias_grad_; }
85
108 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
109
111 [[nodiscard]] bool supports_lrp_rule(LRPRule) const override { return true; }
112
113 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override {
114 return {{"weight", {&kernel_, &kernel_grad_}}, {"bias", {&bias_, &bias_grad_}}};
115 }
116
117
119 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
120
121protected:
130 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
131
132private:
133 int64_t in_channels_;
134 int64_t out_channels_;
135 int64_t kernel_h_;
136 int64_t kernel_w_;
137 int64_t stride_;
138 int64_t padding_;
139 DeviceBackend* backend_;
140 Tensor kernel_; // shape (out_channels, in_channels, kernel_h, kernel_w)
141 Tensor bias_; // shape (out_channels,)
142 Tensor kernel_grad_;
143 Tensor bias_grad_;
144 Tensor last_input_; // (N, in_channels, H, W)
145 Tensor last_im2col_; // (N, patch_size, out_h*out_w), cached for backward
146 Tensor last_pre_bias_output_; // (N, out_channels, out_h, out_w), cached for LRP
147 int64_t last_out_h_ = 0;
148 int64_t last_out_w_ = 0;
149 bool has_forwarded_ = false;
150
151 [[nodiscard]] ConvGeometry geometry() const {
152 const auto s = static_cast<size_t>(stride_), p = static_cast<size_t>(padding_);
153 return {static_cast<size_t>(kernel_h_), static_cast<size_t>(kernel_w_), s, s, p, p};
154 }
155};
156
157} // namespace pulsatrix
2D convolution, batched (input/output are rank-4: N x channels x H x W) – migrated from the original ...
Definition conv2d_module.hpp:31
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition conv2d_module.hpp:119
Tensor backward(const Tensor &grad_output) override
Computes gradients w.r.t. input, kernel, and bias.
std::vector< NamedParamRef > named_parameters() override
This module's trainable parameters, each with its hierarchical name – the one place a module declares...
Definition conv2d_module.hpp:113
OpType op_type() const override
Conv per charter's closed OpType set.
Definition conv2d_module.hpp:67
const Tensor & kernel() const
Definition conv2d_module.hpp:81
const Tensor & kernel_grad() const
Definition conv2d_module.hpp:83
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
Epsilon-rule LRP relevance propagation.
void set_kernel(const std::vector< float > &values)
Vector overload for runtime-sized sources – see Tensor's own vector ctor.
const Tensor & bias_grad() const
Definition conv2d_module.hpp:84
bool supports_lrp_rule(LRPRule) const override
Implements every LRPRule.
Definition conv2d_module.hpp:111
int64_t stride() const
Step between windows along both axes.
Definition conv2d_module.hpp:50
Conv2DModule(int64_t in_channels, int64_t out_channels, int64_t kernel_h, int64_t kernel_w, DeviceBackend *backend, int64_t stride=1, int64_t padding=0)
Constructs a conv layer with zero-initialized kernel/bias.
Tensor forward_impl(const Tensor &input) override
The actual forward computation (im2col + gemm + per-channel bias add).
int64_t padding() const
Zero padding on every side, along both axes.
Definition conv2d_module.hpp:52
void set_bias(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.
void set_kernel(std::initializer_list< float > values)
Overwrites the kernel buffer – test/initialization use only.
const Tensor & bias() const
Definition conv2d_module.hpp:82
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
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
Window geometry for DeviceBackend::im2col / col2im_add: kernel size, stride and zero padding per axis...
Definition device_backend.hpp:171
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57