首页 / 文章 / 推测性解码与提前终止:加速自回归解码

推测性解码与提前终止:加速自回归解码

在可满足接受度与准确率要求的前提下,通过先生成草稿再验证以及基于置信度的层退出机制,可降低解码延迟。

2856 词

为何解码优化处于关键路径上

大语言模型推理包含一个针对提示词并行处理的预填充阶段,以及一个依次生成标记的解码阶段。预填充操作可能会使GPU负荷饱和;而解码阶段通常不会,因为每个新标记的生成都依赖于前一个标记。正是这种顺序依赖性导致即便从理论上看浮点运算量足够,交互延迟依然很高。

Prefill phase:
  Input: [token_1, token_2, ..., token_512]  → all 512 tokens processed in parallel
  Matrix shape: [batch, 512, 4096]
  GPU utilization: high — large matrix, full Tensor Core throughput

Decode phase (one step):
  Input: [token_513]  → one token processed
  Matrix shape: [batch, 1, 4096]
  GPU utilization: low - tiny matrix, most CUDA cores idle

解码过程中的注意力结构

在每个解码步骤中,模型都会针对不断扩大的键/值缓存生成查询。随着上下文长度的增加,每步的处理量也会上升,同时由于聊天应用的延迟性能要求,批量处理的机会也受到限制。

One attention head, decode step:

Query Q: [8, 1, 128]   →  8 × 128 = 1,024 elements
Keys  K: [8, 2048, 128] →  8 × 2048 × 128 = 2M elements
Values V: [8, 2048, 128] →  same

Matrix multiply (QKᵀ) per head:
  Shape: [8, 1, 128] × [8, 128, 2048] → [8, 1, 2048]
  FLOPs: 8 × 1 × 128 × 2048 ≈ 2.1M FLOPs per head

For 32 heads: ≈ 67M FLOPs per layer
For 32 layers: ≈ 2.1B FLOPs total per decode step

A100 peak at BF16: ~312 TFLOPS = 312 × 10¹² FLOPs/sec

Time to execute 2.1B FLOPs at 100% utilization:
  2.1 × 10⁹ / 312 × 10¹² ≈ 0.0067 ms

Actual observed decode latency per step: ~10–30ms

Effective compute utilization: < 0.1%

推测性解码:先生成草稿再验证

较小的草稿模型会生成若干候选令牌;目标模型会在一次并行前向传播中对其进行验证,首次被拒绝时接受部分前缀并重新采样。被接受的草稿能提升每个昂贵目标步骤的有效令牌数量。

Without speculative decoding:
  Generate 5 tokens: 5 × 30ms = 150ms

With speculative decoding (K=4):
  Draft 4 tokens: 4 × 1.5ms = 6ms
  1 target verification pass: ~35ms  (slightly longer than decode,
                                       processes K+1=5 positions)
  Expected accepted tokens per round:
    (1 - 0.80^5) / (1 - 0.80) ≈ 3.36 tokens
  Time per round: 6ms + 35ms = 41ms
  Time per token: 41ms / 3.36 ≈ 12.2ms
Speedup: 30ms → 12.2ms ≈ 2.5×

加速效果取决于草稿模型与目标模型之间的匹配程度。结构简单、熵值较低的文本能接受较长的草稿;而出现意外令牌时则会导致验证失败。

import torch
import torch.nn.functional as F

