RWKV 架构集成到 Hugging Face Transformers

TL;DR

RWKV 是一种模仿 Transformer 注意力机制的新型基于 RNN 的架构,现已在 Hugging Face transformers 库中正式支持,为开发者提供能够以 RNN 的速度和内存效率处理超长上下文的开源语言模型。


RWKV 项目概览

RWKV 项目由 Bo Peng(GitHub: BlinkDL)领导,并由活跃的 Discord 社区维护。Stability AI 捐赠了用于训练的 GPU。项目路线图包括性能改进(例如 RWKV.cpp、量化)、可扩展性提升(数据集处理)以及研究扩展,如聊天微调和多模态微调。社区成员可加入官方 Discord 频道参与贡献。


RWKV 如何桥接 RNN 与 Transformer

RNN 的局限性与 Transformer 的优势

  • 传统 RNN 在每个时间步复用相同的权重,这会导致梯度消失问题以及长期记忆不足。LSTM 和 GRU 在一定程度上缓解了这些问题,但仍在处理超长序列时表现不佳。
  • Transformer 通过自注意力并行处理所有 token,使用 query、key、value 投影来计算注意力得分。该设计解决了长期依赖问题,并相较于传统 RNN 加快了训练速度。

RWKV 的混合设计

  • RWKV 受 Apple 的 Attention‑Free Transformer 启发,并被简化为兼容 RNN 的形式。
  • 它保留了 Transformer 风格的嵌入、层归一化(layer‑norm)和因果语言模型头,但将注意力层替换为基于递归的公式,仍具备与自注意力相同的表达能力。
  • 需要使用诸如 TokenShiftSmallInitEmb 等额外技巧(在官方 GitHub README 中有文档),才能使模型达到 GPT 级别的性能。

RWKV 架构的技术亮点

长上下文能力

  • RWKV 能够处理 8 192 token(ctx8192)的上下文窗口,推理速度和内存使用与 1 024 token 模型相同。
  • 实验性的损失曲线表明,增大上下文长度可提升各模型规模的语言模型损失,展示了有效的长期记忆能力。

训练效率

  • 与传统 RNN 不同,RWKV 可以以“线性化 GPT”方式进行训练,支持批次间并行并比传统递归模型收敛更快。
  • 当前的训练流水线已扩展至 14 B 参数,并在持续修复 RWKV‑4 系列的数值稳定性问题。

可用模型检查点

纯语言模型(RWKV‑4)

  • 模型规模从约 170 M 到 14 B 参数不等。
  • 所有模型均在 The Pile 数据集上进行预训练,并与最先进的基准进行对比,表现相当。

指令微调聊天模型(RWKV‑4 Raven)

  • Raven 系列在指令数据集(如 ALPACA、CodeAlpaca、Guanaco、GPT‑4All 和 ShareGPT)上对 RWKV‑4 进行微调。
  • 存在针对不同语言组合(仅英文、英文 + 中文 + 日文等)和不同规模(1.5 B、7 B、14 B)的变体。
  • 所有检查点均托管在 Hugging Face Hub 上的 RWKV 组织下。

使用 🤗 Transformers 与 RWKV

文本生成示例

from transformers import pipeline
model_id = "RWKV/rwkv-4-169m-pile"
pipe = pipeline("text-generation", model=model_id)
print(pipe("In a shocking finding, scientist discovered a herd of dragons...", max_new_tokens=20))

该 pipeline 返回的续写连贯,可与基于 Transformer 的生成器相媲美。

聊天模型(Raven)示例

from transformers import AutoModelForCausalLM, AutoTokenizer
model_id = "RWKV/rwkv-raven-1b5"
model = AutoModelForCausalLM.from_pretrained(model_id).to(0)
tokenizer = AutoTokenizer.from_pretrained(model_id)
prompt = "### Instruction: Tell me about ravens\n### Response:"
inputs = tokenizer(prompt, return_tensors="pt").to(0)
output = model.generate(inputs["input_ids"], max_new_tokens=100)
print(tokenizer.decode(output[0], skip_special_tokens=True))

模型遵循 Alpaca 风格的指令格式,并生成详细的回复。


将原始 RWKV 权重转换为 Hugging Face 格式

transformers 仓库中附带了一个转换脚本(convert_rwkv_checkpoint_to_hf.py)。用户将原始检查点上传至 Hub 仓库后,运行以下命令:

python convert_rwkv_checkpoint_to_hf.py \
  --repo_id RAW_HUB_REPO \
  --checkpoint_file RAW_FILE \
  --output_dir OUTPUT_DIR

添加 --push_to_hub--model_name 参数可直接将转换后的模型上传至 Hub。


未来方向

  • 多语言 RWKV – 正在开展多语言语料库和分词器的工作,以扩展模型的语言覆盖范围。
  • 社区研究 – Discord 频道承载了关于新训练方案、基准测试和架构调优的项目。
  • 压缩与加速 – 由于 RWKV 仅依赖矩阵‑向量运算,非常适合量化(4‑bit/8‑bit)、ONNX 导出以及光子加速器等实验性硬件。与 optimum 库以及 rwkv.cpprwkv-cpp-cuda 等仓库的集成将进一步提升推理速度。

致谢

Hugging Face 团队感谢 Bo Peng、RWKV 社区以及贡献者 Johan Wind(RWKV 博客文章)、ArEnSc(最初的 Transformers PR)、Merve Noyan、Maria Khalusova 和 Pedro Cuenca 对本次集成的审阅与支持。


引用

如果在研究中使用 RWKV,请使用 RWKV‑LM 仓库中提供的 CITATION.cff 文件进行引用。

Sources