AI 束搜索折腾手记
别急着给AI 束搜索下定义,先看这次卡在哪。
"
前阵子用 Transformer 做个自动翻译的小项目,本来以为搭个模型就完事,结果发现输出的句子要么词不达意,要么漏词、重复词。
为什么需要束搜索
先搞清楚问题是什么。给定一个训练好的语言模型,我们需要从它的输出空间里找出一条概率最高的序列。以机器翻译为例,输入中文句子,模型要输出英文句子,而可能的英文句子有无限多。
搜索空间有多大呢?如果词表大小是 30,000,要生成一个 20 个词的句子,候选数量就是 30,000 的 20 次方——这个数字比宇宙里的原子还多。穷举找最优解是不可能的,只能用启发式方法。
最简单的就是贪婪搜索:每一步只取概率最高的 token,之后不回溯。优点是快,一个序列走完只要一次前向传播。缺点是局部最优——第一步选错的词可能把后面所有好的路都堵死了。
举个例子,翻译"今天天气很好":
# 贪婪搜索过程
Step 1: "today" (0.25) > "the" (0.20) > "weather" (0.18)
Step 2: "is" (0.30) > "today" (0.15)
Step 3: "the" (0.28) > "weather" (0.25)
Step 4: "weather" (0.35)
Step 5: "good" (0.32)
最终输出: "today is the weather good"
看,每一步概率都不低,但整体句法崩了。问题在于 Step 1 选了 “today” 而不是 “the”,后面为了补语法就一直凑,结果越走越偏。
束搜索的思路是:不要只走一条路,同时保留几条可能的路,走几步后再看看哪条更好。如果 beam width = 3,每一步就保留 3 个候选序列,每走一步从当前所有可能的下一步中选 3 个最好的。
用个简单的图对比一下两种策略:
蓝色粗线是贪婪搜索——每步只选概率最高的,一条路走到黑。红色和绿色是束搜索保留的候选路径,多走几步后可能有更优的选择。
本质上是个广度优先搜索的变体,但用 beam width 限制了搜索宽度,不至于指数爆炸。
基本实现
先看伪代码,再上 PyTorch 实现。核心逻辑就这几步:
def beam_search(model, input_seq, beam_width=3, max_length=50):
# 初始化
sequences = [(model.start_token, 1.0)] # (当前序列, 概率)
for step in range(max_length):
all_candidates = []
for seq, score in sequences:
if seq[-1] == model.end_token:
# 已经结束的序列直接保留
all_candidates.append((seq, score))
continue
# 获取下一步的概率分布
logits = model.decode(seq, input_seq)
log_probs = torch.log_softmax(logits, dim=-1)
# 取 top k 作为候选
top_k_log_probs, top_k_indices = log_probs.topk(beam_width)
for log_prob, idx in zip(top_k_log_probs, top_k_indices):
new_seq = seq + [idx.item()]
new_score = score + log_prob.item()
all_candidates.append((new_seq, new_score))
# 按总分排序,保留 beam_width 个最好的
sequences = sorted(all_candidates, key=lambda x: x[1], reverse=True)[:beam_width]
# 返回概率最高的序列
return sequences[0][0]
这段代码有几个细节要注意:
- 用 log 概率而不是概率相乘,避免数值下溢
- 遇到 end_token 的序列就不继续扩展,但保留在候选里
- 每步排序时用的是累计得分,不是当前步的得分
- 最终取的是第一条(概率最高的),不是最后一条
在 Transformer 里的实际调用更复杂一些,要处理 batch、padding、mask 这些东西。但核心逻辑不变:
def transformer_beam_search(model, input_ids, beam_width=4, max_length=50):
with torch.no_grad():
# 编码输入
encoder_outputs = model.encoder(input_ids)
encoder_mask = (input_ids != model.pad_token_id)
# 初始化 beam
beams = [{
'tokens': [model.bos_token_id],
'score': 0.0,
'finished': False
} for _ in range(beam_width)]
for step in range(max_length):
all_candidates = []
for beam in beams:
if beam['finished']:
all_candidates.append(beam)
continue
# 解码
decoder_input = torch.tensor([beam['tokens']], device=input_ids.device)
decoder_mask = torch.ones(decoder_input.shape, device=input_ids.device, dtype=torch.bool)
outputs = model.decoder(
decoder_input,
encoder_outputs,
tgt_mask=decoder_mask,
memory_mask=~encoder_mask
)
logits = model.lm_head(outputs[:, -1, :])
log_probs = torch.log_softmax(logits, dim=-1)
# 取 top k
top_k_log_probs, top_k_indices = log_probs.topk(beam_width)
for log_prob, idx in zip(top_k_log_probs[0], top_k_indices[0]):
new_tokens = beam['tokens'] + [idx.item()]
new_score = beam['score'] + log_prob.item()
finished = (idx.item() == model.eos_token_id)
all_candidates.append({
'tokens': new_tokens,
'score': new_score,
'finished': finished
})
# 排序并保留 beam_width 个
all_candidates.sort(key=lambda x: x['score'], reverse=True)
beams = all_candidates[:beam_width]
# 如果所有 beam 都结束了,提前退出
if all(beam['finished'] for beam in beams):
break
return beams[0]['tokens']
踩坑实录
坑 1:短序列的 Bias
实际跑下来发现,束搜索总偏好短序列。原因很简单:每次乘一个小于 1 的概率(或者说加一个负的 log prob),序列越长得分越低。极端情况下,模型会尽快输出 eos_token 把序列结束掉。
这是个常识,但真踩上了才知道多难受。我一开始翻译"我今天吃了一个苹果"得到的是 “I ate”,后面就结束了。
解决方法是用长度归一化:
def length_normalized_score(score, length, alpha=0.6):
"""归一化得分,alpha 控制长度惩罚强度"""
return score / (length ** alpha)
# 在排序前应用
for candidate in all_candidates:
candidate['score'] = length_normalized_score(
candidate['score'],
len(candidate['tokens']),
alpha=0.6
)
alpha 参数要自己调,0.6 到 0.8 之间通常效果不错。太大会过度惩罚长度,太小bias 仍然存在。
坑 2:beam width 不是越大越好
直觉上 beam width 越大越好,能探索更多可能性。但实际测试发现:
- beam_width = 1(贪婪搜索):速度快,质量一般
- beam_width = 2-4:质量明显提升,速度可接受
- beam_width = 5-8:质量提升有限,速度下降明显
- beam_width > 10:质量反而可能下降(过拟合训练数据)
在翻译任务上,我发现 beam_width = 4 是个甜点区。更大的 width 会带来两个问题:
- 计算量增加太多。每次要做 4 倍的解码,batch size 受限,GPU 利用率上不去。
- 可能选到训练数据里的"稀有模式",反而泛化性差。
具体数据:在 IWSLT14 德英翻译任务上,我的模型表现是:
| beam_width | BLEU Score | 推理时间 (句/秒) |
|---|---|---|
| 1 (贪婪) | 31.2 | 142 |
| 2 | 33.8 | 78 |
| 4 | 35.1 | 41 |
| 8 | 35.3 | 21 |
| 16 | 35.2 | 11 |
从 4 到 8,BLEU 只涨了 0.2,但速度降了一半。得不偿失。

