Skip to content
wangzhaode edited this page Jul 25, 2026 · 1 revision

MNN 支持 RWKV7 线性循环模型:技术报告

1. 背景

RWKV7(Receptance Weighted Key Value v7)是一种**线性循环(linear recurrence)**架构的 LLM,与基于 Softmax Attention 的 Transformer 不同,它通过一个递推状态矩阵 S 来压缩历史信息,推理复杂度为 O(T) 而非 O(T²)。本报告记录为 MNN 添加 RWKV7 模型(以 RWKV7-Goose-World2.8-0.1B-HF 为例)端到端导出与推理支持的完整过程。

按照 MNN 的模型分级体系,RWKV7 属于 Tier 6(全新架构):所有 12 层均为非标准的 RWKV7Attention,没有 Softmax Attention,也没有 RoPE 位置编码。

2. 架构分析

2.1 整体结构

组件 HF 路径 说明
Embedding model.embeddings 词嵌入
Decoder × 12 model.layers 每层含 attn + ffn
Final Norm model.norm LayerNorm
LM Head lm_head 线性投影

每个 Decoder 层(RWKV7Block):

residual = pre_norm(x)          # 仅 layer 0 有 pre_norm
h = attn_norm(residual)
attn_out = RWKV7Attention(h)
h = residual + attn_out
residual = h
h = ffn_norm(h)
ffn_out = RWKV7FeedForward(h)
out = residual + ffn_out

2.2 RWKV7Attention 核心计算

# 1. Token shift(有状态):delta_t = x_{t-1} - x_t,x_{-1} = 0
# 2. 六路混合:x?_t = x_t + delta_t * x_?   (? ∈ {r, w, k, v, a, g})
# 3. 投影 + LoRA:
#    r = r_proj(xr)
#    k = k_proj(xk)
#    v = v_proj(xv)
#    w = -0.6065 * sigmoid(w_lora(xw))        # 衰减项
#    a = sigmoid(a_lora(xa))                  # alpha
#    g = g_lora(xg)                           # gate(内部含 sigmoid)
# 4. v_first 跨层机制:
#    layer 0: v_first = v
#    layer i>0: v = lerp(v, v_first, sigmoid(v_lora(xv)))
# 5. kk = l2norm(k * k_k);k = k + k*(a-1)*k_a
# 6. 递推(有状态,S ∈ [B, H, K, V]):
#    proj = kk^T S                      # 使用衰减前的 S
#    S = exp(w)⊙S - (kk⊙a)⊗proj + k⊗v
#    o = S^T r
# 7. GroupNorm(逐 head,eps = head_dim * norm_eps)
# 8. Gate 修正:corr = sum(r*k*r_k, -1) * v;o = (o + corr) * g
# 9. out = o_proj(o)

2.3 RWKV7FeedForward

delta = token_shift(x)
h = x + delta * x_k
h = key(h)          # Linear: D -> 4D
h = relu(h)^2       # sqrelu
out = value(h)      # Linear: 4D -> D

2.4 关键难点

难点 说明
Token shift 有状态 需要保存上一 token 的 hidden state
递推矩阵 S 有状态 S ∈ [B, H, K, V],跨 decode 步保持
v_first 跨层依赖 layer 0 的 v 需流入 layer 1~11
GroupNorm 无 CPU 实现 MNN CPU 后端无 GroupNorm,需融入自定义算子
无 RoPE 不经过标准 Attention 的 rotary 路径

3. 设计方案

3.1 算子划分

将 RWKV7Attention 拆分为两类算子:

  1. 标准运算保留在图中:线性投影(r/k/v/o_proj)、LoRA、六路混合、sqrelu 等,由 MNN 现有算子(MatMul/Add/Mul/Sigmoid/Tanh/ReLU)表达。
  2. 有状态 + 无 CPU 实现的部分融合为自定义算子
    • Token shift → 复用现有 LinearAttention 算子(OpType 305),新增 attn_type="rwkv7_shift"。免费获得其成熟的状态管理(StateCache / onClone / 快照回滚 / 前缀缓存)。
    • 递推核心(l2norm + k_a 更新 + 递推 + GroupNorm + gate 修正)→ 新增专用算子 RWKV7(OpType 307),输入 r/w/k/a/v/g + 参数 k_k/k_a/gn_w/gn_b/r_k,内部维护递推状态 S。

3.2 v_first 跨层连线

v_first 是图内张量(layer 0 的 v_proj 输出),不是跨步状态。参照 gemma4 的 shared_kv_cache 模式,在 LlmModel.forward 中创建共享容器 rwkv7_vfirst,layer 0 写入 v_first,layer 1~11 读取。ONNX 导出时该依赖自然成为图内连线(已验证 layer 0 的 v_proj 输出被 23 个节点消费)。

3.3 Tokenizer

RWKV7 使用字节级 trie 最长匹配 tokenizer(rwkv_vocab_v20230424.txt)。MNN 的 C++ Tiktoken 同样基于 trie 最长匹配,因此复用 TIKTOKEN 格式导出,将词表按 base64 编码写入 tokenizer.txt,并用特殊 token 字符串覆盖对应索引(如 BOS 在索引 0)。

4. 实现清单

