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

Core RWKV-4 time-mixing block (Peng et al. 2023, arXiv:2305.13048), input (N, L, d_model) -> output (N, L, d_model). More...

#include <rwkv_module.hpp>

Inheritance diagram for pulsatrix::RWKVModule:
Collaboration diagram for pulsatrix::RWKVModule:

Public Member Functions

 RWKVModule (int64_t d_model, DeviceBackend *backend)
 Constructs an RWKV time-mixing layer with zero-initialized parameters.
 
Tensor backward (const Tensor &grad_output) override
 Real backpropagation-through-time across the WKV recurrence: accumulates all nine parameter gradients across every timestep into the same buffers via Tensor::accumulate(). Differentiates through the sigmoid receptance gate, through both exp() nonlinearities (e_t and kk_t), through the num/den quotient and through the decayed a/b state carry, plus the three token-shift mixes.
 
OpType op_type () const override
 Recurrent per charter's closed OpType set – a compound accumulate-over-time operation, shared with RNNModule/LSTMModule/GRUModule/MambaModule (no enum change).
 
void set_W_r (std::initializer_list< float > values)
 Overwrites the receptance projection weight (d_model, d_model) – test/initialization use only.
 
void set_W_k (std::initializer_list< float > values)
 Overwrites the key projection weight (d_model, d_model) – 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_o (std::initializer_list< float > values)
 Overwrites the output projection weight (d_model, d_model) – test/initialization use only.
 
void set_w (std::initializer_list< float > values)
 Overwrites the per-channel decay rate w (d_model,), decay = exp(-w) – test/initialization use only.
 
void set_u (std::initializer_list< float > values)
 Overwrites the per-channel current-token bonus u (d_model,) – test/initialization use only.
 
void set_mu_r (std::initializer_list< float > values)
 Overwrites the receptance token-shift mix ratio (d_model,) – test/initialization use only.
 
void set_mu_k (std::initializer_list< float > values)
 Overwrites the key token-shift mix ratio (d_model,) – test/initialization use only.
 
void set_mu_v (std::initializer_list< float > values)
 Overwrites the value token-shift mix ratio (d_model,) – test/initialization use only.
 
void set_W_r (const std::vector< float > &values)
 std::vector overload of set_W_r() – 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().
 
void set_W_o (const std::vector< float > &values)
 std::vector overload of set_W_o().
 
void set_w (const std::vector< float > &values)
 std::vector overload of set_w().
 
void set_u (const std::vector< float > &values)
 std::vector overload of set_u().
 
void set_mu_r (const std::vector< float > &values)
 std::vector overload of set_mu_r().
 
void set_mu_k (const std::vector< float > &values)
 std::vector overload of set_mu_k().
 
void set_mu_v (const std::vector< float > &values)
 std::vector overload of set_mu_v().
 
const Tensor & W_r () const
 The receptance projection weight (d_model, d_model).
 
const Tensor & W_k () const
 The key projection weight (d_model, d_model).
 
const Tensor & W_v () const
 The value projection weight (d_model, d_model).
 
const Tensor & W_o () const
 The output projection weight (d_model, d_model).
 
const Tensor & w () const
 The per-channel decay rate w (d_model,); the applied decay is exp(-w).
 
const Tensor & u () const
 The per-channel current-token bonus u (d_model,).
 
const Tensor & mu_r () const
 The receptance token-shift mix ratio (d_model,).
 
const Tensor & mu_k () const
 The key token-shift mix ratio (d_model,).
 
const Tensor & mu_v () const
 The value token-shift mix ratio (d_model,).
 
const Tensor & W_r_grad () const
 Accumulated gradient w.r.t. W_r.
 
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.
 
const Tensor & W_o_grad () const
 Accumulated gradient w.r.t. W_o.
 
const Tensor & w_grad () const
 Accumulated gradient w.r.t. w.
 
const Tensor & u_grad () const
 Accumulated gradient w.r.t. u.
 
const Tensor & mu_r_grad () const
 Accumulated gradient w.r.t. mu_r.
 
const Tensor & mu_k_grad () const
 Accumulated gradient w.r.t. mu_k.
 
const Tensor & mu_v_grad () const
 Accumulated gradient w.r.t. mu_v.
 
Tensor propagate_relevance (const Tensor &relevance_out, const LRPRuleConfig &config) override
 The original derived LRP rule (see the class-level note): MambaLRP's detach-the-gate technique adapted to the WKV quotient. Detaches the receptance gate r_t and the num/den weights (e_t, kk_t, decay) as constants, then applies the standard weighted-sum epsilon/z-rule to wkv_t = num_t/den_t and to the state carry a_t = decay*a_{t-1} + kk_t*v_t, threading a state- relevance carry backward across t (mirroring MambaModule's r_h_carry), then the output/value projections' own no-bias z-rule and the value token-shift's weighted-sum split.
 
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 WKV recurrence.
 

