news 2026/9/8 19:34:04

Transformers 中的 Jamba:Transformer–Mamba 混合 MoE 架构解析与部署实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Transformers 中的 Jamba:Transformer–Mamba 混合 MoE 架构解析与部署实战

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)混合架构原理、核心配置参数,以及如何通过PipelineAutoModelForCausalLM和 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=8attn_layer_offset=4),其余为 Mamba 层。与此同时,部分前馈块被替换为 MoE(稀疏专家)块,以在不显著增加计算量的前提下扩充模型容量。

官方 Jamba 系列原始检查点由 AI21(ai21labs组织)发布,本文示例使用的ai21labs/AI21-Jamba-Mini-1.6ai21labs/AI21-Jamba-Large-1.6ai21labs/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_period8每隔多少层出现一次标准注意力层
attn_layer_offset4第一个注意力层所在的层下标
expert_layer_period2每隔多少层出现一次专家(MoE)层
expert_layer_offset1第一个专家层所在的层下标

源码通过三个@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=32num_key_value_heads=8,即 4 个 query 头共享 1 组 KV 头;无偏置线性投影(q/k/v/o),并使用旋转位置编码(RoPE)。前向时通过ALL_ATTENTION_FUNCTIONS.get_interface(...)选择当前attn_implementation(如sdpaflash_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)机制,关键点如下:

  1. in_proj门控投影:把hidden_size投影到2 * intermediate_size(其中intermediate_size = mamba_expand * hidden_size,默认扩展因子 2),随后一分为二得到输入流与silu门控。
  2. 深度可分离因果卷积conv1dmamba_d_conv=4为核宽、按 channel 分组卷积。
  3. x_proj输出输入相关的dt / B / C:这是 Mamba 区别于 S4 的"选择性"来源。
  4. 状态离散化与递归/扫描A_log以 S4D 实数方式初始化并取负指数得到矩阵 A,D初始化为全 1(见 modeling_jamba.py 的_init_weights)。
  5. 与普通 Mamba 的关键差异:对dtBC额外施加 RMSNorm(dt_layernorm/b_layernorm/c_layernorm),见 modeling_jamba.py。

前向中 Jamba 支持两种 SSM 路径(modeling_jamba.py):

  • 单步递归(token 级解码):复用上一 token 的recurrent_stateconv_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)含一个无偏置routerhidden_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 提供三种调用方式(PipelineAutoModel),下面是完整可运行版本。

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")

该方案有三个要点值得展开:

  1. llm_int8_skip_modules=["mamba"]是硬性要求:文档 Notes 第一条明确"不要量化 Mamba 块,以免模型性能退化"。因为 SSM 依赖精确的dt/B/C离散化与状态更新,8-bit 误差会被递归放大。
  2. 量化后的特殊处理路径:源码中专门处理了"模型已被量化"的场景——dt_proj改用带 bias 的nn.Linear调用以兼容量化层(见 modeling_jamba.py),说明 Jamba 的量化支持是经过显式适配的。
  3. device_map手动分布:文档按model.layers.N把相邻层平均分配到 0–7 号 GPU;也可改让from_pretrained(..., device_map="auto")自动分配(AutoModel示例即如此)。

六、注意要点:Mamba kernel 的取舍

文档 Notes 给出两条重要运维经验:

  1. 务必使用优化的 Mamba kernel:没有 mamba-ssm 优化 kernel 的情况下运行 Mamba 会显著变慢(文档原话为 latency 显著下降)。官方 Python 回退实现仅用于无 CUDA kernel 或测试环境。
  2. 若确需关闭 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-ssmcausal-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_size65536词表大小
hidden_size4096隐藏维度
intermediate_size14336FFN 中间维度
num_hidden_layers32总层数(随检查点变化,如示例 device map 对应 72 层版本)
num_attention_heads/num_key_value_heads32 / 8注意力头与 KV 头(GQA)
max_position_embeddings262144最长上下文 256K
num_experts16每 MoE 层专家数
num_experts_per_tok2每 token 路由专家数(Top-2)
expert_layer_period/expert_layer_offset2 / 1MoE 层布局规则
attn_layer_period/attn_layer_offset8 / 4注意力层布局规则
mamba_d_state16SSM 状态维度
mamba_d_conv4深度卷积核宽
mamba_expand2Mamba 隐藏扩展因子
mamba_dt_rank"auto"离散化投影秩,"auto"时按ceil(hidden_size / 16)计算(configuration_jamba.py)
mamba_conv_bias/mamba_proj_biasTrue/False卷积/投影偏置开关
use_mamba_kernelsTrue是否使用 CUDA 优化 kernel
use_mambapyFalse无官方 kernel 时的回退:True用 mamba.py,False用朴素逐token实现
use_associative_scanTruetorch.compile+ torch≥2.9 时用 PyTorch 关联扫描替代朴素循环
rms_norm_eps1e-6RMSNorm 数值稳定项
router_aux_loss_coef0.001负载均衡辅助损失系数
output_router_logitsFalse是否输出路由 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_maskposition_idspast_key_valuesuse_cache,返回MoeModelOutputWithPast(modeling_jamba.py)。
  • JambaForCausalLM:因果语言模型头,forwardlabels支持 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/offsetexpert_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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/8 19:32:18

硬件在环(HiL)测试全解析:原理、应用与职业发展指南

1. HiL 测试到底是做什么的:先把这个行当看清楚再说值不值得很多想入行的人一开始听到“HiL 测试”,脑子里冒出来的问题是:这不就是测测硬件吗?跟板卡测试、整机测试有什么区别?其实差别大了。HiL 全称是 Hardware-in-…

作者头像 李华
网站建设 2026/9/8 19:31:38

【C++ 第二阶段:】智能指针与现代 RAII

C 第二阶段:智能指针与现代 RAII 前言 第一阶段我们手动管理过: FILE* → fopen / fclose char* → new[] / delete[]第二阶段的目标不是“把裸指针换成智能指针”,而是建立清晰的资源所有权模型: 谁拥有资源? 谁负责…

作者头像 李华
网站建设 2026/9/8 19:27:20

应用(客户端)开发、框架开发、驱动/系统开发的关系

最近找工作,好多猎头给我推荐的岗位跟我几乎完全不匹配,有鉴于此,特作此文,以供诸猎头参考。编程语言:打个比方,我想写一本书,可以用中文,也可以用英文,还可以用俄文&…

作者头像 李华