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‑based 生成器相當的連貫續寫。

聊天模型(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。


未來方向

  • Multilingual RWKV – 正在開發多語言語料庫與分詞器,以擴展模型的語言覆蓋範圍。
  • Community research – Discord 頻道承載了關於新訓練配方、基準測試與架構調整的專案。
  • Compression & acceleration – 由於 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