pulsatrix
Loading...
Searching...
No Matches
pulsatrix::MultiHeadAttentionModule Class Reference

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>

Inheritance diagram for pulsatrix::MultiHeadAttentionModule:
Collaboration diagram for pulsatrix::MultiHeadAttentionModule:

Public Member Functions

 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).
 
Sub-module access – weight initialization from tests/loaders, and inspection.
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.
 
- Public Member Functions inherited from pulsatrix::Module
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.
 

Protected Member Functions

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.
 

Detailed Description

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.

Constructor & Destructor Documentation

◆ 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_modelModel/embedding width. Input and output are both (N, L, d_model).
num_headsNumber of attention heads; head_dim = d_model / num_heads.
backendBackend to allocate/compute through. Not owned; must outlive this module.
use_ropeApply RoPE to Q and K after the projections (and after QK-Norm).
use_qk_normApply RMSNorm over each head's head_dim features of Q and K.
Exceptions
std::invalid_argumentif 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.

Member Function Documentation

◆ 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_outputGradient 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_errorif called before any forward().
std::invalid_argumentif 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

Where this layer computes, so forward() rejects an input on another device (FND-8).

Reimplemented from pulsatrix::Module.

◆ 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_argumentif 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()

RMSNormModule * pulsatrix::MultiHeadAttentionModule::k_norm ( )
inline

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

Attention per the charter's closed OpType set – see OpType::Attention's own note.

Implements pulsatrix::Module.

◆ 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_outRelevance at this module's output, (N, L, d_model).
configSupplies 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_errorif called before any forward().
std::invalid_argumentif 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()

RMSNormModule * pulsatrix::MultiHeadAttentionModule::q_norm ( )
inline

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

Cascades to every sub-module, the same way SequentialModule does.

Reimplemented from pulsatrix::Module.

◆ 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: