MLX-VLM 推测解码(Speculative Decoding)实现指南:从架构设计到模型接入
2026/9/18 6:46:07 网站建设 项目流程

MLX-VLM 推测解码(Speculative Decoding)实现指南:从架构设计到模型接入

【免费下载链接】mlx-vlmMLX-VLM is a package for inference and fine-tuning of Vision Language Models (VLMs) on your Mac using MLX.项目地址: https://gitcode.com/GitHub_Trending/ml/mlx-vlm

推测解码(speculative decoding)是一种无损(lossless)解码优化:用一个更小的 drafter(草稿模型)一次性提出多个 token,再由目标模型(target)并行验证整块草稿,接受匹配前缀并额外取一个目标 token,拒绝的缓存条目在下一轮前回滚。本文以 docs/speculative-decoding.md 为主线,结合 MLX-VLM 仓库中mlx_vlm/speculative/的实际实现,完整讲解 DFlash / MTP / EAGLE-3 三类草稿架构的职责划分、缓存事务与精确验证约束、新模型接入流程,以及验证与调优方法。读完本文,你将掌握如何在 MLX-VLM 中接入新的 drafter、理解"逐 token 精确等价"的硬性要求,并学会用合成测试与真实权重验证推测解码的正确性与吞吐收益。

推测解码的核心流程

MLX-VLM 的推测解码遵循标准范式,但有一个关键差异:验证与回滚都通过共享的普通前向(ordinary forward)与缓存事务完成,而非为每个模型单独实现一套验证器。整体流程可抽象为:

target prefill and hidden capture ↓ draft block → target block verification → accept prefix + target token ↑ ↓ └──────────── cache rollback ──────────┘

具体到每一轮(round loop):

  1. Prefill 与隐藏态捕获:目标模型对 prompt 做 prefill,同时按 drafter 的需求暴露隐藏状态(hidden states)。
  2. Draft(草稿):drafter 基于当前隐藏态与上一个 token(bonus token)自回归地提出bs-1个候选 token。
  3. Verify(验证):将[bonus] + draft_tokens拼接后送入目标模型一次前向,得到每个位置的 logits。
  4. Walk / Accept(接受):逐位置比较草稿 token 与目标 greedy(或采样)token,接受最长匹配前缀,并在第一个不匹配处补一个目标 token(bonus token)——这就是_speculative_walk的语义。
  5. Commit / Rollback(提交 / 回滚):接受accepted + 1个位置,拒绝的部分从缓存中回滚,进入下一轮。

从源码看,这一"walk"逻辑在 common.py 中实现为_speculative_walk(单序列)与_speculative_walk_batch(批量)。批量版本通过mx.argmax(mismatches)向量化地找出每行第一个不匹配位置,再配合mx.take_along_axis取出 bonus token,将"接受决策"打包为一次tolist()跨越 Python/图边界。

代码归属:谁负责什么

推测解码横跨生成循环、drafter 架构、缓存与算子层。原文档给出了一张职责表,结合源码可精确对应到具体文件:

区域职责源码位置
generate/ar.pyPrefill、目标缓存创建、推测分发生成主循环,generate_step中按draft_kind调用run_speculative_rounds
speculative/drafters/Checkpoint 加载、drafter 架构、目标兼容性按模型家族划分的子包,如dflash/mtp/eagle3/
speculative/dflash.pyDFlash、DFlash2、DSpark 轮循环_dflash_rounds/_dflash_rounds_batch
speculative/mtp.py原生与 assistant MTP 轮循环_mtp_rounds/_mtp_rounds_batch
speculative/eagle3.pyEAGLE-3 轮循环_eagle3_rounds/_eagle3_rounds_batch
speculative/common.py普通验证前向、接受判定、采样状态、批量护栏verify_forward_speculative_walk*_SpeculativeSamplerRNG
models/<family>/language.py正常模型遍历与隐藏态捕获通过return_hidden/capture_layer_ids暴露
models/cache.py有界时序状态保留与接受态选择缓存类的start_speculation/commit_speculation/abort_speculation
models/linear.pyswitch_layers.pyfast_ops.py共享的解码等价算子与融合内核DECODE_BLOCK_SIZE = 8定义于此
speculative/ops/既有 Qwen 验证算子linear.py
speculative/cache_state.py推测缓存事务的提交与中止SpeculativeCacheTransaction