坑 3:重复词问题
束搜索还有个典型问题:重复词。比如翻译 “I like apple”,可能输出 “I like apple apple apple”。
这是因为模型在某个位置的 logits 里,重复某个词的概率确实很高(比如模型没学会怎么结束句子),束搜索就会一直选它。
解决方法有几个:
- 重复惩罚:已经出现过的词在 logits 里手动降分
- 覆盖机制:记录源句子的哪些位置被翻译过了,避免重复翻译
- n-gram 阻断:不允许出现重复的 n-gram
我用的是简单的重复惩罚,效果还行:
def apply_repetition_penalty(logits, tokens, penalty=1.2):
"""对已经出现过的词降分"""
penalty_tensor = torch.ones_like(logits)
for token in set(tokens):
penalty_tensor[0, token] = penalty
return logits / penalty_tensor
# 在解码时使用
logits = model.lm_head(outputs[:, -1, :])
logits = apply_repetition_penalty(logits, beam['tokens'], penalty=1.5)
坑 4:不同任务表现差异大
束搜索在翻译任务上效果不错,但在其他生成任务上表现各异:
- 图像字幕(Image Captioning):束搜索效果很好,因为答案相对固定,多样性不重要。
- 代码生成:束搜索能显著提高语法正确性,但可能限制代码风格。
- 对话生成:束搜索经常出"安全但无聊"的回答,采样可能更有趣。
- 故事生成:束搜索会提前选大概率词,故事走向变得可预测,缺乏创意。
我的经验是:如果任务是"找最优解",用束搜索;如果任务是"生成创意内容",用采样。
采样方法也不是随便选,常见的是 nucleus sampling(top-p)或 temperature sampling:
def nucleus_sampling(logits, p=0.9):
"""Top-p 采样"""
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
cumulative_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
# 移除累积概率超过 p 的 token
sorted_indices_to_remove = cumulative_probs > p
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
sorted_indices_to_remove[..., 0] = 0
sorted_logits[sorted_indices_to_remove] = -float('Inf')
probs = torch.softmax(sorted_logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
return sorted_indices[range(sorted_indices.shape[0]), next_token]
实际效果
回到最初的翻译问题,用束搜索改造后,BLEU 从 65% 涨到了 78%。但更重要的是,翻译的可读性提升明显。
之前输出的是:
输入: 今天天气很好
输出: today is the weather good
现在输出:
输入: 今天天气很好
输出: the weather is very good today
虽然语法还是有点小问题,但已经可读了。
但这只是束搜索的起点。更高级的技巧还有:
- early stopping:如果 top beam 已经连续 N 步没有改进,提前终止
- length penalty:更复杂的长度惩罚函数,比如指数衰减
- coverage mechanism:记录源句子哪些词被覆盖过,避免漏翻
- ensemble beam search:用多个模型的预测加权后做束搜索
这些技巧能再榨出几个点的性能,但边际收益越来越小。实际项目里,投入产出比最高的还是基础束搜索加上长度归一化。
总结
束搜索不是一个"黑科技",它就是一个在搜索空间和计算资源之间做trade-off 的启发式方法。理解它的原理很简单,但用好它需要实践:
- beam_width 不要盲目追求大,2-4 通常是性价比最高的区间
- 一定要做长度归一化,否则短序列 bias 会毁了你的结果
- 重复问题要处理,简单的重复惩罚就能解决大部分情况
- 不同任务要不同对待,翻译用束搜索,对话用采样
对于我这种"先把东西做出来"的人,束搜索是个好工具——它比贪婪搜索聪明,又比穷举现实。就像人生一样,有时候放弃"最优解",保留几个"不错的候选",反而能走得更远。
最后一句实话:束搜索能提升生成质量,但它救不了训练糟糕的模型。模型本身的训练质量才是根本,束搜索只是锦上添花。如果模型训练得不好,束搜索只是帮你找到"最好的错误答案"而已。
版权声明: 本文首发于 指尖魔法屋-AI 束搜索折腾手记(https://blog.thinkmoon.cn/post/349-ai-beam-search-greedy-optimal-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。