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, )关键配置说明:
| 配置项 | 值 | 说明 |
|---|---|---|
name | UnifiedMTPGlmMoeDsaForCausalLM | 架构标识名,用于 Pipeline 自动匹配 |
example_repo_ids | ["zai-org/GLM-5.2-FP8"] | 在 docs/max/models.mdx 的模型表格中,GLM-5.1 条目下同样列出了zai-org/GLM-5.2、zai-org/GLM-5.2-FP8等 Model ID |
default_encoding | float8_e4m3fn | 默认权重量化编码 |
supported_encodings | float4_e2m1fnx2,float8_e4m3fn,bfloat16 | 支持的量化精度 |
multi_gpu_supported | True | 支持多 GPU 分布式部署 |
default_weights_format | safetensors | HuggingFace safetensors 格式 |
weight_adapters | convert_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在初始化中:
num_draft_steps:从speculative_config.num_speculative_tokens读取每次生成的草稿 token 数量,默认为 1。AcceptanceSampler:初始化接受采样器,支持宽松接受(relaxed acceptance)——在 thinking 阶段,若use_relaxed_acceptance_for_thinking=True,则relaxed_topk和relaxed_delta参数放宽拒绝条件,加速 thinking 阶段的 token 吞吐。该采样器定义在 max/nn/sampling/rejection_sampler.py 中。target = DeepseekV3_2(config):初始化稀疏 V3.2 目标模型,emit_last_token_logits = False以抑制最后一个 token 的 logits 输出。merger = RaggedTokenMerger:初始化 ragged token 合并器,用于拼接用户输入 token 和草稿 token。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_bitmask、wait_payload和device_bitmask_scratch对特定 token 位置施加掩码约束(结构化输出约束)。 accept_and_pick_next_tokens执行贪心拒绝采样,返回:num_accepted_draft_tokens:被接受的草稿 token 数量recovered:被拒绝位置恢复的 tokenbonus: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 中,继承自_UnifiedSpecDecodeModelMixin和Glm5_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),且mla和indexer分开管理:
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方法:
- 实例化
UnifiedMTPGlm5_2模型 - 权重共享:将 draft 的
embed_tokens和lm_head别名为 target 的对应层(strict=False加载时跳过已共享的 key) - 通过
nn_model.input_types(kv_params)构建图输入类型签名 - 用
Graph上下文管理器构建完整的glm5_2_with_mtp_graph计算图 - 从输入中解包四组 KV 缓存(
target_mla、target_indexer、draft_mla、draft_indexer) - 提取采样超参数(
seed、temperature、top_k、max_k、top_p、min_top_p)和in_thinking_phase标志 - 调用
nn_model(...)完成前向传播,绑定图输出
输入 batching
UnifiedMTPGlm5_2BatchProcessor定义在 batch_processor.py 中,继承自DeepseekV3BatchProcessor。它在标准 batch 输入基础上扩展了draft_tokens字段(初始化为None),由 overlap pipeline 在实际执行时填充。
输入结构定义在UnifiedMTPGlm5_2Inputs中(继承UnifiedSpecDecodeInputs和DeepseekV3Inputs),其buffers属性除了父类输入外,额外包含in_thinking_phase标志位。
权重适配器
convert_with_mtp_state_dict
定义在 weight_adapters.py 中,负责将 HuggingFace safetensors 格式的 checkpoint 转换为 MAX 内部格式。权重 key 映射规则如下:
- 常规层:通过
DEEPSEEK_SAFETENSOR_MAP(来自deepseekV3.weight_adapters)将 HuggingFace key 转换为 MAX key - 丢弃 KV 缩放因子:跳过以
.k_scale或.v_scale结尾的 key(MAX 从独立配置路径读取 KV 缓存缩放) - 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.* | 目标模型权重 |
- 权重共享:
embed_tokens和shared_head.head(lm_head)仅在 target 前缀下保存一份,draft 通过模块别名共享,state_dict()自动去重。
Draft 配置创建
_create_draft_config方法(model.py)从 draft 的 state_dict 推导配置:
- 验证 NextN 层存在
decoder_layer.self_attn.kv_a_layernorm.weight - 基于 target 的基础配置创建
DeepseekV3_2NextNConfig - 设置
indexer_types = [](空调度):让单层 MTP 的 indexer 保持满计算,避免引用 target 的 78 层调度方案 - 检测 NVFP4 量化的子树范围(
.mlp.experts.和.self_attn.),将 MTP 层索引添加到正确的量化层集合中 - 如果 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_tokens | int \| None | None | 每次生成的草稿 token 数量 |
num_speculative_tokens_per_batch_size | list[VerifyWidthRange] \| None | None | 按 batch size 分级的草稿数量调度 |
synthetic_acceptance_rate | float \| None | None | 合成接受率(0.0~1.0),用于绕过真实分布模拟 |
use_relaxed_acceptance_for_thinking | bool | False | thinking 阶段是否启用宽松接受 |
relaxed_topk | int | — | 宽松接受的 top-k 范围(需>= 1) |
relaxed_delta | float | — | 宽松接受的 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),仅供参考