pulsatrix
Loading...
Searching...
No Matches
multihead_attention_module.hpp
Go to the documentation of this file.
1
6#pragma once
7
8#include <optional>
9#include <memory>
10#include <vector>
11
13#include "pulsatrix/module.hpp"
17
18namespace pulsatrix {
19
53public:
66 MultiHeadAttentionModule(int64_t d_model, int64_t num_heads, DeviceBackend* backend, bool use_rope = true,
67 bool use_qk_norm = false);
68
86 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
87
89 [[nodiscard]] OpType op_type() const override { return OpType::Attention; }
90
128 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
129
132 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override;
133
135 void set_training(bool training) override;
136
137 [[nodiscard]] int64_t d_model() const { return d_model_; }
138 [[nodiscard]] int64_t num_heads() const { return num_heads_; }
139 [[nodiscard]] int64_t head_dim() const { return head_dim_; }
140 [[nodiscard]] bool uses_rope() const { return use_rope_; }
141 [[nodiscard]] bool uses_qk_norm() const { return use_qk_norm_; }
142
145 [[nodiscard]] LinearModule& q_proj() { return q_proj_; }
146 [[nodiscard]] LinearModule& k_proj() { return k_proj_; }
147 [[nodiscard]] LinearModule& v_proj() { return v_proj_; }
148 [[nodiscard]] LinearModule& out_proj() { return out_proj_; }
150 [[nodiscard]] RMSNormModule* q_norm() { return q_norm_.get(); }
152 [[nodiscard]] RMSNormModule* k_norm() { return k_norm_.get(); }
154
158 [[nodiscard]] const Tensor& last_attention_weights() const { return last_attn_; }
159
160
162 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
163
164protected:
172 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
173
174private:
175 int64_t d_model_;
176 int64_t num_heads_;
177 int64_t head_dim_;
178 bool use_rope_;
179 bool use_qk_norm_;
180 DeviceBackend* backend_;
181
182 LinearModule q_proj_;
183 LinearModule k_proj_;
184 LinearModule v_proj_;
185 LinearModule out_proj_;
186 SoftmaxModule softmax_;
187 // Held by pointer rather than std::optional purely so the "disabled" state costs nothing
188 // and needs no move/copy of a Module subclass (Module declares a virtual destructor,
189 // which suppresses implicit move construction -- std::optional's in-place paths would
190 // work but the pointer is unambiguous).
191 std::unique_ptr<RoPEModule> q_rope_;
192 std::unique_ptr<RoPEModule> k_rope_;
193 std::unique_ptr<RMSNormModule> q_norm_;
194 std::unique_ptr<RMSNormModule> k_norm_;
195
196 // Forward caches. Q/K/V are the post-QK-Norm, post-RoPE, head-split values -- exactly
197 // the operands the two matmuls' backward and Eq. 15 rules need.
198 int64_t last_N_ = 0;
199 int64_t last_L_ = 0;
200 Tensor last_q_;
201 Tensor last_k_;
202 Tensor last_v_;
203 Tensor last_scores_raw_;
204 Tensor last_attn_;
205 Tensor last_context_;
206 bool has_forwarded_ = false;
207};
208
209} // namespace pulsatrix
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.
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
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
softmax(Q @ K^T / sqrt(head_dim)) @ V, multi-head, with optional RoPE and optional QK-Norm....
Definition multihead_attention_module.hpp:52
RMSNormModule * q_norm()
Q's QK-Norm sub-module, or nullptr when use_qk_norm is false.
Definition multihead_attention_module.hpp:150
const Tensor & last_attention_weights() const
Cached attention weights of the last forward, (N, num_heads, L, L) – the softmax output....
Definition multihead_attention_module.hpp:158
OpType op_type() const override
Attention per the charter's closed OpType set – see OpType::Attention's own note.
Definition multihead_attention_module.hpp:89
std::vector< NamedParamRef > named_parameters() override
Every sub-module's parameters, flattened – Q/K/V/O weights and biases, plus the two QK-Norm gammas wh...
Tensor forward_impl(const Tensor &input) override
Runs the 9-step pipeline: project -> split heads -> (QK-Norm) -> (RoPE) -> scores -> softmax -> conte...
void set_training(bool training) override
Cascades to every sub-module, the same way SequentialModule does.
LinearModule & q_proj()
Definition multihead_attention_module.hpp:145
int64_t num_heads() const
Definition multihead_attention_module.hpp:138
Tensor backward(const Tensor &grad_output) override
Gradient w.r.t. this module's input; sub-module parameter gradients accumulate inside those sub-modul...
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
LRP relevance propagation, composed from the sub-modules' own rules plus AttnLRP's Eq....
bool uses_rope() const
Definition multihead_attention_module.hpp:140
LinearModule & v_proj()
Definition multihead_attention_module.hpp:147
LinearModule & k_proj()
Definition multihead_attention_module.hpp:146
RMSNormModule * k_norm()
K's QK-Norm sub-module, or nullptr when use_qk_norm is false.
Definition multihead_attention_module.hpp:152
bool uses_qk_norm() const
Definition multihead_attention_module.hpp:141
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition multihead_attention_module.hpp:162
int64_t head_dim() const
Definition multihead_attention_module.hpp:139
int64_t d_model() const
Definition multihead_attention_module.hpp:137
MultiHeadAttentionModule(int64_t d_model, int64_t num_heads, DeviceBackend *backend, bool use_rope=true, bool use_qk_norm=false)
Constructs a multi-head attention block with zero-initialized projections.
LinearModule & out_proj()
Definition multihead_attention_module.hpp:148
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
Softmax over the tensor's last dimension, applied independently to every "row" (every fixed combinati...
Definition softmax_module.hpp:20
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Dense/fully-connected layer – the reference Module implementation.
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
@ Attention
Multi-head (scaled dot-product) attention – Phase 3's MultiHeadAttentionModule.
RMS normalization layer (Zhang & Sennrich, 2019) – this codebase's first normalization Module,...
Rotary Position Embedding – fixed per-position pair rotation, epsilon-rule LRP.
Rank-agnostic softmax over the last axis, with AttnLRP's Eq. 13 DTD relevance rule.
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57