Detailed Description

Core RWKV-4 time-mixing block (Peng et al. 2023, arXiv:2305.13048), input (N, L, d_model) -> output (N, L, d_model).

Per timestep t (1-based below), batch row b, channel d in [0, d_model), with x_0 = 0 (the zero-initialized "previous token" for the token-shift at t = 1) and a_0 = b_0 = 0 (the zero-initialized running WKV numerator/denominator state): xr_t[b,d] = mu_r[d]*x_t[b,d] + (1-mu_r[d])*x_{t-1}[b,d] (and xk_t/xv_t alike) r_t[b,d] = sigmoid( sum_e xr_t[b,e]*W_r[e,d] ) (Linear, NO bias) k_t[b,d] = sum_e xk_t[b,e]*W_k[e,d] (Linear, NO bias) v_t[b,d] = sum_e xv_t[b,e]*W_v[e,d] (Linear, NO bias) e_t[b,d] = exp(u[d] + k_t[b,d]) (current-token bonus) num_t[b,d] = a_{t-1}[b,d] + e_t[b,d]*v_t[b,d] den_t[b,d] = b_{t-1}[b,d] + e_t[b,d] wkv_t[b,d] = num_t[b,d] / den_t[b,d] decay[d] = exp(-w[d]); kk_t[b,d] = exp(k_t[b,d]) a_t[b,d] = decay[d]*a_{t-1}[b,d] + kk_t[b,d]*v_t[b,d] b_t[b,d] = decay[d]*b_{t-1}[b,d] + kk_t[b,d] o_t[b,d] = sum_e ( r_t[b,e]*wkv_t[b,e] ) * W_o[e,d] (Linear, NO bias)