4.1 Schema

  • schema/default/MNN.fbs:新增 RWKV7 = 307 OpType、RWKV7Param 表(num_heads / head_k_dim / head_v_dim / group_norm_eps),并注册进 OpParameter union。

4.2 Python 导出

文件 改动
utils/model_mapper.py 注册 rwkv7 映射(config/model/decoder/linear_attention/mlp)
utils/transformers.py 新增 RWKV7AttentionRWKV7Mlp;Decoder 处理 pre_norm;create_linear_attention 工厂注册
utils/custom_op.py 新增 FusedRWKV7 自定义算子
utils/mnn_converter.py 新增 rebuild_rwkv7,将 ONNX 自定义算子转为原生 RWKV7 算子
utils/config.py 注册 RWKV7 配置字段(norm_eps、各 low_rank_dim 等)
utils/model.py v_first 跨层容器连线
utils/tokenizer.py RWKV 字节级 tokenizer 导出(TIKTOKEN 格式)
llmexport.py LoRA 嵌套 Linear 替换为 FakeLinear;tokenizer_file 按 tokenizer.json 存在与否自动选择

4.3 C++ 推理

文件 改动
source/backend/cpu/CPURWKV7.cpp/hpp(新增) 递推核心实现:l2norm + k_a 更新 + 递推 + GroupNorm + gate 修正;含状态管理、快照回滚、前缀缓存持久化
source/backend/cpu/CPULinearAttention.cpp/hpp 新增 rwkv7_shift(token shift),onResize/onExecute 增加 rwkv7_shift 分支
source/backend/cpu/CPUOPRegister.cpp 注册 CPURWKV7Creator
source/shape/ShapeAttention.cpp 新增 RWKV7SizeComputer
source/shape/ShapeRegister.cpp 注册 RWKV7SizeComputer

5. 验证

5.1 数学正确性(Python 侧)

编写纯 PyTorch 参考实现(不依赖 Triton),与 HF/fla 模型逐层对比:

  • 定位并修正了 g_lora 的激活位置(sigmoid 在两层 Linear 之间,而非末尾)。
  • 修正后逐层 max abs diff < 0.02(layer 0 完全一致),首 token 与 HF forward() 一致(300 = "A")。
  • MNN 导出模型的 Python test path 生成连贯文本,与参考实现前 11 个 token 完全一致。

注意:本环境中 HF 自带的 generate() 因 fla 的 Cache 与当前 transformers 版本不兼容而损坏(输出乱码),故以纯 PyTorch 参考实现作为对齐基准。

5.2 状态管理正确性

  • 有状态 decode(prefill 一次 + 逐 token 保持状态)与全量重跑逐 token 完全一致。
  • C++ 侧验证:prefill 时 S 正确清零,decode 时 S 从 prefill 状态正确延续。

5.3 C++ 推理

  • fp16 模型(--quant_bit 16)生成连贯且切题的文本,首 token 与参考一致。
  • 4-bit 量化对 0.1B 小模型过于激进,输出退化;建议使用 fp16

6. 踩坑记录

  1. HF generate() 损坏:fla 的 Cache 抽象与当前 transformers 不兼容,use_cache=True 直接报 Can't instantiate abstract class FLALayer。不能以 generate() 输出为基准。
  2. g_lora 激活位置g_lora = Linear -> sigmoid -> Linear(sigmoid 是内部激活),最初误放在末尾导致 layer 0 之后全部发散。
  3. LoRA 权重丢失:LoRA 是嵌套 nn.Sequentiallora.0/lora.2),unload_paramnamed_children() 不会替换它们为 FakeLinear,导致权重走 ONNX 常量路径被置零。修复后 LoRA 输出恢复正常。
  4. MNNConvert 需重编译:schema 新增 OpType 后,必须重新构建 MNNConvert,否则自定义算子会残留为 Extra 节点。
  5. 算子注册是显式调用:CPU 算子和 Shape 的注册函数需在 CPUOPRegister.cpp / ShapeRegister.cpp 中显式调用,否则运行时报 Don't support type [RWKV7]
  6. tokenizer_file 选择:RWKV7 无 tokenizer.json,导出的是 tokenizer.txt,config.json 中需相应配置,否则 C++ 加载失败。

7. 使用方法

# 1. 构建(需重新编译 MNNConvert 与 llm_demo)
mkdir -p build && cd build
cmake .. -DMNN_BUILD_LLM=ON -DMNN_LOW_MEMORY=ON -DMNN_SUPPORT_TRANSFORMER_FUSE=ON
make -j$(nproc) MNNConvert llm_demo

# 2. 导出(推荐 fp16)
cd transformers/llm/export
python llmexport.py \
    --path /path/to/RWKV7-Goose-World2.8-0.1B-HF \
    --export mnn --quant_bit 16 \
    --dst_path ./RWKV7-MNN

# 3. 推理
echo "What is a large language model?" > prompt.txt
./build/llm_demo ./RWKV7-MNN/config.json prompt.txt

8. 后续工作

  • 为 RWKV7 添加 GPU 后端(CUDA/OpenCL/Vulkan)的 RWKV7rwkv7_shift 实现。
  • 优化递推核心的多线程并行度(当前按 batch 并行,decode 时 batch=1 退化为串行)。
  • 支持更大规模的 RWKV7 模型并补充精度/性能基准。

Clone this wiki locally