MAX 中 GLM-5.2 (DeepSeek-V3.2 Sparse) 的统一 MTP 推测解码架构解析
2026/9/12 16:08:18 网站建设 项目流程

MAX 中 GLM-5.2 (DeepSeek-V3.2 Sparse) 的统一 MTP 推测解码架构解析

【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo

本文讲解 MAX 平台中UnifiedMTPGlm5_2模块的设计与实现。该模块将 DeepSeek-V3.2 稀疏 MoE 目标模型和单层稀疏 NextN Draft 模型、贪心拒绝采样及 prefill shift 融合为单一可编译图结构,是 GLM-5.2 系列(zai-org/GLM-5.2)在 MAX Pipeline 中进行推测解码的核心架构组件。读完本文,你将理解其双层 KV 缓存设计、index_share_for_mtp_iteration优化策略、权重量化适配逻辑以及完整的编译与运行时流程。

架构概述

UnifiedMTPGlm5_2定义在 max/python/max/pipelines/architectures/unified_mtp_glm5_2/unified_mtp_glm5_2.py 中,类似于UnifiedMTPDeepseekV3之于 DeepSeek-V3 的关系,是 V3.2 sparse 对应的 MTP(Multi-Token Prediction)版本。与标准 V3.2 的 MTP 有两个关键的结构差异:

  • 双层稀疏 KV 缓存:目标网络(target)和草稿网络(draft)均使用稀疏 MLA(lightning indexer),因此各自携带一对{mla, indexer}KV 缓存,而非单一的 MLA 缓存。
  • index_share_for_mtp_iteration:草稿网络的 lightning indexer 仅在 step 0 执行一次 top-k 选择,之后各 draft step 通过 gather 已被接受的 token 位置来复用该选择结果,避免重复计算。

该模块继承自Module,将 token merging、V3.2 target 前向传播、贪心拒绝采样和稀疏 draft 前向传播融合为一个端到端的图。

模块结构与注册

架构注册

在 arch.py 中,该架构注册为:

unified_mtp_glm5_2_arch = SupportedArchitecture( name="UnifiedMTPGlmMoeDsaForCausalLM", task=PipelineTask.TEXT_GENERATION, example_repo_ids=["zai-org/GLM-5.2-FP8"], default_encoding="float8_e4m3fn", supported_encodings={"float4_e2m1fnx2", "float8_e4m3fn", "bfloat16"}, multi_gpu_supported=True, pipeline_model=UnifiedMTPGlm5_2Model, tokenizer=GlmTokenizer, context_type=TextContext, default_weights_format=WeightsFormat.safetensors, weight_adapters={WeightsFormat.safetensors: convert_with_mtp_state_dict}, supports_empty_batches=True, requires_max_batch_context_length=True, config=Glm5_1Config, memory_planner=DeepseekV3_2MemoryPlanner, batching=UnifiedMTPGlm5_2BatchProcessor, tool_parser="glm45", reasoning_parser="glm45", default_structured_output_backend="xgrammar", default_structured_output_any_whitespace=True, )

关键配置说明:

配置项说明
nameUnifiedMTPGlmMoeDsaForCausalLM架构标识名,用于 Pipeline 自动匹配
example_repo_ids["zai-org/GLM-5.2-FP8"]在 docs/max/models.mdx 的模型表格中,GLM-5.1 条目下同样列出了zai-org/GLM-5.2zai-org/GLM-5.2-FP8等 Model ID
default_encodingfloat8_e4m3fn默认权重量化编码
supported_encodingsfloat4_e2m1fnx2,float8_e4m3fn,bfloat16支持的量化精度
multi_gpu_supportedTrue支持多 GPU 分布式部署
default_weights_formatsafetensorsHuggingFace safetensors 格式
weight_adaptersconvert_with_mtp_state_dict权重 key 映射适配器

模块依赖

依据 BUILD.bazel 的依赖列表,该模块的核心依赖包括:

  • //max/python/max/nn— 神经网络层基类
  • //max/python/max/pipelines/architectures/deepseekV3— DeepSeek-V3 权重映射
  • //max/python/max/pipelines/architectures/deepseekV3_2— V3.2 目标模型
  • //max/python/max/pipelines/architectures/deepseekV3_2_nextn— NextN Draft 模型
  • //max/python/max/pipelines/architectures/glm5_1— GLM-5.1 基类(Tokenizer、ReasoningParser、ToolParser)
  • //max/python/max/pipelines/speculative— 推测解码配置与统一图操作

UnifiedMTPGlm5_2 前向传播详解

__init__与配置初始化

