ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

Transformers 中 Zamba2 混合架构模型解析:从架构原理到推理实践

Transformers 中 Zamba2 混合架构模型解析:从架构原理到推理实践 Transformers 中 Zamba2 混合架构模型解析从架构原理到推理实践【免费下载链接】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/transformersZamba2 是 Zyphra 训练并开源Apache 2.0的混合大语言模型系列将状态空间模型Mamba2与共享 Transformer 层结合在同一网络中在 modeling_zamba2.py 与 configuration_zamba2.py 中被完整集成进 Hugging Face Transformers。本文以docs/source/en/model_doc/zamba2.md为骨架结合源码级证据帮助读者掌握 Zamba2 的层布局与权重共享机制、Zamba2Config全部关键参数、Zamba2Model/Zamba2ForCausalLM/Zamba2ForSequenceClassification三类 API 的用法以及从官方 checkpoint 加载模型做生成推理的完整路径。一、模型概览谁、何时、以何协议发布Zamba2 是一个大语言模型LLM系列由 Zyphra 训练官方权重以Apache 2.0协议开源。它于 2024-11-22 发表在 Hugging Face Papers 上并于 2025-01-27 由贡献者 pglo 合入 Transformers 仓库。该系列包含三个公开规格Zyphra/Zamba2-1.2BZyphra/Zamba2-2.7BZyphra/Zamba2-7B它们都采用下一个 token 预测任务训练并使用Mistral v0.1 tokenizer词汇表大小 32000对应配置中的vocab_size: int 32000见 configuration_zamba2.py。根据官方文档描述Zamba2-1.2B、Zamba2-2.7B 与 Zamba2-7B 分别基于 2T 与 3T 两档 token 规模完成预训练且这一架构是官方在小规模上一系列消融实验ablations后确定的方案。二、架构核心每 6 个 Mamba 块之后共享一个 Transformer 层Zamba2 属于hybrid混合架构这是理解它一切行为的总开关。原始文档概括为一句Zamba2 uses shared transformer layers after every 6 mamba blocks每 6 个 Mamba 块之后放置一个共享的 Transformer 层。源码把它展开为一张精确的 54 层布局表。2.1 默认层布局layers_block_type在 configuration_zamba2.py 的__post_init__中layers_block_type缺省时按如下模式自动生成self.layers_block_type ( [linear_attention] ([linear_attention] * 5 [hybrid]) * 7 [linear_attention] * 4 [hybrid] [linear_attention] * 3 [hybrid] [linear_attention] * 2 )这里有两类 token源码 L126-L127 有明确注释布局 token真实含义linear_attention独立的 Mamba2 解码层Zamba2MambaDecoderLayerhybrid一个共享 Transformer 层 一个线性投影 一个 Mamba2 层的组合Zamba2HybridLayer统计该默认布局可得共 54 层其中 45 层为 Mamba2 线性注意力层、9 层为 hybrid 层hybrid_layer_ids见 L140会把这 9 个 hybrid 层的位置索引抽取出来供后续权重共享使用。正因为每个重复单元是「5 个线性注意力层 1 个 hybrid 层」而 hybrid 层内部还包含一个 Mamba 层所以宏观上呈现「每经过 6 个 Mamba 块出现一次共享 Transformer」的节律。从源码结构看num_hidden_layers54、hidden_size2560、num_attention_heads32、mamba_d_state64、mamba_expand2等默认值与 Zamba2-2.7B 规格对齐配置类上方的auto_docstring(checkpointZyphra/Zamba2-2.7B)也印证了这一点。小规格测试时一般会覆写num_hidden_layers3等字段以加速见 test_modeling_zamba2.py。2.2 Hybrid 层内部两个拼接是灵魂真正的混合发生在单个 hybrid 层的数据流里其载体是Zamba2HybridLayermodeling_zamba2.pytransformer_hidden_states self.shared_transformer( hidden_states, original_hidden_statesoriginal_hidden_states, ... ) transformer_hidden_states self.linear(transformer_hidden_states) hidden_states self.mamba_decoder( hidden_states, transformer_hidden_statestransformer_hidden_states, ... )关键机制有以下四点均可在源码中逐一找到落点特征维拼接在Zamba2AttentionDecoderLayer.forwardmodeling_zamba2.py中注意力层的输入是「上一层 Mamba 输出」与「词嵌入原始输出original_hidden_states」沿最后一维拼接得到的2 * hidden_size向量hidden_states torch.concatenate([hidden_states, original_hidden_states], dim-1) hidden_states self.input_layernorm(hidden_states) hidden_states, _ self.self_attn(hidden_states, ...)Zamba2Attention类的 docstringL219-L233对此有专门解释输入维度为attention_hidden_size 2 * hidden_size这一设计源自 Zamba 论文 Fig. 2通过残差拼接使共享注意力始终能看到最新 token 的原始嵌入避免共享层因位置反复出现而信息衰减。共享权重 非共享 Adapter共享 Transformer 层在每个 hybrid 位置被重复使用为抵消权重共享带来的表达能力损失模型的 q/k/v 投影与 MLP 的 gate/up 投影上额外挂了非共享的低秩适配器Adapters官方在注释里明确指出其形式上等同 LoRA但用于 base model而非微调产物。use_shared_attention_adapterTrue时Zamba2Attention会为每个 hybrid 前向位置生成一套nn.Sequential(nn.Linear(hidden, rank), nn.Linear(rank, hidden))的低秩分支L262-L287前向时把适配器输出加到原始投影结果上L306-L310Zamba2MLP的gate_up_proj_adapter_list同理modeling_zamba2.py。adapter_rank默认 128。权重捆绑周期在Zamba2Model.get_layers()modeling_zamba2.py中模型按num_mem_blocks周期复用共享 Transformer 权重。源码注释描述得很直白# Zamba ties Hybrid module weights by repeating blocks after every # num_mem_blocks. So if num_mem_blocks2, the blocks looks like # [1, 2, 1, 2, 1, 2] where all ones share the same set of weights.即num_mem_blocks1时 9 个 hybrid 位置共享同一套 Transformer 权重通过_tied_weights_keys把layers.{i}.shared_transformer映射到首个源模块num_mem_blocks2时则有两套周期交替的权重。测试文件中专门保留了num_mem_blocks2的官方 checkpoint 回归测试test_modeling_zamba2.py用于验证权重绑定周期与 block 编号的对应关系。这正是在不增加模型规模的前提下增加计算量的设计共享层被反复前向且同一共享层在不同位置拥有各自的 adapter 分支。Transformer 输出回流 Mambahybrid 层末尾共享 Transformer 输出经linear投影后作为transformer_hidden_states与 Mamba 层的输入相加再进入下一段 Mamba源码 L992-L996 引用 Zamba 论文 eq. (6)。整个 block 仍保留标准的 residual 结构。2.3 注意力与位置编码的几个特殊参数Zamba2 的共享注意力层在标准 MHA 之上做了定制Zamba2Attentionmodeling_zamba2.py缩放因子因为注意力输入维是2 * hidden_sizehead 维head_dim attention_hidden_size // num_attention_heads 2 * hidden_size / num_attention_heads代码把原始sqrt(head_dim)缩放改成了scaling (head_dim / 2) ** -0.5L250并保留注释说明这一改动。RoPE 按需开启use_mem_ropeFalse默认时共享注意力层不注入旋转位置编码位置信息主要由 SSM 的时序状态承担设为True时则调用Zamba2RotaryEmbedding生成 cos/sin 并施加到 q/kL316-L318。模型只有在use_mem_ropeTrue时才会构造self.rotary_embL1130-L1135。双掩码Zamba2Model.forward会同时生成两类掩码——linear_attention层用create_recurrent_attention_mask递归注意力掩码full_attentionhybrid 内的 Transformer用create_causal_mask标准因果掩码见 modeling_zamba2.py。这是混合架构前向最需要留意的参数语义差异。三、Mamba2 Mixerchunked scan 与流式状态更新的双路径Zamba2MambaMixermodeling_zamba2.py实现了文档所述架构中的 Mamba2 状态空间主体其内部执行分为四步输入投影与切分in_proj把hidden_size映射到intermediate_size conv_dim n_mamba_heads其中conv_dim intermediate_size 2 * n_groups * mamba_d_state随后拆成 gate、B/C、dt 三部分L758-L760。因果深度卷积conv1d是groupsconv_dim的深度可分离卷积核宽mamba_d_conv4输出接 SiLU。SSM 扫描前向时按chunk_size默认 256把序列切成块做块内 块间联合的分块扫描mamba2_chunk_scanL539-L634当缓存了历史状态且当前序列长度为 1自回归逐 token 生成时切换为递归单步状态更新路径mamba2_selective_state_updateL477-L536。门控归一化 输出投影Zamba2RMSNormGatedgroup RMSNorm × SiLU(gate)之后经过out_proj回到hidden_size。与纯 Mamba 不同Zamba2 的门控归一化放在 scan 之后对intermediate_size做分组归一化这是 Mamba2 架构的标配norm_before_gate等开关均出现在与 kernel 的接口中。模型把 SSM 视为有状态statefulZamba2PreTrainedModel上声明了_is_stateful True、_supports_flash_attn / _supports_flex_attn / _supports_sdpa True、supports_gradient_checkpointing Truemodeling_zamba2.py推理缓存的载体就是通用的DynamicCache在use_cache and past_key_values is None时自动创建L1167-L1168其中保存各 Mamba 层的卷积状态与递归状态。causal_conv1d/mamba_ssm相关的扫描与状态更新算子通过 transformers 的 kernel 注册机制use_kernel_func_from_hub_with_fallback从外部 kernel 源加载加速实现加载失败或训练状态下则自动回退到仓库内自带的 PyTorch 参考实现这些辅助函数定义在 modeling_zamba2.py。四、环境准备与快速开始4.1 版本要求官方文档明确Zamba2 要求transformers 4.48.0。pip install transformers4.48.0源码侧模型位于src/transformers/models/zamba2/其中modeling_zamba2.py文件头注明该文件由modular_zamba2.py自动生成、不可手工修改改动必须落到 modular 源文件仓库 CI 会强制校验一致性。若要在源码内核对实现三个文件各司其职configuration_zamba2.py配置类与默认值modeling_zamba2.py全部模块实现modular_zamba2.pymodular 源文件可编辑的事实源。4.2 直接从官方 checkpoint 做生成推理文档给出的标准推理范式无需任何自定义注册AutoModelForCausalLM与AutoTokenizer会自动按model_typezamba2路由到对应实现from transformers import AutoModelForCausalLM, AutoTokenizer tokenizer AutoTokenizer.from_pretrained(Zyphra/Zamba2-7B) model AutoModelForCausalLM.from_pretrained(Zyphra/Zamba2-7B, device_mapauto) input_text What factors contributed to the fall of the Roman Empire? input_ids tokenizer(input_text, return_tensorspt).to(model.device) outputs model.generate(**input_ids, max_new_tokens100) print(tokenizer.decode(outputs[0]))几个值得说明的要点device_mapauto依赖accelerate完成权重分发对于 7B 级别模型建议同时配合 4-bit 量化如BitsAndBytesConfig在显存受限环境下使用——测试文件 test_modeling_zamba2.py 中已导入BitsAndBytesConfig用于此类场景。生成阶段prepare_inputs_for_generation会把logits_to_keep强制设为配置里的num_logits_to_keepmodeling_zamba2.py默认只计算最后一个 prompt token 的 logits从而显著压低长序列解码的内存峰值。有状态推理依赖DynamicCache因此不要关闭use_cache否则每个 Mamba 层都要从零重建卷积/递归状态既慢又费内存。五、Zamba2Config全部关键参数速查Zamba2Config继承PreTrainedConfig声明了model_type zamba2并把旧式命名自动映射到新字段layer_types - layers_block_type、head_dim - attention_head_dimconfiguration_zamba2.py。下表汇总了配置类 docstring 与类属性中的全部可调参数默认值取自 configuration_zamba2.py参数默认值含义vocab_size32000词表大小对应 Mistral v0.1 tokenizermax_position_embeddings4096最大位置编码长度use_long_contextTrue时自动扩展为 16384hidden_size2560隐藏层宽度num_hidden_layers54总解码层数Mamba 层 hybrid 层layers_block_typeNone每层的类型序列linear_attention/hybrid为None时按 2.1 节规则自动生成mamba_d_state64SSM 状态维度mamba_d_conv4因果卷积核宽mamba_expand2Mamba 内部通道扩展倍数决定intermediate_size与mamba_headdimmamba_ngroups1Mamba2 演化矩阵的 group 数n_mamba_heads8Mamba2 演化矩阵 head 数time_step_min0.001dt 初始化的下限对数均匀采样time_step_max0.1dt 初始化的上限time_step_floor1e-4dt 的 clamp 下限use_mamba_kernelsTrue是否使用快速 Mamba kerneluse_conv_biasTruemixer 的卷积层是否加 biaschunk_size256序列切块大小Mamba2 分块扫描的粒度use_mem_eff_pathFalse是否使用融合的 conv1dscan 路径add_bias_linearFalse各线性层是否加 biasintermediate_sizeNoneMLP 中间维度None时在__post_init__取4 * hidden_sizehidden_actgelu共享 Transformer MLP 的激活函数num_attention_heads32共享注意力 head 数num_key_value_headsNoneKV head 数None时取num_attention_headsattention_dropout0.0注意力 dropoutnum_mem_blocks1不共享的 Transformer 块个数即权重捆绑周期见 2.2 节use_shared_attention_adapterFalse是否在共享注意力的 q/k/v 投影上启用非共享低秩适配器adapter_rank128共享 MLP 与共享注意力中适配器的秩use_mem_ropeFalse是否在共享注意力层加入 RoPErope_parametersNoneRoPE 参数如rope_type、rope_theta传给Zamba2RotaryEmbeddinginitializer_range0.02权重初始化范围rms_norm_eps1e-5RMSNorm 的 epsilonuse_cacheTrue是否返回/使用 KV 与 SSM 状态缓存num_logits_to_keep1生成时仅计算最后 N 个 prompt logits长序列省显存的关键pad_token_id0padding token idbos_token_id1句首 token ideos_token_id2句末 token iduse_long_contextFalse启用上下文扩展版 Zamba改写 RoPE 并把max_position_embeddings拉到 16384tie_word_embeddingsTrue是否捆绑词嵌入与 lm_head 权重__post_init__阶段会派生几个关键派生量configuration_zamba2.pyintermediate_size 4 * hidden_size若未显式指定attention_hidden_size 2 * hidden_size共享注意力的拼接输入维attention_head_dim attention_hidden_size // num_attention_headsmamba_headdim (mamba_expand * hidden_size) // n_mamba_headskv_channels hidden_size // num_attention_headsnum_query_groups num_attention_heads。配置类的标准用法源码 docstring 中的官方示例是直接从配置构建一个小模型from transformers import Zamba2Model, Zamba2Config # 初始化一个 Zamba2-2.7B 风格的配置 configuration Zamba2Config() # 基于该配置初始化模型 model Zamba2Model(configuration) # 读取模型实际生效的配置 configuration model.config六、三类公开模型 API 及其前向语义zamba2模块通过__init__.py的惰性加载暴露四个符号Zamba2Config、Zamba2Model、Zamba2ForCausalLM、Zamba2ForSequenceClassification导出清单见 modeling_zamba2.py。6.1 Zamba2Model纯主干模型nn.Embedding词表 32000→ 按layers_block_type动态构建的层列表 →final_layernorm。前向签名支持input_ids/inputs_embeds二选一、attention_mask、position_ids、past_key_values、use_cache返回BaseModelOutputWithPastlast_hidden_state与past_key_values。注意两点进入第一个层前会original_hidden_states torch.clone(inputs_embeds)该克隆贯穿所有层作为共享 Transformer 的拼接输入modeling_zamba2.py。未开启use_mem_rope时不会构造rotary_embposition_embeddings为None。6.2 Zamba2ForCausalLMZamba2Model 无 bias 的lm_head同时继承GenerationMixin因此model.generate可用。_tied_weights_keys {lm_head.weight: model.embed_tokens.weight}保证输出层与词嵌入在tie_word_embeddingsTrue时共享权重。其核心前向特性是裁剪 logitsslice_indices slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep logits self.lm_head(hidden_states[:, slice_indices, :])配合labels即可直接做语言建模训练/微调与官方 docstring 内嵌示例一致from transformers import AutoTokenizer, Zamba2ForCausalLM model Zamba2ForCausalLM.from_pretrained(Zyphra/Zamba2-7B-v1) tokenizer AutoTokenizer.from_pretrained(Zyphra/Zamba2-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]6.3 Zamba2ForSequenceClassification在主干之上加一个线性分类头score并按因果模型的惯例取序列最后一个 token 的 logits做分类。源码明确了两点语义若配置了pad_token_id默认 0会通过掩码找到每行最后一个非 padding token未配置pad_token_id且 batch 1 时直接抛出ValueError因为无法判定最后一个有效 token只允许 batch1。若以inputs_embeds而非input_ids喂入模型无法感知 padding同样退化为取每行最后一个位置源码会打印警告。num_labels1时计算 MSE 回归损失否则计算交叉熵分类损失见 modeling_zamba2.py 与 docstring。七、验证与回归保障从测试理解使用边界模型能力的可信度取决于测试覆盖。tests/models/zamba2/下的Zamba2ModelTester以num_hidden_layers3的小配置驱动完整测试矩阵test_modeling_zamba2.py覆盖主干前向create_and_check_model、因果语言模型create_and_check_for_causal_lm、序列分类create_and_check_for_sequence_classification分块预填充chunked prefill一致性验证分别在 CPU 与 GPU 上检查同一输入的分块路径与整体路径输出一致L400-L408解码器past_key_values大输入回归create_and_check_decoder_model_past_large_inputsnum_mem_blocks2官方 checkpoint 的权重绑定顺序回归L662 起防止共享 Transformer 被绑到错误的源层。这些测试从工程角度印证了本文第 2、3 节描述的架构要点Mamba2 分块扫描路径必须与逐 token 递归路径数值等价hybrid 权重周期必须与官方发布权重一致stateful 缓存必须能被DynamicCache正确读写。八、模型卡片、问题反馈与许可证官方模型卡片model card与权重发布使用以下标识符可在 Hugging Face 上检索Zyphra/Zamba2-1.2BZyphra/Zamba2-2.7BZyphra/Zamba2-7B文档同时注明模型权重以Apache 2.0开源若遇到模型输出相关问题或希望参与社区讨论可在对应模型的 Hugging Face Discussions 区域发起。Zamba2 在 Transformers 中的实现被作为标准社区贡献维护所有核心实现均可直接阅读本仓库源码modeling_zamba2.py、configuration_zamba2.py、modular_zamba2.py与测试test_modeling_zamba2.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),仅供参考
返回列表