Hugging Face Transformers 4.45.0 动态投机解码
Hugging Face 与 Intel Labs 开发了 动态投机解码,这是一种通过根据任务可提升至 2.7 倍 的文本生成加速方法。该技术自 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 中辅助解码的默认模式。只需在 generate 方法中传入 assistant_model 即可启用:
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 数量的方差很大,说明固定的前瞻值在最大化推理速度方面并不理想。