def speculative_decode(target_model, draft_model, input_ids,
                       max_new_tokens, K=4, temperature=1.0):
    """
    Conceptual speculative decoding loop.
    Real implementations handle KV cache management across both models.
    """
    generated = input_ids.clone()
    while generated.shape[1] - input_ids.shape[1] < max_new_tokens:
        # --- Draft phase ---
        draft_tokens = []
        draft_probs = []
        draft_input = generated.clone()
        for _ in range(K):
            with torch.no_grad():
                draft_logits = draft_model(draft_input).logits[:, -1, :]
            q = F.softmax(draft_logits / temperature, dim=-1)
            token = torch.multinomial(q, num_samples=1)
            draft_tokens.append(token)
            draft_probs.append(q)
            draft_input = torch.cat([draft_input, token], dim=1)
        # --- Verify phase: one target forward pass over all K+1 positions ---
        verify_input = torch.cat([generated] + draft_tokens, dim=1)
        with torch.no_grad():
            target_logits = target_model(verify_input).logits
        # target_logits[:, -K-1:, :] covers all K draft positions + bonus
        # --- Accept/reject ---
        accepted = 0
        for i in range(K):
            p = F.softmax(target_logits[:, -(K+1)+i, :] / temperature, dim=-1)
            q = draft_probs[i]
            token = draft_tokens[i]
            # Acceptance probability
            accept_prob = torch.min(
                torch.ones_like(p.gather(1, token)),
                p.gather(1, token) / (q.gather(1, token) + 1e-9)
            )
            if torch.rand(1) < accept_prob:
                generated = torch.cat([generated, token], dim=1)
                accepted += 1
            else:
                # Sample corrected token and stop this round
                corrected_dist = F.relu(p - q)
                corrected_dist = corrected_dist / corrected_dist.sum(dim=-1, keepdim=True)
                corrected_token = torch.multinomial(corrected_dist, num_samples=1)
                generated = torch.cat([generated, corrected_token], dim=1)
                break
        else:
            # All K accepted - take bonus token
            bonus_logits = target_logits[:, -1, :]
            p_bonus = F.softmax(bonus_logits / temperature, dim=-1)
            bonus_token = torch.multinomial(p_bonus, num_samples=1)
            generated = torch.cat([generated, bonus_token], dim=1)
    return generated

尽可能选择同一系列的草稿模型,即经过精简或量化的简化版本,并根据实际使用数据而非公共博客示例来衡量接受率。

当草稿出现偏差时

如果草稿的分布发生偏移,验证成功率会急剧下降,此时你虽支付了草稿模型成本却收效甚微。需持续监控验证成功率,一旦下降则应恢复使用普通解码方式。

提前终止:在有把握时停止深度层处理

某些架构允许在置信度较高时在中间层终止处理,从而节省对“简单”标记的运算资源。

Layer distribution of exits:Exit at layers 1–8  (very easy tokens like punctuation, articles):  15%
Exit at layers 9–16 (medium tokens, common continuations):           35%
Exit at layers 17–24 (harder tokens, named entities, numbers):       30%
Exit at layers 25–32 (full computation required):                    20%Weighted average layers executed:
  0.15 × 6 + 0.35 × 12 + 0.30 × 20 + 0.20 × 32
  = 0.90 + 4.20 + 6.00 + 6.40
  = 17.5 layers averageSpeedup vs always running 32 layers:
  32 / 17.5 ≈ 1.83×

需谨慎评估其对准确率的影响:提前终止会以速度换取质量,并且对校准参数十分敏感。

推测解码与提前终止

推测解码通过第二个模型来改变标记生成流程;而提前终止则是在同一个模型内调整处理深度。二者针对不同的瓶颈问题,有时可与量化、连续批处理以及KV缓存分页技术结合使用。

Combined stack example:

Target: 7B model, BF16, FlashAttention, PagedAttention
  → Model: ~14 GB, memory-efficient attention, paged KV cache

Draft: 70M model, INT4, FlashAttention
  → Model: ~35 MB, near-zero memory overhead

Speculative decode:
  → with K=4, α=0.80
  → ~2.5× token generation speedup on long outputs

Full stack speedup vs FP32 no-optimization baseline:
  - Quantization: 2–2.3× throughput
  - FlashAttention: 20–40% attention latency reduction
  - Speculative decoding: 2–2.5× decode speedup (on eligible requests)

Combined: 5–8× improvement in end-to-end tokens/second

何时适用

当初步结果的一致性较高且目标步骤是延迟的主要来源时,推测解码较为有效;而当初步结果很少匹配或生成初始结果的开销超过其带来的收益时,该方法则失效。提前终止在许多标记都比较简单且准确率要求允许的情况下效果较好;但对于难度较高的标记或校准不准确的置信度值,该方法则无法发挥作用。

实际应用视角

