pulsatrix
Loading...
Searching...
No Matches
batch_norm_module.hpp
Go to the documentation of this file.
1
7#pragma once
8
9#include <optional>
10#include <initializer_list>
11#include <vector>
12
13#include "pulsatrix/module.hpp"
14
15namespace pulsatrix {
16
17class BatchNormFold;
18
42class BatchNormModule : public Module {
43public:
59 BatchNormModule(int64_t num_channels, DeviceBackend* backend, DeviceType device,
60 float eps = 1e-6f, float momentum = 0.1f);
61
64 BatchNormModule(int64_t num_channels, DeviceBackend* backend);
65
75 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
76
78 [[nodiscard]] OpType op_type() const override { return OpType::Normalization; }
79
81 void set_gamma(std::initializer_list<float> values);
83 void set_beta(std::initializer_list<float> values);
85 void set_gamma(const std::vector<float>& values);
87 void set_beta(const std::vector<float>& values);
88
90 [[nodiscard]] const Tensor& running_mean() const { return running_mean_; }
92 [[nodiscard]] const Tensor& running_var() const { return running_var_; }
97 void set_running_mean(const std::vector<float>& values);
103 void set_running_var(const std::vector<float>& values);
104
105 [[nodiscard]] const Tensor& gamma() const { return gamma_; }
106 [[nodiscard]] const Tensor& beta() const { return beta_; }
107 [[nodiscard]] const Tensor& gamma_grad() const { return gamma_grad_; }
108 [[nodiscard]] const Tensor& beta_grad() const { return beta_grad_; }
109
120 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
121
122 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override {
123 return {{"weight", {&gamma_, &gamma_grad_}}, {"bias", {&beta_, &beta_grad_}}};
124 }
125
126
128 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return gamma_.device(); }
129
130protected:
131 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
132
133private:
134 friend class BatchNormFold;
135
136 // Which statistics the cached forward used, so backward() matches it even if the mode
137 // changes in between.
138 enum class Mode { Batch, Running, Folded };
139
140 int64_t num_channels_;
141 float eps_;
142 float momentum_;
143 DeviceBackend* backend_;
144 Tensor gamma_; // shape (num_channels,)
145 Tensor beta_; // shape (num_channels,)
146 Tensor gamma_grad_;
147 Tensor beta_grad_;
148 Tensor last_input_; // (N, num_channels, H, W)
149 Tensor last_xhat_; // (N, num_channels, H, W)
150 Tensor last_std_; // (num_channels,) one std per channel, on the module's device
151 Tensor running_mean_; // (num_channels,)
152 Tensor running_var_; // (num_channels,)
153 bool has_forwarded_ = false;
154 Mode last_mode_ = Mode::Batch;
155 bool folded_ = false; // set by BatchNormFold: forward/backward become the identity
156};
157
158} // namespace pulsatrix
While alive, merges bn's affine map into conv's weights and makes bn an exact identity; on destructio...
Definition batch_norm_fold.hpp:28
y_{n,c,h,w} = gamma_c * (x_{n,c,h,w} - mu_c)/std_c + beta_c, mu_c/std_c computed per channel c over e...
Definition batch_norm_module.hpp:42
const Tensor & beta() const
Definition batch_norm_module.hpp:106
BatchNormModule(int64_t num_channels, DeviceBackend *backend, DeviceType device, float eps=1e-6f, float momentum=0.1f)
Constructs a BatchNorm layer with zero-initialized gamma and beta.
void set_running_var(const std::vector< float > &values)
Overwrites the running variance, e.g. when loading a pretrained model.
void set_gamma(std::initializer_list< float > values)
Overwrites the per-channel gamma buffer – test/initialization use only.
void set_beta(std::initializer_list< float > values)
Overwrites the per-channel beta buffer – test/initialization use only.
Tensor backward(const Tensor &grad_output) override
Computes the gradient w.r.t. this module's input, and accumulates gamma's/ beta's gradients internall...
void set_beta(const std::vector< float > &values)
Vector overload for runtime-sized sources – see Tensor's own vector ctor.
const Tensor & running_var() const
Per-channel running variance, used in eval mode. Shape (num_channels).
Definition batch_norm_module.hpp:92
std::vector< NamedParamRef > named_parameters() override
This module's trainable parameters, each with its hierarchical name – the one place a module declares...
Definition batch_norm_module.hpp:122
const Tensor & gamma_grad() const
Definition batch_norm_module.hpp:107
const Tensor & beta_grad() const
Definition batch_norm_module.hpp:108
OpType op_type() const override
Normalization per charter's closed OpType set.
Definition batch_norm_module.hpp:78
BatchNormModule(int64_t num_channels, DeviceBackend *backend)
On backend's own device (backend->device()), default eps. Previously the device defaulted to Cpu rega...
const Tensor & gamma() const
Definition batch_norm_module.hpp:105
void set_gamma(const std::vector< float > &values)
Vector overload for runtime-sized sources – see Tensor's own vector ctor.
Tensor forward_impl(const Tensor &input) override
The actual forward computation. Called by forward() after precondition checks.
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition batch_norm_module.hpp:128
void set_running_mean(const std::vector< float > &values)
Overwrites the running mean, e.g. when loading a pretrained model.
const Tensor & running_mean() const
Per-channel running mean, used in eval mode. Shape (num_channels).
Definition batch_norm_module.hpp:90
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
Identity-rule LRP relevance propagation (AttnLRP, Achtibat et al. 2024).
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
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
One collated batch: one stacked Tensor per Sample field position.
Definition collate.hpp:17
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57