-
Notifications
You must be signed in to change notification settings - Fork 2.4k
rwkv7
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 位置编码。
| 组件 | 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
# 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)
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
| 难点 | 说明 |
|---|---|
| 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 路径 |
将 RWKV7Attention 拆分为两类算子:
- 标准运算保留在图中:线性投影(r/k/v/o_proj)、LoRA、六路混合、sqrelu 等,由 MNN 现有算子(MatMul/Add/Mul/Sigmoid/Tanh/ReLU)表达。
-
有状态 + 无 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。
-
Token shift → 复用现有
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 个节点消费)。
RWKV7 使用字节级 trie 最长匹配 tokenizer(rwkv_vocab_v20230424.txt)。MNN 的 C++ Tiktoken 同样基于 trie 最长匹配,因此复用 TIKTOKEN 格式导出,将词表按 base64 编码写入 tokenizer.txt,并用特殊 token 字符串覆盖对应索引(如 BOS 在索引 0)。
-
schema/default/MNN.fbs:新增RWKV7 = 307OpType、RWKV7Param表(num_heads / head_k_dim / head_v_dim / group_norm_eps),并注册进 OpParameter union。
| 文件 | 改动 |
|---|---|
utils/model_mapper.py |
注册 rwkv7 映射(config/model/decoder/linear_attention/mlp) |
utils/transformers.py |
新增 RWKV7Attention、RWKV7Mlp;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 存在与否自动选择 |
| 文件 | 改动 |
|---|---|
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
|
编写纯 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 参考实现作为对齐基准。
- 有状态 decode(prefill 一次 + 逐 token 保持状态)与全量重跑逐 token 完全一致。
- C++ 侧验证:prefill 时 S 正确清零,decode 时 S 从 prefill 状态正确延续。
- fp16 模型(
--quant_bit 16)生成连贯且切题的文本,首 token 与参考一致。 - 4-bit 量化对 0.1B 小模型过于激进,输出退化;建议使用 fp16。
-
HF generate() 损坏:fla 的 Cache 抽象与当前 transformers 不兼容,
use_cache=True直接报Can't instantiate abstract class FLALayer。不能以generate()输出为基准。 -
g_lora 激活位置:
g_lora = Linear -> sigmoid -> Linear(sigmoid 是内部激活),最初误放在末尾导致 layer 0 之后全部发散。 -
LoRA 权重丢失:LoRA 是嵌套
nn.Sequential(lora.0/lora.2),unload_param的named_children()不会替换它们为 FakeLinear,导致权重走 ONNX 常量路径被置零。修复后 LoRA 输出恢复正常。 -
MNNConvert 需重编译:schema 新增 OpType 后,必须重新构建
MNNConvert,否则自定义算子会残留为Extra节点。 -
算子注册是显式调用:CPU 算子和 Shape 的注册函数需在
CPUOPRegister.cpp/ShapeRegister.cpp中显式调用,否则运行时报Don't support type [RWKV7]。 -
tokenizer_file 选择:RWKV7 无
tokenizer.json,导出的是tokenizer.txt,config.json 中需相应配置,否则 C++ 加载失败。
# 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- 为 RWKV7 添加 GPU 后端(CUDA/OpenCL/Vulkan)的
RWKV7与rwkv7_shift实现。 - 优化递推核心的多线程并行度(当前按 batch 并行,decode 时 batch=1 退化为串行)。
- 支持更大规模的 RWKV7 模型并补充精度/性能基准。