Decoding Strategies and Output Control
TL;DR · AI 摘要
本文系统解析Transformer模型生成文本的六大核心解码策略,揭示不同算法对输出质量与多样性的影响机制。
核心要点
- 贪婪解码保证输出稳定性但可能牺牲多样性
- 温度采样通过调整softmax温度参数控制输出随机性
- 核采样通过累积概率阈值实现高效多样性控制
结构提纲
按章节快速跳转。
思维导图
用一张图看清主题之间的关系。
查看大纲文本(无障碍 / 无 JS 友好)
- 解码策略与输出控制
- 基础概念
- Logits处理
- 概率转换
- 核心算法
- 贪婪解码
- 温度采样
- 核采样
- 控制机制
- 重复惩罚
- 停止条件
- 束搜索
金句 / Highlights
值得收藏与分享的关键句。
贪婪解码通过argmax操作实现确定性输出,但可能陷入局部最优
温度参数控制softmax的平滑程度,温度值越低输出越确定
核采样通过累积概率分布动态调整候选token数量,兼顾多样性与效率
解码策略与输出控制 - MachineLearningMastery.com
解码策略与输出控制
By
Adrian Tam
on
2026年8月4日
in
0
Share
Post
语言模型不会直接生成文本。相反,它会返回下一个标记的logits。解码算法决定如何将这些logits转换为标记,重复这个过程会生成输出文本。
解码算法会影响模型的行为。贪婪解码是确定性和稳定的,但可能会显得单调。采样会引入一些随机性,这可能会生成更多样化的文本,但也可能产生错误。束搜索对于某些受限任务可能有用,但通常不是聊天式生成的最佳默认选项。输出约束可以让模型生成JSON或在特定标记处停止。
在本章中,你将学习:
- 贪婪解码
- 温度采样
- Top-k和核采样
- 重复惩罚
- 停止条件
- 束搜索
- 结构化输出约束
让我们开始吧。
解码策略与输出控制 图片由Claudio Testa拍摄。部分权利保留。
概述
本章分为九个部分;它们是:
- 从模型中读取logits
- Top-$k$采样
- 核采样
从模型中读取logits
模型为输入序列中的每个位置返回一个logits向量。对于生成,通常只使用最后一个位置,因为它预测下一个标记。
以下示例使用Hugging Face的transformers库和一个小型GPT-2风格模型。检查点足够小,适合本地实验,但同样的逻辑也适用于更大的模型。
import torch from transformers import AutoModelForCausalLM, AutoTokenizer model_name = "sshleifer/tiny-gpt2" tokenizer = AutoTokenizer.from_pretrained(model_name) model = AutoModelForCausalLM.from_pretrained(model_name) model.eval() prompt = "A language model is" input_ids = tokenizer(prompt, return_tensors="pt").input_ids with torch.no_grad(): outputs = model(input_ids) logits = outputs.logits next_token_logits = logits[:, -1, :] print(next_token_logits.shape)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
import
torch
from
transformers
AutoModelForCausalLM
,
AutoTokenizer
model_name
=
"sshleifer/tiny-gpt2"
tokenizer
.
from_pretrained
(
)
model
eval
prompt
"A language model is"
input_ids
return_tensors
"pt"
with
no_grad
:
outputs
logits
next_token_logits
[
-
]
shape
输出形状是:
[批次大小, 词汇表大小]
logits不是概率。要将logits转换为概率,使用softmax:
probs = torch.softmax(next_token_logits, dim=-1)
probs
softmax
dim
然而,你通常不需要显式计算概率。贪婪解码只需要最大logit的索引,这与概率最高的标记相同。
next_token = next_token_logits.argmax(dim=-1, keepdim=True) print(tokenizer.decode(next_token[0]))
next_token
argmax
keepdim
True
decode
这是最简单的解码策略。
贪婪解码
贪婪解码总是选择得分最高的标记。一个完整的贪婪解码函数可以写成如下形式:
torch.no_grad() def greedy_decode(model, tokenizer, prompt, max_new_tokens=30): input_ids = tokenizer(prompt, return_tensors="pt").input_ids for _ in range(max_new_tokens): outputs = model(input_ids) next_token_logits = outputs.logits[:, -1, :] next_token = next_token_logits.argmax(dim=-1, keepdim=True) input_ids = torch.cat([input_ids, next_token], dim=1) if next_token.item() == tokenizer.eos_token_id: break return tokenizer.decode(input_ids[0], skip_special_tokens=True)
def
greedy_decode
max_new_tokens
30
for
_
range
cat
if
item
==
eos_token_id
break
return
skip_special_tokens
贪婪解码具有确定性。给定相同的模型和提示,它会返回相同的输出。这在调试以及需要避免输出差异化的任务中非常有用。
其弱点在于局部最优的token不一定是最佳的延续。贪婪解码可能导致重复、过度选择常见短语,错失更有意思的延续内容。
温度采样
温度采样从通过温度参数对logits进行缩放后得到的概率分布中进行采样。下图展示了温度如何在不改变原始logits的情况下改变概率分布。相同的十个token得分被转换为概率三次:温度为0.5时一次,温度为1时一次,温度为2时一次。
相同的logits在不同温度下会产生不同的token概率分布。较低温度会使概率集中在最高得分token上,较高温度则会将概率分散到更多token上。
采样过程会从模型的概率分布中随机选择下一个token。温度参数控制该分布的尖锐程度或平坦程度。给定logits $\mathbf{z}$ 和温度 $T$,温度采样使用以下公式:
$$ \mathbf{p} = \operatorname{softmax}(\mathbf{z} / T) $$
较低温度会使分布 $\mathbf{p}$ 更尖锐,较高温度则更平坦。当温度趋近于零时,如果有一个token的logits唯一最高,采样行为会类似于贪婪解码。当温度过高时,logits之间的差异变得不那么重要,模型可能会频繁选择低概率token。
使用温度参数的采样循环如下所示:
@torch.no_grad() def temperature_decode(model, tokenizer, prompt, temperature=0.8, max_new_tokens=30): input_ids = tokenizer(prompt, return_tensors="pt").input_ids assert temperature > 0, "temperature must be positive" for _ in range(max_new_tokens): outputs = model(input_ids) # 对下一个token的logits应用温度参数 logits = outputs.logits[:, -1, :] / temperature # 将logits转换为概率并从分布中采样 probs = torch.softmax(logits, dim=-1) next_token = torch.multinomial(probs, num_samples=1) # 将下一个token追加到输入中用于下一次迭代 input_ids = torch.cat([input_ids, next_token], dim=1) if next_token.item() == tokenizer.eos_token_id: break return tokenizer.decode(input_ids[0], skip_special_tokens=True)
19
@
temperature_decode
temperature
0.8
assert
"temperature must be positive"
对下一个token的logits应用温度参数
/
将logits转换为概率并从分布中采样
multinomial
num_samples
将下一个token追加到输入中用于下一次迭代
温度本身并不是一个调节质量的旋钮。它影响的是随机性的程度。合适的取值取决于具体任务。事实抽取类任务通常需要较低的温度值,而头脑风暴和创意写作可能更适合使用较高的温度值。
Top-$k$ 采样
在上图中,以一个包含10个词元的分布作为示例。实际模型的词表可能包含数十万个词元,其中许多词元在特定上下文中概率极低。
Top-$k$ 采样仅保留概率最高的 $k$ 个词元,移除其他所有词元。其主要目的是防止模型采样到概率极低的词元。它不会避免对完整词表计算logits,但会减少实际采样的候选词元数量。
@torch.no_grad() def top_k_sample(logits, k): assert k > 0, "k 必须为正数" assert k <= logits.size(-1), "k 不得超过词表大小" # 获取 top-k 的 logits 及其索引 values, indices = torch.topk(logits, k) # 将 logits 转换为 top-k 候选词元的概率 probs = torch.softmax(values, dim=-1) # 根据概率从 top-k 索引中进行采样 sampled = torch.multinomial(probs, num_samples=1) # 使用收集的 top-k 索引恢复实际词元 ID next_token = indices.gather(-1, sampled) return next_token
top_k_sample
k
"k 必须为正数"
<=
size
"k 不得超过词表大小"
获取 top-k 的 logits 及其索引
values
indices
topk
将 logits 转换为 top-k 候选词元的概率
根据概率从 top-k 索引中进行采样
sampled
使用收集的 top-k 索引恢复实际词元 ID
gather
Top-$k$ 采样易于理解,但它使用固定数量的候选词元。有时模型非常确定,只有少数词元重要。有时许多词元都合理,此时固定 top-$k$ 截断可能不合适。这促使我们采用核心采样方法。
核心采样
核心采样(也称为 top-$p$ 采样)保留最小的词元集合,其累积概率至少为 $p$。例如,当 $p=0.9$ 时,它会保留那些共同占据 90% 概率质量的最可能词元。
torch.no_grad() def top_p_sampling(logits, temperature=1.0, k=0, p=0.9): """ 应用温度缩放、可选的 top-k 过滤和 top-p 过滤。接受一个 logits 张量并返回一个采样得到的词元 ID。 """ assert logits.dim() == 1, "logits 必须是一个一维张量" assert 0 < p <= 1, "p 必须在 (0, 1] 范围内" vocab_size = logits.size(0) # 应用温度缩放 logits = logits / temperature # 可选的 top-k 过滤 if k > 0 and k < vocab_size: topk_vals, topk_idx = torch.topk(logits, k) # 创建一个填充为 -inf 的掩码,将 top-k 的 logits 放在对应位置 new_logits = torch.full_like(logits, float('-inf')) new_logits[topk_idx] = topk_vals logits = new_logits # top-p(核心)过滤 sorted_logits, sorted_indices = torch.sort(logits, descending=True) sorted_probs = torch.softmax(sorted_logits, dim=-1) cumulative_probs = torch.cumsum(sorted_probs, dim=-1) # 需要移除的词元,但至少保留一个词元 remove = cumulative_probs > p remove[1:] = remove[:-1].clone() remove[0] = False sorted_logits = sorted_logits.masked_fill(remove, float('-inf')) # 采样 final_probs = torch.softmax(sorted_logits, dim=-1) sampled = torch.multinomial(final_probs, num_samples=1) next_token = sorted_indices.gather(-1, sampled) return next_token
20
21
22
23
24
25
26
27
28
29
31
32
33
34
35
36
37
38
top_p_sampling
1.0
p
0.9
""
"
应用温度缩放,可选的top-k过滤和top-p过滤。
接受一个1D的logits张量并返回一个采样后的token ID。
"logits必须是一个1D张量"
<
"p必须在(0, 1]范围内"
vocab_size
应用温度
可选的top-k过滤
and
topk_vals
topk_idx
创建一个填充为-inf的掩码,将top-k的logits放在对应位置
new_logits
full_like
float
'-inf'
new
top-p(nucleus)过滤
sorted_logits
sorted_indices
sort
descending
sorted_probs
cumulative_probs
cumsum
需要移除的token,但至少保留一个token
remove
clone
False
masked_fill
采样
final_probs
上述函数结合了温度采样、可选的top-$k$过滤和top-$p$过滤。将这些技术结合使用非常常见。它们的顺序很重要,因为温度缩放和过滤会影响下一个token采样的分布。top-$p$是自适应的:当模型非常确定时可能只保留少量token,而当分布较广时会保留更多token。
重复惩罚
自回归模型可能会陷入循环,其中某些token模式不断重复。添加重复惩罚会降低已经出现过的token的得分,从而降低这些token再次被选中的可能性。
一个简单的实现方式是将正logits除以惩罚系数,将负logits乘以惩罚系数:
@torch.no_grad() def apply_repetition_penalty(logits, generated_ids, penalty=1.1): assert logits.dim() == 2 and logits.size(0) == 1, ( "logits必须具有形状[1, vocab_size]" ) assert generated_ids.dim() == 2 and generated_ids.size(0) == 1, ( "generated_ids必须具有形状[1, sequence_length]" ) assert penalty >= 1.0, "惩罚系数必须至少为1" if penalty == 1.0: return logits logits = logits.clone() token_ids = set(generated_ids[0].tolist()) for token_id in token_ids: token_logit = logits[0, token_id] logits[0, token_id] = torch.where( token_logit > 0, token_logit / penalty, token_logit * penalty, ) return logits
apply_repetition_penalty
generated_ids
penalty
1.1
"logits必须具有形状[1, vocab_size]"
"generated_ids必须具有形状[1, sequence_length]"
=
"惩罚系数必须至少为1"
token_ids
set
tolist
token_id
token_logit
where
token_logit *
该函数设计得较为简单,假设批量大小为1。例如,相同token的多次出现不会增加惩罚系数。调用者需要自行决定generated_ids是否包含提示token、生成的token或两者都有。如果使用重复惩罚与top-$k$或nucleus采样结合,应先应用惩罚。实际生产实现通常处理更大的批量,并可能区分频率惩罚和存在惩罚。
重复惩罚可能有帮助,但也可能影响质量。某些词需要重复。代码、名称、引用和结构化格式通常需要精确重复。仅在重复确实成为问题时才使用此控制。
束搜索
贪婪解码只保留一个候选序列。束搜索保留多个候选。每一步,它用可能的下一个token扩展每个候选,并保留得分最高的序列。
束搜索在存在明确的序列级目标时非常有用,例如在早期的序列到序列系统中的翻译任务。对于开放式聊天生成,束搜索通常会生成通用文本,因为它倾向于选择高概率的延续。
一个最小的束搜索循环如下所示:
@torch.no_grad() def beam_search(model, tokenizer, prompt, num_beams=3, max_new_tokens=20): input_ids = tokenizer(prompt, return_tensors="pt").input_ids beams = [(0.0, input_ids)] # 每次迭代为每个束添加一个标记 for _ in range(max_new_tokens): candidates = [] # 用每个束的num_beams个最高得分的下一个标记扩展每个束 for score, token_ids in beams: outputs = model(token_ids) logits = outputs.logits[:, -1, :] log_probs = torch.log_softmax(logits, dim=-1) values, indices = torch.topk(log_probs, num_beams, dim=-1) for value, token_id in zip(values[0], indices[0]): next_ids = torch.cat([token_ids, token_id.view(1, 1)], dim=1) candidates.append((score + value.item(), next_ids)) # 为下一次迭代只保留最佳的num_beams个候选 beams = sorted( candidates, key=lambda candidate: candidate[0], reverse=True )[:num_beams] # 只返回最佳束作为最终输出 best_score, best_token_ids = beams[0] return tokenizer.decode(best_token_ids[0], skip_special_tokens=True)
beam_search
num_beams
beams
0.0
每次迭代为每个束添加一个标记
candidates
用每个束的num_beams个最高得分的下一个标记扩展每个束
score
log_probs
log_softmax
value
zip
next_ids
view
append
+
为下一次迭代只保留最佳的num_beams个候选
sorted
key
lambda
candidate
reverse
只返回最佳束作为最终输出
best_score
best_token_ids
此实现故意保持简洁。实际实现应通过序列长度对得分进行归一化,处理序列结束标记,并通过使用KV缓存避免重新计算整个前缀。
束搜索计算成本较高:循环会降低生成速度,束的数量会增加内存使用。如果使用四个束,模型需要跟踪四个延续。这与普通采样相比会增加计算量和缓存内存。因此,束搜索通常在大型语言模型服务中被避免。
停止条件
生成过程必须在某个时刻停止。最简单的停止条件是新生成标记的最大数量。另一个常见条件是模型的序列结束标记。通常语言模型的词表中包含一些特殊标记,序列结束标记就是其中之一。
上面的贪心解码示例可以修改为接受任意停止标记:
@torch.no_grad() def greedy_decode_with_stop(model, tokenizer, prompt, stop_token_id, max_new_tokens=30): input_ids = tokenizer(prompt, return_tensors="pt").input_ids for _ in range(max_new_tokens): outputs = model(input_ids) next_token_logits = outputs.logits[:, -1, :] next_token = next_token_logits.argmax(dim=-1, keepdim=True) input_ids = torch.cat([input_ids, next_token], dim=1) if next_token.item() == stop_token_id: break return tokenizer.decode(input_ids[0], skip_special_tokens=True)
greedy_decode_with_stop
stop_token_id
结构化输出约束
某些应用场景需要模型生成特定格式的输出,例如 JSON、SQL 查询或固定列表中的值。一种方法是提示模型并*期望*其遵循格式。更有效的方法是使用约束解码。
其核心思想是屏蔽会导致输出无效的标记。例如,如果输出必须是三个标签中的一个,可以只对这些标签进行评分:
@torch.no_grad() def choose_label(model, tokenizer, prompt, labels): assert labels, "labels must not be empty" # 运行提示以获取词汇表上的logits input_ids = tokenizer(prompt, return_tensors="pt").input_ids outputs = model(input_ids) logits = outputs.logits[:, -1, :] # 假设每个标签在此上下文中恰好是一个标记,对每个标签进行评分 label_scores = [] for label in labels: # 在标签字符串中包含任何必需的前导空白 label_ids = tokenizer.encode(label, add_special_tokens=False) assert len(label_ids) == 1, f"{label!r} must encode to exactly one token" label_scores.append(logits[0, label_ids[0]].item()) best_score, best_label = max(zip(label_scores, labels)) return best_label
choose_label
labels
"labels must not be empty"
运行提示以获取词汇表上的logits
假设每个标签在此上下文中恰好是一个标记,对每个标签进行评分
label_scores
label
在标签字符串中包含任何必需的前导空白
label_ids
encode
add_special_tokens
len
f
"{label!r} must encode to exactly one token"
best_label
max
该示例仅处理提示后编码为单个标记的标签。分词可能依赖于上下文(包括前导空白),因此调用者必须相应地构建标签。多标记标签需要对完整标记序列进行评分或对每个解码步骤进行约束。结构化解码将输出要求转化为标记约束。更高级的系统使用语法、前缀树或有限状态机来决定每一步的合法标记。
约束解码可以提高可靠性,但也可能降低推理速度。系统必须在每一步计算并应用标记掩码。与所有推理技术一样,您应同时评估质量和性能。
进一步阅读
以下是一些可能有用的资源:
- Softmax函数,维基百科。这是了解logits如何转换为概率的有用参考资料。温度采样是对softmax输入的直接修改,在归一化前将$\mathbf{z}$替换为$\mathbf{z}/T$。
- 束搜索,维基百科。该页面将束搜索描述为一种通用的启发式搜索算法。在语言生成中,束搜索会保留多个候选续写,而不仅仅是单个最佳下一个标记。
- Holtzman 等人的论文《神经文本退化的奇怪案例》。该论文解释了为什么最大似然解码方法(如贪婪解码和束搜索)会产生平淡或重复的文本,并介绍了核采样作为开放式生成的实用替代方案。
- 对比解码:将开放式文本生成作为优化问题,作者:李等。该论文提出了一种解码方法,通过比较专家语言模型与小型业余模型的得分差异,优先选择流畅且信息量大的续写内容。
- 无需微调的结构化NLP任务语法约束解码,作者:耿等。该论文探讨了如何利用形式文法约束语言模型的标记选择,使生成输出符合指定结构。
- 从语言模型生成结构化输出:基准测试与研究,作者:耿等。该论文研究了针对JSON模式等结构化输出的约束解码方法,特别适用于需要可靠机器可读输出而非自由文本的场景。
总结
在本章中,你了解到解码是从事务logits中选择标记的过程。贪婪解码是确定性且简单的。温度采样、top-$k$采样和核采样引入了可控的随机性。集束搜索跟踪多个候选方案但会增加推理成本。重复惩罚和停止条件有助于控制输出长度和行为。结构化输出约束可以使模型输出更易于在应用中使用。
在下一章中,你将学习如何衡量推理性能,从而用实际数值而非直觉来比较这些选择。
更多相关内容
- R语言中的逻辑、流程控制与函数
- 如何通过…控制神经网络模型容量
- 如何控制神经网络训练的稳定性…
- 通过创建定向机器…列表来掌握控制权
- 理解LangChain大语言模型输出解析器
- 训练单输出多元线性回归…