输出指标包括:接受率、平均被接受的长度、质量评估差异值、GPU利用率以及p95延迟。还支持通过特性标志进行优化,并保留传统解码方式作为备用。需记住,剩余的瓶颈可能在于内存带宽、网络或客户端渲染——而不仅仅是浮点运算次数——因此在应用各种技术之前务必先进行性能分析。

总结

预填充与解码会对硬件产生不同压力。相较于接受率和准确率曲线,推测性解码与提前终止是实用的优化手段。应将其视为需要通过控制面板管理的正式功能,而非一次性测试用的方案,这样它们才能在真实流量环境中发挥作用。

辅助理解的参考场景

想象一个70亿参数的模型用于聊天服务,每次回复需要50到150个标记。2000个标记的系统提示语虽然加载量较大,但每轮出现频率不高;用户感知到的延迟主要来自解码步骤。通过推测性接受来减少解码步数,或通过提前终止来降低每步的处理量,都能改善用户的体验。

操作检查清单

  • 记录基准的p95值及质量指标。
  • 添加本地运行的草稿模型,以避免网络延迟。
  • 记录接受结果的直方图数据。
  • 通过事实类测试集对比不同方案的质量。
  • 观察在不同设备上运行草稿模型时的CPU/GPU使用平衡情况。
  • 每次更改分词器或进行微调后都要重新测试。

常见误区

使用随机选取且无关的草稿;忽略温度参数对结果接受度的影响;在没有质量检测机制的情况下就根据处理速度宣布胜利;将提前终止与过度量化结合使用,直至产生大量幻觉。每一个问题都是可测量的——首先依靠仪器数据。

面向平台团队的进一步指导

将推理优化功能集中在服务层,这样应用团队就不必各自设计草稿选择机制。通过标题或追踪属性来显示是否采用了推测性解码或提前终止策略。在费用账单中添加优化标签,以便财务部门了解优化效果。预先演练回滚流程。明确记录:推测性解码并不能替代良好的信息检索或优质提示词——它只是在模型已经知道要表达的内容时,降低生成下一个标记的成本而已。

关于组合使用的深入探讨

应谨慎地将连续批处理与推测机制结合使用:草稿长度会影响到调度器的假设。还需将其与KV缓存淘汰策略配合,以避免长对话导致性能下降。同时要结合预填充阶段的提示词缓存机制,这样既不会在解码时出现问题,也不会造成不必要的资源浪费。整体化的服务设计远优于将单一技术直接套用到关键流程中。

实践者常见问题

推测机制会改变答案吗?如果实现正确,它应与目标分布保持一致;可通过配对测试进行验证。提前终止会改变答案吗?是的,因为这种机制会跳过某些处理层——需预估由此产生的差异。可以在CPU上生成草稿吗?有时可以,但需要先进行性能测试。这对设备上的小型模型有意义吗?通常不如对大型模型的作用显著。处理优先级为:先解决批处理和缓存问题,再考虑推测机制,最后在框架支持的情况下使用提前终止功能。

辅助理解的参考场景

想象一个70亿参数的模型用于聊天服务,每次回复需要50到150个标记。2000个标记的系统提示语虽然加载量大,但每轮出现频率不高;用户感知到的延迟主要来自解码步骤。通过推测性接受来减少解码步数,或通过提前终止来降低每步的处理量,都能改善用户的体验。

操作检查清单

  • 记录基准的p95值及质量指标。
  • 添加本地运行的草稿模型,以避免网络延迟。
  • 记录接受结果的直方图数据。
  • 通过事实类测试集对比不同方案的质量。
  • 观察在不同设备上运行草稿模型时的CPU/GPU使用平衡情况。
  • 每次更改分词器或进行微调后都要重新测试。

常见误区

使用随机选取且无关的草稿;忽略温度参数对结果接受度的影响;在没有质量检测机制的情况下就根据处理速度宣布胜利;将提前终止与过度量化手段结合使用,直至出现幻觉现象加剧。每一个问题都是可测量的——首先依靠工具进行检测。

面向平台团队的进一步指导

将推理优化功能集中在服务层,这样应用团队就不必各自独立设计草稿选择机制。通过标题或追踪属性来显示是否采用了推测性解码或提前终止策略。在费用账单中添加优化标签,以便财务部门能够了解优化效果。预先演练回滚流程。明确记录:推测性解码并不能替代优质的检索或提示词设计——它只是在模型已经知道要表达的内容时,降低生成下一个标记的成本而已。

关于组合使用的深入说明

应谨慎地将连续批处理与推测机制结合使用:草稿长度会影响到调度器的假设。还需将其与KV缓存淘汰策略配合,以避免长对话导致性能下降。同时要结合预填充阶段的提示词缓存机制,这样既不会在解码时出现问题,也不会造成不必要的资源浪费。整体化的服务设计远优于将单一技术直接套用到关键流程中。

实践者常见问题

推测机制会改变答案吗?如果实现正确,它应与目标分布保持一致;可通过配对测试进行验证。提前退出会改变答案吗?是的,因为这种机制会跳过某些处理层——需预估由此产生的差异。可以在CPU上生成草稿吗?有时可以,但需要先进行测试。这对设备上的小型模型有意义吗?通常不如对大型模型的作用显著。处理优先级为:先解决批处理和缓存问题,再考虑推测机制,最后在框架支持的情况下使用提前退出功能。

经过验证的数值直觉

假设某个目标步骤的耗时为10毫秒,而某种方案提出5个标记,其中3个标记的平均接受率为60%。只要接受率保持良好,即使考虑到方案带来的额外开销,每个被接受的标记的实际成本仍低于传统单标记步骤。但如果接受率降至约1个标记,该方案就会失效。正是由于这种敏感性,仪表板才比零散的案例更具参考价值。

质量退化检测机制

在全局启用之前,需在实际问答、编程以及拒绝处理等测试场景中使用固定提示词进行测试。在设置精确分布匹配的推测模式时,比较标记完全相同的比例。同时检查是否存在任何系统性偏差。若要提前终止测试,可跟踪分级任务中的胜率以及现有条件下的人类偏好数据。

硬件部署位置

尽可能将草稿和目标版本放在同一个节点上。跨主机处理草稿会增加网络抖动,从而抵消已取得的优化效果。注意内存使用情况:两个模型加上KV缓存可能会让原本能够容纳一个模型的设备出现内存不足。

交互调度

连续批处理服务器必须考虑可变的推测性扩展需求。糟糕的调度机制会导致批次被拆分,进而降低资源利用率。应与服务维护人员协调,不要仅在应用代码中修改相关标志。

剩余的瓶颈问题

尽管解码性能已得到提升,用户仍可能因工具调用、数据检索或客户端端的Markdown处理而等待。需进行端到端跟踪分析。在错误的环节上进行优化只会浪费工程时间。

总结

顺序解码是自回归结构所固有的成本。在特定条件下,推测性解码与提前终止机制能够降低这一成本。应像处理任何生产功能一样严格管理这些机制:设定指标、设置标志、实现回滚,并明确服务平台团队中的责任人员。

可操作的数值直觉

假设每个目标步骤需要10毫秒,某个方案提出5个标记,其中3个标记的平均接受率为60%。只要接受率保持良好,即使考虑到方案设计带来的额外开销,每个被接受的标记的实际成本仍低于传统的单个标记步骤。但如果接受率降至约1个标记,该方案就会失效。正是由于这种敏感性,仪表板数据比零散的案例更具参考价值。

质量退化处理协议

在全局启用之前,需在事实问答、编程及拒绝处理模块中测试固定提示词的效果。当设置为精确匹配分布时,比较令牌完全相同的比例,同时检查是否存在系统性偏差。为便于提前终止测试,需跟踪分级任务中的胜率以及现有条件下的人类偏好数据。

硬件部署

尽可能将草稿模型与目标模型部署在同一节点上。跨主机运行草稿模型会增加网络抖动,从而抵消已取得的优势。注意内存占用:两个模型加上KV缓存可能会让原本能够容纳一个模型的服务器出现内存不足。

交互调度

连续批处理服务器必须考虑推测扩展量的变化。糟糕的调度策略会导致批次碎片化,降低资源利用率。需与服务维护人员协调,切勿仅在应用代码中更改相关设置。

剩余的瓶颈问题

尽管解码效率有所提升,用户仍可能需等待工具调用、数据检索或客户端端的 Markdown 处理。应进行端到端追踪,否则在错误的部分进行优化只会浪费工程时间。

总结回顾

顺序解码是自回归结构带来的固有开销。在可测量的条件下,推测性解码与提前终止机制能够降低这一开销。应像处理任何生产级功能一样严格管理这些机制:设定指标、设置标志、准备回滚方案,并明确服务平台团队中的责任人员。

基于数据的直观理解

假设每个目标处理步骤需要10毫秒时间,而某种方案先生成5个标记符,其中3个的平均接受率为60%。只要接受率保持良好,即使考虑到方案生成带来的额外开销,每个被接受的标记符的实际成本仍低于传统的单个标记符处理方式。但如果接受率降至约1个标记符,该方案就会失去优势。正是由于这种敏感性,数据看板才比零散的案例更有价值。

质量退化检测方案

在全局启用之前,需在实际问答、编程及拒绝处理场景中测试固定提示词的效果。当设置为精确分布匹配时,比较完全相同的token比例,排查是否存在系统性偏差。如需提前终止测试,可跟踪分级任务中的胜率以及现有条件下的人类偏好数据。

硬件部署

尽可能将草稿模型与目标模型部署在同一节点上。跨主机运行草稿模型会增加网络抖动,从而抵消已取得的优化效果。注意内存占用:两个模型加上KV缓存可能会让原本能够容纳一个模型的服务器出现内存不足。

交互调度

连续批处理服务器必须考虑推测扩展量的不确定性。糟糕的调度机制会导致批次碎片化,降低资源利用率。应与服务维护人员协调,切勿仅在应用代码中修改相关参数。

剩余的瓶颈问题

尽管解码效率有所提升,用户仍可能需等待工具调用、数据检索或客户端端的 Markdown 处理。应进行端到端追踪,否则在错误的部分进行优化只会浪费工程时间。

总结回顾

顺序解码是自回归结构带来的固有开销。在可测量的条件下,推测性解码与提前终止机制能够降低这一开销。应像处理任何生产级功能一样严格管理这些机制:设定指标、设置标志、准备回滚方案,并明确服务平台团队中的责任人员。

基于数据的直观理解

假设每个目标处理步骤需要10毫秒时间,而某种方案先生成5个标记符,其中3个的平均接受率为60%。只要接受率保持良好,即使考虑到方案生成带来的额外开销,每个被接受的标记符的实际成本仍低于传统的单个标记符处理方式。但如果接受率降至约1个标记符,该方案就会失去优势。正是由于这种敏感性,数据看板才比口头描述更有效。

质量退化检测方案

在全局启用之前,需在实际问答、编程及拒绝处理场景中运行固定提示词。当设置为精确分布匹配时,比较完全相同的token比例,排查是否存在系统性偏差。如需提前终止测试,可跟踪分级任务中的胜率以及现有条件下的人类偏好数据。

硬件部署

尽可能将草稿模型与目标模型部署在同一节点上。跨主机运行草稿模型会增加网络抖动,从而抵消已取得的优化效果。注意内存占用:两个模型加上KV缓存可能会让原本能够容纳一个模型的服务器出现内存不足。

交互调度

连续批处理服务器必须考虑推测性扩展带来的变化。糟糕的调度机制会导致批次碎片化,降低资源利用率。应与服务维护人员协调,切勿仅在应用代码中修改相关参数。

剩余的瓶颈问题

尽管解码效率有所提升,用户仍可能面临工具调用、数据检索或客户端 Markdown 处理的延迟。需对整个流程进行端到端追踪。在错误的部分进行优化只会浪费工程时间。

总结回顾

顺序解码是自回归结构带来的固有开销。在特定条件下,推测性解码与提前终止机制能够降低这一开销。应像处理其他生产功能一样严格管理这些技术:设定指标、设置标志、实现回滚机制,并明确服务平台团队中的责任人员。