67 bool use_qk_norm =
false);
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_; }
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_;
206 bool has_forwarded_ =
false;
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