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

y1 = x + MHA(RMSNorm(x)), y2 = y1 + SwiGLU(RMSNorm(y1)). Shape (N, L, d_model) -> (N, L, d_model). More...

#include <transformer_block.hpp>

Inheritance diagram for pulsatrix::TransformerBlock:
Collaboration diagram for pulsatrix::TransformerBlock:

Public Member Functions

 TransformerBlock (int64_t d_model, int64_t num_heads, int64_t d_ff, DeviceBackend *backend, bool use_rope=true, bool use_qk_norm=false)
 Constructs a transformer block with zero-initialized sub-module parameters.
 
Tensor backward (const Tensor &grad_output) override
 Gradient w.r.t. this module's input; sub-module parameter gradients accumulate inside norm1_/mha_/norm2_/swiglu_ (reachable through parameters()).
 
OpType op_type () const override
 Elementwise per this module's own op_type() note above.
 
Tensor propagate_relevance (const Tensor &relevance_out, const LRPRuleConfig &config) override
 LRP relevance propagation: two residual epsilon/z-rule splits composed with norm1_'s/mha_'s/norm2_'s/swiglu_'s own propagate_relevance().
 
std::vector< NamedParamRef > named_parameters () override
 norm1_'s, mha_'s, norm2_'s, and swiglu_'s parameters, flattened.
 
void set_training (bool training) override
 Cascades to every sub-module, the same way SequentialModule/MultiHeadAttentionModule do.
 
int64_t d_model () const
 
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.
RMSNormModule & norm1 ()
 
MultiHeadAttentionModule & mha ()
 
RMSNormModule & norm2 ()
 
SwiGLUModule & swiglu ()
 
- 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: norm1 -> attention -> residual add -> norm2 -> SwiGLU -> residual add.
 

Detailed Description

y1 = x + MHA(RMSNorm(x)), y2 = y1 + SwiGLU(RMSNorm(y1)). Shape (N, L, d_model) -> (N, L, d_model).

Composition over inheritance, third exercise of this principle after MultiHeadAttentionModule and SwiGLUModule: norm1_/norm2_/mha_/swiglu_ are real sub-objects driven through their own forward()/backward()/propagate_relevance(). The only hand-written operation is the residual add, using this codebase's own two-term weighted-sum epsilon/z-rule (the same shape LSTMModule's c_t = f_t*c_{t-1}+i_t*g_t and GRUModule's h_t = (1-z_t)*h_{t-1}+z_t*n_t already use, weight fixed at 1 instead of a gate value) – AttnLRP does not itself address residual connections; this is mission_transformer_block.md's own resolution, not a new rule invented ad hoc.

Note
**RMSNormModule, not LayerNormModule**, for stylistic consistency with this block's other LLaMA-family choices (SwiGLU, RoPE, optional QK-Norm) – not a technical requirement; both are AttnLRP Eq. 19 identity-pass-through and would work identically here.
op_type() reuses OpType::Elementwise – the residual add is a plain elementwise binary op with fixed weight 1, the same category RoPEModule and SwiGLUModule's gate multiply already reuse. Not Composite (this module owns real math, the residual split, disqualifying it by MultiHeadAttentionModule's own precedent) and not a new category (unlike attention's genuinely novel cross-position mixing, a residual add is not architecturally novel).
Conservation is dominated by MultiHeadAttentionModule's own known large gap (measured ~49.6% of its own output in mission_multihead_attention.md), propagated through unchanged by the two residual splits (which conserve near-exactly, same argument as SwiGLUModule's diagonal split) and by SwiGLUModule's own near-exact contribution. Measured and decomposed stage-by-stage in tests/transformer_block_test.cpp, not assumed.
See also
cpp_engineering.aDNA's what/context/cpp_tdd/context_tdd_lrp_rule_pattern_taxonomy.md, "Known-Non-Conserving-by-Design Note" – this block's gap is inherited from MultiHeadAttentionModule, not a new instance of the exception.

Constructor & Destructor Documentation

◆ TransformerBlock()

pulsatrix::TransformerBlock::TransformerBlock ( int64_t  d_model,
int64_t  num_heads,
int64_t  d_ff,
DeviceBackend *  backend,
bool  use_rope = true,
bool  use_qk_norm = false 
)

Constructs a transformer block with zero-initialized sub-module parameters.

Parameters
d_modelModel/embedding width.
num_headsNumber of attention heads; forwarded to MultiHeadAttentionModule.
d_ffSwiGLU hidden width; forwarded to SwiGLUModule.
backendBackend to allocate/compute through. Not owned; must outlive this module.
use_ropeForwarded to MultiHeadAttentionModule.
use_qk_normForwarded to MultiHeadAttentionModule.
Exceptions
std::invalid_argumentpropagated from MultiHeadAttentionModule's or SwiGLUModule's own constructors (d_model <= 0, num_heads <= 0, d_model % num_heads != 0, d_ff <= 0, odd head_dim with use_rope) – no redundant re-validation here.

Member Function Documentation

◆ backward()

Tensor pulsatrix::TransformerBlock::backward ( const Tensor &  grad_output)
overridevirtual

Gradient w.r.t. this module's input; sub-module parameter gradients accumulate inside norm1_/mha_/norm2_/swiglu_ (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, same shape.
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::TransformerBlock::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::TransformerBlock::d_model ( ) const
inline

◆ forward_impl()

Tensor pulsatrix::TransformerBlock::forward_impl ( const Tensor &  input)
overrideprotectedvirtual

Runs: norm1 -> attention -> residual add -> norm2 -> SwiGLU -> residual add.

Parameters
input(N, L, d_model), any device.
Returns
(N, L, d_model).
Exceptions
std::invalid_argumentif input's rank < 2 or final dimension != d_model.

Implements pulsatrix::Module.

◆ mha()

MultiHeadAttentionModule & pulsatrix::TransformerBlock::mha ( )
inline

◆ named_parameters()

std::vector< NamedParamRef > pulsatrix::TransformerBlock::named_parameters ( )
overridevirtual

norm1_'s, mha_'s, norm2_'s, and swiglu_'s parameters, flattened.

Reimplemented from pulsatrix::Module.

◆ norm1()

RMSNormModule & pulsatrix::TransformerBlock::norm1 ( )
inline

◆ norm2()

RMSNormModule & pulsatrix::TransformerBlock::norm2 ( )
inline

◆ op_type()

OpType pulsatrix::TransformerBlock::op_type ( ) const
inlineoverridevirtual

Elementwise per this module's own op_type() note above.

Implements pulsatrix::Module.

◆ propagate_relevance()

Tensor pulsatrix::TransformerBlock::propagate_relevance ( const Tensor &  relevance_out,
const LRPRuleConfig &  config 
)
overridevirtual

LRP relevance propagation: two residual epsilon/z-rule splits composed with norm1_'s/mha_'s/norm2_'s/swiglu_'s own propagate_relevance().

Parameters
relevance_outRelevance at this module's output, matching the cached forward shape.
configSupplies the epsilon stabilizer for the residual splits and every composed sub-module rule.
Returns
Relevance at this module's input, same shape as relevance_out.
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.

◆ set_training()

void pulsatrix::TransformerBlock::set_training ( bool  training)
overridevirtual

Cascades to every sub-module, the same way SequentialModule/MultiHeadAttentionModule do.

Reimplemented from pulsatrix::Module.

◆ swiglu()

SwiGLUModule & pulsatrix::TransformerBlock::swiglu ( )
inline

The documentation for this class was generated from the following file: