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 个最好的。

用个简单的图对比一下两种策略:

graph TD Start[开始: <start>] --> Step1[Step 1] Step1 --> T1a[today: 0.25] Step1 --> T1b[the: 0.20] Step1 --> T1c[weather: 0.18] T1a --> Step2a[Step 2] T1b --> Step2b[Step 2] T1c --> Step2c[Step 2] Step2a --> T2a[is: 0.30] Step2a --> T2b[today: 0.15] Step2b --> T2c[weather: 0.28] Step2c --> T2d[is: 0.25] style T1a stroke:#2E86AB,stroke-width:3px style T1b stroke:#E63946,stroke-width:3px style T1c stroke:#E63946,stroke-width:3px style T2a stroke:#2E86AB,stroke-width:3px style T2c stroke:#2E86AB,stroke-width:3px style T2d stroke:#2E86AB,stroke-width:3px

蓝色粗线是贪婪搜索——每步只选概率最高的,一条路走到黑。红色和绿色是束搜索保留的候选路径,多走几步后可能有更优的选择。

本质上是个广度优先搜索的变体,但用 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 会带来两个问题:

  1. 计算量增加太多。每次要做 4 倍的解码,batch size 受限,GPU 利用率上不去。
  2. 可能选到训练数据里的"稀有模式",反而泛化性差。

具体数据:在 IWSLT14 德英翻译任务上,我的模型表现是:

beam_widthBLEU Score推理时间 (句/秒)
1 (贪婪)31.2142
233.878
435.141
835.321
1635.211

从 4 到 8,BLEU 只涨了 0.2,但速度降了一半。得不偿失。

beam width 对性能的影响:BLEU Score(蓝色)和推理速度(红色)的双坐标轴图显示,beam_width=4 是个甜点区——质量提升明显且速度可接受

坑 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 的启发式方法。理解它的原理很简单,但用好它需要实践:

  1. beam_width 不要盲目追求大,2-4 通常是性价比最高的区间
  2. 一定要做长度归一化,否则短序列 bias 会毁了你的结果
  3. 重复问题要处理,简单的重复惩罚就能解决大部分情况
  4. 不同任务要不同对待,翻译用束搜索,对话用采样

对于我这种"先把东西做出来"的人,束搜索是个好工具——它比贪婪搜索聪明,又比穷举现实。就像人生一样,有时候放弃"最优解",保留几个"不错的候选",反而能走得更远。

最后一句实话:束搜索能提升生成质量,但它救不了训练糟糕的模型。模型本身的训练质量才是根本,束搜索只是锦上添花。如果模型训练得不好,束搜索只是帮你找到"最好的错误答案"而已。

版权声明: 本文首发于 指尖魔法屋-AI 束搜索折腾手记https://blog.thinkmoon.cn/post/349-ai-beam-search-greedy-optimal-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!