Hugging Face Transformers 4.45.0 動態投機解碼
動態投機解碼的運作方式
動態投機解碼透過最佳化「投機前瞻」(speculation lookahead, SL) 來提升標準投機解碼的效能——即快速草稿模型在較大的目標模型同步驗證之前所產生的 token 數量。
先前的方法使用固定的 SL 或根據前一次迭代的接受率所制定的啟發式策略,而動態投機解碼則利用助理模型自身的信心來決定何時停止。具體而言,系統會監測每個預測 token 的 logits softmax。若助理模型的信心低於預先定義的 assistant_confidence_threshold,則該迭代的 token 產生會停止,並將序列送至目標模型驗證,即使尚未達到最大 num_assistant_tokens。
效能基準測試
在 RTX 4090 上使用貪婪解碼(temperature = 0)進行的基準測試顯示,動態方法在各種模型配對與任務上始終優於基於啟發式的方式:
| 目標模型 | 草稿(助理)模型 | 任務 | 加速比 - 啟發式 | 加速比 - 動態 |
|---|---|---|---|---|
facebook/opt-6.7b |
facebook/opt-125m |
摘要 | 1.82x | 2.71x |
facebook/opt-6.7b |
facebook/opt-125m |
開放式生成 | 1.23x | 1.59x |
Salesforce/codegen-6B-mono |
Salesforce/codegen-350M-mono |
程式碼生成(Python) | 0.89x | 1.09x |
google/flan-t5-xl |
google/flan-t5-small |
摘要 | 1.18x | 1.31x |
meta-llama/Llama-3.1-8B |
meta-llama/Llama-3.2-1B |
摘要 | 1.00x | 1.52x |
meta-llama/Llama-3.1-8B |
meta-llama/Llama-3.2-1B |
開放式生成 | 1.00x | 1.18x |
meta-llama/Llama-3.1-8B |
meta-llama/Llama-3.2-1B |
程式碼生成(Python) | 1.09x | 1.15x |
這些基準測試的主要發現包括:
- Llama 3.1/3.2 配對:動態方法在摘要任務上取得 1.52 倍的加速,而啟發式方法則未顯示顯著的加速。
- 程式碼生成:對於
codegen-6B-mono,啟發式方法實際上導致速度下降(0.89 倍),而動態方法則提供了加速(1.09 倍)。
在 Transformers 4.45.0 中的實作
動態投機是 Transformers 4.45.0 中助理解碼的預設模式。可透過將 assistant_model 傳入 generate 方法來實作:
from transformers import AutoModelForCausalLM, AutoTokenizer
import torch
prompt = "Alice and Bob"
checkpoint = "EleutherAI/pythia-1.4b-deduped"
assistant_checkpoint = "EleutherAI/pythia-160m-deduped"
device = "cuda" if torch.cuda.is_available() else "cpu"
tokenizer = AutoTokenizer.from_pretrained(checkpoint)
inputs = tokenizer(prompt, return_tensors="pt").to(device)
model = AutoModelForCausalLM.from_pretrained(checkpoint).to(device)
assistant_model = AutoModelForCausalLM.from_pretrained(assistant_checkpoint).to(device)
outputs = model.generate(**inputs, assistant_model=assistant_model)
設定參數
使用者可以在 assistant_model.generation_config 中調整以下參數,以優化特定資料集的效能:
assistant_confidence_threshold:草稿模型停止產生的信心門檻(logits 的 softmax)。num_assistant_tokens:助理模型每次迭代可產生的最大 token 數量。num_assistant_tokens_schedule:可設定為'dynamic'(預設)、'heuristic'或'constant',以切換不同的投機策略。
理論基礎:Oracle 模型
為了驗證動態調整的必要性,研究人員使用了一個「oracle」模型,透過持續產生 token 直至草稿模型與目標模型出現不一致,從而找出每次迭代的最佳投機前瞻。對 MBPP 資料集與 Alpaca 資料集的分析顯示,最佳草稿 token 數量變異很大,證明固定的前瞻值在最大化推論速度方面並非最佳。