AI 模型与平台
Jamba:AI21 Labs 的新型混合 Transformer-Mamba 语言模型
语言模型在近年来取得了快速的进展,Transformer 架构成为自然语言处理的主流。但是,随着模型的规模增加,处理长上下文、内存效率和吞吐量的挑战变得更加明显。
AI21 Labs 推出了 Jamba,一种结合了 Transformer 和 Mamba 架构的混合语言模型。这种混合框架使得 Jamba 在处理长上下文和内存效率方面取得了显著的改进。本文详细介绍了 Jamba 的架构、性能和潜在应用。
Jamba 概述
Jamba 是由 AI21 Labs 开发的混合语言模型,结合了 Transformer 层和 Mamba 层,并集成了 Mixture-of-Experts (MoE) 模块。这种架构使得 Jamba 能够平衡内存使用、吞吐量和性能,成为自然语言处理任务的强大工具。该模型设计为适合单个 80GB GPU,提供高吞吐量和小内存占用,同时保持了最先进的性能。
Jamba 的架构
Jamba 的架构是其能力的基础。它采用了一种新型的混合设计,交替使用 Transformer 层和 Mamba 层,并集成了 MoE 模块以增强模型的容量而不显著增加计算需求。
1. Transformer 层
Transformer 架构已成为现代语言模型的标准,能够高效地处理并行处理和捕获长距离依赖关系。然而,其性能往往受到高内存和计算需求的限制,特别是在处理长上下文时。Jamba 通过集成 Mamba 层来解决这些限制。
2. Mamba 层
Mamba 是一种最近的状态空间模型 (SSM),旨在比传统的 RNN 或甚至 Transformer 更高效地处理长距离关系。Mamba 层特别适用于减少 Transformer 中的 key-value (KV) 缓存的内存占用。通过交替使用 Mamba 层和 Transformer 层,Jamba 降低了整体内存使用,同时保持了高性能,特别是在需要长上下文处理的任务中。
3. Mixture-of-Experts (MoE) 模块
Jamba 中的 MoE 模块引入了一种灵活的方式来扩展模型容量。MoE 允许模型增加可用参数的数量,而不需要在推理时成比例增加活跃参数。在 Jamba 中,MoE 应用于某些 MLP 层,路由机制选择每个令牌激活的顶级专家。这种选择性激活使得 Jamba 能够在保持高效率的同时处理复杂任务。
以下图像演示了混合注意力-Mamba 模型中的诱导头的功能,这是 Jamba 的一个关键特征。在这个例子中,注意力头负责预测诸如“正面”或“负面”等标签,用于情感分析任务。突出显示的单词说明了模型的注意力如何在少数样本中强烈关注标签令牌,特别是在预测最终标签之前的关键时刻。这种注意力机制在模型的上下文学习能力中发挥着至关重要的作用,即模型必须根据给定的上下文和少数样本推断出适当的标签。
通过将 MoE 与注意力-Mamba 混合架构集成,Jamba 在性能方面取得了显著的改进。表格中显示了这种集成的性能改进,Jamba 通过使用 MoE 增加了其容量,而不成比例地增加了计算成本。这在各种基准测试中表现得尤为明显,例如 HellaSwag、WinoGrande 和自然问题 (NQ)。具有 MoE 的模型不仅实现了更高的准确率(例如,在 WinoGrande 上达到 66.0%,而没有 MoE 时为 62.5%),而且在不同领域中展示了更好的对数概率(例如,在 C4 上达到 -0.534)。
关键架构特征
- 层组成: Jamba 的架构由组合了 Mamba 和 Transformer 层的块组成,按照特定的比例(例如,1:7,即每七个 Mamba 层对应一个 Transformer 层)。这种比例经过优化以实现最佳性能和效率。
- MoE 集成: MoE 层每隔几层应用一次,有 16 个专家可用,每个令牌激活两个顶级专家。这种配置使得 Jamba 能够有效地扩展,同时在内存使用和计算效率之间取得平衡。
- 归一化和稳定性: 为了确保训练过程中的稳定性,Jamba 在 Mamba 层中使用了 RMSNorm,这有助于缓解大型激活脉冲等问题,这些问题可能在大规模模型中出现。
Jamba 的性能和基准测试
Jamba 已在广泛的基准测试中进行了评估,展示了其在各个方面的竞争力。以下部分突出了 Jamba 在常见 NLP 基准测试和长上下文评估中的优势。
1. 常见 NLP 基准测试
Jamba 已在多个学术基准测试中进行了评估,包括:
- HellaSwag (10 次采样):Jamba 实现了 87.1% 的性能分数,超过了许多竞争模型。
- WinoGrande (5 次采样):Jamba 得分 82.5%,再次展示了其处理复杂语言推理的能力。
- ARC-Challenge (25 次采样):Jamba 展现出强大的性能,得分 64.4%,反映了其处理具有挑战性的多项选择问题的能力。
在综合基准测试中,如 MMLU (5 次采样),Jamba 实现了 67.4% 的分数,表明其在多样任务中的稳健性。
2. 长上下文评估
Jamba 的一个突出特点是其处理极长上下文的能力。该模型支持最长 256K 令牌的上下文长度,在公开可用模型中排名第一。这种能力通过针对海量数据集的基准测试进行了评估,Jamba 在不同上下文长度(最高 256K 令牌)中展示了出色的检索准确率。
3. 吞吐量和效率
Jamba 的混合架构显著提高了吞吐量,特别是在处理长序列时。

在比较不同模型的吞吐量(每秒令牌数)的测试中,Jamba 一致地超过了其同行,特别是在大批量和长上下文场景中。例如,在 128K 令牌的上下文中,Jamba 的吞吐量是 Mixtral(一个可比模型)的三倍。

