Bamba-9B:推理高效的 Hybrid Mamba2 模型

TL;DR

Bamba-9B 是由 IBM、Princeton、CMU 和 UIUC 發佈的推理高效 Hybrid Mamba2 模型。在 vLLM 中,與標準 Transformer 相比,它展現了 2.5 倍的吞吐量增益和 2 倍的延遲降低,並可立即在 transformers、vLLM、TRL 和 llama.cpp 中使用。

動機

Transformer 的推理受限於隨上下文長度增長的 KV-cache 瓶頸。Hybrid Mamba2 架構保持 KV-cache 大小不變,解決了這個瓶頸。Bamba-9B 在 7B-10B 規模上驗證了 hybrid Mamba2 方法,使用完全開放的數據,並提供可重複的檢查點以鼓勵社群實驗。

在 transformers 中的使用

若要在 🤗 Transformers 函式庫中運行 Bamba-9B,請使用 AutoModelForCausalLM 和 AutoTokenizer 載入模型與分詞器,然後呼叫 generate。範例程式碼:

from transformers import AutoModelForCausalLM, AutoTokenizer
model = AutoModelForCausalLM.from_pretrained("ibm-fms/Bamba-9B")
tokenizer = AutoTokenizer.from_pretrained("ibm-fms/Bamba-9B")
txt = ["Mamba is a snake with following properties"]
inputs = tokenizer(txt, return_tensors='pt', return_token_type_ids=False)
out = model.generate(**inputs, max_new_tokens=64)
print(tokenizer.batch_decode(out, skip_special_tokens=True)[0])

評估

與 SoTA Transformer 模型的比較

Bamba-9B 在 HF OpenLLM v1 排行榜上的平均得分為 62.31,略低於 Meta Llama 3.1 8B (63.51),但在某些指標上高於 Olmo2 7B (66.17) 和 IBM Granite v3 8B (67.47)。在 v2 排行榜上,Bamba-9B 平均為 10.91,低於 Llama 3.1 8B (14.27),但在排除數學和 MMLU 任務後表現相當;屆時平均值接近 Llama 3.1 8B 的 44.68,而 Bamba-9B 為 45.53。

與使用相似 Token 預算訓練的 Transformer 模型比較

Bamba-9B (2.2T tokens) 平均得分為 62.31,優於使用 2T tokens 訓練的 Olmo1.5 7B (55.8)。即使將 2T-token 的 Bamba-9B 檢查點 (59.11) 與 Llama2 7B (53.78) 和 IBM Granite 7B (52.07) 相比,Bamba-9B 仍然更高,這表明在數據量相等的情況下具有競爭力。

與其他 Mamba/Mamba2 模型的比較

Bamba-9B 平均得分為 62.31,在某些指標上高於 NVIDIA Mamba2 Hybrid 8B* (58.78) 和 Zamba 7B (64.36),但低於 Falcon Mamba 7B (65.31)。表格顯示,hybrid Mamba2 模型可以在提供高達 5 倍理論推理效率的同時,交付具競爭力的結果。

推理效率

在 NVIDIA H100 80GB GPU 上使用 vLLM 進行測量(跨越從 1K 到 64K tokens 的 batch sizes 和 sequence lengths),Bamba-9B 的吞吐量比 Meta Llama 3.1 8B 高出達 2.5 倍,延遲低 2 倍。算術強度分析預測,當解碼階段成為記憶體受限時,潛在速度提升可達 5 倍;目前的 vLLM 結果受限於缺乏分塊預填充 (chunked pre-fill) 支援、Transformer 風格的記憶體分配,以及針對 H100 未優化的 Mamba2 核心 (kernels)。

模型架構

Bamba-9B 總共使用 32 層:3 層全注意力 (full-attention) 層和 29 層 Mamba2 層,MLP 擴展因子為 3.5,詞表大小為 128k,使用 RoPE 嵌入和 GQA (8 個 KV-heads, 32 個 heads)。與 NVIDIA hybrid Mamba2 8B 模型相比,Bamba-9B 將注意力層從 4 層減少到 3 層,並增加了 RoPE。

數據

第一階段訓練使用了 Dolma v1.7,隨後使用 FineWeb-edu 和 Cosmopedia 進行了額外的 200B tokens 訓練。所有數據都在使用 Ray 框架的內部 Red Hat OpenShift 叢集上進行了分詞。第一階段的數據混合情況已在部落格文章中視覺化。

預訓練

預訓練分階段進行:在 1.8B/100B tokens 時進行消融實驗,接著使用 Dolma 進行 3B/2T tokens 訓練,隨後進行 9B/2T tokens 訓練,最後使用 FineWeb-edu 和 Cosmopedia 進行 200B tokens 的微調階段。訓練超參數:餘弦學習率排程 (cosine LR schedule),峰值 3e-4,2000 步二次熱身 (quadratic warmup),衰減 0.033,結束學習率 1e-5,AdamW (β1=0.9, β2=0.95),權重衰減 0.1,序列長度 4096,全域 batch size 1.5M tokens,在 IBM Cloud Vela 上使用 192 張 A100 GPU 訓練約 2 個月。由於部署錯誤和硬體故障,由 Autopilot 系統檢測到三次任務中斷。

