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)和因果语言模型头,但将注意力层替换为基于递归的公式,仍具备与自注意力相同的表达能力。
- 需要使用诸如
TokenShift和SmallInitEmb等额外技巧(在官方 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.cpp、rwkv-cpp-cuda等仓库的集成将进一步提升推理速度。
致谢
Hugging Face 团队感谢 Bo Peng、RWKV 社区以及贡献者 Johan Wind(RWKV 博客文章)、ArEnSc(最初的 Transformers PR)、Merve Noyan、Maria Khalusova 和 Pedro Cuenca 对本次集成的审阅与支持。
引用
如果在研究中使用 RWKV,请使用 RWKV‑LM 仓库中提供的 CITATION.cff 文件进行引用。