Jamba 混合架构


文档摘要

Jamba 混合架构 本节摘要:状态空间模型(SSM)与 Transformer 想要不同的东西——Transformer 用二次成本注意力换质量,SSM 用递归换线性时间推理与常量内存但质量滞后。AI21 的 Jamba(2024 年 3 月)与 Jamba 1.5(8 月)把两者放同一模型:每 7 层 Mamba 配 1 层 Transformer,每隔一块上 MoE,256K 上下文装单张 80GB GPU。Mamba-3(ICLR 2026)用复值状态空间与 MIMO 投影收紧 SSM 侧。

Jamba 混合架构

本节摘要:状态空间模型(SSM)与 Transformer 想要不同的东西——Transformer 用二次成本注意力换质量,SSM 用递归换线性时间推理与常量内存但质量滞后。AI21 的 Jamba(2024 年 3 月)与 Jamba 1.5(8 月)把两者放同一模型:每 7 层 Mamba 配 1 层 Transformer,每隔一块上 MoE,256K 上下文装单张 80GB GPU。Mamba-3(ICLR 2026)用复值状态空间与 MIMO 投影收紧 SSM 侧。本节端到端读两架构,解释为何混合配方存活了三年规模化而纯 SSM 与纯 Transformer 长上下文尝试没有——你会算 Jamba 在 256K 的 KV 缓存(8 倍小于纯 Transformer),理解 SSM 的常量内存递归。

学习目标

阅读完本节,你应当能够:

  1. 解释 Jamba 块里的三种原语——Transformer 层、Mamba 层、MoE——以及 1:7:偶 的交织配方。
  2. 高层陈述 SSM 的递归长什么样,以及为何它启用常量内存推理。
  3. 算 Jamba 模型在 256K 上下文的 KV 缓存足迹,对比纯 Transformer 会需的。
  4. 说出 Mamba-3 的三个创新(指数梯形离散化、复值状态更新、MIMO)及各自针对的问题。

一、问题与直觉

注意力对序列长是二次的,状态空间模型是线性的。这差异累积:256K token,Transformer 注意力图每头 650 亿项;SSM 的递归状态不论序列长都是固定大小。纯 SSM 模型(Mamba、Mamba-2)在小规模匹配 Transformer 困惑度,但在状态追踪任务上滞后,在某些上下文检索类别上失败。直觉:SSM 把历史压进固定状态,历史长了信息泄漏;注意力精确记住一切但付二次代价。

显然修法:都用。Transformer 层放精确召回要紧处,SSM 层放其他,调比例。Jamba 是首批生产级规模(52B 总、12B 激活、256K 上下文、单 80GB GPU)发布这混合配方的模型。Jamba 1.5 扩到 398B 总/94B 激活。Mamba-3(ICLR 2026)是当前最佳纯 SSM 基线,可围绕它重建混合。

一页纸的 SSM

状态空间模型通过固定大小状态 h 处理序列 x₁,...,x_N:

h_t = A·h_{t-1} + B·x_t y_t = C·h_t

每步状态经线性动态 A 演化、取输入 B·x_t、发输出 C·h_t。A、B、C 可学。关键性质:算 y_t 只需 h_{t-1}x_t,不需任何更早的 x——内存常量,推理每 token O(1)。建模质量的技巧在 A 的结构:S4(Gu 2021)用高度结构化矩阵训练时可作长卷积高效算;Mamba(Gu、Dao 2023)用数据依赖的 A、B、C(「选择性」部分);Mamba-2(2024)进一步简化;Mamba-3(2026)在特定位重加复杂度。关键:对解码 LLM,SSM 层是注意力层的即插即用替代,有固定大小每层状态而非增长的 KV 缓存。

Jamba 块

Jamba 块按两数交织层:l(注意力与 Mamba 比,Jamba 用 l=8,即每 7 Mamba 配 1 Transformer,7+1=8 层一组);e(MoE 频率,Jamba 用 e=2,每隔一层 MoE)。一块内层序列:M M M M M M M A(7 Mamba + 1 Attention),| M | M | M | M(| 标 MoE)。每 Jamba 块 8 层,4 块深(32 层)得 28 Mamba 与 4 Attention,16 层用 MoE。

