
Hugging Face Transformers 中的 Zamba 混合架构模型Mamba 共享 Transformer 的状态空间语言模型详解【免费下载链接】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导读Zamba 是 Zyphra 团队提出的混合大语言模型架构它把状态空间模型以 Mamba 为代表的线性推理复杂度与 Transformer 的表达能力结合在同一个解码器中其开源权重 Zamba-7B-v1 采用 Apache 2.0 协议发布。本文以 docs/source/en/model_doc/zamba.md 为骨架并结合本仓库的 配置实现、模型实现 与 测试套件完整讲解 Zamba 的架构原理、全部ZambaConfig参数、环境搭建与推理实战让你既能直接跑通 Zamba-7B-v1也能从源码级理解它为什么被称为“带周期共享注意力层的 Mamba”。一、Zamba 模型概述Zamba 是一个通过下一个 Token 预测next-token prediction训练的因果语言模型LLM由 Zyphra 训练并开源权重采用Apache 2.0 许可证。它首次在论文中公开的时间是 2024-05-26并于 2024-10-04 由贡献者 pglo 合入本仓库对应 Transformers 4.46 左右版本线。其核心定位可以概括为一句话它是状态空间模型具体为 Mamba与 Transformer 的混合体hybrid。设计上的两个突出特征是稀疏放置共享注意力层每隔 6 个 Mamba 块放置一个共享的 Transformer 层shared transformer layer。复用现成分词器直接使用 Mistral v0.1 的分词器词汇表 32000无需额外训练分词器。Zyphra 在推出最终方案前进行了多轮小规模消融实验ablations at small scales最终确认了“周期性共享注意力 Mamba 主体”的排列。Zamba-7B-v1 在约 1T tokens 的文本与代码数据上完成预训练。模型权重、社区讨论等可围绕模型标识符Zyphra/Zamba-7B-v1获取官方模型卡也在同名的 Hub 仓库下维护。从工程角度看本仓库把 Zamba 系列实现为标准的PreTrainedModel家族自动注册到AutoModelForCausalLM等入口中见 modeling_auto.py 与 auto_mappings.py同时显式导出ZambaModel、ZambaForCausalLM、ZambaForSequenceClassification、ZambaPreTrainedModel与ZambaConfig。二、架构原理为什么是“Mamba 主体 周期性共享 Transformer”要理解 Zamba先看模型整体的层排列方式。官方模型文档给出的描述是每 6 个 Mamba 块之后放置一个共享 Transformer 层。而在ZambaConfig的默认构建逻辑中这一模式被精确化为下面的列表生成规则见 configuration_zamba.pyif self.layers_block_type is None: self.layers_block_type [ linear_attention, linear_attention, hybrid, ] [ hybrid if i % self.attn_layer_period self.attn_layer_offset else linear_attention for i in range(self.num_hidden_layers - 3) ]即层类型按linear_attention纯 Mamba 块与hybrid混合块两类标注由attn_layer_period与attn_layer_offset两个参数共同决定周期与偏移。以 Zamba-7B-v1 的默认配置num_hidden_layers 76、周期 6、偏移 4计算会得到 13 个hybrid位置测试代码里对注意力输出数量的断言ceil((num_hidden_layers - attn_layer_offset) / attn_layer_period) 1正好与之对应见 test_modeling_zamba.py。2.1 从数据流看懂“共享”与“混合”打开 modeling_zamba.py 的主干类ZambaModel第 820 行起其__init__与forward揭示了 Zamba 的两大核心设计。设计一Transformer 层的权重是全局共享的。在构建 76 层时每个hybrid位置都会包出一个ZambaHybridLayer但它内部持有的shared_transf其实是指向同一份ZambaAttentionDecoderLayer参数for layer_id, layer_type in enumerate(self.layers_block_type): mamba ZambaMambaDecoderLayer(config, layer_idxlayer_id) if layer_type hybrid: linear nn.Linear(self.config.hidden_size, self.config.hidden_size, biasFalse) layers.append(ZambaHybridLayer(ZambaAttentionDecoderLayer(config), linear, mamba))共享机制在权重层面通过_tied_weights_keys正则实现把除首个 hybrid 位置之外的所有layers.*.shared_transf绑定到第一个 hybrid 层因此整网虽有多处注意力实际可学习参数只有一份注意力层 各自的投影/归一化小件。这也是测试中注明output_attentions只产出注意力层份数的原因。设计二混合层内部有“拼接—投影—残差注入”三段流ZambaHybridLayer.forward第 731 行起shared_transf处理当前 Mamba 输出与“原始嵌入”的拼接结果详见下文 2.2紧接一个无偏置的线性层self.linear把 Transformer 输出投影回 hidden sizemamba_decoder收到该投影结果后会先把它加到自己的输入上再做 Mamba 计算即论文 eq.(6) 的残差注入见 ZambaMambaDecoderLayer.forward 的注释最后再做 Mamba 输出的残差连接。这种“周期插入共享注意力 残差回注”的混合方式让信息可以在有限的注意力位置被“复习”与全局整合而大部分层仍享受 Mamba 的线性复杂度状态传递。2.2 注意力输入翻倍original_hidden_states的拼接技巧Zamba 最容易被忽略但非常关键的设计是共享注意力层的输入维度是普通隐藏维度hidden_size的两倍。原因在ZambaAttention的类注释里写得很清楚第 118-125 行输入维度为attention_hidden_size 2 * hidden_sizehead 维度为attention_hidden_size // num_heads。多出来的 2 倍来自输入是original_hidden_states词嵌入输出与前一个 Mamba 层输出的拼接见论文 fig. 2。对应地ZambaAttentionDecoderLayer.forward首先执行hidden_states torch.concatenate([hidden_states, original_hidden_states], dim-1) hidden_states self.input_layernorm(hidden_states)而original_hidden_states正是ZambaModel.forward中对词嵌入输出做的克隆第 879 行它会一路携带到每个注意力位置。也就是说共享注意力每次看到的都是“该位置当前的深层状态 最初的 token 嵌入”让 76 层深网在注意力发生的少数位置仍能直接感知词级原始信息缓解纯 Mamba 堆叠可能带来的信息遗忘。其它来自 Transformer 一侧的实现细节归一化沿用 Llama 系风格ZambaRMSNorm等价 T5LayerNorm并在进入注意力前对两倍宽度输入做 RMSNormMLP 前另有pre_ff_layernorm注意力本身是标准的 MHA支持GQA 式 KV 头复用num_key_value_heads默认 16与num_attention_heads相同repeat_kv逻辑与 Llama 一致缩放因子被调整为(head_dim / 2) ** -0.5即除以sqrt(head_dim/2)与两倍宽度的设计配套见 ZambaAttention注意力后端通过ALL_ATTENTION_FUNCTIONS接口分发配合_supports_flash_attn True、_supports_sdpa True第 786-787 行因此FlashAttention 与 PyTorch SDPA 均可使用MLP 采用 SwiGLU 式门控结构ZambaMLP第 605-618 行激活为 GELU。2.3 多头的 MambaZambaMambaMixerMamba 主体在实现上并非直接复刻MambaMixer而是在其上做了多头化改造。ZambaMambaMixer的 docstring 说明它与 Mamba 原版的两点差异in_proj输出按n_mamba_heads默认 2切成多个头x_proj_weight与dt_proj的权重/偏置每个 Mamba 头各有一套各头独立完成与原始 Mamba 相同的计算直到out_proj之前再把各头预激活拼接起来送入输出投影modeling_zamba.py。单头的内部计算流程第 464 行forward依然是 Mamba 的标准四步门控线性投影in_proj把输入投影到2 × (mamba_expand × hidden_size)拆出hidden_states_B_C与gate因果卷积causal_conv1d核宽mamba_d_conv4沿序列做深度可分离因果卷积可选用silu激活选择性状态空间扫描由输入驱动生成离散化时间步dt与选择性参数B、C对每个 Mamba 头执行 per-head selective scan训练/长序列走全序列扫描mamba_selective_scan增量解码单 Token 时走逐头状态更新mamba_selective_state_update即按ssm_state state * dA dBx递推输出投影out_proj汇合所有 Mamba 头并映射回 hidden size。其中A采用 S4D 实数初始化对数域存储A_logD初始化为 1dt_proj_bias通过 softplus 反函数做逆初始化保证初始时间步落在[time_step_min, time_step_max]区间内且不低于time_step_floor_init_weights。2.4 为什么需要两条 mask因果 mask 与循环 mask由于 Zamba 中同时存在注意力层需要标准因果 mask与 Mamba 层padding 语义不同需要把 padding 状态清零ZambaModel.forward在进入层循环前会一次性构造好两类掩码并封装成字典modeling_zamba.pycausal_mask_mapping { full_attention: create_causal_mask(**mask_kwargs), # 供共享注意力层使用 linear_attention: create_recurrent_attention_mask(**mask_kwargs), # 供 Mamba 层使用 }在层循环中hybrid层把线性注意力掩码交给 Mamba 部分、把全注意力因果掩码交给共享 Transformer 部分同时供ZambaMambaDecoderLayer处理。padding 位置的状态清零还依赖apply_mask_to_padding_states辅助函数第 183-192 行参考 state-spaces/mamba issue #66避免 padding token 的状态污染真实序列。三、ZambaConfig全部配置参数与源码级说明模型配置类位于 configuration_zamba.pymodel_type zamba并提供了两层向后兼容别名映射attribute_map {layer_types: layers_block_type, head_dim: attention_head_dim}即旧式layer_types/head_dim字段会被自动归一化到新字段。下面给出该类在 class 级默认值 与 docstring 中定义的全部参数含 Zamba-7B-v1 实际默认参数默认值含义vocab_size32000词表大小对齐 Mistral v0.1 分词器tie_word_embeddingsTrue是否绑定输入/输出词嵌入hidden_size3712隐藏层维度attention_hidden_sizeNone推导为2 * hidden_size 7424注意力输入维度拼接后翻倍intermediate_size14848MLP 中间维度num_hidden_layers76解码器层数含纯 Mamba 块与混合块num_attention_heads16注意力头数attention_head_dimNone推导为2 * hidden_size // num_attention_heads 464注意力单头维度num_key_value_heads16KV 头数GQAn_mamba_heads2每个 Mamba 层的 Mamba 头数hidden_actgeluTransformer MLP 激活hidden_mamba_actsiluMamba 内部激活initializer_range0.02权重初始化标准差rms_norm_eps1e-5RMSNorm 的 epsilonuse_cacheTrue是否返回/使用 KV 缓存num_logits_to_keep1生成时只计算最后 N 个位置 logits省显存pad_token_id/bos_token_id/eos_token_id0 / 1 / 2特殊 Token 编号max_position_embeddings4096最大位置长度注意力部分attention_dropout0.0注意力 dropoutattn_layer_period6每多少个位置出现一次共享注意力attn_layer_offset4共享注意力在周期内的偏移use_mamba_kernelsTrue是否使用快速 Mamba 内核mamba_d_state16SSM 状态维度 Nmamba_d_conv4因果卷积核宽mamba_expand2内维扩展倍数SSM 内维 expand × hiddenmamba_dt_rankauto时间步投影秩auto 时取ceil(hidden_size / 16) 232time_step_min/time_step_max0.001 / 0.1初始时间步采样区间time_step_floor1e-4时间步下限mamba_conv_biasTrue因果卷积是否带偏置mamba_proj_biasFalseMamba 线性投影是否带偏置layers_block_typeNone自动生成每层类型列表等价旧字段layer_types3.1__post_init__的推导逻辑创建配置时如果没有显式给出部分字段会依据下面三条规则自动补齐configuration_zamba.pyself.attention_hidden_size self.attention_hidden_size or 2 * self.hidden_size self.attention_head_dim self.attention_head_dim or 2 * self.hidden_size // self.num_attention_heads self.mamba_dt_rank math.ceil(self.hidden_size / 16) if self.mamba_dt_rank auto else self.mamba_dt_rank对应 Zamba-7B-v1attention_hidden_size 7424attention_head_dim 7424 // 16 464mamba_dt_rank ceil(3712/16) 232。3.2validate_architecture构造期校验配置类通过strict装饰器启用架构校验其中 validate_architecture 校验if (self.mamba_expand * self.hidden_size) % self.n_mamba_heads ! 0: raise ValueError(intermediate_size should be divisible by n_mamba_heads.)即 SSM 内维必须能被 Mamba 头数整除否则无法均分多头构造ZambaConfig时会直接报错。默认 7424 % 2 0 满足条件。单元测试中的微型配置如hidden_size64、mamba_expand保持默认也遵循此约束见 test_modeling_zamba.py。3.3num_logits_to_keep长序列推理的显存开关这是一个容易被忽略但在长上下文生成中极其重要的参数默认 1。生成时模型只对最后一个 prompt token计算 logits 即可继续采样若对整条长序列都算 logits 会显著增加显存占用。ZambaForCausalLM.prepare_inputs_for_generation会在生成入口自动把logits_to_keep注入为config.num_logits_to_keepmodeling_zamba.py并在forward中用slice(-logits_to_keep, None)裁切只算尾部 logits。四、环境准备与安装Prerequisites运行 Zamba 需要满足以下前提Transformers 版本Zamba 模型自 2024-10-04 合入本仓库官方文档要求使用 4.46.0 及以上版本对应文档中的安装命令为 4.45.0 起建议直接升级到当前仓库对应的最新主线pip install transformers4.45.0Mamba 快速内核强烈建议要运行优化过的 Mamba 实现需要额外安装两个 CUDA 扩展包pip install mamba-ssm causal-conv1d1.2.0CUDA 设备使用快速内核要求模型运行在 CUDA 设备上。不用内核的降级方案不带优化 Mamba 内核也能跑但文档明确提示这不推荐因为会产生显著更高的推理延迟。需要加载模型时显式传use_mamba_kernelsFalse。需要注意use_mamba_kernelsTrue时会强制要求mamba-ssm与causal-conv1d已安装且模块位于 CUDA 设备上否则ZambaConfig会抛ValueError见 configuration_zamba.py 的 docstring。在无内核的纯 PyTorch 路径上模型实现提供了三层后备机制见 modeling_zamba.py优先走mamba_ssm/causal_conv1d的融合内核其次在 PyTorch ≥ 2.9 时可尝试torch的 associative scan 或可选的mambapy.pscan并行扫描最后退化为逐时间步的循环扫描recurrent iteration。这也解释了为什么tests/models/zamba/test_modeling_zamba.py的集成测试统一使用use_mamba_kernelsFalse以便在任何环境复现。五、快速上手推理Quick Start5.1 最小推理示例下面的示例直接取自官方模型文档加载预训练权重Zyphra/Zamba-7B-v1并完成 100 个新 Token 的续写from transformers import AutoModelForCausalLM, AutoTokenizer tokenizer AutoTokenizer.from_pretrained(Zyphra/Zamba-7B-v1) model AutoModelForCausalLM.from_pretrained(Zyphra/Zamba-7B-v1, device_mapauto) input_text A funny prompt would be input_ids tokenizer(input_text, return_tensorspt).to(model.device) outputs model.generate(**input_ids, max_new_tokens100) print(tokenizer.decode(outputs[0]))要点说明Zamba 走标准因果 LM 生成接口继承GenerationMixinAutoModelForCausalLM会根据配置自动解析到本仓库的ZambaForCausalLM自动映射注册见 modeling_auto.py。device_mapauto依赖 accelerate若未安装或想手动控制设备也可先model.to(cuda)再自行把输入搬到model.device。若未安装 Mamba 内核需要改为from_pretrained(Zyphra/Zamba-7B-v1, device_mapauto, use_mamba_kernelsFalse)。num_logits_to_keep默认已为 1长 prompt 场景无需手动设置即可获得省显存收益。5.2 带 batch 与 padding 的生成Zamba 模型是**有状态stateful**的_is_stateful True同时keys_to_ignore_at_inference [past_key_values]、_skip_keys_device_placement [past_key_values]见 ZambaPreTrainedModel。KV/SSM 状态缓存在解码期通过DynamicCache保存。集成测试给出了一套 batch padding 的标准写法test_modeling_zamba.pytokenizer.add_special_tokens({pad_token: [PAD]}) model.resize_token_embeddings(len(tokenizer)) inputs tokenizer( [Hey how are you doing on this lovely evening?, Tell me a story], paddingTrue, return_tensorspt, ).to(model.device) out model.generate(**inputs, do_sampleFalse, max_new_tokens10) output_sentences tokenizer.batch_decode(out)由于 Mamba 的 padding 状态必须显式清零请务必为每个 batch 传入attention_mask上面的 pad token 扩展与resize_token_embeddings是必须的前置步骤。5.3 缓存正确性带缓存 vs 不带缓存的输出一致性作为混合模型Zamba 的增量解码需要同时维护注意力 KV 缓存与 Mamba 的卷积/循环状态。测试test_decoder_model_past_with_large_inputstest_modeling_zamba.py验证了把整段序列一次性前向与先算前半段缓存、再只喂后续 3 个 Token 的增量前向二者在裁剪后的 hidden states 上torch.allclose(atol1e-3)成立。这从测试层面证明了状态缓存路径与完整前向路径的数值一致性。六、面向任务的三种模型封装本仓库为 Zamba 提供了三种开箱即用的封装__all__见 modeling_zamba.pyZambaModel裸的 Transformer 解码器主体。forward接收input_ids/inputs_embeds、attention_mask、position_ids、past_key_values、use_cache等返回BaseModelOutputWithPastlast_hidden_state与可选的past_key_values。主干在forward开头会把词嵌入克隆为original_hidden_states并贯穿整个层循环。feature-extraction pipeline 使用该类。ZambaForCausalLM在ZambaModel之上叠加lm_head无偏置线性层支持labels计算交叉熵损失与model.generate(...)文本生成。_tied_weights_keys会把lm_head.weight与model.embed_tokens.weight绑定因此权重共享后总参数量显著低于同等宽度的纯 Transformer。text-generation pipeline 使用该类。模型文档内置的 docstring 示例给出了标准用法第 958-970 行from transformers import AutoTokenizer, ZambaForCausalLM model ZambaForCausalLM.from_pretrained(Zyphra/Zamba-7B-v1) tokenizer AutoTokenizer.from_pretrained(Zyphra/Zamba-7B-v1) prompt Hey, are you conscious? Can you talk to me? inputs tokenizer(prompt, return_tensorspt) generate_ids model.generate(inputs.input_ids, max_length30) tokenizer.batch_decode(generate_ids, skip_special_tokensTrue, clean_up_tokenization_spacesFalse)[0]ZambaForSequenceClassification叠加一个分类头score nn.Linear(hidden_size, num_labels)与 GPT-2 等因果模型一致取最后一个 token的表示做分类。其行为细节值得注意见 modeling_zamba.py若配置了pad_token_id取每行最右侧的非 padding token因此可以同时兼容左 padding 与右 padding。若未配置pad_token_idbatch 1 会直接报错无法区分 paddingbatch 1 时取行尾 token。若以inputs_embeds传入而非input_ids无法判断 padding 位置同样退化为取行尾。problem_type会根据num_labels与标签 dtype 自动推导为 regression / single-label / multi-label并分别选用 MSE、CrossEntropy 或 BCEWithLogits 损失。测试中的pipeline_model_mappingtest_modeling_zamba.py确认了 Zamba 可接入的 Pipelinefeature-extractionZambaModel、text-generationZambaForCausalLM、text-classification / zero-shotZambaForSequenceClassification。七、支持的注意力实现与训练特性注意力后端_supports_flash_attn True、_supports_sdpa True即既支持 FlashAttention 2也支持 PyTorch 原生 SDPA二者均可通过attn_implementationflash_attention_2/sdpa或模型默认配置启用。与文档头部徽章FlashAttention、SDPA一致。在启用 Flash Attention 时注意 Zamba 的共享注意力采用绑定权重若配合 4-bit 量化需留意 dtype 一致性测试中因此跳过了 FA2 fp32 LN 专项见 test_modeling_zamba.py。梯度检查点supports_gradient_checkpointing TrueMamba 层继承GradientCheckpointingLayer利于长序列微调显存控制。加载与离线相关限制由于混合层类与绑定权重的组合测试中注明 CPU offload 与磁盘 offloadbin / safetensors暂不适用见 test_modeling_zamba.py 的 skip 说明做大规模设备迁移时需留意。八、用测试与真实推理结果验证模型仓库自带了两级测试tests/models/zamba/test_modeling_zamba.py单元/通用测试层ZambaModelTest通过ModelTesterMixin、GenerationTesterMixin、PipelineTesterMixin覆盖三种封装类的前向形状、损失、带缓存增量生成一致性、注意力输出数量ceil((num_hidden_layers - attn_layer_offset)/attn_layer_period) 1、config 通用测试等。测试配置采用迷你尺寸如hidden_size64、attn_layer_offset1、num_hidden_layers5并强制use_mamba_kernelsFalse保证无 CUDA 内核也能运行。集成测试层ZambaModelIntegrationTest使用真实权重Zyphra/Zamba-7B-v1bfloat16 加载给出了可复现的输出基准输入Hey how are you doing on this lovely evening?贪心解码 10 个新 Token 的期望输出为s Hey how are you doing on this lovely evening? I hope you are all doing well. I am同时附带了前 40 个 logits 的逐位参考值test_modeling_zamba.py。如果你在本地复现推理可以把这些期望输出当作“环境是否配置正确”的冒烟测试依据。九、常见问题与注意事项小结版本与内核先确认transformers满足 4.45/4.46 以上的要求追求性能请安装mamba-ssm与causal-conv1d1.2.0并保证 CUDA 环境否则记得显式use_mamba_kernelsFalse并接受更高的推理延迟。device 要求use_mamba_kernelsTrue时内核只在 CUDA 上可用CPU 或 Apple Silicon 环境请关闭该开关。padding 语义Mamba 是循环状态模型padding token 必须通过attention_mask显式清零batch 解码务必传 mask分类任务应预置pad_token_id以获得正确的“最后非 padding token”语义。分词器Zamba 直接使用 Mistral v0.1 分词器词表 32000s为 BOS、/s为 EOS、[PAD]需自行添加无需单独训练。许可与获取模型权重以 Apache 2.0 开源社区讨论与模型卡围绕官方 Zamba-7B-v1 仓库进行模型相关 issue 也可在该仓库讨论区提出。与后续版本的关系本仓库测试注释中多次提到 “Same as zamba2”说明 Zamba 系列架构在同一代码库中持续演进本文描述的具体行为周期共享注意力、多头 Mamba、状态缓存语义等以当前仓库实现为准升级版本时建议回归验证上述测试基准。参考资料仓库内定位官方模型文档docs/source/en/model_doc/zamba.md配置实现src/transformers/models/zamba/configuration_zamba.py模型实现src/transformers/models/zamba/modeling_zamba.py测试套件tests/models/zamba/test_modeling_zamba.pyAuto 类注册src/transformers/models/auto/modeling_auto.py、src/transformers/models/auto/auto_mappings.py【免费下载链接】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),仅供参考