拓冰建站拓冰建站
首页 / 资讯中心 / 正文

MAX 平台 Nemotron-H 混合架构解析:Mamba-2 + NoPE Attention + relu2 MLP 的工程化实现

MAX 平台 Nemotron-H 混合架构解析Mamba-2 NoPE Attention relu2 MLP 的工程化实现【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo导读本文以 MAX 平台Modular Platform包含 MAX 与 Mojo中max.pipelines.architectures.nemotron_h模块为主线深入解析 NVIDIA Nemotron-HNemotron-3 系列这一混合解码器架构在 MAX 推理管线中的完整落地方式如何用hybrid_override_pattern描述稀疏注意力 选择性状态空间SSM 稠密 MLP/MoE的层间混合Mamba-2 SSD chunked scan 与 NoPE GQA attention 如何在一个语言图内共存以及模型opt 每张量静态 FP8 量化如何在 Mamba 与 MLP 投影上按模块生效。读完本文你将掌握 Nemotron-H 在 MAX 中的配置字段、混合层映射规则、状态池conv/SSM slot 池的推理生命周期以及从 HuggingFace 权重到 MAX 图编译的完整适配链路。一、模块定位一个文档桩背后的完整架构包max/python/docs/pipelines.architectures.nemotron_h.rst是 Sphinx autosummary 的模块文档入口其正文由automodule:: max.pipelines.architectures.nemotron_h指令动态生成实际内容全部来自max/python/max/pipelines/architectures/nemotron_h/目录下的六个源码文件。因此本文以该目录为核心研究对象文件职责nemotron_h.pynn.Module 层MLP、MoE、Attention、Mamba-2 mixer、Block、完整解码器model_config.pyNemotronHConfig配置类、混合模式解析、FP8 量化配置构建model.py管线模型NemotronHModel与输入封装NemotronHInputsarch.py架构注册SupportedArchitectureweight_adapters.pyHuggingFace → MAX 权重名称映射与类型修正state_cache.pyGPU 常驻 conv/SSM 状态 slot 池tokenizer.py推理分隔符think//thinktoken id 解析模块顶层__init__.py导出四个公共符号NemotronHConfig、NemotronHInputs、NemotronHModel、nemotron_h_arch。理解这四个符号就等于理解了该架构从配置 → 输入 → 模型 → 注册的完整骨架。二、架构总览一种没有旋转位置编码的混合解码器nemotron_h.py的模块 docstring 明确给出了该架构的数学本质——它与 HuggingFace 的NemotronHForCausalLM其torch_forward逐操作对应翻译为惯用的 MAX 表达Block 结构pre-norm RMSNorm → mixer → residual addresidual_in_fp32False即标准的 pre-norm 残差块。Mamba-2 mixerin_proj → [gate, hidden_states_B_C, dt]对hidden_states_B_C做 depthwise SiLU 卷积SSD chunked scan同时服务 prefill 与 decodegated RMSNormnorm_before_gateFalseout_proj。AttentionGQAGrouped-Query Attention、NoPE无 RoPE 旋转位置编码、无 bias。MLPrelu2 down(relu(up(x))**2)非门控non-gated、无 bias。其中最值得注意的设计是NoPENemotronHAttention的 docstring 指出位置信息经由 SSM 层流动因此注意力层不添加任何位置编码——与 HF 参考实现中position_embeddings未被使用的事实一致。这意味着传统 Transformer 依赖 RoPE 注入的位置信息在 Nemotron-H 中改由 Mamba-2 的选择性状态空间机制承载。相应地arch.py中注册的weight_adapters注释也重申了NoPE: attention adds no rotary embedding。三、混合层模式hybrid_override_pattern的字符映射Nemotron-H 是稀疏注意力 Mamba-2 MLP/MoE的层间混合模型每一层到底放哪种 mixer由 HuggingFace 配置中的hybrid_override_pattern字符串决定。model_config.py中的parse_hybrid_pattern()负责把它解析成逐层类型列表字符层类型说明MmambaMamba-2 选择性状态空间混合器*attentionNoPE GQA 注意力-mlp非门控 relu2 稠密前馈层EmoeNemotron-3 MoE 混合层如 30B-A3B其余任何字符都会抛出ValueError: invalid hybrid_override_pattern character。例如测试文件 test_nemotron_h_fp8_kv_config.py 中使用的M-*M模式表示四层结构Mamba → MLP → Attention → Mamba。NemotronHConfig提供两个派生属性便于下游消费mamba_layer_indices与attention_layer_indices分别返回 Mamba 层和 Attention 层的绝对层号列表。在NemotronH完整解码器的构造函数中每个块根据config.layer_kinds[layer_idx]实例化对应 mixer同时维护两个独立的索引KV 缓存索引只给attention层顺序分配 0、1、2、… 的 KV 缓存切片与绝对层号解耦Mamba 层索引只给mamba层顺序编号用于索引该层的 conv/SSM 状态池。四、NemotronHConfig配置字段全景NemotronHConfigmodel_config.py继承ArchConfigWithStoredKVParams, ArchConfigWithKVCache是一个kw_onlydataclass。它声明了两组编码能力DEFAULT_ENCODING: ClassVar[SupportedEncoding] bfloat16 SUPPORTED_ENCODINGS: ClassVar[set[SupportedEncoding]] { bfloat16, float8_e4m3fn, }4.1 核心维度字段字段含义hidden_size/vocab_size/num_hidden_layers解码器基础维度layer_norm_epsilonRMSNorm 的 epsilonmax_seq_len最大序列长度dtype模型激活/权重主精度实际恒为 bf16见下文 FP8 说明tie_word_embeddings是否共享 embedding 与 lm_head 权重4.2 Attention 字段NoPE GQAnum_attention_heads/num_key_value_headsQ 头数与 KV 头数GQAattention_head_dim注意力头维度attention_bias默认False无 bias。resolve_attention_head_dim()按照 HF 参考NemotronHAttention的语义解析头维度优先取head_dim其次attention_head_dim最后回退到hidden_size // num_attention_heads——这是为了与参考实现逐位对齐。4.3 MLP 字段intermediate_sizeup/down 投影中间维度mlp_hidden_act默认relu2mlp_bias默认False。4.4 MoE 字段仅 Nemotron-3 的E混合层num_experts默认 0、num_experts_per_tok、moe_intermediate_size、moe_shared_expert_intermediate_size、routed_scaling_factor默认 1.0、norm_topk_prob默认 True。这些字段从 HF 配置的n_routed_experts、num_experts_per_tok等键读取并用getattr(..., 0)守卫——纯稠密的 4B/8B 变体没有这些键保持默认值即可不受影响。代码注释特别提醒_tok/_token的命名差异是刻意保留的num_experts_per_tok镜像 HF 配置键名再喂给 MAX 侧的num_experts_per_token参数。4.5 Mamba-2 mixer 字段字段含义mamba_num_heads/mamba_head_dimSSM 头数与每头维度n_groupsSSD scan 的分组数ssm_state_sizeSSM 状态维度 dstateconv_kerneldepthwise 卷积核大小 Kchunk_sizeSSD chunked scan 的 chunk 大小use_conv_bias默认 Truemamba_proj_bias默认 Falsetime_step_limit默认(0.0, inf)两个关键的派生维度属性mamba_intermediate_size mamba_num_heads * mamba_head_dimconv_dim mamba_intermediate_size 2 * n_groups * ssm_state_size即 hidden B C 三段拼接宽度mamba_in_proj_out mamba_intermediate_size conv_dim mamba_num_heads即融合in_proj的完整输出宽度[gate | hidden_states_B_C | dt]。4.6 FP8 层集合字段fp8_mamba_layers、fp8_mlp_layers、fp8_moe_layers三个set[int]分别记录哪些层号上的 in/out_proj、MLP up/down_proj、MoE 专家投影被量化到 FP8is_fp8为汇总标志。它们由populate_fp8_layers(state_dict)根据检查点中的weight_scale键反推——一个 Linear 是 FP8 当且仅当它在检查点中存在weight_scale这恰好是 modeloptexclude_modules列表的精确补集详见第六节。五、四大混合器的实现原理nemotron_h.py定义了五个 nn.Module 层类其中NemotronHBlock按kind分派到三种 mixermamba/attention/moe其中mlp与moe都走 MLP 家族而NemotronH是完整解码器。5.1NemotronHMLP非门控 relu2 稠密层def __call__(self, x: TensorValue) - TensorValue: return self.down_proj(_relu2(self.up_proj(x)))其中_relu2(x) relu(x) ** 2。up_proj维度为hidden → intermediatedown_proj为intermediate → hidden二者都支持mlp_bias与 FP8quant_config有quant_config时权重以float8_e4m3fn存储即_weight_dtype()的返回值。5.2NemotronHExpertMLP与NemotronHMoE128 专家 top-6 1 共享专家Nemotron-330B-A3B 混合体的 MoE 有两大特色非门控专家NemotronHExpertMLP只构建up_proj/down_proj没有gate_proj。它通过is_shardingTrue初始化MLP基类来跳过门控投影的构建并重写sharding_strategy与shard()——因为基类会去分片不存在的gate_proj。张量并行时up_proj按 rowwise 分片、down_proj按 columnwise 分片。Sigmoid top-k 路由器NemotronHMoEGate是 DeepSeek 风格的路由器——sigmoid 门控得分、加性的e_score_correction_bias仅用于选择专家而权重使用加偏前的得分。由于n_group topk_group 1分组受限方案退化为普通 top-k因此只用ops.top_kops.gather二者在 Apple/Metal 上都有原生分支刻意避开了moe_router_group_limited其 warp 集合的WARP_SIZE % group_size约束在 128 专家、n_group 1时失败。NemotronHMoE重写了gate_up_proj属性非门控专家的权重栈只堆叠 up 投影形状为[num_experts, moe_intermediate_size, hidden]。其__call__分两条路径bf16 路径无quant_config直接委托基类MoE的分组矩阵乘路由FP8 权重压缩路径W8A16专家权重以float8_e4m3fn存储送入同一 dtype 泛型的分组 matmulnaive kernel 在加载时将 E4M3 权重拓宽为 fp32 参与累加每个专家的标量weight_scale作为 matmul 后按行的精确去量化因子折叠标量可提出求和符号因此该折叠是精确的而非近似。共享专家则走稠密 FP8 Linear 路径。5.3NemotronHAttentionNoPE GQA关键实现点nemotron_h.py融合 QKV一个qkv_projmatmul 输出q_dim 2*kv_dim宽度再ops.split成 q | k | v。权重适配器把检查点中分离的 q/k/v 权重按 q, then k, then v 顺序拼接成qkv_proj.weight。无旋转K/V 直接store_k_cache_ragged/store_v_cache_ragged写入分页缓存不做任何 RoPE 处理。FP8 KV 缓存支持当kv_params.is_fp8_kv_dtype时q/k/v 先 cast 到缓存 dtype 再入缓存FP8 flash attention 要求 query 与缓存 dtype 一致输出再转回激活 dtype 进入o_proj。scale 采用sqrt(1/head_dim)mask 为CAUSAL_MASK核心计算是flash_attention_raggedragged 前缀注意力。保持 bf16注意力投影始终是 bf16——4B FP8 检查点将其排除在量化之外8B Reasoning 检查点的每张量 FP8 q/k/v/o 由权重适配器在加载时去量化回 bf16。5.4NemotronHMamba2Mixer融合 in_proj 就地状态池这是全模块最复杂的部分其 docstring 声称与 HFNemotronHMamba2Mixer逐操作对齐融合in_proj一个 matmul 输出[gate(intermediate) | hidden_states_B_C(conv_dim) | dt(nheads)]。检查点只有单个in_proj.weight带单一每张量 FP8weight_scale/input_scale融合 FP8 matmul 与三个复刻同一共享 scale 的 matmul 数值等价对精度无影响。实现上有一个关键技巧由于 fused matmul 的行 stride如 17504会让 strided 的gate视图破坏下游 gated group-RMSNorm 的归约对齐known-limitations/strided-split-misaligns-gpu-group-reduce代码先把整个 fused 输出 cast 到 fp32 再 split4 字节 stride 让归约保持对齐hidden_BC/dt则从原始 bf16 输出上 split它们喂给 conv/SSD kernel能容忍 split-view stride。Depthwise SiLU 卷积causal_conv1d_varlen_fwd在channels_lastTrue下处理[N, conv_dim]以slot_idx[b]为槽位就地读写conv_pool消除了卷积两侧的[conv_dim, N]转置注释称这是 prefill 粘合 kernel 的主要开销4k-token 请求在 B200 上约 9.5ms。SSD chunked scanmamba2_ssd_chunk_scan_varlen_fwd_inplace直接从ssm_pool[slot_idx[b]]读初始状态scan 结束后把终态写回同一槽位——图侧完全不需要 gather/scatter_nd/buffer_store 的整池往返注释称这一就地化消除了 B200 上约 30% 的 decode 墙钟时间。SSD kernel 同时服务 prefill 与 decodedecode 即 seqlen-1 的序列。Gated group RMSNorm_gated_group_rmsnorm用单个融合 kernel 复现 HFZamba2RMSNormGated的norm_before_gateFalse语义fp32 中做 silu-gate、对group_size做 group RMSNorm、乘 fp32 norm weight替换原本会低化成 3~4 次串行 GPU 分派的链式操作。5.5NemotronHBlock与完整解码器NemotronHNemotronHBlock是 pre-norm 残差块__call__被设计为不可直接调用抛出RuntimeError块的分派发生在NemotronH.__call__中。NemotronH的完整前向流程为embed → 逐块循环mamba 取conv_pools[mamba_i]/ssm_pools[mamba_i]attention 取kv_collections[0]并传入顺序 KV 索引moe/mlp 无状态直传→ 残差累加 →logits_postprocessfinal RMSNorm lm_head。input_types()定义了语言图的输入顺序tokens, input_row_offsets, return_n_logits, *kv_inputs, slot_idx, *conv_pools, *ssm_pools, has_initial_state。其中 conv 池为模型 dtype 的可变缓冲[max_slots, conv_dim, conv_kernel-1]SSM 池为_ssm_state_dtype()可变缓冲[max_slots, nheads, head_dim, dstate]has_initial_state为[batch]bool全新 prefill 为空decode 为全 True。六、FP8 量化按模块生效的 per-tensor staticNemotron-H 的 FP8 路径与常见的模型级量化不同是**按模块per-module**生效的build_fp8_quant_config()model_config.py扫描检查点只要存在float8_e4m3fn权重就返回一个QuantConfig其input_scale/weight_scale均为ScaleGranularity.TENSOR、ScaleOrigin.STATIC、dtype fp32格式为COMPRESSED_TENSORS_FP8。它刻意绕开通用的mlp_quantized_layers/attn_quantized_layers机制那些硬编码的self_attn/mlp.{gate,up,down}命名不匹配 Nemotron 的backbone.layers.{i}.mixer.*。from_hf()有一个反直觉但关键的处理即使解析出的编码是float8_e4m3fn模型 dtype 仍强制为bfloat16——FP8 只作用于特定 Linear若把 embedding 等全部声明为 fp8 会因 dtype 不匹配而无法加载 bf16 检查点张量。FP8 在NemotronH构造时按层分配仅当层号落在config.fp8_mamba_layers/fp8_mlp_layers/fp8_moe_layers内时该块的 mixer 才拿到quant_configattention、conv1d、各类 norm 与lm_head始终留在 bf16。weight_adapters.py补充了 FP8 检查点的精确行为F8_E4M3 权重原样保留weight_scale/input_scalecast 到 fp32被排除的模块lm_head、第 [11,16,23,31] 层的 mamba in/out_proj、所有 conv1d因无 scale 张量而保持 bf16。8B Reasoning 检查点的每张量 FP8 注意力投影则在加载时去量化回 bf16fp8_e4m3fn_to_float32 应用标量weight_scale其k_scale/v_scale等 KV 缓存 scale 被消费/丢弃。6.1 FP8 KV 缓存默认规则construct_kv_params()实现了一条参考配置对齐规则当解析出的编码为float8_e4m3fn且用户未显式指定kv_cache_format时KV 缓存 dtype 默认取float8_e4m3fn对齐 vLLM 的--kv-cache-dtype fp8参考显式覆盖与非 FP8 模型则保留解析出的 dtype。同时KV 缓存只为 Attention 层分配——其num_layers参数是模式中*的个数而非num_hidden_layers。该规则由 test_nemotron_h_fp8_kv_config.py 在纯 CPU 上验证无 GPU 即可运行。七、权重适配器从NemotronHForCausalLM到 MAX 命名空间weight_adapters.py 的convert_nemotron_h_state_dict完成命名映射与 dtype/形状修正前缀改写剥离backbone.前缀backbone.embeddings→embed_tokensbackbone.norm_f→norm_fbackbone.layers.N→blocks.N同时兼容 transformers 参考实现的model.前缀。Mamba 相关conv1d 权重保持三维[dim, 1, K]A_log/D/dt_bias每头标量与 gatednorm.weightcast 到 fp32。MoE 相关路由门控权重mixer.gate.weight→mixer.gate.gate_score.weight对齐 MAXMoEGate的嵌套结构e_score_correction_biascast 到 fp32路由/共享专家 up/down 投影 1:1 映射。融合 QKVblocks.{i}.mixer.{q,k,v}_proj.weight按 q、k、v 顺序拼接为qkv_proj.weighto_proj独立保留。八、状态池与推理生命周期conv/SSM 的 slot 管理由于 Mamba-2 的循环状态无法从 token 前缀重建Nemotron-H 在推理时必须显式管理每请求的 SSM 状态。NemotronHStateCachestate_cache.py在 GPU 上预分配两类可变缓冲池conv_pool[l][max_slots, conv_dim, conv_kernel-1]模型 dtype由causal_conv1d_varlen_fwd在slot_idx[batch_item]槽位就地改写ssm_pool[l][max_slots, nheads, head_dim, dstate]fp32Apple GPU 上为 bf16——仅存储精度scan 始终在 fp32 寄存器中累加由mamba2_ssd_chunk_scan_varlen_fwd_inplace同槽位就地读写。生命周期与 qwen3_5 的GatedDeltaNetStateCache同构claim(request_id)注册请求并清零槽位 →slot_idx_for()写入槽位索引 →model.execute消费池与索引两个 inplace kernel 直接改池图没有状态输出→release(request_id)释放槽位。NemotronHModel还实现了SupportsSSMStateWarmup接口release_warmup_state()在设备图捕获graph capture暖机的每个(batch_size, cache_length)探测后释放暖机槽位防止池被暖机扫描耗尽。此外_has_initial_state_prealloc恒为全 True——请求槽位在 claim 时清零因此加载零初始状态与从零开始的 prefill 等价无需单独的 prefill/decode 双图。由于 SSM 循环状态不可从 token 前缀重建arch.py明确要求required_arguments{enable_prefix_caching: False}禁用前缀缓存并声明multi_gpu_supportedFalse单 GPU。这是理解该架构部署边界的关键约束。九、架构注册、推理与工具调用解析arch.py中的nemotron_h_arch SupportedArchitecture(...)完成架构注册nameNemotronHForCausalLMtaskTEXT_GENERATION示例仓库nvidia/NVIDIA-Nemotron-3-Nano-4B-FP8与nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-FP8默认权重格式 safetensorsbf16 为默认编码支持float8_e4m3fn复用 Qwen3.5 的解析器Nemotron-3 的聊天模板是 Qwen 格式——生成提示中预填think\n隐式开启推理、显式/think关闭在之前的助手轮次回填think/think工具调用渲染为tool_call/function.../parameter...块。因此reasoning_parserqwen3_5、tool_parserqwen3_5并在此处导入 qwen3_5 的解析模块以完成惰性注册。NemotronHTokenizertokenizer.py实现ReasoningPipelineTokenizer协议在初始化时通过resolve_single_special_token解析think//think的 token id供重叠管线overlap pipeline的思考模式追踪直接读取。十、模块级测试与验证仓库为 Nemotron-H 提供了多个集成测试佐证上述实现test_nemotron_h_fp8_kv_config.pyCPU 上验证 FP8 KV 缓存 dtype 选择规则test_nemotron_h_state_warmup.py 与 test_attention_fp8_kv_gpu.py状态暖机与 GPU 上 FP8 KV 注意力路径。从代码注释可见验证严格遵循先最小化冒烟mini-smoke再完整 serve的顺序——例如融合 in_proj 的 stridedgate对齐问题先在真实几何fused 17504、group_size 960的最小化 GPU 复现上确认无CUDA_ERROR_MISALIGNED_ADDRESS再在 FP8 量化路径上完整验证后才声称可服务。结语Nemotron-H 在 MAX 中的落地是架构创新 × 工程优化的典型样本hybrid_override_pattern用四个字符表达层间混合NoPE 设计把位置信息让渡给 SSMMamba-2 的 conv/SSM 状态以就地 slot 池的形式融入单一语言图省去约 30% 的 decode 墙钟与约 9.5ms 的 prefill 转置开销FP8 则以按模块、W8A16 权重压缩、标量 scale 精确折叠的方式与 bf16 主体共存。本文涉及的配置字段、混合模式规则与状态池生命周期均可直接作为在 MAX 中加载、编译与推理 Nemotron-H 系列检查点4B FP8 / 30B-A3B FP8 / 8B Reasoning的实践参考。【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门