數據載入器 (Data loader)

發佈的狀態化 (stateful) 數據載入器提供檢查點續傳、自動重新縮放、零開銷洗牌串流 (shuffled streaming)、無對等點流量的非同步分散式操作、動態數據混合與即時分詞,且為 PyTorch 原生、模組化且可擴展。它已在數百個訓練任務中通過實戰測試,並與 Torch Titan 整合。

量化

使用帶有 llm-compressor 的 FMS Model Optimizer 框架,Bamba-9B 檢查點被量化為 fp8,導致精度損失微乎其微:OpenLLM v1 平均值從 62.31 降至 61.5 (-0.1),v2 平均值從 10.91 降至 10.04 (-0.9)。vLLM 中的 fp8 推理啟用功能正等待 Mamba2 層的 kernel 更新。

上下文長度擴展

將 LongRoPE 應用於全注意力層可以擴展 Bamba-9B 的上下文長度。初步的 PhoneBook 檢索測試顯示,在無需微調的情況下,擴展後的模型在長度達 16K tokens 時優於基礎版 Bamba-9B、Llama2-7B 和 Llama3-8B,並能達到 Llama3.1-8B 的性能。在 32K tokens 時,Llama3.1-8B 領先。

總結

Bamba-9B 是由 IBM、Princeton、CMU 和 UIUC 開發的 hybrid Mamba2 模型,使用 2.2T 開放數據訓練,在 vLLM 中比 Llama 3.1 8B 實現了 2.5 倍的吞吐量和 2 倍的延遲增益,可立即在 transformers、vLLM、TRL 和 llama.cpp 中使用,並附帶訓練、微調和擴展預訓練方案以及一個狀態化數據載入器。

未來工作

計劃包括在更多數據上進行持續預訓練、使用社群建議的混合數據進行 SFT、使用 Tulu-3、Orca-AgentInstruct 和 Daring-Anteater 數據集進行監督式微調、在 vLLM 中啟用分塊預填充和適當的記憶體分配、增加用於更快推理的 fp8 kernels、應用 torch.compile 和 fp8 訓練,以及將上下文長度擴展到 1M+ tokens。

貢獻者

數據收集與策劃:AllenAI (Dolma) 和 Hugging Face (FineWeb-edu, Cosmopedia)。 數據預處理:IBM 團隊成員 Tuan Hoang Trong, Syed Zawad, Jay Gala, Ryan Gordon,使用 IBM Data Prep Kit。 模型架構:Tri Dao (Princeton), Albert Gu (CMU), Linsong Chu (IBM), Davis Wertheimer (IBM), Minjia Zhang (UIUC), Mudhakar Srivatsa (IBM), Raghu Ganti (IBM)。 模型訓練:IBM 團隊成員 Linsong Chu, Divya Kumari, Davis Wertheimer, Raghu Ganti, Dakshi Agrawal。 模型微調:IBM 團隊成員 Sukriti Sharma, Anh Uong (透過 TRL)。 模型推理:IBM 與社群貢獻者 Fabian Lim, Antoni Viros i Martin, Adnan Hoque, Jamie Yang, Nelson Nimura Gonzalez, Joshua Rosenkranz, Nick Hill, Gabe Goodhart。 量化:IBM 團隊成員 Naigang Wang, Charlie Liu。 評估:由 Yotam Perlitz, Ofir Arviv, Michal Shmueli-Scheuer, Haoechen Shen, Minjia Zhang (UIUC) 領導的 IBM 評估團隊。 領導層致謝:Priya Nagpurkar, David Cox, Sriram Raghavan, Aya Soffer, Ruchir Puri, Mukesh Khare。 社群感謝:Pablo Montalvo-Leroux, Aritra Roy Gosthipaty, Vaibhav Srivastav (Hugging Face), Stas Bekman (Contextual AI), Tyler Michael Smith (Neural Magic)。 同時感謝 Meta PyTorch、AllenAI 和 Hugging Face 的開源貢獻。

附錄:算術強度

附錄推導了注意力機制與 Bamba 模型的計算與記憶體方程式,顯示 Bamba-9B 在解碼階段的記憶體優勢對於長序列 (>16K tokens) 而言,可產生高達 5 倍於 Llama 的加速。目前的 vLLM 測量值(2.5 倍吞吐量和 2 倍延遲)受限於缺乏分塊預填充、Transformer 風格的記憶體分配,以及 H100 上未優化的 Mamba2 kernels。

Sources