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