Transformers 中的 Jamba:Transformer–Mamba 混合 MoE 架构解析与部署实战
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
本篇技术指南以 Hugging Face Transformers 官方 Jamba 模型文档为主体,结合本仓库 Jamba 源码实现、配置类与测试套件,系统讲解 Jamba 的块式(blocks-and-layers)混合架构原理、核心配置参数,以及如何通过Pipeline、AutoModelForCausalLM和 bitsandbytes 量化在多 GPU 环境下完成文本生成与部署。读完后你将掌握 Jamba 的架构设计逻辑,并能直接复现完整的推理、量化与对话式生成方案。
一、Jamba 是什么:一次 Transformer 与状态空间模型的架构融合
Jamba 是 AI21 Labs 提出的混合 Transformer–Mamba 专家混合(MoE)语言模型,官方文档将其定位为:同时继承两大模型家族优势——Transformer 的建模性能,以及 Mamba 这类状态空间模型(SSM)的高推理效率与超长上下文能力(最长 256K tokens)。其参数量覆盖约 52B 到 398B 的总规模区间。
Jamba 的架构核心是一种blocks-and-layers(块与层)的组织方式:每个 Jamba 块内部要么包含一个自注意力层、要么包含一个 Mamba 层,其后总是跟一个多层感知机(MLP)前馈块;整体上每 8 层中约有 1 层是标准 Transformer 注意力层(默认attn_layer_period=8、attn_layer_offset=4),其余为 Mamba 层。与此同时,部分前馈块被替换为 MoE(稀疏专家)块,以在不显著增加计算量的前提下扩充模型容量。
官方 Jamba 系列原始检查点由 AI21(ai21labs组织)发布,本文示例使用的ai21labs/AI21-Jamba-Mini-1.6、ai21labs/AI21-Jamba-Large-1.6、ai21labs/AI21-Jamba-1.5-Large均为该系列可加载权重。
1.1 混合架构的动机
- Transformer 优势:全局建模能力强,成熟的注意力实现(SDPA、FlashAttention-2)在长程依赖上表现稳定。
- Mamba(SSM)优势:线性复杂度的状态递归,支持超长序列与更快解码;但由于没有注意力这种“内容寻址”机制,单独使用仍有局限。
- MoE 的作用:在不放大每个 token 的稠密计算量的前提下,通过多个专家 + Top-k 路由扩充参数容量(capacity)。
仓库中JambaPreTrainedModel同时声明了_supports_flash_attn = True与_supports_sdpa = True(见 modeling_jamba.py),意味着混合架构中的注意力部分可以无缝切换 SDPA / FlashAttention-2 等实现。
二、层布局机制:period 与 offset 如何拼装混合网络
JambaConfig(configuration_jamba.py)用两组"周期 + 偏移"参数精确描述混合网络布局,这是理解 Jamba 结构的关键入口:
| 配置项 | 默认值 | 含义 |
|---|---|---|
attn_layer_period | 8 | 每隔多少层出现一次标准注意力层 |
attn_layer_offset | 4 | 第一个注意力层所在的层下标 |
expert_layer_period | 2 | 每隔多少层出现一次专家(MoE)层 |
expert_layer_offset | 1 | 第一个专家层所在的层下标 |
源码通过三个@property将布局物化为逐层序列(见 configuration_jamba.py):
@property def layers_block_type(self): return [ "attention" if i % self.attn_layer_period == self.attn_layer_offset else "mamba" for i in range(self.num_hidden_layers) ] @property def layers_num_experts(self): return [ self.num_experts if i % self.expert_layer_period == self.expert_layer_offset else 1 for i in range(self.num_hidden_layers) ]即:
layers_block_type:把num_hidden_layers(默认 32)逐层判定为"attention"或"mamba";layers_num_experts:对应层 FFN 的专家数——命中 offset/period 规则时取num_experts(默认 16),否则为 1(即普通稠密 MLP)。
模型主模块JambaModel.__init__正是依照layers_block_type逐层实例化对应 decoder 层(见 modeling_jamba.py):
decoder_layers = [] for i in range(config.num_hidden_layers): layer_class = ALL_DECODER_LAYER_TYPES[config.layers_block_type[i]] decoder_layers.append(layer_class(config, layer_idx=i)) self.layers = nn.ModuleList(decoder_layers)其中ALL_DECODER_LAYER_TYPES = {"attention": JambaAttentionDecoderLayer, "mamba": JambaMambaDecoderLayer}(见 modeling_jamba.py)。两种 decoder 层结构对称:各自执行输入 RMSNorm →(注意力 或 Mamba 混合器)→ 残差 →pre_ff_layernorm→ FFN/MoE → 残差。
2.1 架构合法性校验
由于 offset/period 是排列规则的"分母",JambaConfig 提供了严格校验(configuration_jamba.py):
attn_layer_offset必须小于attn_layer_period;expert_layer_offset必须小于expert_layer_period。
否则直接抛出ValueError。对应地,测试套件在JambaConfigTester中验证了非法偏移组合(例如attn_layer_offset=4, attn_layer_period=4)会触发校验失败(test_modeling_jamba.py),保证自定义架构时不会产生"找不到锚点层"的无效配置。
三、三种核心子层:注意力、Mamba 混合器与 MoE
3.1 标准注意力层:GQA + RoPE,统一注意力函数接口
JambaAttention(modeling_jamba.py)采用多查询分组注意力(GQA):默认num_attention_heads=32、num_key_value_heads=8,即 4 个 query 头共享 1 组 KV 头;无偏置线性投影(q/k/v/o),并使用旋转位置编码(RoPE)。前向时通过ALL_ATTENTION_FUNCTIONS.get_interface(...)选择当前attn_implementation(如sdpa、flash_attention_2)对应的注意力函数,回退实现是eager_attention_forward(modeling_jamba.py),内部完成 KV 头复制、缩放点积、softmax(float32 精度)与 dropout。这也解释了文档徽标中同时支持 FlashAttention 与 SDPA 的原因。
3.2 Mamba 混合器:选择性 SSM 的完整数据流
JambaMambaMixer(modeling_jamba.py)实现了 Mamba 论文中“选择性状态空间”(selective SSM)机制,关键点如下:
in_proj门控投影:把hidden_size投影到2 * intermediate_size(其中intermediate_size = mamba_expand * hidden_size,默认扩展因子 2),随后一分为二得到输入流与silu门控。- 深度可分离因果卷积:
conv1d以mamba_d_conv=4为核宽、按 channel 分组卷积。 x_proj输出输入相关的dt / B / C:这是 Mamba 区别于 S4 的"选择性"来源。- 状态离散化与递归/扫描:
A_log以 S4D 实数方式初始化并取负指数得到矩阵 A,D初始化为全 1(见 modeling_jamba.py 的_init_weights)。 - 与普通 Mamba 的关键差异:对
dt、B、C额外施加 RMSNorm(dt_layernorm/b_layernorm/c_layernorm),见 modeling_jamba.py。
前向中 Jamba 支持两种 SSM 路径(modeling_jamba.py):
- 单步递归(token 级解码):复用上一 token 的
recurrent_state与conv_state,走mamba_selective_state_update,这是自回归生成阶段的关键路径; - 整序列扫描:对新序列执行
mamba_selective_scan,支持三种后端——官方 mamba-ssm CUDA kernel、mamba.py 的并行前缀扫描(pscan,受use_mambapy控制)、以及 PyTorch 的torch._higher_order_ops.associative_scan(仅在torch.compile追踪且torch >= 2.9.0时启用,配置use_associative_scan=True)。
这些 kernel 回调均通过@use_kernel_func_from_hub_with_fallback(...)注入(modeling_jamba.py),其 Python 实现内还包含对 padding token 的状态遮蔽(apply_mask_to_padding_states,见 modeling_jamba.py),避免 padding 污染 SSM 状态。
3.3 MoE:Top-2 路由 + 3D 专家张量 + 负载均衡损失
JambaSparseMoeBlock(modeling_jamba.py)含一个无偏置router(hidden_size → num_experts),每个 token 经 softmax 后取 Top-k(默认num_experts_per_tok=2)的权重与索引;JambaExperts(modeling_jamba.py)将专家权重组织为 3D 参数张量(num_experts × dim × dim),按命中专家的 token 分组计算,gate_up_proj先算 silu 门控,再用down_proj回投,最后乘 Top-k 权重并index_add_聚合。相比“padding 到满容量”的朴素 MoE,这种块稀疏写法不会丢弃 token,也不浪费容量填充。- 训练时可通过
output_router_logits=True输出各层路由 logits,并由load_balancing_loss_func(modeling_jamba.py)按 Switch Transformer 的负载均衡损失(式 4–6)逐层累计、O(seq_len × num_experts) 内存峰值地归一化计算;JambaForCausalLM中按router_aux_loss_coef(默认 0.001)叠加到主损失(modeling_jamba.py)。
3.4 双掩码机制
由于混合网络中两种层形态不同,JambaModel.forward为两类层分别构造掩码:注意力层使用create_causal_mask(因果注意力掩码),Mamba 层使用create_recurrent_attention_mask(递归掩码),再按config.layer_types("full_attention"/"linear_attention")分发到每一层(见 modeling_jamba.py)。这是"Transformer + SSM 能否正确拼接"的工程关键。
四、文本生成实战:三种官方示例
本仓库文档 jamba.md 提供三种调用方式(Pipeline与AutoModel),下面是完整可运行版本。
4.1 方式一:Pipeline(最快上手)
# 若使用 CUDA 上的 Mamba,请先安装优化 kernel: # !pip install mamba-ssm causal-conv1d>=1.2.0 from transformers import pipeline pipe = pipeline( task="text-generation", model="ai21labs/AI21-Jamba-Mini-1.6", device=0, ) pipe("Plants create energy through a process known as")4.2 方式二:AutoModel+ SDPA + 静态缓存
from transformers import AutoModelForCausalLM, AutoTokenizer tokenizer = AutoTokenizer.from_pretrained("ai21labs/AI21-Jamba-Large-1.6") model = AutoModelForCausalLM.from_pretrained( "ai21labs/AI21-Jamba-Large-1.6", device_map="auto", attn_implementation="sdpa", # 或 "flash_attention_2" ) input_ids = tokenizer("Plants create energy through a process known as", return_tensors="pt").to(model.device) output = model.generate(**input_ids, cache_implementation="static") print(tokenizer.decode(output[0], skip_special_tokens=True))这里的attn_implementation="sdpa"与源码_supports_sdpa = True相印证;cache_implementation="static"则利用 Jamba 是状态型模型(_is_stateful = True)且_can_compile_fullgraph = True(modeling_jamba.py)的特性,为torch.compile全图编译场景服务。
4.3 使用对话模板(Chat Template)
Jamba 各检查点随附 chat template,可直接将多轮messages结构编码为模型输入(这也是文档量化示例中使用的方式):
messages = [ {"role": "system", "content": "You are an ancient oracle who speaks in cryptic but wise phrases, always hinting at deeper meanings."}, {"role": "user", "content": "Hello!"}, ] input_ids = tokenizer.apply_chat_template(messages, add_generation_prompt=True, return_tensors='pt').to(model.device) outputs = model.generate(input_ids, max_new_tokens=216) conversation = tokenizer.decode(outputs[0], skip_special_tokens=True) assistant_response = conversation.split(messages[-1]['content'])[1].strip() print(assistant_response)官方文档中该示例(对 8-bit 量化模型)输出类似:Seek and you shall find. The path is winding, but the journey is enlightening. What wisdom do you seek from the ancient echoes?
五、大模型量化部署:8-bit 权重量化与多 GPU 平铺
Jamba 系列参数量大,单卡显存难以容纳。文档指出:量化以更低精度表示权重,能显著降低大模型的内存负担,可用 Transformers 的多种量化后端(参见量化总览)。下面是文档给出的完整 8-bit 示例(使用 bitsandbytes):
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig quantization_config = BitsAndBytesConfig( load_in_8bit=True, llm_int8_skip_modules=["mamba"], # 关键:跳过 Mamba 模块的量化 ) # 把 72 层网络平均铺到 8 张 GPU 的 device map(节选开头几项示意) device_map = { 'model.embed_tokens': 0, 'model.layers.0': 0, ... 'model.layers.8': 0, 'model.layers.9': 1, ... 'model.layers.71': 7, 'model.final_layernorm': 7, 'lm_head': 7, } model = AutoModelForCausalLM.from_pretrained( "ai21labs/AI21-Jamba-Large-1.6", attn_implementation="flash_attention_2", quantization_config=quantization_config, device_map=device_map, ) tokenizer = AutoTokenizer.from_pretrained("ai21labs/AI21-Jamba-Large-1.6")该方案有三个要点值得展开:
llm_int8_skip_modules=["mamba"]是硬性要求:文档 Notes 第一条明确"不要量化 Mamba 块,以免模型性能退化"。因为 SSM 依赖精确的dt/B/C离散化与状态更新,8-bit 误差会被递归放大。- 量化后的特殊处理路径:源码中专门处理了"模型已被量化"的场景——
dt_proj改用带 bias 的nn.Linear调用以兼容量化层(见 modeling_jamba.py),说明 Jamba 的量化支持是经过显式适配的。 device_map手动分布:文档按model.layers.N把相邻层平均分配到 0–7 号 GPU;也可改让from_pretrained(..., device_map="auto")自动分配(AutoModel示例即如此)。
六、注意要点:Mamba kernel 的取舍
文档 Notes 给出两条重要运维经验:
- 务必使用优化的 Mamba kernel:没有 mamba-ssm 优化 kernel 的情况下运行 Mamba 会显著变慢(文档原话为 latency 显著下降)。官方 Python 回退实现仅用于无 CUDA kernel 或测试环境。
- 若确需关闭 kernel:在
from_pretrained传入use_mamba_kernels=False:
import torch from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained( "ai21labs/AI21-Jamba-1.5-Large", use_mamba_kernels=False, device_map="auto", )这一配置语义见 configuration_jamba.py:use_mamba_kernels默认为True,仅在安装了mamba-ssm与causal-conv1d、且 Mamba 模块运行在 CUDA 设备上时才可用;若置为True而环境不满足,会抛出ValueError。测试环境中的小型 Jamba tester 也因此显式设置use_mamba_kernels=False(见 test_modeling_jamba.py)以脱离 GPU kernel 依赖运行数值与梯度测试。
若模型经device_map做流水线并行(部分 Mamba 模块落在 CPU),kernel 可能不可用,此时需按上述方式显式关闭或迁移设备。
七、JambaConfig:完整参数表与自定义布局
JambaConfig的字段与默认值集中在 configuration_jamba.py,关键项汇总如下:
| 参数 | 默认值 | 说明 |
|---|---|---|
vocab_size | 65536 | 词表大小 |
hidden_size | 4096 | 隐藏维度 |
intermediate_size | 14336 | FFN 中间维度 |
num_hidden_layers | 32 | 总层数(随检查点变化,如示例 device map 对应 72 层版本) |
num_attention_heads/num_key_value_heads | 32 / 8 | 注意力头与 KV 头(GQA) |
max_position_embeddings | 262144 | 最长上下文 256K |
num_experts | 16 | 每 MoE 层专家数 |
num_experts_per_tok | 2 | 每 token 路由专家数(Top-2) |
expert_layer_period/expert_layer_offset | 2 / 1 | MoE 层布局规则 |
attn_layer_period/attn_layer_offset | 8 / 4 | 注意力层布局规则 |
mamba_d_state | 16 | SSM 状态维度 |
mamba_d_conv | 4 | 深度卷积核宽 |
mamba_expand | 2 | Mamba 隐藏扩展因子 |
mamba_dt_rank | "auto" | 离散化投影秩,"auto"时按ceil(hidden_size / 16)计算(configuration_jamba.py) |
mamba_conv_bias/mamba_proj_bias | True/False | 卷积/投影偏置开关 |
use_mamba_kernels | True | 是否使用 CUDA 优化 kernel |
use_mambapy | False | 无官方 kernel 时的回退:True用 mamba.py,False用朴素逐token实现 |
use_associative_scan | True | torch.compile+ torch≥2.9 时用 PyTorch 关联扫描替代朴素循环 |
rms_norm_eps | 1e-6 | RMSNorm 数值稳定项 |
router_aux_loss_coef | 0.001 | 负载均衡辅助损失系数 |
output_router_logits | False | 是否输出路由 logits |
例如自定义"每 4 层放 1 个注意力层、偏移为 1、共 12 层"的轻量布局:
from transformers import JambaConfig, JambaForCausalLM config = JambaConfig( vocab_size=32000, hidden_size=768, num_hidden_layers=12, attn_layer_period=4, attn_layer_offset=1, expert_layer_period=2, expert_layer_offset=0, num_experts=8, num_experts_per_tok=2, use_mamba_kernels=False, # CPU/无 kernel 环境必须关闭 ) model = JambaForCausalLM(config)注意 offset 必须严格小于对应 period,否则构造时校验会直接报错。
八、公开 API 一览
Jamba 在 Transformers 中对外暴露的类(注册于init.py)包括:
JambaConfig:配置类(本文第七节)。JambaModel:主干模型,forward支持input_ids/inputs_embeds(二者必须且只能给一个)、attention_mask、position_ids、past_key_values、use_cache,返回MoeModelOutputWithPast(modeling_jamba.py)。JambaForCausalLM:因果语言模型头,forward的labels支持 LM loss(-100 忽略),logits_to_keep控制只计算最后 N 个位置的 logits,output_router_logits控制辅助损失输出(modeling_jamba.py),并混入GenerationMixin以支持generate。JambaForSequenceClassification:基于GenericForSequenceClassification实现的序列分类头(modeling_jamba.py)。
三者的forward完整签名可在对应 autodoc 中查阅。此外,modeling_jamba.py由 modular_jamba.py 自动生成(文件头有明确声明),若要修改实现应改动 modular 源文件。
结语
Jamba 的价值在于用工程手段把两种范式的层"安全地焊在一起":attn_layer_period/offset与expert_layer_period/offset定义了宏观拓扑,JambaMambaMixer的选择性扫描与 kernel 回退保证了 SSM 部分既快又可控,MoE 的块稀疏专家张量与负载均衡损失在扩容的同时维持训练稳定。在推理部署侧,请始终牢记两条铁律:不要量化 Mamba 块(用llm_int8_skip_modules=["mamba"]显式跳过),非 CUDA 环境务必设use_mamba_kernels=False。本文全部代码示例均可直接运行,配合 Jamba 官方文档 与 测试套件 即可完成从理解到上手的闭环。
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考