draft_kind 决定轮循环,而非 checkpoint 架构

一个容易混淆的点是:draft_kind选择的是"轮循环"(round loop),而不是 drafter 的 checkpoint 架构。drafter 的 HFmodel_typedflash/mtp/eagle3的映射定义在 drafters/init.py:

KNOWN_DRAFTER_KINDS = {"dflash", "mtp", "eagle3"} DRAFTER_KIND_BY_MODEL_TYPE = { "deepseek_v4_mtp": "mtp", "deepseek_v4_dspark": "dflash", "dspark": "dflash", "gemma4_dspark": "dflash", "eagle3": "eagle3", "gemma4_assistant": "mtp", "glm5_next_mtp": "mtp", "qwen3_5_mtp": "mtp", "laguna": "dflash", "muse_glimmer_assistant": "dflash", "qwen3_dspark": "dflash", # ... } DEFAULT_DRAFTER_KIND = "dflash"

resolve_drafter_kind(drafters/init.py)负责在调用方未传kind时从config.jsonmodel_type自动探测;若调用方显式传入的kind与 drafter 的model_type冲突,则自动覆盖并告警(例如用户把--draft-model指向gemma4_assistantcheckpoint 却忘了--draft-kind mtp,系统会替用户纠正而不是深陷draft_block中报错)。

validate_drafter_compatibility(drafters/init.py)则坚持"用架构与 config 字段判断,而不是仓库名",从而保证量化后的 MLX 转换与本地 checkpoint 同样被接受;对 MTP 类 drafter 还会比对backbone_hidden_size/target_hidden_size与目标模型的hidden_size,不匹配直接抛错。

添加一个新模型(drafter)

原文档给出 6 步接入流程,结合源码可进一步落实每一步的落点:

  1. 先勘察真实config.json与张量名。接入前必须确定:目标层(target layers)、隐藏态输入、block-size 语义、缓存归属与量化方式。从drafters/<family>/config.py可以看到每个 drafter 都有一段 config 归一化逻辑,专门处理 HF 权重命名与 MLX 格式的差异。

  2. speculative/drafters/<family>/下新增或复用 drafter,实现:config 归一化、checkpoint 清洗(sanitization)、draft_block、缓存重置、目标兼容性检查。典型结构包含__init__.pyconfig.py、模型主体.py,部分家族(如qwen3_5_mtpglm5_next_mtpdeepseek_v4_mtp)还带split.py负责从目标权重拆分 drafter。

  3. 注册model_type到正确的draft_kind(见上文映射表),优先使用架构与 config 字段而非仓库名判断。

  4. 让目标 prefill 返回 drafter 所需的隐藏态。DFlash 与 EAGLE-3 通常按capture_layer_ids捕获若干指定层;MTP 通常消费最终隐藏态,并可能共享目标 K/V。这一点在 utils.py 的speculative_prefill_kwargs中直接体现:

    if draft_kind == "mtp": return {"return_hidden": True, "return_shared_kv": True} if draft_kind == "eagle3": return {"capture_layer_ids": _eagle3_capture_layer_ids(drafter)} if draft_kind == "dflash": return {"capture_layer_ids": list(drafter.config.target_layer_ids)}
  5. 使用模型的普通前向 + 共享缓存事务。这是本仓库设计哲学的核心:有状态算子通过时间缓存接口写穿(write through),使得回滚无需对层做第二套实现;数值分派放在普通共享算子中,使用同一份权重与同一模块实例;接受判定与历史保留策略不得侵入模型层。

  6. 先写合成契约测试,再用真实目标与 drafter checkpoint 验证,然后才可宣称支持。

共享验证前向:verify_forward

GLM 与 DeepSeek 使用共享运行时的verify_forward(common.py)。它的做法是:

  • 开一个缓存事务(start_speculative_cache);
  • 把验证输入按DECODE_BLOCK_SIZE(在 linear.py 中定义为8)切成短块,逐块调用普通 model 前向
  • 若只有一块,直接返回;多块则沿序列轴拼接 logits 与 hidden_states,并返回事务对象。

MTP 请求最终隐藏态,DFlash 请求其配置的捕获层;轮循环负责采样、提交与中止。GLM 与 DeepSeek 都不需要专门的 speculative verifier 或 rollback 方法——普通前向本身就是验证器。

普通前向通过return_hiddencapture_layer_ids暴露隐藏态;当捕获态在送入 LM head 前需要架构特定的归一化时,提供一个普通的logits_from_hidden方法即可。drafter 自行负责 reshape 自己的输入。

原文档特别强调:迁移期内既有模型的 hooks 仍受支持,但新模型应使用普通前向/缓存契约;仓库中不存在 target 注册表、复制的模型视图或参数包装层。

精确验证(Exact Verification):为何普通前向≠重复单 token 解码

推测解码的正确性底线是与自回归解码逐 token 等价。原文档明确指出一个微妙陷阱:普通的多 token 目标前向并不自动等价于重复的单 token 解码,原因包括:

  • 内核分派可能改变浮点运算顺序(kernel dispatch can change floating-point order);
  • Mamba、gated-delta、卷积或旋转缓存(rotating caches)需要一个针对接受位置的显式状态
  • 在 argmax 平局附近,极小的数值差异就可能改变生成的序列。

因此验证必须以重复的普通单 token 前向作为参照基准。models/linear.pyswitch_layers.py等中的短块线性、专家与超连接算子保留了约简顺序;因果注意力保留逐位置缓存顺序;图像 prefill 掩码保留完整可见性规则。运行时会按DECODE_BLOCK_SIZE切分更大的验证块,但整个验证块共享一个缓存事务;大的普通 prefill 调用仍走批量执行路径。

验证必须满足的硬性条件

  • 产生与自回归解码相同的 greedy 目标 token;
  • 让每个缓存都推进过整个验证块;
  • accepted + 1个 token 处恢复有状态缓存;
  • 正确处理零接受、部分接受、完全接受三种情况;
  • 要么支持逐行(per-row)批量回滚,要么要求统一接受(uniform acceptance);
  • 分块 prefill 时保留完整的必需 prompt 隐藏态;
  • 匹配所有受支持的权重格式,包括量化输出头(quantized output heads)

原文档还强调:融合 argmax 与自定义 Metal 内核是优化手段而非正确性捷径,在等价性被证明之前必须保留全 logits 回退路径(full-logit fallback)。

接受判定的实现细节

  • Greedy 路径_speculative_walk直接比较草稿 token 与目标 greedy token,取第一个不匹配位置;批量版_speculative_walk_batch向量化完成。
  • 采样路径:DFlash 的_sample_dflash_target_walk逐位置计算 log-probs 并调用 sampler;MTP 的_speculative_walk_deferred_greedy则是延迟投影——只在到达拒绝点前按需把目标隐藏态投影为 logits,避免为注定被丢弃的位置浪费 LM head 计算。
  • 批量统一接受_requires_uniform_batch_acceptance(common.py)检查 drafter 或目标模型是否声明requires_uniform_batch_acceptance。这是一个重要的兼容性细节:即使 drafter 未声明,目标模型(如 qwen3_5 配合 qwen3_5_mtp)若使用矩形批量 KV 缓存(单一_idx),就无法表示参差接受(ragged accepts),否则会留下幽灵零键被注意力读到(源码注释中引用了 issue #1962),因此目标模型也可能独立要求统一接受。

采样 RNG 隔离

_SpeculativeSamplerRNG(common.py)负责保持目标与 drafter 的采样 RNG 流相互独立:draft 调用前保存目标 RNG 状态、恢复 drafter 的 RNG 状态;draft 完成后反向恢复。这样目标模型的采样分布不会因推测解码的存在而改变——这正是"无损"在采样场景下的保证。

线性注意力的时间缓存(Temporal Caches)

推测解码的缓存回滚是正确性的核心。ArraysCache拥有有界历史(bounded history),覆盖当前验证窗口;普通解码只保留最新状态。关键在于:同一个前向调用可以在缓存事务内运行,无需模型专属 checkpoint 或循环层替换。

共享 gated-delta 算子

共享的 gated-delta 算子接受cachecache_index,委托给cache.update_recurrent——后者仅在需要历史时才从内核请求中间状态cache.update_window存储卷积或 token 窗口并保留其时序视图。两种方法都支持一次整块更新或容量内的多次小块更新。

事务语义

SpeculativeCacheTransaction(cache_state.py)是这一设计的具体实现:

  • Commit:按每行接受的数量保留对应输入位置,cache.commit_speculation(lengths, generation)选择各行的接受态;
  • Abort:恢复起始状态,cache.abort_speculation(generation)释放历史;
  • 超容量更新与不完整历史被拒绝(reject)。

新的循环算子只需要一次性地支持共享的状态生产契约,使用它的模型无需知道 MTP 或草稿长度的存在;数值投影与注意力分派属于普通算子的一部分,与缓存历史保留相互独立。

旋转缓存(Rotating Cache)的特殊处理

对于旋转缓存,_RotatingCacheTransaction(cache_state.py)的做法是:记录传入的 KV,而服务缓存保持其原生布局;提交时若各行接受数不同,则 abort 后按各行的接受长度重新prepare、按原始顺序重放(replay)接受的更新、再finalize。原文档强调:仅仅恢复游标(cursor)在驱逐或旋转后是不够的——事务必须保留服务窗口与传入 KV,然后按原始顺序重放被接受的更新。

验证与维护:如何证明它是对的

第一层:廉价合成测试

从配置归一化、严格权重加载、捕获层顺序、块大小、采样状态与缓存回滚等"廉价测试"起步。核心断言是:在零接受、部分接受、完全接受三种情况下,比较缓存状态与下一个目标 token,覆盖单序列与批量生成。仓库中 test_cache.py 的test_pooling_cache_speculative_commit_matches_prefix_replaytest_batch_pooling_cache_speculative_commit_matches_ragged_prefixes正是这类"提交与前缀重放一致"的契约测试。

第二层:真实 checkpoint 贪婪生成对拍

用真实 checkpoint 的贪婪生成与同一自回归基线对拍,要求跨多个 prompt 与上下文长度逐 token 相等(token-for-token equality),然后才进入基准测试。

基准测试纪律

  • 所有变体在同一进程内、共享同一基线上测量;
  • 报告中位解码吞吐与接受率(median decode throughput and acceptance);
  • 拒绝那种只提升孤立内核、却不改善端到端解码的改动。

运行时统计由_record_speculative_roundformat_speculative_stats(common.py)支持,例如输出形如:

Speculative decoding: 2.34 accepted tokens/round (1.34 accepted drafts/round, 67.0% of drafted, avg draft 2.00) over 500 rounds

回归触发条件

当改动涉及目标层、缓存、量化、采样、批处理或分块 prefill时,必须重跑合成回滚测试与真实模型精确性测试。新的 checkpoint 布局或架构标签应视为兼容性变更,而不是"既有适配器仍然适用"的证据。

旋转缓存的边界测试

对旋转缓存,要越过窗口边界、跨多次接受/拒绝轮次测试。逐行独立测试批量:即使 batch=1 测试通过,中间状态内核的错误 stride 也可能在批量时写出其输出分配之外。

运行时块大小(Runtime Block Sizes)与实测取舍

块大小决定每轮提出多少草稿 token,直接影响验证成本与接受收益的平衡。原文档给出了两个具体案例:

GLM-5-Next:从原生 MTP 深度扩展

GLM-5-Next 从原生 MTP 深度起步,当共享接受策略观察到可靠接受率时,可以扩展到3-token 块(2 个提案)——这降低了 GLM-5.3-Flash 贪婪 batch-one 与 batch-four 检查中的轮开销。但批量随机采样默认保持 2-token 上限,因为在该负载下额外提案无法偿还其验证成本。显式的runtime_block_size--draft-block-size设置优先。

DeepSeek-V4 DSpark:默认 2-token 块

DeepSeek-V4 DSpark 默认2-token 块(1 个提案),其普通算子保留解码算术与物理注意力窗口顺序。更大的块可能"花得比省得多"(larger blocks can cost more than their accepted tokens save)。文档明确指出:当前精确路径(exact path)在实测的 M3 Ultra 负载上仍慢于基线,吞吐优先时应使用普通解码。此外,挂载 drafter 后文本 prefill 仍保持分块,捕获的特征跨块保留;图像跨度仍使用模型的整图 prefill 策略。

块大小的自适应选择在源码中有多处实现:DFlash 的_dflash_next_block_size(dflash.py)基于最近 8 轮的接受率动态调整——接受率低于 0.30 或均值低于 2.0 时快速回退,接受率 ≥0.85 且满命中率 ≥0.75 时才逐步增长回配置上限;MTP 的_effective_mtp_block_size(mtp.py)则要求最近 32 轮中"达到配置深度"的命中率 ≥0.65 才允许超出配置深度。这些启发式都遵循同一原则:配置深度是上限,收益不达预期就快速回退

命令行接入与配置入口

推测解码通过三个 CLI 参数启用(定义于 generate/dispatch.py):

  • --draft-model:drafter 路径或 HF id(例如z-lab/Qwen3.5-4B-DFlash);
  • --draft-kinddflash/mtp/eagle3,默认从 drafter 的 HFmodel_type自动探测;
  • --draft-block-size:覆盖 drafter 配置的块大小。

加载路径(dispatch.py)依次执行:load_drafter(自动探测或覆盖kind)→validate_drafter_compatibility(不兼容则告警并禁用推测路径)→ 把draft_model/draft_kind/draft_block_size注入generate_step。生成结束后可通过format_speculative_stats(draft_model)打印接受率统计。

在代码中直接调用时,入口是 utils.py 的run_speculative_rounds(同步生成)与run_speculative_server_rounds(服务端批量、支持 continuous batching),它们按draft_kind分派到dflash/mtp/eagle3各自的轮循环(get_speculative_rounds_batch是批量分派表)。Prefill 阶段的隐藏态捕获由SpeculativePrefill跨 chunk 累积拼接,保证分块 prefill 下 drafter 所需的完整隐藏态不丢失。

总结

MLX-VLM 的推测解码设计有三个可独立引用的要点:

  1. 普通前向即验证器:通过共享缓存事务与verify_forward,GLM、DeepSeek 等模型无需专用 verifier 或 rollback 方法,接受判定、采样状态与批量护栏全部收敛在speculative/common.py,新架构只需实现普通前向契约。
  2. 精确等价是硬约束:多 token 前向 ≠ 重复单 token 解码,必须以普通单 token 前向为参照,覆盖零/部分/完全接受与旋转缓存回滚,量化头与批量参差接受(或统一接受钳制)都必须显式处理。
  3. 块大小是运行时决策draft_kind只决定轮循环,块大小由配置、runtime_block_size与基于接受率的自适应策略共同决定,且必须以"端到端中位吞吐与接受率"而非孤立内核指标来评估。

对于希望在 Apple Silicon 上获得无损解码加速、或向 MLX-VLM 接入新推测解码架构的开发者,建议的下一步是:阅读 speculative/drafters/README.md 与 generate/ar.py 的调用链,参照 test_cache.py 的契约测试模式为新架构建立合成测试,再用真实 checkpoint 对拍基线,最后在同一进程内完成吞吐与接受率的基准对比。

【免费下载链接】mlx-vlmMLX-VLM is a package for inference and fine-tuning of Vision Language Models (VLMs) on your Mac using MLX.项目地址: https://gitcode.com/GitHub_Trending/ml/mlx-vlm

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

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

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

立即咨询