为何 1:7 比

AI21 跑消融:什么注意力-Mamba 比给最佳每参数困惑度 长上下文召回?太多注意力(1:1):质量升但内存速度降;太少(1:15):内存好但上下文检索失败;甜点 1:7 或 1:8。直觉:Transformer 层处理精确召回与状态追踪,Mamba 层处理廉价的大宗处理。

位置编码

Mamba 层本身位置感知(经递归)。原 Mamba 混合里的注意力层不用 RoPE——SSM 层提供位置信息。Jamba 1.5 给注意力层加 RoPE 做更长上下文泛化,这是基于长上下文评估的经验后修。

内存预算

Jamba-1 形状(32 层:28 Mamba + 4 Attention,隐藏 4096,32 注意力头):

  • KV 缓存(仅注意力层):2×4×32×128×256K×2 = 8.4 GB 在 256K BF16。仅 4 注意力层贡献。
  • SSM 状态:每 token 前缀 28×hidden×state_size,但这是每层固定大小,不随序列长缩放。典型 Mamba 状态每特征 16,隐藏 4096:28×4096×16×2 = 3.7 MB 总。

对比纯 Transformer 32 层同隐藏全 MHA 32 头:2×32×32×128×256K×2 = 128 GB,KV 缓存 8 倍缩减。即使对多数 2024 模型用的 GQA(8)基线(2×32×8×128×256K×2 = 32 GB),Jamba 的 1:7 混合(16GB)仍 2 倍小。这就是 AI21 说的「256K 上下文装单 80GB GPU」——全 MHA 纯 Transformer 的 KV 缓存装不下,即使 GQA 基线也没空间装权重激活,Jamba 有。

Mamba-3:2026 纯 SSM 基线

Mamba-3(ICLR 2026,arXiv:2603.15569)在纯 SSM 侧引三创新:(1) 指数梯形离散化,替 Mamba-2 的欧拉法离散化为更具表达力的递归,卷积式操作作用在核递归内的状态输入上;(2) 复值状态更新,之前 Mamba 把状态矩阵从复(S4)降到实对角(Mamba)再到缩放单位(Mamba-2),Mamba-3 重加复值——等价于状态上的数据依赖旋转编码,恢复之前实值简化丢失的状态追踪能力;(3) 多输入多输出(MIMO)投影,替每特征标量投影用矩阵值投影,提升建模能力与推理时硬件利用而不增 decode 延迟。1.5B 参数下,Mamba-3 平均下游精度比 Gated DeltaNet 高 0.6 点,MIMO 变体再加 1.2 点共 1.8 点增益;同状态大小下 Mamba-3 用一半状态匹配 Mamba-2。

何时用混合

  • 长上下文、内存紧:Jamba 类混合,SSM 层控 KV 缓存,注意力层保精确召回。
  • 状态追踪任务重:需 Mamba-3 的复值状态或更多注意力层。
  • 短上下文、追求极致质量:纯 Transformer,GQA/MHA 的 KV 缓存在短上下文不痛。

二、从零实现:层混合计算器

不重写 SSM/Transformer,而是写个计算器:给定层数、l 比、e 频率,算各类型层数、KV 缓存、SSM 状态。

def jamba_layer_mix(total_layers, l=8, e=2): num_blocks = total_layers // l mamba_per_block = l - 1; attn_per_block = 1 total_mamba = num_blocks * mamba_per_block total_attn = num_blocks * attn_per_block moe_layers = total_layers // e return {"mamba": total_mamba, "attention": total_attn, "moe": moe_layers} # Jamba-1 32 层 l=8 e=2: 28 Mamba, 4 Attention, 16 MoE def jamba_kv_cache(num_attn_layers, num_heads, head_dim, seq_len, bytes_per=2): return 2 * num_attn_layers * num_heads * head_dim * seq_len * bytes_per / 1e9 # Jamba-1 256K: 4 注意力层 → 8.4 GB;对比纯 Transformer 32 层 → 128 GB def ssm_state(num_mamba_layers, hidden, state_per_feature=16, bytes_per=2): return num_mamba_layers * hidden * state_per_feature * bytes_per / 1e6 # MB,不随序列长