UnifiedMTPGlm5_2.__init__接受三个核心配置:

def __init__( self, config: DeepseekV3_2Config, # 目标模型配置 draft_config: DeepseekV3_2NextNConfig | None, # 草稿模型配置 speculative_config: SpeculativeConfig | None, # 推测解码配置 enable_structured_output: bool = False, ) -> None

在初始化中:

  1. num_draft_steps:从speculative_config.num_speculative_tokens读取每次生成的草稿 token 数量,默认为 1。
  2. AcceptanceSampler:初始化接受采样器,支持宽松接受(relaxed acceptance)——在 thinking 阶段,若use_relaxed_acceptance_for_thinking=True,则relaxed_topkrelaxed_delta参数放宽拒绝条件,加速 thinking 阶段的 token 吞吐。该采样器定义在 max/nn/sampling/rejection_sampler.py 中。
  3. target = DeepseekV3_2(config):初始化稀疏 V3.2 目标模型,emit_last_token_logits = False以抑制最后一个 token 的 logits 输出。
  4. merger = RaggedTokenMerger:初始化 ragged token 合并器,用于拼接用户输入 token 和草稿 token。
  5. draft = DeepseekV3_2NextN(draft_config):初始化单层稀疏 NextN 草稿模型。

__call__前向流程

前向传播分为 Step 0 和后续步骤,整体流程如下:

阶段 1:Token 合并
merged_tokens, merged_offsets, host_merged_offsets = merge_tokens_and_host_offsets( self.merger, tokens, input_row_offsets, draft_tokens, host_input_row_offsets, )

将用户输入的tokens和预先准备的draft_tokens按 batch 合并,同时合并对应的 row offsets。这一步由RaggedTokenMerger完成。

阶段 2:目标模型前向
target_outputs = self.target( merged_tokens, signal_buffers, target_mla_kv, target_indexer_kv, # 目标模型的双层 KV 缓存 return_n_logits, merged_offsets, host_merged_offsets, data_parallel_splits, batch_context_lengths, ep_inputs, )

目标模型使用两级缓存

  • target_mla_kv:MLA(Multi-head Latent Attention)的 KV 缓存
  • target_indexer_kv:Lightning Indexer 的 KV 缓存(用于稀疏注意力)

返回(logits, offsets, hs_0..hs_{n-1}),其中 hidden states 为 ALL_NORMALIZED 模式(已在层内完成最终归一化)。

阶段 3:拒绝采样与 Bitmask
effective_bitmasks = apply_overlap_bitmask( pinned_bitmask, wait_payload, device_bitmask_scratch, num_steps=draft_tokens.shape[1], device=device0, ) num_accepted_draft_tokens, recovered, bonus, next_tokens = accept_and_pick_next_tokens( self.acceptance_sampler, draft_tokens, logits, seed=seed[0], temperature=temperature, top_k=top_k, max_k=max_k, top_p=top_p, min_top_p=min_top_p, in_thinking_phase=in_thinking_phase, token_bitmasks=effective_bitmasks, )
  • 如果enable_structured_output=True,则通过pinned_bitmaskwait_payloaddevice_bitmask_scratch对特定 token 位置施加掩码约束(结构化输出约束)。
  • accept_and_pick_next_tokens执行贪心拒绝采样,返回:
    • num_accepted_draft_tokens:被接受的草稿 token 数量
    • recovered:被拒绝位置恢复的 token
    • bonus:bonus token(从目标分布中额外采样)
    • next_tokens:下一轮的输入 token
阶段 4:Draft Step 0(带 index_share 初始化)
self.draft.return_hidden_states = ReturnHiddenStates.ALL self.draft.return_logits = ReturnLogits.VARIABLE self.draft.emit_last_token_logits = False # 抑制 lm_head 的大词汇表投影 draft_outputs = self.draft( shifted_corrected, hidden_states, signal_buffers, draft_mla_kv, draft_indexer_kv, return_n_logits, merged_offsets_per_dev, host_merged_offsets, data_parallel_splits, batch_context_lengths, ep_inputs, prev_topk_indices=None, reuse_prev_topk=False, )

Step 0 的特殊之处:

  • 设置return_hidden_states = ALL,返回所有层的 hidden states(draft 仅 1 层,返回全部)
  • 设置return_logits = VARIABLE,返回每个 token 位置的 logits(用于计算 draft argmax)
  • 设置emit_last_token_logits = False抑制最后一个 token 的 lm_head 投影(该位置不参与 draft 自回归)
  • prev_topk_indices=None/reuse_prev_topk=False:在此步骤中计算lightning indexer 的 top-k,并保存供后续复用

Step 0 输出布局为:(logits, offsets, hs[n], topk[n])

阶段 5:Draft 后续步骤(复用 top-k)

在进入循环之前,切换 draft 的返回模式:

self.draft.return_hidden_states = ReturnHiddenStates.LAST_PER_DEVICE self.draft.return_logits = ReturnLogits.LAST_TOKEN self.draft.emit_last_token_logits = True

同时,切换 draft MLA 缓存的分发元数据:

draft_mla_kv = [ replace(kv, max_prompt_length=one, attention_dispatch_metadata=kv.draft_attention_dispatch_metadata, mla_num_partitions=kv.draft_mla_num_partitions, ) for kv in draft_mla_kv ]

index_share_for_mtp_iteration核心优化:在 Draft Step 0 中已经计算了 lightning indexer 的 top-k 选择结果step0_topk,后续步骤通过gather_accepted_hidden_states收集已接受位置的 top-k 索引后,在迭代中作为prev_topk_indices传入,并设置reuse_prev_topk=True跳过重复的 top-k 计算

reuse_topk = gather_accepted_hidden_states( step0_topk, merged_offsets=merged_offsets, merged_offsets_per_dev=merged_offsets_per_dev, num_accepted=num_accepted_draft_tokens, num_draft_tokens=draft_tokens.shape[1], data_parallel_degree=..., data_parallel_splits=..., signal_buffers=..., device=device0, split_prefix="mtp_topk", ) for step in range(1, self.num_draft_steps): step_outputs = self.draft( next_draft_tokens, draft_hs, signal_buffers, step_mla_kv, draft_indexer_kv, draft_return_n_logits, decode_offsets_per_dev, host_decode_offsets, data_parallel_splits, batch_context_lengths, ep_inputs, prev_topk_indices=reuse_topk, # 复用 step 0 的 top-k reuse_prev_topk=True, # 跳过重复计算 split_prefix=f"mtp_draft_step{step}", )

每个后续步骤的缓存mla_cache_lengths_per_dev会递增+1,而 indexer 缓存长度保持不变(因为 indexer 仅在 step 0 参与)。

阶段 6:输出组装
if len(all_draft_tokens) > 1: new_token = ops.stack(all_draft_tokens, axis=-1) else: new_token = ops.unsqueeze(all_draft_tokens[0], -1) return (num_accepted_draft_tokens, next_tokens, new_token)

最终返回三元组(被接受的草稿数量, 下一轮主 token, 拼接的新 draft tokens)

PipelineModel 编译与运行时

UnifiedMTPGlm5_2Model

该 PipelineModel 定义在 model.py 中,继承自_UnifiedSpecDecodeModelMixinGlm5_1Model

权重加载

_load_state_dict从 checkpoint 解析target.*draft.*前缀:

self._draft_state_dict = { k[len("draft."):]: v for k, v in raw_state_dict.items() if k.startswith("draft.") } # 某些 checkpoint 共享 shared_head_norm 与 final norm if ("shared_head_norm.weight" not in self._draft_state_dict and "target.norm.weight" in raw_state_dict): self._draft_state_dict["shared_head_norm.weight"] = raw_state_dict["target.norm.weight"]
KV 缓存树

_create_model_config构建嵌套的{target: {mla, indexer}, draft: {mla, indexer}}KV 缓存树。draft 的缓存仅有 1 层(num_layers=1),且mlaindexer分开管理:

draft_kv = MultiKVCacheParams.from_params({ "mla": replace(target_mla_params, num_layers=1), "indexer": replace(target_indexer_params, num_layers=1), }) self.kv_params = MultiKVCacheParams.from_params( {"target": target_kv, "draft": draft_kv} )
分布式专家并行(EP)

_init_distributed_runtime处理专家并行初始化。对于 NVFP4 量化检查点,其 MTP 层的 routed experts 以 bf16 精度存储(无.weight_scale),因此 draft 的 EP 分发精度必须从 NVFP4 提升到 bf16:

draft_moe_dispatches_bf16 = not _subtree_quantized( self._draft_state_dict, ".mlp.experts." ) if draft_moe_dispatches_bf16: ep_alloc_config = replace(model_config.ep_config, dispatch_dtype=DType.bfloat16, dispatch_quant_config=None, fused_shared_expert=model_config.n_shared_experts == 1, )
图编译

_build_graph_for_compile方法:

  1. 实例化UnifiedMTPGlm5_2模型
  2. 权重共享:将 draft 的embed_tokenslm_head别名为 target 的对应层(strict=False加载时跳过已共享的 key)
  3. 通过nn_model.input_types(kv_params)构建图输入类型签名
  4. Graph上下文管理器构建完整的glm5_2_with_mtp_graph计算图
  5. 从输入中解包四组 KV 缓存(target_mlatarget_indexerdraft_mladraft_indexer
  6. 提取采样超参数(seedtemperaturetop_kmax_ktop_pmin_top_p)和in_thinking_phase标志
  7. 调用nn_model(...)完成前向传播,绑定图输出

输入 batching

UnifiedMTPGlm5_2BatchProcessor定义在 batch_processor.py 中,继承自DeepseekV3BatchProcessor。它在标准 batch 输入基础上扩展了draft_tokens字段(初始化为None),由 overlap pipeline 在实际执行时填充。

输入结构定义在UnifiedMTPGlm5_2Inputs中(继承UnifiedSpecDecodeInputsDeepseekV3Inputs),其buffers属性除了父类输入外,额外包含in_thinking_phase标志位。

权重适配器

convert_with_mtp_state_dict

定义在 weight_adapters.py 中,负责将 HuggingFace safetensors 格式的 checkpoint 转换为 MAX 内部格式。权重 key 映射规则如下:

  1. 常规层:通过DEEPSEEK_SAFETENSOR_MAP(来自deepseekV3.weight_adapters)将 HuggingFace key 转换为 MAX key
  2. 丢弃 KV 缩放因子:跳过以.k_scale.v_scale结尾的 key(MAX 从独立配置路径读取 KV 缓存缩放)
  3. MTP 层重映射:MTP 层在 checkpoint 中位于layers.{num_hidden_layers}.索引处,映射到draft.*。特定子模块映射如下:
Checkpoint Key 前缀MAX Key 路径说明
layers.N.shared_head.norm.draft.shared_head_norm.共享头部归一化
layers.N.enorm.draft.enorm.专家归一化
layers.N.hnorm.draft.hnorm.头部归一化
layers.N.eh_proj.draft.eh_proj.专家头部投影
layers.N.其他draft.decoder_layer.解码器层(self-attention / MLP)
layers.N.前缀target.*目标模型权重
  1. 权重共享embed_tokensshared_head.headlm_head)仅在 target 前缀下保存一份,draft 通过模块别名共享,state_dict()自动去重。

Draft 配置创建

_create_draft_config方法(model.py)从 draft 的 state_dict 推导配置:

  1. 验证 NextN 层存在decoder_layer.self_attn.kv_a_layernorm.weight
  2. 基于 target 的基础配置创建DeepseekV3_2NextNConfig
  3. 设置indexer_types = [](空调度):让单层 MTP 的 indexer 保持满计算,避免引用 target 的 78 层调度方案
  4. 检测 NVFP4 量化的子树范围(.mlp.experts..self_attn.),将 MTP 层索引添加到正确的量化层集合中
  5. 如果 draft 的 routed experts 未量化(NVFP4 下 MTP 层为 bf16),则修正 EP 分发配置为bfloat16

使用方式

要使用该架构在 MAX 中加载 GLM-5.2-FP8 模型并启用 MTP 推测解码,需要在 Pipeline 配置中指定:

python -m max.pipelines.run \ --model-path zai-org/GLM-5.2-FP8 \ --speculative-config.num-speculative-tokens 3 \ --speculative-config.use-relaxed-acceptance-for-thinking \ --speculative-config.relaxed-topk 5 \ --speculative-config.relaxed-delta 0.1

关键推测解码配置项说明(源自 max/python/max/pipelines/speculative/config.py):

配置项类型默认值说明
num_speculative_tokensint \| NoneNone每次生成的草稿 token 数量
num_speculative_tokens_per_batch_sizelist[VerifyWidthRange] \| NoneNone按 batch size 分级的草稿数量调度
synthetic_acceptance_ratefloat \| NoneNone合成接受率(0.0~1.0),用于绕过真实分布模拟
use_relaxed_acceptance_for_thinkingboolFalsethinking 阶段是否启用宽松接受
relaxed_topkint宽松接受的 top-k 范围(需>= 1
relaxed_deltafloat宽松接受的 delta 阈值(0.0~1.0)

支持的量化编码可通过--dtype float8_e4m3fn(默认)、--dtype bfloat16--dtype float4_e2m1fnx2指定。多 GPU 训练可通过--tensor-parallel-size N启用。

该架构已默认注册到 MAX Pipeline 中,通过架构名UnifiedMTPGlmMoeDsaForCausalLM自动匹配zai-org/GLM-5.2-FP8等模型 ID,无需额外的手动注册步骤。

【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询