pulsatrix
Loading...
Searching...
No Matches
rms_norm_module.hpp
Go to the documentation of this file.
1
6#pragma once
7
8#include <optional>
9#include <initializer_list>
10#include <vector>
11
12#include "pulsatrix/module.hpp"
13
14namespace pulsatrix {
15
32class RMSNormModule : public Module {
33public:
46 RMSNormModule(int64_t num_features, DeviceBackend* backend, DeviceType device,
47 float eps = 1e-6f);
48
51 RMSNormModule(int64_t num_features, DeviceBackend* backend);
52
61 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
62
64 [[nodiscard]] OpType op_type() const override { return OpType::Normalization; }
65
67 void set_gamma(std::initializer_list<float> values);
68
70 void set_gamma(const std::vector<float>& values);
71
72 [[nodiscard]] const Tensor& gamma() const { return gamma_; }
73 [[nodiscard]] const Tensor& gamma_grad() const { return gamma_grad_; }
74
86 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
87
88 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override {
89 return {
90 {"weight", {&gamma_, &gamma_grad_}},
91 };
92 }
93
94
96 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return gamma_.device(); }
97
98protected:
99 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
100
101private:
102 int64_t num_features_;
103 float eps_;
104 DeviceBackend* backend_;
105 Tensor gamma_;
106 Tensor gamma_grad_;
107 Tensor last_input_;
108 Tensor last_rms_; // (N,) one rms per batch row, on the module's device
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
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
y_{n,i} = gamma_i * x_{n,i} / rms(x_n), rms(x_n) = sqrt(mean_i(x_{n,i}^2) + eps), computed independen...
Definition rms_norm_module.hpp:32
RMSNormModule(int64_t num_features, DeviceBackend *backend)
On backend's own device (backend->device()), default eps. Previously the device defaulted to Cpu rega...
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
Identity-rule LRP relevance propagation (AttnLRP, Achtibat et al. 2024).
RMSNormModule(int64_t num_features, DeviceBackend *backend, DeviceType device, float eps=1e-6f)
Constructs an RMSNorm layer with zero-initialized gamma.
void set_gamma(std::initializer_list< float > values)
Overwrites the gamma 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 gradient internally.
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 rms_norm_module.hpp:96
std::vector< NamedParamRef > named_parameters() override
This module's trainable parameters, each with its hierarchical name – the one place a module declares...
Definition rms_norm_module.hpp:88
void set_gamma(const std::vector< float > &values)
Vector overload for runtime-sized sources – see Tensor's own vector ctor.
OpType op_type() const override
Normalization per charter's closed OpType set.
Definition rms_norm_module.hpp:64
const Tensor & gamma() const
Definition rms_norm_module.hpp:72
const Tensor & gamma_grad() const
Definition rms_norm_module.hpp:73
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
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57