Orthrus-Qwen3: 利用基於擴散的投機解碼加速 LLM 推論

自回歸 (AR) Transformer——大多數現代大型語言模型 (LLMs) 的架構——面臨的挑戰在於其序列性質。逐一生成 token 是計算成本高昂的過程,這往往成為實際應用中的瓶頸。為了縮小此差距,研究人員探索了投機解碼 (speculative decoding),即由一個較小、較快的「草稿」模型預測多個 token,再由較大的「目標」模型並行驗證它們。

Orthrus-Qwen3 透過將可訓練的擴散注意力模組直接整合到凍結的 Qwen3 主幹網路的每一層中,為這個問題引入了一種新穎的方法。與傳統的投機解碼不同,Orthrus 不需要外部的草稿模型或獨立的 KV cache,從而在不犧牲準確度的情況下實現了顯著的加速。

Orthrus 的運作原理:擴散注意力機制

Orthrus 的核心在於透過在每一層中注入可訓練的擴散注意力模組來修改標準的 AR Transformer。其關鍵創新在於基礎模型(凍結的 AR Transformer)與擴散頭 (diffusion head) 被設計為共享單一的 KV cache。

該過程分為兩個階段:

  1. 擴散傳遞 (The Diffusion Pass):擴散頭會並行投影 $K=32$ 個 token。這是一個單步去噪過程,用於預測多個潛在的下一個 token。
  2. 驗證傳遞 (The Verification Pass):接著,AR 頭會在第二次傳遞中驗證這些 token。它會接受與基礎模型原始輸出分佈一致的最長匹配前綴。

由於基礎模型的權重是凍結的,其輸出分佈在理論上與原始 Qwen3 模型完全相同。這確保了使用者獲得的相同品質的回答,但速度顯著提升。

性能基準測試與優勢

根據作者所述,與傳統的擴散語言模型 (diffusion LMs) 以及現有的投機解碼方法相比,Orthrus-Qwen3 在吞吐量和效率方面都提供了實質性的改進。

吞吐量與速度

  • 每次前向傳播的 token 數 (TPF):Orthrus 實現了高達 7.8 倍的 TPF,在 MATH-500 基準測試中,實際時鐘速度 (wall-clock speedup) 約為 6 倍。

與擴散語言模型的比較

傳統的擴散語言模型(如 Dream, Fast-dLLM-v2, 和 Mercury)通常會修改模型的基礎權重以實現並行生成。這往往會導致準確度的損失。例如,Fast-dLLM-v2 在 MATH-500 上出現了 11 分的下降。相比之下,Orthrus 凍結了主幹網路,確保準確度與 Qwen3-8B 完全一致。

與投機解碼的比較

與 EAGLE-3 和 DFlash 等方法相比,Orthrus 提供了幾項架構上的優勢:

  • 無需外部草稿模型:不需要初始化或同步另一個獨立模型,這消除了首字生成時間 (TTFT) 的懲罰。
  • 記憶體效率:KV 開銷極小,維持在 $O(1)$(約為 4.5 MiB 平整值)。
  • 更高的接受率:在 MATH-500 基準測試中,Orthrus 達到了 11.7 的接受長度,而 DFlash 為 7.9,EAGLE-3 為 3.5。

訓練與實作細節

Orthrus 的實作在訓練需求方面非常高效。僅有 16% 的參數需要訓練,且該模型在 8x H200 GPU 上使用不到 1B 個 token 進行了 24 小時的訓練。

研究人員發現,KL 蒸餾 (KL distillation) 在提高接受率方面比交叉熵 (CE) 表現更好,而單步去噪 (single-step denoising) 過程(6.35 TPF) 比多步去噪 (3.53 TPF) 表現更佳。

限制與考量因素

雖然 Orthrus-Qwen3 的結果令人印象深刻,但該模型目前受限於凍結的基礎模型。這意味著它會繼承原始 Qwen3 的所有偏見、幻覺和知識缺口。此外,目前的評估僅限於 Qwen3 以及貪婪/拒絕採樣方法。

隨著社群對這項工作的潛蹤性進行討論,對於將其應用於其他模型(例如 DeepSeek-V3 或用於本地 LLM 執行的量化 GGUF 版本)的興趣也日益增加。如果成功移植,這將能顯從根本上降低高規模 AI 提供商和本地愛好者的延遲與擁塞。

Sources