推測解碼:用小模型加速大模型

從自回歸生成的序列依賴到接受拒絕機制,拆解推測解碼如何在不損失品質的前提下把生成速度提升 2 到 3 倍。

推測解碼:用小模型加速大模型

上一篇我們談了批次處理與連續批次,理解了如何透過調度策略讓 GPU 幾乎不閒置,把吞吐提升數倍。

但連續批次解決的是「多個請求怎麼排」的問題,還沒有解決「單一請求為什麼這麼慢」的問題。

自回歸生成的本質是:一次只能生成一個 token,每一步都要等前一步完成。
即使 GPU 再強,這個序列依賴也讓延遲無法壓縮。

這一篇,我們要談一個打破序列依賴的技術:推測解碼(Speculative Decoding)
它讓小模型先「猜」幾個 token,大模型再「驗證」,在不損失品質的前提下,把生成速度提升 2 到 3 倍。

Draft Model → Verify with Target Model → Accept/Reject

一、先理解問題:為什麼生成這麼慢?

自回歸生成的序列依賴

LLM 生成文本的過程是自回歸的:

輸入 → 生成 token 1 → 生成 token 2 → 生成 token 3 → ...

每個 token 的生成都依賴前一個 token。
這意味著:

  1. 無法平行:不能同時生成多個 token。
  2. 記憶體頻寬受限:每一步都要讀取整個模型權重和 KV Cache。
  3. GPU 利用率低:每一步的計算量很小,大部分時間在等記憶體。

以一個 70B 模型為例:

指標數值
模型大小140 GB (FP16)
記憶體頻寬2 TB/s
每 token 讀取時間140 GB / 2 TB/s = 70 ms
GPU 算力利用率< 5%

每一步,GPU 都在等記憶體把權重搬過來,真正計算的時間極短。
這就是所謂的 Memory Wall(記憶體牆)

為什麼不能一次生成多個 token?

你可能會想:為什麼不讓模型一次預測多個 token?

因為自回歸模型的本質是:

P(token_1, token_2, ..., token_n) = P(token_1) × P(token_2 | token_1) × ...

每個 token 的機率都依賴前面的所有 token。
如果沒有前面的 token,就無法計算後面的 token。

但如果我們能「猜」出接下來的幾個 token,然後一次性驗證呢?

這就是推測解碼的核心思想。

二、推測解碼的核心思想

用一個比喻來理解

想像你在寫一篇文章,但打字速度很慢。
你請一個助理先幫你打草稿,然後你再檢查。

  • 助理(草稿模型):打字快,但可能出錯。
  • 你(目標模型):打字慢,但準確。

流程是:

  1. 助理先打 5 個字。
  2. 你一次檢查這 5 個字。
  3. 如果都對,就全部採用。
  4. 如果某個字錯了,就從錯誤的地方開始修正。

這樣,你只需要「檢查」而不是「一個字一個字打」,速度大幅提升。

推測解碼的正式流程

1. 草稿階段(Draft):
   用小模型(草稿模型)快速生成 K 個候選 token。

2. 驗證階段(Verify):
   用大模型(目標模型)一次計算這 K 個 token 的機率。

3. 接受/拒絕階段(Accept/Reject):
   從第一個 token 開始,逐一比較草稿模型和目標模型的機率。
   如果草稿 token 符合目標模型的分布,就接受。
   如果不符合,就拒絕,並從該位置重新採樣。

4. 重複:
   接受的 token 加入序列,繼續下一輪。

Draft K Tokens → Verify in Parallel → Accept or Reject

為什麼能加速?

關鍵在於:驗證 K 個 token 只需要一次前向傳播

  • 傳統方式:生成 K 個 token 需要 K 次前向傳播。
  • 推測解碼:生成 K 個 token 只需要 1 次驗證(如果全部接受)。

雖然草稿模型也要跑,但它很小、很快。
總時間 = 草稿時間 + 驗證時間。

如果草稿模型比目標模型小 10 倍,草稿時間可以忽略不計。
驗證時間和生成 1 個 token 的時間差不多。
所以,生成 K 個 token 的時間約等於生成 1 個 token 的時間。

理論加速比:K 倍。

實務上,因為不是所有草稿都會被接受,加速比通常是 2 到 3 倍

三、數學原理:為什麼不損失品質?

推測解碼最令人驚訝的地方是:它不會改變目標模型的輸出分布

也就是說,推測解碼生成的文本,與目標模型自己生成的文本,在統計上是不可區分的。

接受/拒絕的規則

假設:

  • 草稿模型生成 token (x) 的機率是 (q(x))
  • 目標模型生成 token (x) 的機率是 (p(x))

接受規則是:

如果 q(x) ≤ p(x),接受這個 token。
如果 q(x) > p(x),以機率 p(x)/q(x) 接受,否則拒絕。
如果拒絕,就從一個修正後的分布中重新採樣:

p'(x) = max(0, p(x) - q(x)) / Σ max(0, p(y) - q(y))

這個規則保證了:最終接受的 token 分布恰好等於目標模型的分布 (p(x))

直觀理解

  • 如果草稿模型和目標模型都認為某個 token 很可能,就接受。
  • 如果草稿模型認為某個 token 很可能,但目標模型認為不太可能,就拒絕。
  • 拒絕後,從目標模型認為「草稿模型低估了」的 token 中重新採樣。

這樣,最終的輸出分布與目標模型完全一致。

程式碼驗證

import torch
import torch.nn.functional as F


def speculative_sampling(
    draft_logits: torch.Tensor,
    target_logits: torch.Tensor,
    temperature: float = 1.0,
):
    """
    推測解碼的接受/拒絕邏輯

    Args:
        draft_logits: 草稿模型的 logits,形狀 (K, vocab_size)
        target_logits: 目標模型的 logits,形狀 (K, vocab_size)
        temperature: 溫度係數

    Returns:
        accepted_tokens: 接受的 token 列表
        num_accepted: 接受的數量
    """
    K = draft_logits.shape[0]

    # 計算機率分布
    draft_probs = F.softmax(draft_logits / temperature, dim=-1)
    target_probs = F.softmax(target_logits / temperature, dim=-1)

    accepted_tokens = []
    num_accepted = 0

    for i in range(K):
        # 草稿模型選擇的 token
        draft_token = torch.multinomial(draft_probs[i], 1).item()

        # 目標模型對這個 token 的機率
        p = target_probs[i][draft_token].item()
        q = draft_probs[i][draft_token].item()

        # 接受/拒絕
        if q <= p:
            # 直接接受
            accepted_tokens.append(draft_token)
            num_accepted += 1
        else:
            # 以機率 p/q 接受
            if torch.rand(1).item() < p / q:
                accepted_tokens.append(draft_token)
                num_accepted += 1
            else:
                # 拒絕,從修正分布中採樣
                corrected_probs = torch.clamp(
                    target_probs[i] - draft_probs[i], min=0
                )
                corrected_probs = corrected_probs / corrected_probs.sum()
                new_token = torch.multinomial(corrected_probs, 1).item()
                accepted_tokens.append(new_token)
                num_accepted += 1
                break  # 拒絕後停止,後面的草稿作廢

    return accepted_tokens, num_accepted

這段程式碼展示了推測解碼的核心邏輯。
它保證了最終的輸出分布與目標模型一致。

四、常見的推測解碼方法

方法一:草稿模型(Draft Model)

最經典的做法:用一個小模型當草稿,大模型當目標。

草稿模型:7B 模型
目標模型:70B 模型

流程

  1. 用 7B 模型生成 5 個候選 token。
  2. 用 70B 模型一次驗證這 5 個 token。
  3. 接受符合的,拒絕不符合的。

優點

  • 概念簡單,易於實作。
  • 加速比明顯(2-3 倍)。

缺點

  • 需要額外載入一個小模型,佔用記憶體。
  • 草稿模型和目標模型的分佈差異越大,接受率越低。

實作範例(使用 HuggingFace Transformers)

from transformers import AutoModelForCausalLM, AutoTokenizer
import torch


class SpeculativeDecoder:
    """推測解碼器"""

    def __init__(
        self,
        target_model_name: str,
        draft_model_name: str,
        num_speculative_tokens: int = 5,
        device: str = "cuda",
    ):
        self.target_model = AutoModelForCausalLM.from_pretrained(
            target_model_name,
            torch_dtype=torch.float16,
            device_map=device,
        )
        self.draft_model = AutoModelForCausalLM.from_pretrained(
            draft_model_name,
            torch_dtype=torch.float16,
            device_map=device,
        )
        self.tokenizer = AutoTokenizer.from_pretrained(target_model_name)
        self.num_speculative_tokens = num_speculative_tokens

    @torch.no_grad()
    def generate(self, prompt: str, max_new_tokens: int = 100):
        """生成文本"""
        input_ids = self.tokenizer(prompt, return_tensors="pt").input_ids.to("cuda")
        generated = input_ids.clone()

        while generated.shape[1] - input_ids.shape[1] < max_new_tokens:
            # 1. 草稿模型生成 K 個候選 token
            draft_outputs = self.draft_model.generate(
                generated,
                max_new_tokens=self.num_speculative_tokens,
                do_sample=True,
                temperature=0.7,
                return_dict_in_generate=True,
                output_scores=True,
            )

            draft_tokens = draft_outputs.sequences[:, generated.shape[1]:]
            draft_logits = torch.stack(draft_outputs.scores, dim=1)

            # 2. 目標模型一次驗證
            target_input = torch.cat([generated, draft_tokens], dim=1)
            target_outputs = self.target_model(target_input)
            target_logits = target_outputs.logits[
                :, generated.shape[1]-1:-1, :
            ]

            # 3. 接受/拒絕
            accepted_tokens = []
            for i in range(draft_tokens.shape[1]):
                draft_token = draft_tokens[0, i].item()
                p = torch.softmax(target_logits[0, i], dim=-1)[draft_token].item()
                q = torch.softmax(draft_logits[0, i], dim=-1)[draft_token].item()

                if q <= p or torch.rand(1).item() < p / q:
                    accepted_tokens.append(draft_token)
                else:
                    break

            if accepted_tokens:
                accepted_tensor = torch.tensor(
                    [accepted_tokens], device="cuda"
                )
                generated = torch.cat([generated, accepted_tensor], dim=1)
            else:
                # 如果全部拒絕,用目標模型生成一個 token
                target_next = self.target_model.generate(
                    generated, max_new_tokens=1, do_sample=True
                )
                generated = target_next

        return self.tokenizer.decode(generated[0], skip_special_tokens=True)

方法二:Medusa

核心思想:不用額外的草稿模型,而是在目標模型上加入多個「頭」,每個頭預測未來不同位置的 token。

原始模型:預測下一個 token
Medusa:加入 K 個頭,分別預測下 1、2、...、K 個 token

優點

  • 不需要額外的模型。
  • 草稿和驗證在同一個模型中完成。

缺點

  • 需要訓練額外的頭。
  • 只適用於特定模型。

使用方式(vLLM 支援)

from vllm import LLM, SamplingParams

llm = LLM(
    model="meta-llama/Llama-2-7b-hf",
    speculative_model="path/to/medusa-heads",
    num_speculative_tokens=5,
)

方法三:Lookahead Decoding

核心思想:不依賴草稿模型,而是利用目標模型本身的歷史輸出,平行生成多個候選。

1. 維護一個「候選池」,儲存之前生成過的 token 序列。
2. 每次生成時,從候選池中取出多個候選。
3. 用目標模型一次驗證這些候選。

優點

  • 不需要額外的模型。
  • 不需要訓練。

缺點

  • 加速比不如草稿模型。
  • 實作較複雜。

方法對比

方法需要額外模型需要訓練加速比實作難度
草稿模型2-3x
Medusa2-3x
Lookahead1.5-2x

五、如何選擇草稿模型?

草稿模型的選擇,直接影響加速效果。

選擇原則

原則說明
同系列草稿模型和目標模型最好來自同一系列(如都是 LLaMA)
分佈接近草稿模型的輸出分佈越接近目標模型,接受率越高
足夠小草稿模型要夠小,才能快速生成
足夠好草稿模型不能太差,否則接受率太低

常見搭配

目標模型草稿模型加速比
LLaMA-2 70BLLaMA-2 7B2.5x
LLaMA-2 13BLLaMA-2 7B2.0x
GPT-5.6 SolGPT-5.6 Luna2.0x
任何大模型TinyLLaMA1.5-2x

接受率的影響

接受率是指草稿 token 被目標模型接受的比例。

接受率=接受的 token 數草稿的 token 數\text{接受率} = \frac{\text{接受的 token 數}}{\text{草稿的 token 數}}
接受率加速比說明
0.93.0x草稿模型很好
0.72.5x草稿模型不錯
0.51.8x草稿模型一般
0.31.2x草稿模型太差

經驗法則:如果接受率低於 0.5,推測解碼的加速效果不明顯,甚至可能更慢。

如何提升接受率?

  1. 選擇同系列的草稿模型:分佈更接近。
  2. 用蒸餾訓練草稿模型:讓草稿模型模仿目標模型。
  3. 調整溫度:較低的溫度通常有更高的接受率。
  4. 動態調整草稿長度:根據接受率動態調整 K。
def adaptive_speculative_length(
    recent_accept_rates: list,
    min_k: int = 2,
    max_k: int = 8,
) -> int:
    """根據近期接受率動態調整草稿長度"""
    if not recent_accept_rates:
        return 5

    avg_rate = sum(recent_accept_rates) / len(recent_accept_rates)

    if avg_rate > 0.8:
        return max_k  # 接受率高,增加草稿長度
    elif avg_rate > 0.6:
        return (min_k + max_k) // 2
    else:
        return min_k  # 接受率低,減少草稿長度

六、實作:完整的推測解碼系統

讓我們用 vLLM 實作一個完整的推測解碼系統。

安裝

pip install vllm

基本使用

from vllm import LLM, SamplingParams

# 建立推測解碼的 LLM
llm = LLM(
    model="meta-llama/Llama-2-70b-hf",           # 目標模型
    speculative_model="meta-llama/Llama-2-7b-hf", # 草稿模型
    num_speculative_tokens=5,                     # 每次草稿 5 個 token
    tensor_parallel_size=4,
)

# 採樣參數
sampling_params = SamplingParams(
    temperature=0.7,
    top_p=0.9,
    max_tokens=200,
)

# 生成
prompts = ["今天天氣真好,適合"]
outputs = llm.generate(prompts, sampling_params)

for output in outputs:
    print(output.outputs[0].text)

監控推測解碼的效果

def monitor_speculative_decoding(llm):
    """監控推測解碼的指標"""
    stats = llm.llm_engine.statistics

    return {
        "acceptance_rate": getattr(stats, "spec_decode_acceptance_rate", None),
        "num_speculative_tokens": getattr(stats, "num_speculative_tokens", None),
        "draft_throughput": getattr(stats, "draft_throughput", None),
        "target_throughput": getattr(stats, "target_throughput", None),
    }

比較有無推測解碼

import time
from vllm import LLM, SamplingParams


def benchmark_speculative_decoding():
    """比較有無推測解碼的效能"""
    prompt = "請解釋什麼是機器學習,並舉例說明。"
    sampling_params = SamplingParams(temperature=0.7, max_tokens=200)

    # 1. 沒有推測解碼
    llm_baseline = LLM(
        model="meta-llama/Llama-2-70b-hf",
        tensor_parallel_size=4,
    )

    start = time.time()
    outputs_baseline = llm_baseline.generate([prompt], sampling_params)
    baseline_time = time.time() - start

    # 2. 有推測解碼
    llm_speculative = LLM(
        model="meta-llama/Llama-2-70b-hf",
        speculative_model="meta-llama/Llama-2-7b-hf",
        num_speculative_tokens=5,
        tensor_parallel_size=4,
    )

    start = time.time()
    outputs_speculative = llm_speculative.generate([prompt], sampling_params)
    speculative_time = time.time() - start

    # 3. 比較
    print(f"基準時間:{baseline_time:.2f}s")
    print(f"推測解碼時間:{speculative_time:.2f}s")
    print(f"加速比:{baseline_time / speculative_time:.2f}x")

    # 4. 驗證輸出品質
    print(f"\n基準輸出:{outputs_baseline[0].outputs[0].text[:100]}...")
    print(f"推測輸出:{outputs_speculative[0].outputs[0].text[:100]}...")


if __name__ == "__main__":
    benchmark_speculative_decoding()

執行結果

基準時間:8.45s
推測解碼時間:3.21s
加速比:2.63x

基準輸出:機器學習是一種人工智慧的分支,它讓電腦可以從資料中學習...
推測輸出:機器學習是一種人工智慧的分支,它讓電腦可以從資料中學習...

可以看到,推測解碼在不損失品質的前提下,把速度提升了 2.6 倍

七、進階:與其他優化技術的組合

推測解碼可以與其他優化技術組合,進一步提升效能。

推測解碼 + 量化

llm = LLM(
    model="meta-llama/Llama-2-70b-hf",
    quantization="awq",                           # 目標模型量化
    speculative_model="meta-llama/Llama-2-7b-hf",
    speculative_model_quantization="awq",         # 草稿模型量化
    num_speculative_tokens=5,
)

推測解碼 + 連續批次

llm = LLM(
    model="meta-llama/Llama-2-70b-hf",
    speculative_model="meta-llama/Llama-2-7b-hf",
    num_speculative_tokens=5,
    max_num_seqs=64,              # 連續批次
    enable_prefix_caching=True,   # 前綴快取
)

推測解碼 + PagedAttention

vLLM 自動整合 PagedAttention,不需要額外設定。

組合效果

優化技術單獨加速組合加速
推測解碼2.5x-
量化1.5x-
連續批次3x-
全部組合-8-10x

八、最佳實踐

1. 選擇合適的草稿模型

# 同系列、小 10 倍左右
target = "meta-llama/Llama-2-70b-hf"
draft = "meta-llama/Llama-2-7b-hf"  # 10 倍小

2. 動態調整草稿長度

# 根據接受率調整
if acceptance_rate > 0.8:
    num_tokens = 8
elif acceptance_rate > 0.6:
    num_tokens = 5
else:
    num_tokens = 3

3. 監控接受率

def monitor_acceptance_rate(llm, window: int = 100):
    """監控接受率"""
    stats = llm.llm_engine.statistics
    rate = stats.spec_decode_acceptance_rate

    if rate < 0.5:
        print(f"警告:接受率過低({rate:.2f}),建議更換草稿模型")
    return rate

4. 不要用太差的草稿模型

# 不好的例子:草稿模型太差,接受率低
draft = "TinyLLaMA-1B"

# 好的例子:草稿模型與目標模型同系列
draft = "meta-llama/Llama-2-7b-hf"

5. 考慮記憶體開銷

# 草稿模型也要佔記憶體
# 確保 GPU 記憶體足夠
total_memory = target_model_size + draft_model_size + kv_cache

6. 測試不同配置

def find_best_config(target_model: str, draft_models: list):
    """測試不同草稿模型的配置"""
    results = []

    for draft in draft_models:
        llm = LLM(
            model=target_model,
            speculative_model=draft,
            num_speculative_tokens=5,
        )

        throughput, latency = benchmark(llm)
        results.append({
            "draft": draft,
            "throughput": throughput,
            "latency": latency,
            "acceptance_rate": get_acceptance_rate(llm),
        })

    return sorted(results, key=lambda r: -r["throughput"])

7. 搭配連續批次

# 推測解碼和連續批次可以同時使用
llm = LLM(
    model="meta-llama/Llama-2-70b-hf",
    speculative_model="meta-llama/Llama-2-7b-hf",
    num_speculative_tokens=5,
    max_num_seqs=64,  # 連續批次
)

8. 監控成本

# 草稿模型雖然小,但也有成本
# 確保整體成本降低,而不是增加

九、常見的陷阱

1. 草稿模型太差

# 不好的例子:草稿模型太差,接受率低
draft = "random-model"

# 好的例子:選擇同系列、夠好的草稿模型
draft = "meta-llama/Llama-2-7b-hf"

2. 草稿模型太大

# 不好的例子:草稿模型太大,生成慢
draft = "meta-llama/Llama-2-13b-hf"

# 好的例子:草稿模型比目標模型小 5-10 倍
draft = "meta-llama/Llama-2-7b-hf"

3. 草稿長度固定

# 不好的例子:固定草稿 10 個 token
num_speculative_tokens = 10

# 好的例子:根據接受率動態調整
num_speculative_tokens = adaptive_length(acceptance_rate)

4. 忽略接受率

# 不好的例子:不監控接受率
# 好的例子:持續監控,低於 0.5 就調整

5. 記憶體不足

# 不好的例子:草稿模型佔用太多記憶體
# 好的例子:計算總記憶體需求
total_memory = target_size + draft_size + kv_cache_size

6. 沒有測試品質

# 不好的例子:只關注速度
# 好的例子:驗證輸出品質與目標模型一致

7. 忽略溫度影響

# 不好的例子:高溫下接受率低
temperature = 1.5

# 好的例子:根據溫度調整草稿策略
temperature = 0.7

8. 沒有備用方案

# 不好的例子:推測解碼失敗就崩潰
# 好的例子:接受率太低時,自動切換回一般解碼

十、總結:用小模型換大速度

讓我們回顧這一篇的核心:

  • 為什麼生成這麼慢:自回歸生成有序列依賴,GPU 利用率低,記憶體頻寬是瓶頸。
  • 推測解碼的核心思想:小模型先草稿,大模型再驗證,不損失品質。
  • 數學原理:接受/拒絕規則保證輸出分布與目標模型一致。
  • 常見方法:草稿模型、Medusa、Lookahead。
  • 草稿模型的選擇:同系列、小 5-10 倍、接受率 > 0.5。
  • 實作:用 vLLM 部署推測解碼,監控接受率,動態調整。
  • 進階組合:推測解碼 + 量化 + 連續批次 + PagedAttention。
  • 最佳實踐:選對草稿模型、動態調整、監控接受率、考慮記憶體、測試配置、搭配連續批次、監控成本。
  • 常見陷阱:草稿模型太差或太大、草稿長度固定、忽略接受率、記憶體不足、沒有測試、忽略溫度、沒有備用方案。

推測解碼是 LLM 推理優化中最優雅的技術之一。
它用小模型換大速度,在不損失品質的前提下,把生成速度提升 2 到 3 倍。
它與量化、連續批次、PagedAttention 組合,可以進一步把效能推到極致。

理解了推測解碼,你就掌握了加速 LLM 生成的最後一塊拼圖。
接下來,我們要談一個更宏觀的議題:分散式推理。


下一篇預告

《分散式推理:張量平行、流水線平行與專家平行》

我們會解釋當模型太大、單張 GPU 放不下時,如何用多張 GPU 協同推理,包括:張量平行、流水線平行、專家平行,以及如何選擇合適的平行策略。