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

TL;DR

Bamba-9B 是由 IBM、普林斯顿大学、卡内基梅隆大学 (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 size 和序列长度),Bamba-9B 的吞吐量比 Meta Llama 3.1 8B 高出 2.5 倍,延迟降低了 2 倍。算术强度分析预测,当解码阶段成为内存受限时,可能会实现 5 倍的加速;目前的 vLLM 结果受限于缺乏分块预填充 (chunked pre-fill) 支持、Transformer 风格的内存分配以及针对 H100 未优化的 Mamba2 内核。

模型架构

Bamba-9B 总共使用 32 层:3 层全注意力 (full-attention) 层和 29 层 Mamba2 层,MLP 扩展因子为 3.5,词表大小为 128k,使用 RoPE 嵌入和 GQA (8 个 KV 头,32 个头)。与 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 系统检测到三次作业中断。

数据加载器

发布的有状态数据加载器提供检查点可恢复、自动缩放、零开销洗牌流式传输、无对等点流量的异步分布式操作、动态数据混合和即时分词,并且是 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 层的内核更新。

上下文长度扩展

将 LongRoPE 应用于全注意力层可以扩展 Bamba-9B 的上下文长度。初步的 PhoneBook 检索测试显示,在无需微调的情况下,扩展后的模型在高达 16K tokens 时表现优于基础版 Bamba-9B、Llama2-7B 和 Llama3-8B,并能达到 Llama3.1-8B 的性能。在 32K tokens 时,Llama3.1-8B 领先。

总结

Bamba-9B 是由 IBM、普林斯顿大学、CMU 和 UIUC 开发的 hybrid Mamba2 模型,基于 2.2T 开放 token 训练,在 vLLM 中比 Llama 3.1 8B 实现 2.5 倍的吞吐量和 2 倍的延迟增益,可立即在 transformers、vLLM、TRL 和 llama.cpp 中使用,并附带训练、微调和扩展预训练方案以及一个有状态的数据加载器。

未来工作

计划包括在更多数据上进行持续预训练,使用社区建议的混合数据进行 SFT,使用 Tulu-3、Orca-AgentInstruct 和 Daring-Anteater 数据集进行监督微调,在 vLLM 中启用分块预填充和适当的内存分配,添加用于更快推理的 fp8 内核,应用 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 内核。

Sources