使用 Jamba:Python
对于渴望尝试 Jamba 的开发人员和研究人员,AI21 Labs 已在 Hugging Face 等平台上提供了该模型,使其可用于广泛的应用。以下代码片段演示了如何加载和使用 Jamba 生成文本:
<p>from transformers import AutoModelForCausalLM, AutoTokenizer</p>
<p>model = AutoModelForCausalLM.from_pretrained("ai21labs/Jamba-v0.1")
tokenizer = AutoTokenizer.from_pretrained("ai21labs/Jamba-v0.1")</p>
<p>input_ids = tokenizer("最近的超级碗 LVIII", return_tensors='pt').to(model.device)["input_ids"]</p>
<p>outputs = model.generate(input_ids, max_new_tokens=216)</p>
print(tokenizer.batch_decode(outputs))
该简单脚本加载 Jamba 模型和分词器,根据给定的输入提示生成文本,并打印生成的输出。
微调 Jamba
Jamba 被设计为一个基础模型,这意味着它可以被微调以适应特定的任务或应用。微调允许用户将模型适应特定领域,提高在专用任务上的性能。以下示例展示了如何使用 PEFT 库微调 Jamba:
import torch
from datasets import load_dataset
from trl import SFTTrainer, SFTConfig
from peft import LoraConfig
from transformers import AutoTokenizer, AutoModelForCausalLM, TrainingArguments
<p>tokenizer = AutoTokenizer.from_pretrained("ai21labs/Jamba-v0.1")
model = AutoModelForCausalLM.from_pretrained(
"ai21labs/Jamba-v0.1", device_map='auto', torch_dtype=torch.bfloat16)</p>
<p>lora_config = LoraConfig(r=8,
target_modules=[
"embed_tokens","x_proj", "in_proj", "out_proj", # mamba
"gate_proj", "up_proj", "down_proj", # mlp
"q_proj", "k_proj", "v_proj"
# attention],
task_type="CAUSAL_LM", bias="none")</p>
<p>dataset = load_dataset("Abirate/english_quotes", split="train")
training_args = SFTConfig(output_dir="./results",
num_train_epochs=2,
per_device_train_batch_size=4,
logging_dir='./logs',
logging_steps=10, learning_rate=1e-5, dataset_text_field="quote")
trainer = SFTTrainer(model=model, tokenizer=tokenizer, args=training_args,
peft_config=lora_config, train_dataset=dataset,
)
trainer.train()
该代码片段微调 Jamba 在一组英文引语的数据集上,调整模型的参数以更好地适应特定任务的文本生成。
部署和集成
AI21 Labs 已通过各种平台和部署选项使 Jamba 家族广泛可用:
- 云平台:
- 在主要云提供商的 Google Cloud Vertex AI、Microsoft Azure 和 NVIDIA NIM (NVDA ) 上可用。
- 即将在 Amazon Bedrock、Databricks Marketplace 和 Snowflake Cortex 上推出。
- AI 开发框架:
- 与流行框架如 LangChain 和 LlamaIndex(即将推出)集成。
- AI21 Studio:
- 通过 AI21 自有的开发平台直接访问。
- Hugging Face:
- 模型可供下载和实验。
- 本地部署:
- 针对具有特定安全或合规需求的组织的私有部署选项。
- 定制解决方案:
- AI21 为企业客户提供定制模型和微调服务。
开发者友好功能
Jamba 模型具有多个内置功能,使其对开发人员特别有吸引力:
- 函数调用:轻松将外部工具和 API 集成到您的 AI 工作流中。
- 结构化 JSON 输出:直接从自然语言输入生成干净、可解析的数据结构。
- 文档对象消化:高效地处理和理解复杂的文档结构。
- RAG 优化:内置功能以增强检索增强生成管道。
这些功能,加上模型的长上下文窗口和高效处理能力,使 Jamba 成为广泛开发场景中的多功能工具。
伦理考虑和负责任的 AI
虽然 Jamba 的能力令人印象深刻,但以负责任的 AI 心态来对待其使用至关重要。AI21 Labs 强调了几个重要点:
- 基础模型性质:Jamba 1.5 模型是预训练的基础模型,没有特定的对齐或指令微调。
- 缺乏内置安全措施:模型没有内置的审查机制。
- 谨慎部署:在生产环境或面向最终用户使用 Jamba 之前,应实施额外的适应和安全措施。
- 数据隐私:在使用基于云的部署时,应注意数据处理和合规要求。
- 偏见意识:像所有大型语言模型一样,Jamba 可能会反映其训练数据中的偏见。用户应意识到这一点,并实施适当的缓解措施。
通过考虑这些因素,开发人员和组织可以负责任、合乎道德地利用 Jamba 的能力。
AI 开发的新篇章?
AI21 Labs 推出的 Jamba 家族标志着大型语言模型演进的重要里程碑。通过结合 Transformer 和状态空间模型的优势,集成 Mixture-of-Experts 技术,并突破上下文长度和处理速度的界限,Jamba 为 AI 应用开辟了新的可能性。随着 AI 社区继续探索和在此创新架构基础上进行建设,我们可以期待在模型效率、长上下文理解和实际 AI 部署方面取得进一步的进展。Jamba 家族不仅代表了一组新模型,也代表了对大型 AI 系统设计和实施的潜在转变。
随着 AI 社区继续探索和在此创新架构基础上进行建设,我们可以期待在模型效率、长上下文理解和实际 AI 部署方面取得进一步的进展。Jamba 家族代表了对大型 AI 系统设计和实施的潜在转变,这可能会开启 AI 开发的新篇章。














