 的統(tǒng)一 MTP 推測(cè)解碼架構(gòu)解析)
MAX 中 GLM-5.2 (DeepSeek-V3.2 Sparse) 的統(tǒng)一 MTP 推測(cè)解碼架構(gòu)解析【免費(fèi)下載鏈接】mojoThe Modular Platform (includes MAX Mojo)項(xiàng)目地址: https://gitcode.com/GitHub_Trending/mo/mojo本文講解 MAX 平臺(tái)中UnifiedMTPGlm5_2模塊的設(shè)計(jì)與實(shí)現(xiàn)。該模塊將 DeepSeek-V3.2 稀疏 MoE 目標(biāo)模型和單層稀疏 NextN Draft 模型、貪心拒絕采樣及 prefill shift 融合為單一可編譯圖結(jié)構(gòu)是 GLM-5.2 系列zai-org/GLM-5.2在 MAX Pipeline 中進(jìn)行推測(cè)解碼的核心架構(gòu)組件。讀完本文你將理解其雙層 KV 緩存設(shè)計(jì)、index_share_for_mtp_iteration優(yōu)化策略、權(quán)重量化適配邏輯以及完整的編譯與運(yùn)行時(shí)流程。架構(gòu)概述UnifiedMTPGlm5_2定義在 max/python/max/pipelines/architectures/unified_mtp_glm5_2/unified_mtp_glm5_2.py 中類似于UnifiedMTPDeepseekV3之于 DeepSeek-V3 的關(guān)系是 V3.2 sparse 對(duì)應(yīng)的 MTPMulti-Token Prediction版本。與標(biāo)準(zhǔn) V3.2 的 MTP 有兩個(gè)關(guān)鍵的結(jié)構(gòu)差異雙層稀疏 KV 緩存目標(biāo)網(wǎng)絡(luò)target和草稿網(wǎng)絡(luò)draft均使用稀疏 MLAlightning indexer因此各自攜帶一對(duì){mla, indexer}KV 緩存而非單一的 MLA 緩存。index_share_for_mtp_iteration草稿網(wǎng)絡(luò)的 lightning indexer 僅在 step 0 執(zhí)行一次 top-k 選擇之后各 draft step 通過 gather 已被接受的 token 位置來復(fù)用該選擇結(jié)果避免重復(fù)計(jì)算。該模塊繼承自Module將 token merging、V3.2 target 前向傳播、貪心拒絕采樣和稀疏 draft 前向傳播融合為一個(gè)端到端的圖。模塊結(jié)構(gòu)與注冊(cè)架構(gòu)注冊(cè)在 arch.py 中該架構(gòu)注冊(cè)為unified_mtp_glm5_2_arch SupportedArchitecture( nameUnifiedMTPGlmMoeDsaForCausalLM, taskPipelineTask.TEXT_GENERATION, example_repo_ids[zai-org/GLM-5.2-FP8], default_encodingfloat8_e4m3fn, supported_encodings{float4_e2m1fnx2, float8_e4m3fn, bfloat16}, multi_gpu_supportedTrue, pipeline_modelUnifiedMTPGlm5_2Model, tokenizerGlmTokenizer, context_typeTextContext, default_weights_formatWeightsFormat.safetensors, weight_adapters{WeightsFormat.safetensors: convert_with_mtp_state_dict}, supports_empty_batchesTrue, requires_max_batch_context_lengthTrue, configGlm5_1Config, memory_plannerDeepseekV3_2MemoryPlanner, batchingUnifiedMTPGlm5_2BatchProcessor, tool_parserglm45, reasoning_parserglm45, default_structured_output_backendxgrammar, default_structured_output_any_whitespaceTrue, )關(guān)鍵配置說明配置項(xiàng)值說明nameUnifiedMTPGlmMoeDsaForCausalLM架構(gòu)標(biāo)識(shí)名用于 Pipeline 自動(dòng)匹配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 IDdefault_encodingfloat8_e4m3fn默認(rèn)權(quán)重量化編碼supported_encodingsfloat4_e2m1fnx2,float8_e4m3fn,bfloat16支持的量化精度multi_gpu_supportedTrue支持多 GPU 分布式部署default_weights_formatsafetensorsHuggingFace safetensors 格式weight_adaptersconvert_with_mtp_state_dict權(quán)重 key 映射適配器模塊依賴依據(jù) BUILD.bazel 的依賴列表該模塊的核心依賴包括//max/python/max/nn— 神經(jīng)網(wǎng)絡(luò)層基類//max/python/max/pipelines/architectures/deepseekV3— DeepSeek-V3 權(quán)重映射//max/python/max/pipelines/architectures/deepseekV3_2— V3.2 目標(biāo)模型//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— 推測(cè)解碼配置與統(tǒng)一圖操作UnifiedMTPGlm5_2 前向傳播詳解__init__與配置初始化UnifiedMTPGlm5_2.__init__接受三個(gè)核心配置def __init__( self, config: DeepseekV3_2Config, # 目標(biāo)模型配置 draft_config: DeepseekV3_2NextNConfig | None, # 草稿模型配置 speculative_config: SpeculativeConfig | None, # 推測(cè)解碼配置 enable_structured_output: bool False, ) - None在初始化中num_draft_steps從speculative_config.num_speculative_tokens讀取每次生成的草稿 token 數(shù)量默認(rèn)為 1。AcceptanceSampler初始化接受采樣器支持寬松接受relaxed acceptance——在 thinking 階段若use_relaxed_acceptance_for_thinkingTrue則relaxed_topk和relaxed_delta參數(shù)放寬拒絕條件加速 thinking 階段的 token 吞吐。該采樣器定義在 max/nn/sampling/rejection_sampler.py 中。target DeepseekV3_2(config)初始化稀疏 V3.2 目標(biāo)模型emit_last_token_logits False以抑制最后一個(gè) token 的 logits 輸出。merger RaggedTokenMerger初始化 ragged token 合并器用于拼接用戶輸入 token 和草稿 token。draft DeepseekV3_2NextN(draft_config)初始化單層稀疏 NextN 草稿模型。__call__前向流程前向傳播分為 Step 0 和后續(xù)步驟整體流程如下階段 1Token 合并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和預(yù)先準(zhǔn)備的draft_tokens按 batch 合并同時(shí)合并對(duì)應(yīng)的 row offsets。這一步由RaggedTokenMerger完成。階段 2目標(biāo)模型前向target_outputs self.target( merged_tokens, signal_buffers, target_mla_kv, target_indexer_kv, # 目標(biāo)模型的雙層 KV 緩存 return_n_logits, merged_offsets, host_merged_offsets, data_parallel_splits, batch_context_lengths, ep_inputs, )目標(biāo)模型使用兩級(jí)緩存target_mla_kvMLAMulti-head Latent Attention的 KV 緩存target_indexer_kvLightning Indexer 的 KV 緩存用于稀疏注意力返回(logits, offsets, hs_0..hs_{n-1})其中 hidden states 為 ALL_NORMALIZED 模式已在層內(nèi)完成最終歸一化。階段 3拒絕采樣與 Bitmaskeffective_bitmasks apply_overlap_bitmask( pinned_bitmask, wait_payload, device_bitmask_scratch, num_stepsdraft_tokens.shape[1], devicedevice0, ) num_accepted_draft_tokens, recovered, bonus, next_tokens accept_and_pick_next_tokens( self.acceptance_sampler, draft_tokens, logits, seedseed[0], temperaturetemperature, top_ktop_k, max_kmax_k, top_ptop_p, min_top_pmin_top_p, in_thinking_phasein_thinking_phase, token_bitmaskseffective_bitmasks, )如果enable_structured_outputTrue則通過pinned_bitmask、wait_payload和device_bitmask_scratch對(duì)特定 token 位置施加掩碼約束結(jié)構(gòu)化輸出約束。accept_and_pick_next_tokens執(zhí)行貪心拒絕采樣返回num_accepted_draft_tokens被接受的草稿 token 數(shù)量recovered被拒絕位置恢復(fù)的 tokenbonusbonus token從目標(biāo)分布中額外采樣next_tokens下一輪的輸入 token階段 4Draft 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_indicesNone, reuse_prev_topkFalse, )Step 0 的特殊之處設(shè)置return_hidden_states ALL返回所有層的 hidden statesdraft 僅 1 層返回全部設(shè)置return_logits VARIABLE返回每個(gè) token 位置的 logits用于計(jì)算 draft argmax設(shè)置emit_last_token_logits False抑制最后一個(gè) token 的 lm_head 投影該位置不參與 draft 自回歸prev_topk_indicesNone/reuse_prev_topkFalse在此步驟中計(jì)算lightning indexer 的 top-k并保存供后續(xù)復(fù)用Step 0 輸出布局為(logits, offsets, hs[n], topk[n])階段 5Draft 后續(xù)步驟復(fù)用 top-k在進(jìn)入循環(huán)之前切換 draft 的返回模式self.draft.return_hidden_states ReturnHiddenStates.LAST_PER_DEVICE self.draft.return_logits ReturnLogits.LAST_TOKEN self.draft.emit_last_token_logits True同時(shí)切換 draft MLA 緩存的分發(fā)元數(shù)據(jù)draft_mla_kv [ replace(kv, max_prompt_lengthone, attention_dispatch_metadatakv.draft_attention_dispatch_metadata, mla_num_partitionskv.draft_mla_num_partitions, ) for kv in draft_mla_kv ]index_share_for_mtp_iteration核心優(yōu)化在 Draft Step 0 中已經(jīng)計(jì)算了 lightning indexer 的 top-k 選擇結(jié)果step0_topk后續(xù)步驟通過gather_accepted_hidden_states收集已接受位置的 top-k 索引后在迭代中作為prev_topk_indices傳入并設(shè)置reuse_prev_topkTrue跳過重復(fù)的 top-k 計(jì)算。reuse_topk gather_accepted_hidden_states( step0_topk, merged_offsetsmerged_offsets, merged_offsets_per_devmerged_offsets_per_dev, num_acceptednum_accepted_draft_tokens, num_draft_tokensdraft_tokens.shape[1], data_parallel_degree..., data_parallel_splits..., signal_buffers..., devicedevice0, split_prefixmtp_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_indicesreuse_topk, # 復(fù)用 step 0 的 top-k reuse_prev_topkTrue, # 跳過重復(fù)計(jì)算 split_prefixfmtp_draft_step{step}, )每個(gè)后續(xù)步驟的緩存mla_cache_lengths_per_dev會(huì)遞增1而 indexer 緩存長(zhǎng)度保持不變因?yàn)?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)最終返回三元組(被接受的草稿數(shù)量, 下一輪主 token, 拼接的新 draft tokens)。PipelineModel 編譯與運(yùn)行時(shí)UnifiedMTPGlm5_2Model該 PipelineModel 定義在 model.py 中繼承自_UnifiedSpecDecodeModelMixin和Glm5_1Model。權(quán)重加載_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構(gòu)建嵌套的{target: {mla, indexer}, draft: {mla, indexer}}KV 緩存樹。draft 的緩存僅有 1 層num_layers1且mla和indexer分開管理draft_kv MultiKVCacheParams.from_params({ mla: replace(target_mla_params, num_layers1), indexer: replace(target_indexer_params, num_layers1), }) self.kv_params MultiKVCacheParams.from_params( {target: target_kv, draft: draft_kv} )分布式專家并行EP_init_distributed_runtime處理專家并行初始化。對(duì)于 NVFP4 量化檢查點(diǎn)其 MTP 層的 routed experts 以 bf16 精度存儲(chǔ)無.weight_scale因此 draft 的 EP 分發(fā)精度必須從 NVFP4 提升到 bf16draft_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_dtypeDType.bfloat16, dispatch_quant_configNone, fused_shared_expertmodel_config.n_shared_experts 1, )圖編譯_build_graph_for_compile方法實(shí)例化UnifiedMTPGlm5_2模型權(quán)重共享將 draft 的embed_tokens和lm_head別名為 target 的對(duì)應(yīng)層strictFalse加載時(shí)跳過已共享的 key通過nn_model.input_types(kv_params)構(gòu)建圖輸入類型簽名用Graph上下文管理器構(gòu)建完整的glm5_2_with_mtp_graph計(jì)算圖從輸入中解包四組 KV 緩存target_mla、target_indexer、draft_mla、draft_indexer提取采樣超參數(shù)seed、temperature、top_k、max_k、top_p、min_top_p和in_thinking_phase標(biāo)志調(diào)用nn_model(...)完成前向傳播綁定圖輸出輸入 batchingUnifiedMTPGlm5_2BatchProcessor定義在 batch_processor.py 中繼承自DeepseekV3BatchProcessor。它在標(biāo)準(zhǔn) batch 輸入基礎(chǔ)上擴(kuò)展了draft_tokens字段初始化為None由 overlap pipeline 在實(shí)際執(zhí)行時(shí)填充。輸入結(jié)構(gòu)定義在UnifiedMTPGlm5_2Inputs中繼承UnifiedSpecDecodeInputs和DeepseekV3Inputs其buffers屬性除了父類輸入外額外包含in_thinking_phase標(biāo)志位。權(quán)重適配器convert_with_mtp_state_dict定義在 weight_adapters.py 中負(fù)責(zé)將 HuggingFace safetensors 格式的 checkpoint 轉(zhuǎn)換為 MAX 內(nèi)部格式。權(quán)重 key 映射規(guī)則如下常規(guī)層通過DEEPSEEK_SAFETENSOR_MAP來自deepseekV3.weight_adapters將 HuggingFace key 轉(zhuǎn)換為 MAX key丟棄 KV 縮放因子跳過以.k_scale或.v_scale結(jié)尾的 keyMAX 從獨(dú)立配置路徑讀取 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.專家頭部投影l(fā)ayers.N.其他draft.decoder_layer.解碼器層self-attention / MLP非layers.N.前綴target.*目標(biāo)模型權(quán)重權(quán)重共享embed_tokens和shared_head.headlm_head僅在 target 前綴下保存一份draft 通過模塊別名共享state_dict()自動(dòng)去重。Draft 配置創(chuàng)建_create_draft_config方法model.py從 draft 的 state_dict 推導(dǎo)配置驗(yàn)證 NextN 層存在decoder_layer.self_attn.kv_a_layernorm.weight基于 target 的基礎(chǔ)配置創(chuàng)建DeepseekV3_2NextNConfig設(shè)置indexer_types []空調(diào)度讓單層 MTP 的 indexer 保持滿計(jì)算避免引用 target 的 78 層調(diào)度方案檢測(cè) NVFP4 量化的子樹范圍.mlp.experts.和.self_attn.將 MTP 層索引添加到正確的量化層集合中如果 draft 的 routed experts 未量化NVFP4 下 MTP 層為 bf16則修正 EP 分發(fā)配置為bfloat16使用方式要使用該架構(gòu)在 MAX 中加載 GLM-5.2-FP8 模型并啟用 MTP 推測(cè)解碼需要在 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關(guān)鍵推測(cè)解碼配置項(xiàng)說明源自 max/python/max/pipelines/speculative/config.py配置項(xiàng)類型默認(rèn)值說明num_speculative_tokensint \| NoneNone每次生成的草稿 token 數(shù)量num_speculative_tokens_per_batch_sizelist[VerifyWidthRange] \| NoneNone按 batch size 分級(jí)的草稿數(shù)量調(diào)度synthetic_acceptance_ratefloat \| NoneNone合成接受率0.0~1.0用于繞過真實(shí)分布模擬use_relaxed_acceptance_for_thinkingboolFalsethinking 階段是否啟用寬松接受relaxed_topkint—寬松接受的 top-k 范圍需 1relaxed_deltafloat—寬松接受的 delta 閾值0.0~1.0支持的量化編碼可通過--dtype float8_e4m3fn默認(rèn)、--dtype bfloat16或--dtype float4_e2m1fnx2指定。多 GPU 訓(xùn)練可通過--tensor-parallel-size N啟用。該架構(gòu)已默認(rèn)注冊(cè)到 MAX Pipeline 中通過架構(gòu)名UnifiedMTPGlmMoeDsaForCausalLM自動(dòng)匹配zai-org/GLM-5.2-FP8等模型 ID無需額外的手動(dòng)注冊(cè)步驟。【免費(fèi)下載鏈接】mojoThe Modular Platform (includes MAX Mojo)項(xiàng)目地址: https://gitcode.com/GitHub_Trending/mo/mojo創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考