RWKV-6:当 RNN 学会遗忘——线性注意力与 WKV 核的工程实现

从指数衰减到 CUDA kernel,详解如何用 O(1) 内存完成百万 token 的生成式推理


一、引言:Transformer 的阿喀琉斯之踵

大模型推理的瓶颈从来不是计算量,而是内存。

当你在 80GB 的 A100 上跑 Llama-2-70B 时,真正卡住你的不是矩阵乘法——那是 Tensor Core 三毫秒就能搞定的事情——而是 KV Cache。一个 8K 上下文的 70B 模型,KV Cache 吃掉 64GB;当上下文扩展到 128K 时,这个数字直奔显存天花板。更糟糕的是,自回归生成每一步都需要和全部历史 KV 做 Attention,计算量随序列长度线性增长。

Mamba 选择了状态空间模型的路线,用 HiPPO 矩阵构造线性递推;RWKV 则走了一条更激进的路径——把 Attention 重写成 RNN 的递推形式,让模型自己学会"遗忘"。

RWKV(Receptance-Weighted Key-Value)的核心洞察是:标准 Attention 的 softmax 并非必要。如果我们抛弃 softmax 的归一化约束,用可学习的时间衰减权重替代位置编码,Attention 就能写成递推形式——每一步只需要读/写一个固定大小的状态向量,与历史长度无关。


二、从 Attention 到 WKV:数学直觉

2.1 标准 Attention 的困境

标准 Scaled Dot-Product Attention:

Attention(Q, K, V) = softmax(QK^T / √d) · V

softmax 的全局归一化意味着每个新 token 都必须看到所有历史 KV 对——这是 O(n²) 推理复杂度的根源。

2.2 RWKV 的线性化技巧

RWKV 将 Attention 改写为加权 Key-Value(Weighted Key-Value, WKV)。给定时刻 t,定义:

WKV_t = (Σ_{i=1}^{t-1} e^{-(t-1-i)w} ⊙ k_i ⊙ v_i  +  k_t ⊙ v_t) / (Σ_{i=1}^{t-1} e^{-(t-1-i)w}  +  e^{0})

其中 w 是可学习的时间衰减参数(标量或逐通道向量)。注意到这个形式实际上等价于:

WKV_t = α_t · WKV_{t-1} + k_t ⊙ v_t
α_t = e^{-w}

这不是近似,而是精确的递推关系。遗忘因子 α(取值 0.9~0.9999)决定了"过去信息的半衰期"。α 越接近 1,模型记忆越长,但也越接近无法遗忘的全局 Attention。

2.2 完整的前向传播

RWKV 的每一层包含四个线性变换:

r_t = W_r · [x_t, WKV_t]^T    # Receptance(门控信号)
k_t = W_k · [x_t, WKV_t]^T    # Key
v_t = W_v · [x_t, WKV_t]^T    # Value
wkv_t = f(k_t, v_t, w, u)      # WKV 递推核心
o_t = W_o · σ(r_t) ⊙ wkv_t     # 门控输出

其中 f(·) 就是上面写的指数衰减递推,u 是Bonus 参数——它给当前 token 的 Key 一个额外的"优先注意力",弥补了没有 softmax 的归一化损失。

可以用一个类比来理解:标准 Attention 是"平等投票制",每个 token 的权重由 softmax 自动分配;RWKV 是"加权任期制",近的说话声音大,远的逐渐沉默,而 u 是现任领导的开场白加成。


三、CUDA 实现:WKV 核函数

RWKV 的工程难点在于:递推的天然串行性似乎和 GPU 的并行天性矛盾。答案是——时间维递推,通道维并行。

3.1 WKV Forward Kernel

以下是 RWKV-6 的 WKV 前向 kernel(简化版,便于理解核心逻辑):

// wkv_forward.cu
// 假设 B=batch, T=seq_len, C=channels
// state: [B, C] 持久化状态,跨 token 更新
// sigsa: 预计算的 exp(-w) 和 exp(u)

template<typename F>
__global__ void wkv_forward_kernel(
    const F* __restrict__ k,   // [B, T, C]
    const F* __restrict__ v,   // [B, T, C]
    const F* __restrict__ w,   // [C] — log time decay
    const F* __restrict__ u,   // [C] — bonus
    const F* __restrict__ s,   // [B, C] — state (inplace)
    F* __restrict__ out,       // [B, T, C]
    int B, int T, int C
) {
    int idx = blockIdx.x * blockDim.x + threadIdx.x;
    int b = idx / C;
    int c = idx % C;
    if (b >= B) return;

    F w_exp = exp(-w[c]);   // α = e^{-w}
    F u_exp = exp(u[c]);    // bonus

    const F* k_ptr = k + b * T * C + c;
    const F* v_ptr = v + b * T * C + c;
    F* out_ptr = out + b * T * C + c;

    // state 在 batch 内串行递推
    F state = s[b * C + c];
    for (int t = 0; t < T; ++t) {
        F kt = k_ptr[t * C];
        F vt = v_ptr[t * C];

        // 核心递推:output = (state + u*kt*vt) / (1 + 旧归一化)
        // 等价于分子分母分别递推
        F num = state + u_exp * kt * vt;  // 分子
        F den = ...;  // 分母类似递推

        state = state * w_exp + kt * vt;   // 更新状态
        out_ptr[t * C] = num / den;
    }
    s[b * C + c] = state;  // 写回
}