Note
Scope cut, mirroring every prior recurrent module's own: this is the time-mixing (WKV) block only. A full RWKV layer additionally has a separate channel-mixing feedforward block, not built here – the same "core mechanism only, not the full published block" convention RNNModule/LSTMModule/GRUModule/MambaModule each follow.
Numerical scope cut: the recurrence above is RWKV's direct, unstabilized running- sum form. RWKV's reference implementation additionally carries a running maximum M_t purely to keep exp() from overflowing float on very long sequences – a pure numerical-conditioning device that does not change the mathematical function being computed. It is deliberately omitted here (safe and exact at this project's test- scale sequence lengths); it would have to be added back before production-scale long-sequence use. Same class of documented simplification as RNNModule's zero h_0 and MambaModule's Euler-approximated Bbar.
a_0 = b_0 = 0 and x_0 = 0, zero-initialized and not learnable – the same scope cut every prior recurrent module makes for its own initial state.
Device-generic (GPU-native-kernels Mission 6): forward, backward and propagate_relevance run entirely through DeviceBackend – the projections through gemm/gemm_ex, the token shift, the WKV recurrence (exp, sigmoid, the a_t/b_t carry), its BPTT and its LRP through DeviceBackend::ssm_pass (one lane per (batch, channel), sequential over time) – so the module runs on CPU, CUDA and HIP with no host round-trip.
LRP rule – original derivation (2026-09-27, operator-directed reassessment following campaign_exai_dl_library_phase6_modern_architectures's Decision Point 2 and RetNet's own resolved derivation). The mission's a-priori hypothesis was that the WKV num/den quotient would inherit SoftmaxModule's non-conserving DTD- approximation shape; fresh re-derivation (not a re-run of the same guess) found otherwise. Unrolling: num_t = a_{t-1} + e_t*v_t and den_t = b_{t-1} + e_t, with a_t = decay*a_{t-1} + kk_t*v_t (kk_t = exp(k_t), no bonus) – so wkv_t = num_t/den_t is, at every step, a two-term weighted sum of a_{t-1} and v_t with weights 1/den_t and e_t/den_t (which sum to exactly 1 by construction), and a_t is itself a two-term weighted sum of a_{t-1} and v_t with weights decay and kk_t. This is structurally MambaModule's own h_t = Abar_t*h_{t-1} + Bbar_t*x_t shape, not softmax's cross-normalizing shape – so the same MambaLRP technique applies: detach e_t, kk_t, decay (and the receptance gate r_t) as constants (they are already data-dependent gates in Mamba's Abar/Bbar/C, detached there for the same reason), and apply the standard weighted-sum epsilon/z-rule to the two surviving weighted-sum nodes (wkv_t's quotient and a_t's carry), threading a state-relevance carry backward across t exactly the way MambaModule's own r_h_carry does. Consequently w_r_, mu_r_, w_k_, mu_k_, u_ never appear in propagate_relevance() at all – receptance and the key path are consumed only through their cached, detached forward values (last_r_, last_e_, last_kk_), mirroring MambaModule's own identical treatment of w_delta_/bias_delta_/w_b_/w_c_. w_ (the decay rate) is the one exception: decay = exp(-w_) is recomputed from the live parameter rather than cached, so it is technically referenced – but only to reconstruct a detached forward value used as a fixed weight, the same role every other detached gate plays, not to compute w_'s own relevance share (there is no such concept for any weight tensor in this codebase's LRP rules). None of these six parameters ever receives a bilinear-split or weighted-sum relevance share of its own. Verified on a hand-worked single-channel, L=3 example before implementation (see the mission's Completion Summary): the composed rule conserves exactly there (mod the usual epsilon stabilizers), so RWKVModule – like RetNetModule, and unlike SoftmaxModule/MultiHeadAttentionModule – belongs in lrp_conservation_test.cpp's AllModuleTypeCases(). See the campaign's Decision Point 2 addendum for the full outcome, distinguishing this resolved case from RetNet's independently-resolved (and approximation-free) one.

Constructor & Destructor Documentation

◆ RWKVModule()

pulsatrix::RWKVModule::RWKVModule ( int64_t  d_model,
DeviceBackend *  backend 
)

Constructs an RWKV time-mixing layer with zero-initialized parameters.

Parameters
d_modelModel/channel dimension (also the input and output feature dimension).
backendBackend to allocate/compute through. Not owned; must outlive this module.
Exceptions
std::invalid_argumentif d_model <= 0 – external boundary (construction arguments can originate from the Python bindings with no upstream validation).

Member Function Documentation

◆ backward()

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

Real backpropagation-through-time across the WKV recurrence: accumulates all nine parameter gradients across every timestep into the same buffers via Tensor::accumulate(). Differentiates through the sigmoid receptance gate, through both exp() nonlinearities (e_t and kk_t), through the num/den quotient and through the decayed a/b state carry, plus the three token-shift mixes.

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::RWKVModule::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::RWKVModule::forward_impl ( const Tensor &  input)
overrideprotectedvirtual

The actual forward computation – per-timestep tied-weight WKV 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.

◆ mu_k()

const Tensor & pulsatrix::RWKVModule::mu_k ( ) const
inline

The key token-shift mix ratio (d_model,).

◆ mu_k_grad()

const Tensor & pulsatrix::RWKVModule::mu_k_grad ( ) const
inline

Accumulated gradient w.r.t. mu_k.

◆ mu_r()

const Tensor & pulsatrix::RWKVModule::mu_r ( ) const
inline

The receptance token-shift mix ratio (d_model,).

◆ mu_r_grad()

const Tensor & pulsatrix::RWKVModule::mu_r_grad ( ) const
inline

Accumulated gradient w.r.t. mu_r.

◆ mu_v()

const Tensor & pulsatrix::RWKVModule::mu_v ( ) const
inline

The value token-shift mix ratio (d_model,).

◆ mu_v_grad()

const Tensor & pulsatrix::RWKVModule::mu_v_grad ( ) const
inline

Accumulated gradient w.r.t. mu_v.

◆ named_parameters()

std::vector< NamedParamRef > pulsatrix::RWKVModule::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::RWKVModule::op_type ( ) const
inlineoverridevirtual

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

Implements pulsatrix::Module.

◆ propagate_relevance()

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

The original derived LRP rule (see the class-level note): MambaLRP's detach-the-gate technique adapted to the WKV quotient. Detaches the receptance gate r_t and the num/den weights (e_t, kk_t, decay) as constants, then applies the standard weighted-sum epsilon/z-rule to wkv_t = num_t/den_t and to the state carry a_t = decay*a_{t-1} + kk_t*v_t, threading a state- relevance carry backward across t (mirroring MambaModule's r_h_carry), then the output/value projections' own no-bias z-rule and the value token-shift's weighted-sum split.

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.
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
w_r_, mu_r_, w_k_, mu_k_, u_ never appear below at all; w_ appears only to reconstruct the detached decay value, not to compute its own relevance share – see the class-level note. Only w_v_, w_o_, mu_v_ (and w_ for decay) participate.
Device-generic (GPU-native-kernels Mission 6) – see the class-level note.
Conserves near-exactly (measured in rwkv_module_test.cpp), gated only by the usual epsilon stabilizers – NOT a known non-conserving approximation like SoftmaxModule's Eq. 13.

Implements pulsatrix::Module.

◆ set_mu_k() [1/2]

void pulsatrix::RWKVModule::set_mu_k ( const std::vector< float > &  values)

std::vector overload of set_mu_k().

◆ set_mu_k() [2/2]

void pulsatrix::RWKVModule::set_mu_k ( std::initializer_list< float >  values)

Overwrites the key token-shift mix ratio (d_model,) – test/initialization use only.

◆ set_mu_r() [1/2]

void pulsatrix::RWKVModule::set_mu_r ( const std::vector< float > &  values)

std::vector overload of set_mu_r().

◆ set_mu_r() [2/2]

void pulsatrix::RWKVModule::set_mu_r ( std::initializer_list< float >  values)

Overwrites the receptance token-shift mix ratio (d_model,) – test/initialization use only.

◆ set_mu_v() [1/2]

void pulsatrix::RWKVModule::set_mu_v ( const std::vector< float > &  values)

std::vector overload of set_mu_v().

◆ set_mu_v() [2/2]

void pulsatrix::RWKVModule::set_mu_v ( std::initializer_list< float >  values)

Overwrites the value token-shift mix ratio (d_model,) – test/initialization use only.

◆ set_u() [1/2]

void pulsatrix::RWKVModule::set_u ( const std::vector< float > &  values)

std::vector overload of set_u().

◆ set_u() [2/2]

void pulsatrix::RWKVModule::set_u ( std::initializer_list< float >  values)

Overwrites the per-channel current-token bonus u (d_model,) – test/initialization use only.

◆ set_w() [1/2]

void pulsatrix::RWKVModule::set_w ( const std::vector< float > &  values)

std::vector overload of set_w().

◆ set_w() [2/2]

void pulsatrix::RWKVModule::set_w ( std::initializer_list< float >  values)

Overwrites the per-channel decay rate w (d_model,), decay = exp(-w) – test/initialization use only.

◆ set_W_k() [1/2]

void pulsatrix::RWKVModule::set_W_k ( const std::vector< float > &  values)

std::vector overload of set_W_k().

◆ set_W_k() [2/2]

void pulsatrix::RWKVModule::set_W_k ( std::initializer_list< float >  values)

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

◆ set_W_o() [1/2]

void pulsatrix::RWKVModule::set_W_o ( const std::vector< float > &  values)

std::vector overload of set_W_o().

◆ set_W_o() [2/2]

void pulsatrix::RWKVModule::set_W_o ( std::initializer_list< float >  values)

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

◆ set_W_r() [1/2]

void pulsatrix::RWKVModule::set_W_r ( const std::vector< float > &  values)

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

◆ set_W_r() [2/2]

void pulsatrix::RWKVModule::set_W_r ( std::initializer_list< float >  values)

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

◆ set_W_v() [1/2]

void pulsatrix::RWKVModule::set_W_v ( const std::vector< float > &  values)

std::vector overload of set_W_v().

◆ set_W_v() [2/2]

void pulsatrix::RWKVModule::set_W_v ( std::initializer_list< float >  values)

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

◆ u()

const Tensor & pulsatrix::RWKVModule::u ( ) const
inline

The per-channel current-token bonus u (d_model,).

◆ u_grad()

const Tensor & pulsatrix::RWKVModule::u_grad ( ) const
inline

Accumulated gradient w.r.t. u.

◆ w()

const Tensor & pulsatrix::RWKVModule::w ( ) const
inline

The per-channel decay rate w (d_model,); the applied decay is exp(-w).

◆ w_grad()

const Tensor & pulsatrix::RWKVModule::w_grad ( ) const
inline

Accumulated gradient w.r.t. w.

◆ W_k()

const Tensor & pulsatrix::RWKVModule::W_k ( ) const
inline

The key projection weight (d_model, d_model).

◆ W_k_grad()

const Tensor & pulsatrix::RWKVModule::W_k_grad ( ) const
inline

Accumulated gradient w.r.t. W_k.

◆ W_o()

const Tensor & pulsatrix::RWKVModule::W_o ( ) const
inline

The output projection weight (d_model, d_model).

◆ W_o_grad()

const Tensor & pulsatrix::RWKVModule::W_o_grad ( ) const
inline

Accumulated gradient w.r.t. W_o.

◆ W_r()

const Tensor & pulsatrix::RWKVModule::W_r ( ) const
inline

The receptance projection weight (d_model, d_model).

◆ W_r_grad()

const Tensor & pulsatrix::RWKVModule::W_r_grad ( ) const
inline

Accumulated gradient w.r.t. W_r.

◆ W_v()

const Tensor & pulsatrix::RWKVModule::W_v ( ) const
inline

The value projection weight (d_model, d_model).

◆ W_v_grad()

const Tensor & pulsatrix::RWKVModule::W_v_grad ( ) const
inline

Accumulated gradient w.r.t. W_v.


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