1. 束搜索算法基础解析
束搜索(Beam Search)是序列生成任务中广泛使用的启发式搜索算法,在机器翻译、语音识别、文本摘要等场景表现优异。与贪心搜索每次只保留一个最优候选不同,束搜索会在每个时间步保留k个最有可能的候选序列(称为束宽),通过这种折中方案平衡计算成本和结果质量。
1.1 算法核心流程
典型束搜索实现包含以下关键步骤:
- 初始化:将起始符号(如
<start>)作为初始序列,计算其对数概率(通常为0) - 扩展候选:
- 对当前所有候选序列,预测下一个词的概率分布
- 对每个候选序列,保留概率最高的k个扩展结果
- 筛选存活:
- 合并所有候选序列的扩展结果
- 全局保留总概率最高的k个序列
- 终止判断:
- 当任何候选序列生成结束符号(如
<end>)时,将其移出候选池 - 当所有候选序列都终止或达到最大长度时停止搜索
- 当任何候选序列生成结束符号(如
python复制def beam_search(decoder, start_token, k=5, max_len=100):
candidates = [([start_token], 0.0)] # (sequence, log_prob)
for _ in range(max_len):
new_candidates = []
for seq, score in candidates:
if seq[-1] == EOS_TOKEN: # 已终止的序列不再扩展
new_candidates.append((seq, score))
continue
next_probs = decoder.predict(seq) # 获取下一个词的概率分布
top_k = next_probs.argsort()[-k:] # 选择top-k候选
for token in top_k:
new_seq = seq + [token]