三、框架对比

Jamba 与 Jamba 1.5 开权重,在 vLLM、SGLang 里支持,推理时 SSM 层用常量内存递归、注意力层用 KV 缓存。对比纯 Transformer(Llama、Mistral):256K 上下文 Jamba 内存远低(16GB vs 32~128GB),但短上下文质量可能略低(28 层是 Mamba 而非注意力)。对比纯 SSM(Mamba-2、Mamba-3):Jamba 的注意力层保精确召回与状态追踪,纯 SSM 在这些任务上滞后。Mamba-3 尚未大规模用于生产混合,但是下一代 Jamba 级模型 SSM 侧的显然候选。

四、可复用产物

本节产出 outputs/prompt-hybrid-arch-selector.md——一个提示,给定目标上下文长、内存预算、任务类型(召回重 vs 大宗处理),推荐 SSM-Transformer 比、是否上 MoE、用 Mamba-2 还是 Mamba-3,算 KV 缓存与 SSM 状态,给出部署建议。

五、练习

  1. (Easy) 用层混合计算器算不同比(1:1、1:7、1:15)下 32 层模型的注意力与 Mamba 层数,标出 KV 缓存差异。

  2. (Medium) 算 Jamba-1(28 Mamba+4 Attention)vs 纯 Transformer(32 Attention)vs GQA(8)基线在 256K 的 KV 缓存,验证 Jamba 的 8 倍 / 2 倍缩减。

  3. (Medium) 实现 SSM 递归 h_t = A·h_{t-1}+B·x_t 的玩具版,验证状态大小不随序列长增长(常量内存),对比 KV 缓存的线性增长。

  4. (Hard) 消融实验:训三个小模型(纯注意力、1:7 混合、纯 SSM)在小语料上,对比困惑度与「针在草堆」召回,验证 1:7 甜点。

  5. (Hard) 实现复值状态更新(Mamba-3 核心):把 SSM 的 A 从实对角改复值,在状态追踪任务(如复制、翻转字符串)上对比实值版,验证复值恢复状态追踪能力。

本节要点回顾

  1. SSM 线性,Transformer 二次:256K 注意力图每头 650 亿项,SSM 递归状态固定大小不论序列长。
  2. SSM 常量内存推理:h_t=A·h_{t-1}+B·x_t,算 y_t 只需上一步状态与当前输入,O(1) 每 token。
  3. Jamba 1:7:偶:每 7 Mamba 配 1 Transformer,每隔一层 MoE,甜点平衡召回与内存。
  4. Transformer 处理精确召回:SSM 压历史进固定状态会泄漏,注意力层补状态追踪与精确检索。
  5. 256K 装单 80GB GPU:Jamba KV 缓存 8.4GB(仅 4 注意力层)+ SSM 状态 3.7MB,纯 Transformer 全 MHA 要 128GB。
  6. KV 缓存 8 倍缩减:对比纯 Transformer 全 MHA;即使对比 GQA(8)基线仍 2 倍小。
  7. SSM 状态不随序列长:每层固定大小,与 KV 缓存的线性增长对比是 SSM 的核心优势。
  8. Mamba-3 三创新:指数梯形离散化、复值状态(数据依赖旋转)、MIMO 投影,恢复状态追踪、提效率。
  9. Mamba-3 用半状态匹配 Mamba-2:复值状态的表达力红利。
  10. 何时用混合:长上下文内存紧用 Jamba 类,短上下文追极致质量用纯 Transformer,状态追踪重需 Mamba-3 复值或更多注意力。

下一节,异步推理与 Hogwild!:多 worker 共享一个 KV 缓存,无需微调就能涌现协作,推理并行的新维度。


发布者: 作者: Rohit Gupta 转发
评论区 (0)
U