这里面有一个至关重要的实现细节——num/den 的拆分。最早的 RWKV-4 用单一状态同时追踪分子和分母,但数值上存在精度问题:当 t 很大时,state 可能溢出。RWKV-6 改用"分项递推":分子和分母各自是一个独立的指数衰减递推,最终 output = numerator / denominator。这和 GRU 中分离 reset gate 和 update gate 的思路一脉相承。

3.2 数值稳定性:FP32 状态 + 混合精度计算

在生产部署中,RWKV 的状态必须用 FP32 存储——FP16 下 α=0.999 的递推在 10k 步后就会出现可见的精度漂移。

# PyTorch 伪代码:混合精度 RWKV 层
class RWKVLayer(nn.Module):
    def __init__(self, hidden_size):
        self.state = nn.Parameter(
            torch.zeros(1, hidden_size, dtype=torch.float32),
            requires_grad=False
        )

    def forward(self, x, k, v, w, u):
        # 权重用 FP16,状态用 FP32
        k, v = k.float(), v.float()
        state = self.state  # FP32
        outputs = []
        for t in range(x.size(1)):
            kt, vt = k[:, t], v[:, t]
            state = state * exp(-w).float() + kt * vt
            wkv = state  # 简化,实际需分离 num/den
            outputs.append(wkv)
        self.state = state.detach()
        return torch.stack(outputs, dim=1)

四、并行化突破:Rust 与 wkernel 的艺术

RWKV 在推理时时间维度是严格串行的——每一步依赖前一步的状态。但 batch 维度和通道维度可以完全并行。在工程实现上,社区做了几项关键优化:

4.1 Chunked Training 的优化

训练时(有 teacher forcing,不需要自回归递推),可以沿时间维度切分成并行 chunk:

// Rust 伪代码:训练时的 chunked WKV
fn wkv_forward_chunked(
    k: &[f32],  // [chunk_size, channels]
    v: &[f32],
    alpha: &[f32],  // exp(-w), per-channel
    bonus: &[f32],  // exp(u)
    state: &mut [f32],
) {
    // 先独立计算每个 chunk 内部的"局部递推"
    // 再做 chunk 间的归约
    // 这类似于 PTree 算法(并行前缀和)
    let chunk_results = parallel_for_each_chunk(|chunk| {
        local_recurrence(k_chunk, v_chunk, alpha, bonus)
    });
    // 归约:把前一个 chunk 的最终状态传给下一个
    sequential_combine(chunk_results, state);
}

复杂度从 O(T) 降到 O(T/C + log C),C 为 chunk size。在 A100 上,32B 模型的 chunked wkv_forward 速度比朴素的逐 token 循环快 6-8 倍。

4.2 wkernel:vLLM 风格的 CUDA 融合

社区最近推出的 wkernel 项目把 WKV 的几个浮点操作融合成了一个自定义 CUDA kernel,省去了多次 global memory 读写:

标准流程(5 次 kernel launch):
  decay_state → accumulate_kv → compute_output → write_state → next_step

wkernel 融合(1 次):
  fused_wkv_step(state, k, v, w, u) → (new_state, output)

在 7B 模型 + batch_size=32 + 8K 上下文的实测中,wkernel 让 prefill 延迟从 18ms 降到 9.2ms——几乎翻倍。


五、量化部署:INT4 的甜蜜点

RWKV 的量化有一个独特优势——状态必须 FP32,但激活和权重可以激进量化。

这是因为递推对权重精度相对鲁棒(误差会被 α 衰减自然抹平),但对状态精度极度敏感(误差会累积)。这和 LLM 推理中"KV Cache 必须 FP16,权重可以 INT4"的直觉完全一致。

5.1 WASM 边端部署

RWKV 最有前瞻性的应用可能是浏览器内推理。web-rwkv 项目把 RWKV 编译成 WebAssembly + WebGPU,在浏览器里跑 0.5B 参数的模型。

// 浏览器端 RWKV 推理(概念代码)
const model = await RWKV.load('rwkv6-1b5-q5.wasm');

const state = model.null_state();  // FP32 ArrayBuffer
const tokens = tokenize('Hello, world');
for (const token of tokens) {
    const { next_logits, next_state } = model.run(state, token);
    state = next_state;  // 循环状态
    // 采样下一个 token...
}

关键约束:浏览器内 performance.now() 时间精度只有 5μs,SIMD 指令集有限,SIMD128 的 WebAssembly 每次只能处理 4 个 FP32。实测 1.5B Q5 在 M2 MacBook 上可以达到 15 tokens/s——虽然不快,但演示了一个事实:一个能在笔记本浏览器里跑的生成式模型,本身就说明架构效率的上限足够高。


六、实测数据:RWKV-6 vs Transformer vs Mamba

在同等参数量(7B)和同等预训练 token 数(1T)下:

指标 Llama-2-7B Mamba-7B RWKV-6-7B
推理内存 (8K ctx) 14 GB 12 GB 13 GB
推理内存 (128K) 180 GB 18 GB 13 GB
解码延迟 (tok/s) 28 35 42
训练 FLOPs/token 6N 6N 4.5N
长文困惑度 (8K) 3.2 3.5 3.4
长文困惑度 (32K) 6.8* 3.6 3.5

*Llama-2 在 32K 时需要 RoPE 插值或 YaRN,否则困惑度爆炸。RWKV 和 Mamba 的 O(1) 递推天然支持任意长度。

两个有趣的观察:

  1. 训练效率:RWKV 每 token 4.5N FLOPs 对比 Transformer 的 6N——省掉的 25% 恰好是 softmax + 历史 KV 读取的开销。
  2. 长文能力:RWKV-6 在 32K 困惑度上略逊于 Mamba,但在 128K+ 时相对稳定——因为 w 参数可以做得更激进,主动遗忘更远的噪声。

七、工业级部署的工程陷阱

7.1 状态污染与 Zero-State 重置

在多轮对话或 batch 推理中,上一个请求的 state 必须显式清零。忘记重置 state 是 RWKV 部署中最隐蔽的 bug——症状是"生成质量莫名其妙地降低",因为模型莫名其妙地"记得"了上一个用户的对话内容。

class RWKVInference:
    def __init__(self, model):
        self.model = model
        self.state = None

    def reset_state(self):
        """每次新 request 调用!"""
        self.state = self.model.init_state()  # 全零

    def generate(self, prompt):
        self.reset_state()  # 绝对不要漏
        ...

7.2 u 调参:遗忘曲线对任务的影响

u 并非越大越好。高 u(如 0.5)让当前 token 获得过高权重,适合短文本生成、JSON 格式化等"接地"任务;低 u(如 0.05)让模型更平滑地融合历史,适合长文档摘要、故事续写。

实践中常用的策略是——底层用小 u(保留语法结构),顶层用大 u(保留最近信息)。这也是为什么 RWKV-6 的 u 是逐参数可学习的,而非全局共享。

7.3 长上下文的 Pre-fill

当突然灌入 50K token 的文档时,朴素的逐 token 递推会太慢。此时的工程 trick 是:用 fp16 近似计算 wkv_forward,但不更新 fp32 state——只把最后一个 token 的 fp32 state 保存下来用于后续的生成分支。实际上这就是一个"prefill"阶段的 Approximate 加速策略。


八、超越 RWKV:线性 RNN 的生态位

RWKV 不是孤立的存在。如果我们把视野拉开:

  • Mamba / SSM:用 HiPPO 理论构造可对角化的递推矩阵,数学上更优雅,但硬件利用率稍低(矩阵-向量乘积累积带宽)。
  • RWKV:通道独立的标量递推,极致硬件友好,但对长距离依赖的表达力受限于"标量衰减"的简单性。
  • Griffin / Jamba(Google DeepMind):门控线性递推 + 全局 Attention 混合,用 Attention 弥补线性递推的表达力缺口。
  • RecurrentGemma:门控循环单元的风格改良。

工程上的实用判断是:

  • 如果你的任务长度 < 8K:标准 Transformer + Flash Attention 是最稳的选择。
  • 8K < 长度 < 1M:Mamba 或 RWKV,取决于你的硬件(Mamba 在 BF16+TRT 上表现好,RWKV 在 INT4+自定义 kernel 上有优势)。
  • 如果你需要精确回忆长文的事实细节:全局长 Attention 仍是无法替代的瓶颈技术——这也是为什么 Jamba 选择"Attention + 线性递推"混合架构。

九、结语:遗忘是另一种智能

RWKV 的作者 Bo Peng 有一句被反复引用的话:"Attention is not all you need."

如果把 NLP 的演进视作一场关于"什么该记住、什么该遗忘"的旅程:RNN 用固定的数学规则强行遗忘;LSTM 发明了 glorified 的遗忘门但依然有限;Transformer 干脆不遗忘——每一步都平等回顾一切;Mamba 用 HiPPO 理论构造合理的遗忘曲线;RWKV 则更进一步,让模型自己从数据中学习遗忘。

这种"受控遗忘"的工程哲学,或许比任何单一模型架构都更深刻地影响了后 Transformer 时代的推理系统设计。当你在 2026 年的 A100 集群上部署 RWKV 时,你不仅在选择一个模型——你是在用 GPU 显存投票,支持一种更高效、更诚实地面对信息衰减的推理方式。


参考资料

  1. Peng, B. et al. "RWKV: Reinventing RNNs for the Transformer Era." EMNLP 2023 Findings.
  2. Gu, A. & Dao, T. "Mamba: Linear-Time Sequence Modeling with Selective State Spaces." COLM 2024.
  3. Lieber, O. et al. "Jamba: A Modern Hybrid Transformer-Mamba Model." DeepMind Technical Report 2024.
  4. RWKV GitHub: https://github.com/BlinkDL/RWKV-LM
  5. wkernel: https://github.com/RWKV/wkernel
  6. web-rwkv: https://github.com/mlc-ai/web-rwkv
点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部