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

Core RetNet retention block (Sun et al. 2023, arXiv:2307.08621), recurrent mode, input (N, L, d_model) -> output (N, L, d_model). More...

#include <retnet_module.hpp>

Inheritance diagram for pulsatrix::RetNetModule:
Collaboration diagram for pulsatrix::RetNetModule:

Public Member Functions

 RetNetModule (int64_t d_model, int64_t key_dim, float gamma, DeviceBackend *backend)
 Constructs a retention layer with zero-initialized parameters.
 
Tensor backward (const Tensor &grad_output) override
 Real backpropagation-through-time across the retention recurrence: accumulates all three parameter gradients across every timestep into the same buffers via Tensor::accumulate(). Threads a (key_dim, d_model) state-gradient accumulator backwards through the gamma-decayed carry, then backprops the three no-bias linear projections onto the shared grad_input slot.
 
OpType op_type () const override
 Recurrent per charter's closed OpType set – a compound accumulate-over-time operation, shared with RNNModule/LSTMModule/GRUModule/MambaModule/RWKVModule (no enum change).
 
void set_W_Q (std::initializer_list< float > values)
 Overwrites the query projection weight (d_model, key_dim) – test/initialization use only.
 
void set_W_K (std::initializer_list< float > values)
 Overwrites the key projection weight (d_model, key_dim) – test/initialization use only.
 
void set_W_V (std::initializer_list< float > values)
 Overwrites the value projection weight (d_model, d_model) – test/initialization use only.
 
void set_W_Q (const std::vector< float > &values)
 std::vector overload of set_W_Q() – for callers building values programmatically.
 
void set_W_K (const std::vector< float > &values)
 std::vector overload of set_W_K().
 
void set_W_V (const std::vector< float > &values)
 std::vector overload of set_W_V().
 
const Tensor & W_Q () const
 The query projection weight (d_model, key_dim).
 
const Tensor & W_K () const
 The key projection weight (d_model, key_dim).
 
const Tensor & W_V () const
 The value projection weight (d_model, d_model).
 
const Tensor & W_Q_grad () const
 Accumulated gradient w.r.t. W_Q.
 
const Tensor & W_K_grad () const
 Accumulated gradient w.r.t. W_K.
 
const Tensor & W_V_grad () const
 Accumulated gradient w.r.t. W_V.
 
float gamma () const
 The fixed retention decay. A hyperparameter, not a parameter – it has no gradient and is absent from parameters().
 
Tensor propagate_relevance (const Tensor &relevance_out, const LRPRuleConfig &config) override
 The original derived LRP rule (see the class-level note): unrolls the retention recurrence into Y = G @ V with G[t,s] = gamma^(t-s)*(Q_t.K_s) for s <= t (0 above the diagonal – no relevance ever reaches a future key), applies AttnLRP Eq. 15's bilinear split twice (once for Y = G @ V, once for QK = Q @ K^T) with an exact constant-scale pass-through for gamma^(t-s) composed in between, then the three no-bias projections' standard weighted-connection epsilon/z-rule.
 
std::vector< NamedParamRef > named_parameters () override
 This module's trainable parameters, each with its hierarchical name – the one place a module declares its parameters (roadmap FND-1).
 
std::optional< DeviceType > compute_device () const override
 Where this layer computes, so forward() rejects an input on another device (FND-8).
 
- 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).
 
virtual void set_training (bool training)
 Sets this module's training/eval mode. Defaults to training (matches every mainstream framework's Module default).
 
bool is_training () const
 Whether this module is currently in training mode.
 

Protected Member Functions

Tensor forward_impl (const Tensor &input) override
 The actual forward computation – per-timestep tied-weight retention recurrence.
 

Detailed Description

Core RetNet retention block (Sun et al. 2023, arXiv:2307.08621), recurrent mode, input (N, L, d_model) -> output (N, L, d_model).

Per timestep t (1-based below), batch row b, key index i in [0, key_dim), value index j in [0, d_model), with S_0[b,i,j] = 0 (the zero-initialized retention state) and a fixed scalar decay gamma: Q_t[b,i] = sum_e x_t[b,e]*W_Q[e,i] (Linear, NO bias) K_t[b,i] = sum_e x_t[b,e]*W_K[e,i] (Linear, NO bias) V_t[b,j] = sum_e x_t[b,e]*W_V[e,j] (Linear, NO bias) S_t[b,i,j] = gamma*S_{t-1}[b,i,j] + K_t[b,i]*V_t[b,j] o_t[b,j] = sum_i( Q_t[b,i]*S_t[b,i,j] )

