softmax(Q @ K^T / sqrt(head_dim)) @ V, multi-head, with optional RoPE and optional QK-Norm. Shape (N, L, d_model) -> (N, L, d_model).
More...
#include <multihead_attention_module.hpp>
|
| | 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.
|
| |
| Tensor | backward (const Tensor &grad_output) override |
| | Gradient w.r.t. this module's input; sub-module parameter gradients accumulate inside those sub-modules (reachable through parameters()).
|
| |
| OpType | op_type () const override |
| | Attention per the charter's closed OpType set – see OpType::Attention's own note.
|
| |
| 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. 15 bilinear ("uniform") rule for the two batched matmuls.
|
| |
| 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 when enabled. RoPE and softmax contribute none.
|
| |
| void | set_training (bool training) override |
| | Cascades to every sub-module, the same way SequentialModule does.
|
| |
| int64_t | d_model () const |
| |
| int64_t | num_heads () const |
| |
| int64_t | head_dim () const |
| |
| bool | uses_rope () const |
| |
| bool | uses_qk_norm () const |
| |
| const Tensor & | last_attention_weights () const |
| | Cached attention weights of the last forward, (N, num_heads, L, L) – the softmax output. Exposed because "what did each head attend to" is the single most-asked explainability question about this module.
|
| |
| std::optional< DeviceType > | compute_device () const override |
| | Where this layer computes, so forward() rejects an input on another device (FND-8).
|
| |
|
| LinearModule & | q_proj () |
| |
| LinearModule & | k_proj () |
| |
| LinearModule & | v_proj () |
| |
| LinearModule & | out_proj () |
| |
| RMSNormModule * | q_norm () |
| | Q's QK-Norm sub-module, or nullptr when use_qk_norm is false.
|
| |
| RMSNormModule * | k_norm () |
| | K's QK-Norm sub-module, or nullptr when use_qk_norm is false.
|
| |
| virtual | ~Module ()=default |
| |
| Tensor | forward (const Tensor &input) |
| | Runs this module's forward computation.
|
| |
| std::pair< Tensor, NodeId > | forward_traced (const Tensor &input, NodeId input_node, ComputationGraph &graph, Autograd &autograd) |
| | Runs forward() while also registering a ComputationGraph node (tagged with this module's op_type(), parented to input_node) and wiring an Autograd backward function that reuses this module's own backward() – the opt-in traced/explainable path, per Phase 2 Mission 0.
|
| |
| virtual bool | supports_lrp_rule (LRPRule rule) const |
| | Whether propagate_relevance() implements rule (no silent fallback: callers such as ExplainerContext::relevance_pass() throw rather than run a module on a rule it does not implement).
|
| |
| virtual std::vector< ParamRef > | parameters () |
| | This module's trainable parameters and their gradients, for an optimizer to update uniformly across module types.
|
| |
| void | set_requires_grad (bool requires_grad, const std::string &prefix="") |
| | Freezes (false) or unfreezes (true) parameters by name (roadmap FND-2).
|
| |
| bool | is_training () const |
| | Whether this module is currently in training mode.
|
| |
|
| Tensor | forward_impl (const Tensor &input) override |
| | Runs the 9-step pipeline: project -> split heads -> (QK-Norm) -> (RoPE) -> scores -> softmax -> context -> merge heads -> output projection.
|
| |
softmax(Q @ K^T / sqrt(head_dim)) @ V, multi-head, with optional RoPE and optional QK-Norm. Shape (N, L, d_model) -> (N, L, d_model).
Composition over inheritance. Every sub-operation that already exists as a shipped, tested Module in this codebase is held here as a real member object and driven through its own forward()/backward()/propagate_relevance():
LinearModule q_proj_/k_proj_/v_proj_/out_proj_ (all d_model -> d_model),
RoPEModule q_rope_/k_rope_ (only when use_rope),
RMSNormModule q_norm_/k_norm_ (only when use_qk_norm, head_dim-sized gamma),
SoftmaxModule softmax_. None of their math is reimplemented here. What is implemented here is only what has no existing module: the two batched matmuls (Q@K^T, Attn@V), the 1/sqrt(head_dim) scale, and the head split/merge permutation – with their own forward, backward and LRP rule.
- Note
- Two RoPE instances and two QK-Norm instances, not one shared each. Both of those module types cache their own forward input/output for use by
backward() and propagate_relevance(). A single shared instance applied to Q and then to K would leave only K's activations in the cache, so the subsequent Q backward/relevance pass would silently redistribute through K's numbers. Separate instances are the only correct choice given those modules' caching contract (and, for QK-Norm, it also matches the usual published formulation, which gives Q and K independent gammas).
-
QK-Norm gamma is initialized to 1.0, not to
RMSNormModule's own zero default. A zero gamma would make QK-Norm annihilate Q and K entirely, so use_qk_norm=true on a freshly constructed module would produce uniform attention regardless of the input – a silently degenerate configuration rather than a neutral default. Every other parameter here keeps this codebase's zero-init convention.
-
**
use_rope=true requires an even head_dim** – RoPEModule rotates adjacent feature pairs and rejects an odd dimension. Checked here, before the sub-module is constructed, so the error message names the real cause.
◆ MultiHeadAttentionModule()
| pulsatrix::MultiHeadAttentionModule::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.
- Parameters
-
| d_model | Model/embedding width. Input and output are both (N, L, d_model). |
| num_heads | Number of attention heads; head_dim = d_model / num_heads. |
| backend | Backend to allocate/compute through. Not owned; must outlive this module. |
| use_rope | Apply RoPE to Q and K after the projections (and after QK-Norm). |
| use_qk_norm | Apply RMSNorm over each head's head_dim features of Q and K. |
- Exceptions
-
| std::invalid_argument | if d_model <= 0, num_heads <= 0, d_model % num_heads != 0, or use_rope with an odd head_dim – external boundary (constructor arguments can originate from Phase 5's Python bindings with no upstream validation), same convention as every other module. |
◆ backward()
| Tensor pulsatrix::MultiHeadAttentionModule::backward |
( |
const Tensor & |
grad_output | ) |
|
|
overridevirtual |
Gradient w.r.t. this module's input; sub-module parameter gradients accumulate inside those sub-modules (reachable through parameters()).
- Parameters
-
| grad_output | Gradient w.r.t. this module's output, shape (N, L, d_model) matching the cached forward shape. |
- Returns
- Gradient w.r.t. this module's input, shape
(N, L, d_model) – the sum of the three paths (through Q, through K, through V) back into the shared input.
- Note
- The only hand-written gradient math here is the two batched matmuls (
dA = dC @ B^T, dB = A^T @ dC, per (n, h) slice) and the head split/merge inverses (pure data movement). Everything else is a real sub-module backward() call. Verified against central finite differences over the input and every sub-module parameter in tests/multihead_attention_module_test.cpp.
- Exceptions
-
| std::logic_error | if called before any forward(). |
| std::invalid_argument | if grad_output's shape differs from the cached forward shape. |
- Note
- Device-generic: runs on Cpu, Cuda or Hip tensors (GPU-native-kernels Mission 2).
Implements pulsatrix::Module.
◆ compute_device()
| std::optional< DeviceType > pulsatrix::MultiHeadAttentionModule::compute_device |
( |
| ) |
const |
|
inlineoverridevirtual |
◆ d_model()
| int64_t pulsatrix::MultiHeadAttentionModule::d_model |
( |
| ) |
const |
|
inline |
◆ forward_impl()
| Tensor pulsatrix::MultiHeadAttentionModule::forward_impl |
( |
const Tensor & |
input | ) |
|
|
overrideprotectedvirtual |
Runs the 9-step pipeline: project -> split heads -> (QK-Norm) -> (RoPE) -> scores -> softmax -> context -> merge heads -> output projection.
- Parameters
-
| input | (N, L, d_model), any device. |
- Returns
(N, L, d_model).
- Exceptions
-
| std::invalid_argument | if input is not rank-3 with a final dimension of d_model. |
Implements pulsatrix::Module.
◆ head_dim()
| int64_t pulsatrix::MultiHeadAttentionModule::head_dim |
( |
| ) |
const |
|
inline |
◆ k_norm()
K's QK-Norm sub-module, or nullptr when use_qk_norm is false.
◆ k_proj()
| LinearModule & pulsatrix::MultiHeadAttentionModule::k_proj |
( |
| ) |
|
|
inline |
◆ last_attention_weights()
| const Tensor & pulsatrix::MultiHeadAttentionModule::last_attention_weights |
( |
| ) |
const |
|
inline |
Cached attention weights of the last forward, (N, num_heads, L, L) – the softmax output. Exposed because "what did each head attend to" is the single most-asked explainability question about this module.
◆ named_parameters()
| std::vector< NamedParamRef > pulsatrix::MultiHeadAttentionModule::named_parameters |
( |
| ) |
|
|
overridevirtual |
Every sub-module's parameters, flattened – Q/K/V/O weights and biases, plus the two QK-Norm gammas when enabled. RoPE and softmax contribute none.
Reimplemented from pulsatrix::Module.
◆ num_heads()
| int64_t pulsatrix::MultiHeadAttentionModule::num_heads |
( |
| ) |
const |
|
inline |
◆ op_type()
| OpType pulsatrix::MultiHeadAttentionModule::op_type |
( |
| ) |
const |
|
inlineoverridevirtual |
◆ out_proj()
| LinearModule & pulsatrix::MultiHeadAttentionModule::out_proj |
( |
| ) |
|
|
inline |
◆ propagate_relevance()
| Tensor pulsatrix::MultiHeadAttentionModule::propagate_relevance |
( |
const Tensor & |
relevance_out, |
|
|
const LRPRuleConfig & |
config |
|
) |
| |
|
overridevirtual |
LRP relevance propagation, composed from the sub-modules' own rules plus AttnLRP's Eq. 15 bilinear ("uniform") rule for the two batched matmuls.
For O = A @ B per (n, h) slice, Eq. 15 (Achtibat et al. 2024) splits each output's relevance evenly between the two operands – hence the factor 2 in the denominator, which this codebase's additive epsilon rule (LinearModule/RNNModule/RoPEModule) does not have: R_A[i,j] += sum_k (A[i,j]*B[j,k] / (2*O[i,k] + eps*sign(O[i,k]))) * R_O[i,k] R_B[j,k] += sum_i (A[i,j]*B[j,k] / (2*O[i,k] + eps*sign(O[i,k]))) * R_O[i,k] Applied once to context = Attn @ V and once to scores_raw = Q @ K^T (there B is K^T, so the resulting R_B is transposed back into K's layout).
- Parameters
-
| relevance_out | Relevance at this module's output, (N, L, d_model). |
| config | Supplies the epsilon stabilizer for Eq. 15 and for the projections' epsilon rule. |
- Returns
- Relevance at this module's input,
(N, L, d_model) – summed over the Q/K/V paths.
- Note
- This composed rule does not conserve relevance, and is deliberately excluded from tests/lrp_conservation_test.cpp's AllModuleTypeCases() for the same reason SoftmaxModule is: the softmax step in the middle of the pipeline is AttnLRP Eq. 13, a first-order DTD approximation with a known residual "hidden bias term" (see softmax_module.hpp). Eq. 15 is itself only conserving in the sense that the two operands' shares sum back to
R_O – and only when the pipeline feeding it conserves. The composed block's actual measured gap is reported by tests/multihead_attention_module_test.cpp's dedicated measurement test rather than asserted away or forced into the shared tolerance.
- See also
- cpp_engineering.aDNA's what/context/cpp_tdd/context_tdd_lrp_rule_pattern_taxonomy.md, "Known-Non-Conserving-by-Design Note".
- Note
- The
1/sqrt(head_dim) scale is a positive constant factor, under which the epsilon rule is exactly the identity (x*c/(c*x) == 1), so relevance at the scaled scores equals relevance at the raw product and Eq. 15 is applied against the cached raw Q @ K^T – no separate scale step, and no scale-dependent relevance.
- Exceptions
-
| std::logic_error | if called before any forward(). |
| std::invalid_argument | if relevance_out's shape differs from the cached forward shape. |
- Note
- Device-generic: runs on Cpu, Cuda or Hip tensors (GPU-native-kernels Mission 3).
Implements pulsatrix::Module.
◆ q_norm()
Q's QK-Norm sub-module, or nullptr when use_qk_norm is false.
◆ q_proj()
| LinearModule & pulsatrix::MultiHeadAttentionModule::q_proj |
( |
| ) |
|
|
inline |
◆ set_training()
| void pulsatrix::MultiHeadAttentionModule::set_training |
( |
bool |
training | ) |
|
|
overridevirtual |
◆ uses_qk_norm()
| bool pulsatrix::MultiHeadAttentionModule::uses_qk_norm |
( |
| ) |
const |
|
inline |
◆ uses_rope()
| bool pulsatrix::MultiHeadAttentionModule::uses_rope |
( |
| ) |
const |
|
inline |
◆ v_proj()
| LinearModule & pulsatrix::MultiHeadAttentionModule::v_proj |
( |
| ) |
|
|
inline |
The documentation for this class was generated from the following file: