推測解碼:用小模型加速大模型
從自回歸生成的序列依賴到接受拒絕機制,拆解推測解碼如何在不損失品質的前提下把生成速度提升 2 到 3 倍。
推測解碼:用小模型加速大模型
上一篇我們談了批次處理與連續批次,理解了如何透過調度策略讓 GPU 幾乎不閒置,把吞吐提升數倍。
但連續批次解決的是「多個請求怎麼排」的問題,還沒有解決「單一請求為什麼這麼慢」的問題。
自回歸生成的本質是:一次只能生成一個 token,每一步都要等前一步完成。
即使 GPU 再強,這個序列依賴也讓延遲無法壓縮。
這一篇,我們要談一個打破序列依賴的技術:推測解碼(Speculative Decoding)。
它讓小模型先「猜」幾個 token,大模型再「驗證」,在不損失品質的前提下,把生成速度提升 2 到 3 倍。
一、先理解問題:為什麼生成這麼慢?
自回歸生成的序列依賴
LLM 生成文本的過程是自回歸的:
輸入 → 生成 token 1 → 生成 token 2 → 生成 token 3 → ...
每個 token 的生成都依賴前一個 token。
這意味著:
- 無法平行:不能同時生成多個 token。
- 記憶體頻寬受限:每一步都要讀取整個模型權重和 KV Cache。
- 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,然後一次性驗證呢?
這就是推測解碼的核心思想。
二、推測解碼的核心思想
用一個比喻來理解
想像你在寫一篇文章,但打字速度很慢。
你請一個助理先幫你打草稿,然後你再檢查。
- 助理(草稿模型):打字快,但可能出錯。
- 你(目標模型):打字慢,但準確。
流程是:
- 助理先打 5 個字。
- 你一次檢查這 5 個字。
- 如果都對,就全部採用。
- 如果某個字錯了,就從錯誤的地方開始修正。
這樣,你只需要「檢查」而不是「一個字一個字打」,速度大幅提升。
推測解碼的正式流程
1. 草稿階段(Draft):
用小模型(草稿模型)快速生成 K 個候選 token。
2. 驗證階段(Verify):
用大模型(目標模型)一次計算這 K 個 token 的機率。
3. 接受/拒絕階段(Accept/Reject):
從第一個 token 開始,逐一比較草稿模型和目標模型的機率。
如果草稿 token 符合目標模型的分布,就接受。
如果不符合,就拒絕,並從該位置重新採樣。
4. 重複:
接受的 token 加入序列,繼續下一輪。
為什麼能加速?
關鍵在於:驗證 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 模型
流程:
- 用 7B 模型生成 5 個候選 token。
- 用 70B 模型一次驗證這 5 個 token。
- 接受符合的,拒絕不符合的。
優點:
- 概念簡單,易於實作。
- 加速比明顯(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 | 低 |
| Medusa | 否 | 是 | 2-3x | 中 |
| Lookahead | 否 | 否 | 1.5-2x | 高 |
五、如何選擇草稿模型?
草稿模型的選擇,直接影響加速效果。
選擇原則
| 原則 | 說明 |
|---|---|
| 同系列 | 草稿模型和目標模型最好來自同一系列(如都是 LLaMA) |
| 分佈接近 | 草稿模型的輸出分佈越接近目標模型,接受率越高 |
| 足夠小 | 草稿模型要夠小,才能快速生成 |
| 足夠好 | 草稿模型不能太差,否則接受率太低 |
常見搭配
| 目標模型 | 草稿模型 | 加速比 |
|---|---|---|
| LLaMA-2 70B | LLaMA-2 7B | 2.5x |
| LLaMA-2 13B | LLaMA-2 7B | 2.0x |
| GPT-5.6 Sol | GPT-5.6 Luna | 2.0x |
| 任何大模型 | TinyLLaMA | 1.5-2x |
接受率的影響
接受率是指草稿 token 被目標模型接受的比例。
| 接受率 | 加速比 | 說明 |
|---|---|---|
| 0.9 | 3.0x | 草稿模型很好 |
| 0.7 | 2.5x | 草稿模型不錯 |
| 0.5 | 1.8x | 草稿模型一般 |
| 0.3 | 1.2x | 草稿模型太差 |
經驗法則:如果接受率低於 0.5,推測解碼的加速效果不明顯,甚至可能更慢。
如何提升接受率?
- 選擇同系列的草稿模型:分佈更接近。
- 用蒸餾訓練草稿模型:讓草稿模型模仿目標模型。
- 調整溫度:較低的溫度通常有更高的接受率。
- 動態調整草稿長度:根據接受率動態調整 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 協同推理,包括:張量平行、流水線平行、專家平行,以及如何選擇合適的平行策略。