Note
Scope cut – no rotary position encoding. The paper applies an xpos/RoPE-style rotation Theta_t to Q_t/K_t. RoPE already exists here as its own composable module (RoPEModule, Phase 3 Mission 2), exactly the way MultiHeadAttentionModule optionally composes a separate RoPEModule instance rather than baking rotation into the attention math. This module does not wire one in either – that composition is a future mission's job, once this module's own core recurrence is proven (mirroring how MultiHeadAttentionModule itself came after RoPEModule existed standalone). Q_t/K_t here are plain linear projections, unrotated.
Scope cut – single head only. The paper's own per-head formula is complete and citable standalone; multi-head is an orthogonal expressivity choice (like MultiHeadAttentionModule's head-splitting), not a structural requirement. Consistent with RNNModule/LSTMModule/GRUModule/MambaModule/RWKVModule all being single-layer, single-mechanism cores.
Scope cut – no separate output projection. V's own dimension is d_model, so o_t is already this module's output shape; the same "no unnecessary extra projection" simplification MambaModule made by not needing a W_O. (RWKVModule does need one, because its r_t*wkv_t intermediate isn't already the output shape.)
gamma is a fixed constructor hyperparameter, a plain float, NOT a learned Tensor – that is what the paper specifies (unlike Mamba's A, which is learned). No gradient is accumulated for it. 0 < gamma < 1 is expected but deliberately NOT enforced: the same "no stability constraint enforced" disposition MambaModule takes for A.
S_0 = 0, zero-initialized and not learnable – the same scope cut every prior recurrent module makes for its own initial state.
There is no nonlinearity anywhere in this recurrence (it is purely linear in x through the projections and bilinear in Q/K/V through the state), so unlike MambaModule/RWKVModule there is no elementwise nonlinearity to fuse.
Device-generic (GPU-native-kernels Mission 6): forward, backward and propagate_relevance run entirely through DeviceBackend – the projections through gemm/gemm_ex, the retention-state recurrence and its BPTT through DeviceBackend::ssm_pass (one lane per (batch, channel), sequential over time), and the LRP chain through ssm_pass plus lrp_bilinear_matmul – so the module runs on CPU, CUDA and HIP with no host round-trip.
LRP rule – original derivation (2026-09-27, operator-directed follow-on to campaign_exai_dl_library_phase6_modern_architectures's Decision Point 2, which found no published citable rule for RetNet – a literature-search result, not a mathematical impossibility). Unrolling the recurrence: with S_0 = 0, S_t[b,i,j] = sum_{s=0}^{t} gamma^(t-s) * K_s[b,i]*V_s[b,j], so Y_t[b,j] = sum_i Q_t[b,i]*S_t[b,i,j] = sum_{s=0}^{t} gamma^(t-s) * (Q_t[b,:].K_s[b,:]) * V_s[b,j] – a **causal, gamma-decay-gated, attention-shaped weighted sum Y = G @ V with G[t,s] = gamma^(t-s)*(Q_t.K_s) for s <= t (else 0). This is structurally identical to MultiHeadAttentionModule's own context = Attn @ V / scores = Q @ K^T shape, so propagate_relevance() reuses that module's AttnLRP Eq. 15 bilinear-split rule (duplicated locally, same per-module-owns-its-helpers convention as MambaModule's local stabilize()) twice – once for Y = G @ V, once for QK = Q @ K^T – composed with one exact constant-scale identity pass-through for the gamma^(t-s) factor (which, unlike Mamba's Abar/Bbar or RWKV's decay/kk, is a genuine fixed constructor hyperparameter here, not detached-as-if-constant – so this step introduces zero approximation, not even MambaLRP's kind). Verified by hand on a worked d_model=2, key_dim=2, L=3 example before implementation (see the mission's Completion Summary and retnet_module_test.cpp's conservation test); the three composed steps (two Eq. 15 splits, each conserving exactly via its factor-2 denominator, plus the exact pass-through, plus the three no-bias linear projections' own standard z-rule) conserve near-exactly, gated only by the usual epsilon stabilizers – so RetNetModule is in lrp_conservation_test.cpp's AllModuleTypeCases(), unlike RWKVModule/SoftmaxModule/MultiHeadAttentionModule. See the campaign's Decision Point 2 addendum for the full outcome.

Constructor & Destructor Documentation

◆ RetNetModule()

pulsatrix::RetNetModule::RetNetModule ( int64_t  d_model,
int64_t  key_dim,
float  gamma,
DeviceBackend *  backend 
)

Constructs a retention layer with zero-initialized parameters.

Parameters
d_modelModel dimension (also the value dimension and the output dimension).
key_dimQuery/key projection dimension.
gammaFixed retention decay applied to the carried state each timestep. Not learned, and not constrained – see the class-level note.
backendBackend to allocate/compute through. Not owned; must outlive this module.
Exceptions
std::invalid_argumentif d_model <= 0 or key_dim <= 0 – external boundary (construction arguments can originate from the Python bindings with no upstream validation).

Member Function Documentation

◆ backward()

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

Real backpropagation-through-time across the retention recurrence: accumulates all three parameter gradients across every timestep into the same buffers via Tensor::accumulate(). Threads a (key_dim, d_model) state-gradient accumulator backwards through the gamma-decayed carry, then backprops the three no-bias linear projections onto the shared grad_input slot.

Parameters
grad_outputGradient w.r.t. this module's output. Must be (N, L, d_model) matching the most recent forward() call's output shape.
Returns
Gradient w.r.t. this module's input, shape (N, L, d_model).
Exceptions
std::logic_errorif forward() has never been called.
std::invalid_argumentif grad_output's shape doesn't match the cached forward output shape.
Note
Device-generic (GPU-native-kernels Mission 6) – see the class-level note.

Implements pulsatrix::Module.

◆ compute_device()

std::optional< DeviceType > pulsatrix::RetNetModule::compute_device ( ) const
inlineoverridevirtual

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

Reimplemented from pulsatrix::Module.

◆ forward_impl()

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

The actual forward computation – per-timestep tied-weight retention recurrence.

Exceptions
std::invalid_argumentif input isn't rank-3 (N, L, d_model), or its last dimension doesn't match d_model.
Note
Device-generic (GPU-native-kernels Mission 6) – see the class-level note.

Implements pulsatrix::Module.

◆ gamma()

float pulsatrix::RetNetModule::gamma ( ) const
inline

The fixed retention decay. A hyperparameter, not a parameter – it has no gradient and is absent from parameters().

◆ named_parameters()

std::vector< NamedParamRef > pulsatrix::RetNetModule::named_parameters ( )
inlineoverridevirtual

This module's trainable parameters, each with its hierarchical name – the one place a module declares its parameters (roadmap FND-1).

Returns
{name, {value, grad}} entries pointing directly at this module's own members, in a fixed order. Names are unique within the module tree. Default: empty (a parameterless module like ReluModule needs no override).
Note
Override this, not parameters(): saving, loading, freezing by name and optimizer parameter groups all key on these names.

Reimplemented from pulsatrix::Module.

◆ op_type()

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

Recurrent per charter's closed OpType set – a compound accumulate-over-time operation, shared with RNNModule/LSTMModule/GRUModule/MambaModule/RWKVModule (no enum change).

Implements pulsatrix::Module.

◆ propagate_relevance()

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

The original derived LRP rule (see the class-level note): unrolls the retention recurrence into Y = G @ V with G[t,s] = gamma^(t-s)*(Q_t.K_s) for s <= t (0 above the diagonal – no relevance ever reaches a future key), applies AttnLRP Eq. 15's bilinear split twice (once for Y = G @ V, once for QK = Q @ K^T) with an exact constant-scale pass-through for gamma^(t-s) composed in between, then the three no-bias projections' standard weighted-connection epsilon/z-rule.

Parameters
relevance_outRelevance at this module's output. Must be (N, L, d_model) matching the most recent forward() call's output shape.
configSupplies the epsilon stabilizer used throughout (both Eq. 15 calls and the three projections' z-rule denominators).
Returns
Relevance at this module's input, shape (N, L, d_model).
Exceptions
std::logic_errorif forward() has never been called.
std::invalid_argumentif relevance_out's shape doesn't match the cached forward output shape.
Note
Device-generic (GPU-native-kernels Mission 6) – see the class-level note.
Conserves near-exactly (measured in retnet_module_test.cpp), gated only by the usual epsilon stabilizers – see the class-level note for why every composed step is exact or near-exact. Unlike SoftmaxModule/MultiHeadAttentionModule/ RWKVModule, this rule is NOT a known non-conserving approximation.

Implements pulsatrix::Module.

◆ set_W_K() [1/2]

void pulsatrix::RetNetModule::set_W_K ( const std::vector< float > &  values)

std::vector overload of set_W_K().

◆ set_W_K() [2/2]

void pulsatrix::RetNetModule::set_W_K ( std::initializer_list< float >  values)

Overwrites the key projection weight (d_model, key_dim) – test/initialization use only.

◆ set_W_Q() [1/2]

void pulsatrix::RetNetModule::set_W_Q ( const std::vector< float > &  values)

std::vector overload of set_W_Q() – for callers building values programmatically.

◆ set_W_Q() [2/2]

void pulsatrix::RetNetModule::set_W_Q ( std::initializer_list< float >  values)

Overwrites the query projection weight (d_model, key_dim) – test/initialization use only.

◆ set_W_V() [1/2]

void pulsatrix::RetNetModule::set_W_V ( const std::vector< float > &  values)

std::vector overload of set_W_V().

◆ set_W_V() [2/2]

void pulsatrix::RetNetModule::set_W_V ( std::initializer_list< float >  values)

Overwrites the value projection weight (d_model, d_model) – test/initialization use only.

◆ W_K()

const Tensor & pulsatrix::RetNetModule::W_K ( ) const
inline

The key projection weight (d_model, key_dim).

◆ W_K_grad()

const Tensor & pulsatrix::RetNetModule::W_K_grad ( ) const
inline

Accumulated gradient w.r.t. W_K.

◆ W_Q()

const Tensor & pulsatrix::RetNetModule::W_Q ( ) const
inline

The query projection weight (d_model, key_dim).

◆ W_Q_grad()

const Tensor & pulsatrix::RetNetModule::W_Q_grad ( ) const
inline

Accumulated gradient w.r.t. W_Q.

◆ W_V()

const Tensor & pulsatrix::RetNetModule::W_V ( ) const
inline

The value projection weight (d_model, d_model).

◆ W_V_grad()

const Tensor & pulsatrix::RetNetModule::W_V_grad ( ) const
inline

Accumulated gradient w.r.t. W_V.


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