跳转至

大模型手撕-核心题

实现贪心解码函数 greedy_decode(logits),逐步选取概率最大的 token 拼接成序列。

1.1 原理

贪心解码(Greedy Decoding)是自回归生成中最简单、最确定性的策略。在每一个生成步骤,模型输出一个形状为 (vocab_size,) 的 logits 张量,经过 softmax 后得到概率分布。贪心解码直接选择概率最高的 token 作为本步的输出,将其拼接到已生成序列的末尾,然后该新序列作为下一步的输入继续生成,直到遇到结束符(EOS)或达到预设的最大长度。

贪心解码不进行任何随机采样,因此生成结果是完全确定的(假设模型固定且无随机性 dropout 等)。其优点是速度极快,无需维护候选集;缺点是缺乏全局最优性,容易陷入局部最优,在长文本生成中常导致重复和退化(如循环输出相同短语)。

1.2 实现

import torch

def greedy_decode(model, input_ids, max_length=50, eos_token_id=None):
    """
    贪心解码自回归生成。

    Args:
        model: 自回归语言模型,调用 model(input_ids) 返回 logits (batch_size, seq_len, vocab_size)。
               假设输入 batch_size=1,已处理 prompt。
        input_ids: (1, prompt_len) 初始提示词 token ID 张量。
        max_length: 最大生成总长度(包括 prompt)。
        eos_token_id: 结束符 token ID,遇到则停止。

    Returns:
        generated: 生成的完整 token ID 列表(包括 prompt)。
    """
    model.eval()
    generated = input_ids.tolist()[0]  # 转为Python列表
    with torch.no_grad():
        for _ in range(max_length - len(generated)):
            # 将当前序列转为tensor输入模型
            inputs = torch.tensor([generated], device=input_ids.device)
            outputs = model(inputs)
            # 取最后一个位置的logits
            next_logits = outputs[0, -1, :]  # shape: (vocab_size,)
            # 贪心选择最大概率的token
            next_token_id = torch.argmax(next_logits).item()
            generated.append(next_token_id)
            # 如果遇到eos则停止
            if eos_token_id is not None and next_token_id == eos_token_id:
                break
    return generated

1.3 关键细节

  • KV Cache 优化:上述实现每次都将整个历史序列重新输入模型,计算量随生成长度平方增长。生产环境中必须使用 KV Cache:在每步仅输入最新 token,利用缓存的历史 Key/Value 张量进行注意力计算。实现 KV Cache 需要模型支持 past_key_values 参数。这里为了清晰展示解码逻辑,采用简化形式,但面试中应主动指出此优化。

  • 模型调用约定:通常自回归模型的 forward 接受 input_idspast_key_values,返回 logits 和新的 past_key_values。上述函数可扩展为接受 past_key_values 参数并返回更新后的缓存。

  • EOS 处理:遇到 EOS 时应当立即停止,避免生成无意义的后续 token。同时应确保生成的序列中不会因为 EOS 后的 token 而被截断。

  • Tensor 操作:使用 torch.argmax 直接获取索引,高效简洁。

1.4 贪心解码的局限与改进方向

  • 无法产生多样性输出。

  • 容易陷入重复循环,尤其在开放域对话中。

  • 对于需要探索的任务,通常与温度缩放、Top-k/Top-p 结合使用,或在束搜索中作为最终评分基准。


实现基础束搜索 beam_search(logits_fn, k, max_len),每步扩展并保留 k 条最优序列。

2.1 原理

束搜索(Beam Search)在每一步维护一个大小为 k(束宽)的候选序列集合,而不是仅保留一个最佳序列。具体流程:

  1. 初始化:起始序列为 prompt,得分(对数概率和)为 0。

  2. 对当前束中的每一条序列,模型给出下一步的 logits,计算所有可能的下一个 token 的对数概率。

  3. 将所有 (当前序列, 下一个token) 的组合视为新候选,共 k * |V| 个(实际常只取每个序列的 top-k 候选以加速),每个候选的得分为 原序列得分 + log P(token|sequence)

  4. 按得分降序排列,保留得分最高的 k 个候选作为下一轮的束。

  5. 重复直到满足终止条件:所有序列都生成了 EOS 或达到最大长度。

  6. 最终从所有已完成和未完成的序列中选择得分最高的一条作为输出。

束搜索通过维护多条路径,能在一定程度上避免贪心的局部最优,找到全局更好的序列。

2.2 实现

def beam_search(model, input_ids, beam_width=3, max_length=50, eos_token_id=None):
    """
    基础束搜索。

    Args:
        model: 自回归模型,返回 logits。
        input_ids: (1, prompt_len) 初始提示词。
        beam_width (k): 束宽。
        max_length: 最大生成长度(包括prompt)。
        eos_token_id: EOS token ID。

    Returns:
        best_sequence: (list) 得分最高的token序列。
    """
    device = input_ids.device
    prompt = input_ids[0].tolist()
    # 束中的每个元素为 (tokens_list, cumulative_log_prob)
    beams = [(prompt, 0.0)]
    finished = []  # 已完成的序列

    for step in range(max_length - len(prompt)):
        if not beams:  # 所有束都已完成
            break
        all_candidates = []
        # 为每个束计算下一步 logits
        for tokens, score in beams:
            if eos_token_id is not None and tokens[-1] == eos_token_id:
                finished.append((tokens, score))
                continue
            input_tensor = torch.tensor([tokens], device=device)
            outputs = model(input_tensor)
            next_logits = outputs[0, -1, :]  # (vocab_size,)
            log_probs = torch.log_softmax(next_logits, dim=-1)
            # 为效率,可仅考虑 top-k 个候选,此处取全部
            topk_log_probs, topk_indices = torch.topk(log_probs, beam_width, dim=-1)
            for i in range(beam_width):
                next_token = topk_indices[i].item()
                new_score = score + topk_log_probs[i].item()
                new_tokens = tokens + [next_token]
                all_candidates.append((new_tokens, new_score))

        if not all_candidates:
            break

        # 按得分降序排序,保留 beam_width 个最佳
        ordered = sorted(all_candidates, key=lambda x: x[1], reverse=True)
        beams = ordered[:beam_width]

    # 合并已完成和剩余束
    final_candidates = finished + beams
    # 选取得分最高的
    best_tokens, best_score = max(final_candidates, key=lambda x: x[1])
    return best_tokens

2.3 重要说明

  • 计算效率:上述实现每步对束中的每条序列独立调用模型,没有利用 batch。实际工程中应将束中所有序列 batch 化:将 beams 中所有 tokens 堆叠为 (beam_width, cur_len) 的张量,一次前向得到所有序列的下一步 logits,然后扩展、排序、选择。这能大幅提升 GPU 利用率。

  • 对数概率:使用 log_softmax 将乘法转换为加法,避免浮点下溢。

  • EOS 处理:当某序列生成 EOS 时,将其移入 finished 列表,并停止对其扩展,但不影响其他序列。

  • 束宽选择:beam_width 通常取 3~10,过大计算量剧增,且可能引入更多噪音。

2.4 局限性

  • 倾向于产生较短序列(因为对数概率为负,累加导致长序列得分低),需要结合长度惩罚。

  • 多样性不足:多个束可能收敛到相似的高分序列,缺乏差异性。


实现带长度惩罚的束搜索,每条序列得分除以 len^α 后排序。

3.1 原理

长度惩罚(Length Penalty)用于缓解束搜索偏好短序列的问题。对于序列长度 L(通常只计生成部分),修正后的得分 score' = score / L^α,其中 α 是长度惩罚系数(常为 0.6~1.0)。排序时使用修正得分,使得长度适中的序列更可能胜出。

3.2 实现

在基础束搜索的基础上,修改排序逻辑和最终选择逻辑。

def beam_search_len_penalty(model, input_ids, beam_width=3, max_length=50,
                            eos_token_id=None, alpha=0.7):
    device = input_ids.device
    prompt_len = input_ids.shape[1]
    prompt = input_ids[0].tolist()
    beams = [(prompt, 0.0)]  # (tokens, raw_score)
    finished = []

    for step in range(max_length - prompt_len):
        if not beams:
            break
        all_candidates = []
        for tokens, score in beams:
            if eos_token_id is not None and tokens[-1] == eos_token_id:
                finished.append((tokens, score))
                continue
            input_tensor = torch.tensor([tokens], device=device)
            outputs = model(input_tensor)
            next_logits = outputs[0, -1, :]
            log_probs = torch.log_softmax(next_logits, dim=-1)
            topk_log_probs, topk_indices = torch.topk(log_probs, beam_width)
            for i in range(beam_width):
                next_token = topk_indices[i].item()
                new_score = score + topk_log_probs[i].item()
                new_tokens = tokens + [next_token]
                all_candidates.append((new_tokens, new_score))

        if not all_candidates:
            break

        # 计算长度惩罚得分并排序
        def len_pen_score(tokens, raw_score):
            gen_len = len(tokens) - prompt_len
            return raw_score / (gen_len ** alpha) if gen_len > 0 else raw_score

        # 为每个候选添加惩罚得分
        scored_candidates = [(tokens, raw_score, len_pen_score(tokens, raw_score))
                             for tokens, raw_score in all_candidates]
        # 按惩罚得分降序
        scored_candidates.sort(key=lambda x: x[2], reverse=True)
        beams = [(tokens, raw_score) for tokens, raw_score, _ in scored_candidates[:beam_width]]

    # 同样用长度惩罚选择最终最佳
    final = finished + beams
    best = max(final, key=lambda x: len_pen_score(x[0], x[1]))
    return best[0]

3.3 参数 α 的选择

  • α = 0:等同于无惩罚。

  • α = 1.0:强惩罚,倾向最短序列。

  • 常用值 0.6~0.8,在翻译、摘要等任务中微调。

3.4 注意事项

  • 长度计算方式:应只计算生成部分的长度,不包括 prompt。否则 prompt 长度影响所有序列的惩罚,不公平。

  • 惩罚应在最终选择时也应用,而不仅仅是每步排序。因为最终比较的是不同长度的已完成序列。


实现束搜索的提前终止逻辑,所有束生成 EOS 时停止,并正确处理已完成序列。

4.1 原理

提前终止(Early Stopping)是束搜索的标准优化。当束中所有活跃序列都已完成(即最后一个 token 都是 EOS),继续生成无意义,应立即终止。同时,已完成序列需要被安全存放,并在最终评选中参与竞争。

实现的关键是维护 beamsfinished 两个列表:每步将遇到 EOS 的序列移入 finished,并从 beams 中移除;当 beams 为空时终止循环。

4.2 实现

在基础束搜索代码中,我们已实现了提前终止的核心逻辑。以下是更健壮的版本,明确处理活跃束与完成束的分离。

def beam_search_early_stop(model, input_ids, beam_width=3, max_length=50, eos_token_id=None):
    device = input_ids.device
    prompt = input_ids[0].tolist()
    prompt_len = len(prompt)
    # active_beams 存储 (tokens, score)
    active_beams = [(prompt, 0.0)]
    finished_beams = []

    for step in range(max_length - prompt_len):
        if not active_beams:
            break
        next_candidates = []
        for tokens, score in active_beams:
            # 注意:这里假设 tokens 的最后一个不会是EOS,因为我们在上一轮已经移除了
            input_tensor = torch.tensor([tokens], device=device)
            outputs = model(input_tensor)
            next_logits = outputs[0, -1, :]
            log_probs = torch.log_softmax(next_logits, dim=-1)
            topk_log_probs, topk_indices = torch.topk(log_probs, beam_width)
            for i in range(beam_width):
                next_token = topk_indices[i].item()
                new_score = score + topk_log_probs[i].item()
                new_tokens = tokens + [next_token]
                if eos_token_id is not None and next_token == eos_token_id:
                    finished_beams.append((new_tokens, new_score))
                else:
                    next_candidates.append((new_tokens, new_score))

        # 从扩展后的候选(不含EOS结束的)中选出 beam_width 个继续
        next_candidates.sort(key=lambda x: x[1], reverse=True)
        active_beams = next_candidates[:beam_width]

    # 将所有序列合并,选取最佳(如果希望,可应用长度惩罚等)
    all_sequences = finished_beams + active_beams
    all_sequences.sort(key=lambda x: x[1], reverse=True)
    best_tokens, _ = all_sequences[0]
    return best_tokens

4.3 边界情况

  • 如果 prompt 本身已经以 EOS 结尾?不应发生,但可预先检查。

  • 如果某束刚刚生成 EOS,但在 next_candidates 中已被剔除,则 active_beams 数量减少,这是正确的。

  • 可能出现一种情况:所有活跃束都扩展后,next_candidates 数量少于 beam_width,此时 active_beams 会缩减,甚至变为 0,循环退出。

4.4 优化建议

  • 批量前向:将活跃束中的所有 tokens 整理为一批,一次前向,可大幅提速。

5.1 原理

多样束搜索(Diverse Beam Search, DBS)解决普通束搜索多样性不足的问题。它将 k 个束分成 g 组,每组独立进行束搜索,同时引入组间差异惩罚:若某个组选择的 token 与之前组(按顺序)的选择相同,则对该 token 的得分施加惩罚。这使得各组被迫探索不同的生成方向。

典型做法:

  • k 分成 g 组,每组有 k/g 个束。

  • 在第 t 步,依次处理每组:对于组 i,计算所有候选的原始得分,然后对于候选中的每个 token,如果该 token 与之前组(0..i-1)在同样位置 t 生成的 token 重复,则减去惩罚值 δ(或乘以惩罚因子)。

  • 惩罚强度 δ 控制多样性程度。

5.2 实现

def diverse_beam_search(model, input_ids, beam_width=5, groups=2, diversity_strength=0.5,
                        max_length=50, eos_token_id=None):
    """
    diverse beam search
    Args:
        beam_width: 总束宽,需被groups整除。
        groups: 分组数。
        diversity_strength: 惩罚强度(得分减去该值乘以重复次数)。
    """
    assert beam_width % groups == 0
    group_size = beam_width // groups
    prompt = input_ids[0].tolist()
    prompt_len = len(prompt)
    device = input_ids.device

    # 为每组初始化束
    group_beams = [[(prompt, 0.0)] for _ in range(groups)]
    finished = [[] for _ in range(groups)]

    for step in range(max_length - prompt_len):
        # 存储本步各组选出的token,用于后续组计算惩罚
        chosen_tokens_per_group = []  # list of list of token ids
        new_group_beams = []
        for g in range(groups):
            candidates = []
            # 收集本组当前束的所有扩展
            for tokens, score in group_beams[g]:
                if eos_token_id is not None and tokens[-1] == eos_token_id:
                    finished[g].append((tokens, score))
                    continue
                input_t = torch.tensor([tokens], device=device)
                outputs = model(input_t)
                next_logits = outputs[0, -1, :]
                log_probs = torch.log_softmax(next_logits, dim=-1)
                # 为加速,取 topk*2 以便惩罚后仍有足够候选
                topk_log_probs, topk_indices = torch.topk(log_probs, beam_width * 2)
                for i in range(topk_indices.shape[0]):
                    next_token = topk_indices[i].item()
                    new_score = score + topk_log_probs[i].item()
                    # 多样性惩罚:检查是否与前面组选中的token重复
                    penalty = 0.0
                    for prev_g in range(g):
                        # 本步中前面组已决定的token(如果是EOS则不考虑)
                        if step < len(chosen_tokens_per_group[prev_g]):
                            if chosen_tokens_per_group[prev_g][step] == next_token:
                                penalty += diversity_strength
                    new_score -= penalty
                    new_tokens = tokens + [next_token]
                    candidates.append((new_tokens, new_score))

            if not candidates:
                new_group_beams.append([])
                chosen_tokens_per_group.append([])
                continue

            # 选择本组得分最高的 group_size 个束,并记录本步选择的token(用于后续组)
            candidates.sort(key=lambda x: x[1], reverse=True)
            selected = candidates[:group_size]
            new_group_beams.append(selected)
            # 记录本组本步选出的所有token(注意可能由于有多个束,我们只对束内最高分的token?通常是对每个束记录其最后token)
            # 更标准的DBS是记录本组所有束的最后token,后续组若生成相同token就惩罚。
            group_step_tokens = [seq[-1] for seq, _ in selected]
            chosen_tokens_per_group.append(group_step_tokens)

        group_beams = new_group_beams
        # 如果所有组的活跃束都为空,则结束
        if all(len(b) == 0 for b in group_beams):
            break

    # 合并所有组的所有序列,选取最佳
    all_final = []
    for g in range(groups):
        all_final.extend(finished[g])
        for tokens, score in group_beams[g]:
            if not (eos_token_id and tokens[-1] == eos_token_id):  # 未完成的也加入
                all_final.append((tokens, score))
    best = max(all_final, key=lambda x: x[1])
    return best[0]

5.3 说明

  • 上述实现中,多样性惩罚是基于“前面组在同一时间步生成的 token”进行的。更复杂的实现可以考虑整个序列历史。

  • 惩罚强度 diversity_strength 需要调优,太大可能导致生成不连贯,太小则多样性不足。

  • 在 NLP 任务中,多样束搜索常用于生成多个不同图像描述或对话回复。


实现 Top‑k 采样 top_k_sampling(logits, k, temperature)。

6.1 原理

Top‑k 采样在生成时截断概率分布,只保留概率最高的 k 个 token,将其余 token 的概率置零,然后重新归一化,从该截断分布中随机采样。这能有效避免选取到低概率的无关 token(噪声),同时保持一定的随机性和多样性。温度参数用于在截断前缩放 logits,控制分布的尖锐度。

6.2 实现

def top_k_sampling(logits, k, temperature=1.0):
    """
    Args:
        logits: (vocab_size,) 浮点张量。
        k: 保留的候选 token 数量。
        temperature: 温度参数,>1 更平滑,<1 更尖锐。
    Returns:
        sampled_token_id: int
    """
    # 温度缩放
    if temperature != 1.0:
        logits = logits / temperature
    # 获取 top-k 的 logits 及其索引
    topk_logits, topk_indices = torch.topk(logits, k, dim=-1)
    # 创建一个新的 logits 张量,填充 -inf
    masked_logits = torch.full_like(logits, float('-inf'))
    masked_logits.scatter_(0, topk_indices, topk_logits)
    # softmax 得到概率
    probs = torch.softmax(masked_logits, dim=-1)
    # 采样
    token_id = torch.multinomial(probs, num_samples=1).item()
    return token_id

6.3 参数指南

  • k:通常取 10~100。值太小会导致输出重复、缺乏多样性;太大则引入噪音,可能生成不连贯文本。

  • temperature:T=1 无缩放;T<1 使高概率 token 更突出(接近贪心);T>1 使分布更平坦,增加低概率 token 被选中的机会。

6.4 与贪心的比较

  • 贪心解码等价于 T=0 且 k=1 的极限情况。

  • Top‑k 通过保留 k 个候选并随机采样,打破了贪心解码的确定性,能产生更多样化的输出。


实现 Top‑p(Nucleus)采样 top_p_sampling(logits, p, temperature)。

7.1 原理

Top‑p 采样(核采样)不固定候选数量,而是动态选择最小集合,使得这些 token 的累积概率达到阈值 p。具体步骤:将 logits 按概率降序排序,累加概率,找到第一个使得累积概率超过 p 的位置,保留该位置及之前的所有 token,将其余 token 的概率置零,重新归一化后采样。这种方法可以适应不同置信度的分布,避免 Top‑k 有时截断过多或过少的弊端。

7.2 实现

def top_p_sampling(logits, p, temperature=1.0):
    """
    Args:
        logits: (vocab_size,) 张量。
        p: 累积概率阈值,通常 0.9~0.95。
        temperature: 温度缩放。
    Returns:
        token_id: int
    """
    if temperature != 1.0:
        logits = logits / temperature
    # 按降序排序
    sorted_logits, sorted_indices = torch.sort(logits, descending=True)
    # 计算累积概率
    sorted_probs = torch.softmax(sorted_logits, dim=-1)
    cumsum_probs = torch.cumsum(sorted_probs, dim=-1)
    # 找到需要截断的位置:累积概率 > p 的第一个索引,但要确保至少保留一个token
    # 方法:创建掩码,将超过p的部分设为True
    # 我们希望保留 cumsum_probs <= p 的部分,以及刚好越过p的那个token(为了至少达到p)
    # 常用技巧:将第一个超过p的位置保留,所以将其掩码置为False
    mask = cumsum_probs > p
    # 将第一个超过p的位置的mask设为False(即保留该token)
    mask[1:] = mask[:-1].clone()
    mask[0] = False
    # 将mask对应的logits设为 -inf
    sorted_logits[mask] = float('-inf')
    # 重新映射回原始顺序
    filtered_logits = torch.full_like(logits, float('-inf'))
    filtered_logits.scatter_(0, sorted_indices, sorted_logits)
    # 采样
    probs = torch.softmax(filtered_logits, dim=-1)
    token_id = torch.multinomial(probs, num_samples=1).item()
    return token_id

7.3 与 Top‑k 的对比

  • Top‑k 固定候选数量,对分布变化不敏感:当模型非常确定时(少数 token 占绝大部分概率),Top‑k 仍保留 k 个,可能引入噪音;当模型不确定时(概率分散),Top‑k 可能遗漏重要的长尾 token。

  • Top‑p 动态调整候选集大小,更灵活,通常在多样性和质量间取得更好平衡。


实现温度缩放函数,并将其与贪心解码结合输出单个 token。

8.1 温度缩放原理

温度 T 对 logits 进行线性缩放:scaled_logits = logits / T。当 T > 1 时,概率分布趋于均匀,采样更随机;T < 1 时,分布更陡峭,接近 one-hot(极端时 T→0 退化为 argmax)。温度缩放在贪心解码中直接改变 logits,然后取 argmax 等同于先缩放再取最大值。因此可封装一个温度缩放函数,然后用 argmax 或采样。

8.2 实现

def temperature_scale(logits, temperature):
    """对 logits 进行温度缩放"""
    if temperature == 0:
        raise ValueError("Temperature must be > 0; for deterministic decoding use argmax directly.")
    return logits / temperature

def temperature_greedy_decode(logits, temperature):
    """结合温度缩放和贪心解码,返回一个token"""
    scaled_logits = temperature_scale(logits, temperature)
    token_id = torch.argmax(scaled_logits).item()
    return token_id

8.3 解析

  • temperature 非常小(如 0.1),输出近乎确定性,但可能仍受浮点精度影响。

  • 通常不建议在贪心中使用温度缩放,因为贪心本意是取最大值,温度缩放并不改变 argmax 的结果(单调变换,argmax 不变)。严格来说,argmax(softmax(logits/T)) 等价于 argmax(logits),因为 softmax 是单调递增函数,温度缩放不改变相对顺序。因此温度缩放结合贪心解码不会改变输出!这是一个常见的误解,面试中需明确指出。如果希望温度影响贪心,那是不可能的;要引入随机性必须使用采样。所以这里的结合实际是指“如果贪心解码,温度缩放无效”。我们可以讨论这一点,并说明正确的结合方式:温度缩放在采样策略中有效,而贪心不是采样。

修正:面试中我们应指出“温度缩放不会改变 argmax 的结果”,因此贪心+温度是无意义的。如果要求输出单个 token 且希望温度影响结果,应当使用温度缩放 + 随机采样,而不是贪心。所以题目可能是在考察对温度的理解。我们可以实现温度缩放的函数,然后说明其与采样结合,而不是贪心。

根据题意“实现温度缩放函数,并将其与贪心解码结合输出单个 token”,或许出题人希望我们结合的是:先温度缩放,然后贪心(取argmax)。但如上分析,argmax 不变。我们可以指出这一点,然后提供温度采样函数。

此处我实现温度缩放函数,并用采样来展示温度的影响,并讨论贪心为何不受影响。

def temperature_sampling(logits, temperature):
    """温度缩放后随机采样一个token"""
    scaled = logits / temperature
    probs = torch.softmax(scaled, dim=-1)
    return torch.multinomial(probs, num_samples=1).item()

8.4 结论

  • 温度缩放是采样策略的预处理步骤,不是贪心解码的组成部分。

  • 贪心解码永远选择概率最高的 token,不受温度影响。


实现组合 Top‑k + Top‑p 采样,先 Top‑k 再 Top‑p 截断后采样。

9.1 原理

组合采样结合了 Top‑k 和 Top‑p 的优点:先用 Top‑k 丢弃长尾噪声,再在剩余高质量候选上应用 Top‑p 动态截断,以应对不同的分布形状。步骤如下:

  1. 温度缩放 logits。

  2. 选出 top-k 个最高概率的 logits,其余置 -inf

  3. 在截断后的 logits 上应用 Top‑p 过滤。

  4. softmax 后采样。

9.2 实现

def top_k_top_p_sampling(logits, k, p, temperature=1.0):
    """
    先 Top-k 再 Top-p,最后采样。
    """
    if temperature != 1.0:
        logits = logits / temperature

    # 第一步:Top‑k 过滤
    topk_logits, topk_indices = torch.topk(logits, k, dim=-1)
    # 创建 masked logits
    k_masked = torch.full_like(logits, float('-inf'))
    k_masked.scatter_(0, topk_indices, topk_logits)

    # 第二步:在 Top‑k 结果上应用 Top‑p
    sorted_logits, sorted_indices = torch.sort(k_masked, descending=True)
    sorted_probs = torch.softmax(sorted_logits, dim=-1)
    cumsum = sorted_probs.cumsum(dim=-1)
    # 确定截断掩码
    mask = cumsum > p
    mask[1:] = mask[:-1].clone()
    mask[0] = False
    sorted_logits[mask] = float('-inf')

    # 重建原始顺序的 logits
    final_logits = torch.full_like(logits, float('-inf'))
    final_logits.scatter_(0, sorted_indices, sorted_logits)

    # 采样
    probs = torch.softmax(final_logits, dim=-1)
    token = torch.multinomial(probs, 1).item()
    return token

9.3 参数推荐

  • k: 50~200,视词汇表大小而定。

  • p: 0.9~0.95。

  • temperature: 0.7~1.0。

该组合在实际应用(如 GPT 系列)中被广泛采用,能很好地平衡流畅性与多样性。


实现重复惩罚函数,支持频率惩罚和存在惩罚两种模式。

10.1 原理

重复惩罚用于减少模型生成重复文本。通过在生成每个 token 前,对已出现过的 token 施加惩罚,降低其被再次选中的概率。两种常见模式:

  • 频率惩罚(Frequency Penalty):根据 token 在已生成序列中出现的次数按比例惩罚,出现越多,惩罚越重。公式:logits[t] -= freq[t] * penalty

  • 存在惩罚(Presence Penalty):只要 token 出现过至少一次,就施加一个固定惩罚,与出现次数无关。公式:logits[t] -= exists[t] * penalty

通常 penalty 为正值(如 0.1~1.0),直接减少 logits。

10.2 实现

def apply_repetition_penalty(logits, generated_ids, penalty, mode='frequency'):
    """
    Args:
        logits: (vocab_size,) 当前步的 logits。
        generated_ids: list of int,已生成的 token ID 序列(不包括当前)。
        penalty: 惩罚系数,通常 0.1~1.0。
        mode: 'frequency' 或 'presence'。
    Returns:
        penalized_logits: 同形状张量。
    """
    if not generated_ids or penalty == 0.0:
        return logits

    # 统计频率或存在性
    if mode == 'frequency':
        # 计算每个token的出现次数
        counts = torch.bincount(torch.tensor(generated_ids, device=logits.device),
                                minlength=logits.shape[0])
        # 对出现过的token施加惩罚:logits -= counts * penalty
        # 注意:我们希望出现次数多的惩罚更重,因此直接按次数乘以系数
        penalized = logits - counts.float() * penalty
    elif mode == 'presence':
        # 只要出现过就惩罚固定值
        unique_ids = set(generated_ids)
        mask = torch.zeros(logits.shape[0], dtype=torch.bool, device=logits.device)
        mask[list(unique_ids)] = True
        penalized = logits - mask.float() * penalty
    else:
        raise ValueError("mode must be 'frequency' or 'presence'")
    return penalized

10.3 使用方式

在每一步解码前,先对 logits 调用此函数施加惩罚,然后再进行 softmax 和采样或贪心。注意惩罚不应应用于 EOS token?通常 EOS 也需要适当惩罚以防止提前结束,但也可以选择对特殊 token 免除惩罚。

10.4 注意事项

  • 惩罚强度过大可能导致输出不连贯或语法错误。

  • 通常结合 temperature 和其他采样策略使用。

  • 实现时需注意 generated_ids 应包含 prompt 部分的 token 吗?通常只对生成部分的重复进行惩罚,但有时 prompt 中的高频词也可能导致生成时被惩罚,可根据需要决定。


实现硬约束解码,禁止生成与已有序列末尾形成重复 n-gram 的 token。

11.1 原理

硬约束解码(Hard N-gram Blocking)用于防止模型生成重复的连续片段。在每一步生成时,检查当前序列末尾是否存在一个长度为 n 的片段,它在已生成的文本中出现过(尤其是最近的输出),如果下一个 token 会导致即将形成的 n-gram 已经存在,则将该 token 的概率强制置零(logits 设为 -inf)。这能有效避免模型陷入重复循环,提升生成文本的多样性。

11.2 实现

def ngram_blocking_processor(logits, generated_ids, n=3):
    """
    将导致形成重复 n-gram 的 token 的 logits 置为 -inf。

    Args:
        logits: (vocab_size,) 当前步的 logits。
        generated_ids: list of int,已生成的 token ID 序列。
        n: n-gram 长度,通常 3 或 4。

    Returns:
        processed_logits: 同形状张量。
    """
    if len(generated_ids) < n - 1:
        # 序列太短,无法形成 n-gram,不做约束
        return logits

    # 取当前序列末尾的 (n-1) 个 token 作为前缀
    prefix = tuple(generated_ids[-(n-1):])  # 长度 n-1
    # 扫描已生成序列中所有长度为 n 的 n-gram,记录 (prefix, next_token) 对
    blocked_tokens = set()
    for i in range(len(generated_ids) - n + 1):
        if tuple(generated_ids[i:i+n-1]) == prefix:
            # 找到了相同前缀的 n-gram,记录它的最后一个 token
            next_tok = generated_ids[i+n-1]
            blocked_tokens.add(next_tok)

    # 如果末尾序列本身不足 n-1,但我们可以只检查最后的 (n-1) 长度历史
    # 将 blocked_tokens 中的 token 的 logits 置为 -inf
    processed_logits = logits.clone()
    for tok in blocked_tokens:
        processed_logits[tok] = float('-inf')
    return processed_logits

11.3 使用方式

将此函数作为 LogitsProcessor 插入解码循环的 softmax 之前,通常与采样策略配合:

logits = model(...)
logits = ngram_blocking_processor(logits, generated_ids, n=3)
# 然后进行温度缩放、Top-k/Top-p、采样

11.4 注意事项

  • n 值不宜太小(如2),否则会过度限制常见短语的正常生成;也不宜太大(如5+),则几乎不生效。

  • 这种硬约束可能在某些任务中导致生成不连贯,因此常与软惩罚(如重复惩罚)结合,或仅在检测到严重重复时启用。

  • 实现中 prefix(n-1) 个 token,这是为了检查“如果下一个 token 是 X,将形成一个新的 n-gram 是否已存在”。所以扫描的是整个序列中已有的所有长度为 n 的片段,并记录以相同前缀结尾的那些 token。

  • 注意:我们只应检查生成部分(不含 prompt)的重复,但具体取决于任务,通常对整个序列(含 prompt)检查也是可行的,因为 prompt 中也可能有需要避免的重复模式。


实现 min‑p 采样,过滤概率低于 min_p * max_prob 的 token 后归一化采样。

12.1 原理

Min‑p 采样是一种动态的截断策略,由 Nguyen et al. 提出。它计算当前概率分布中的最大概率值 max_prob,然后将所有概率低于 min_p * max_prob 的 token 丢弃(概率置零),再重新归一化进行采样。与 Top‑p 不同,它不累积概率,而是基于最大概率的比例来截断。当分布非常尖锐时,只有少数 token 保留;当分布平坦时,保留较多 token。这使其适应性更强,尤其在创意性和连贯性之间取得较好平衡。

12.2 实现

def min_p_sampling(logits, min_p=0.1, temperature=1.0):
    """
    Args:
        logits: (vocab_size,) 张量。
        min_p: 最小概率比例,通常 0.05~0.2。
        temperature: 温度缩放。
    Returns:
        sampled_token_id: int
    """
    if temperature != 1.0:
        logits = logits / temperature
    probs = torch.softmax(logits, dim=-1)
    max_prob = probs.max()
    # 计算阈值
    threshold = min_p * max_prob
    # 将低于阈值的 logits 置为 -inf
    filtered_logits = logits.clone()
    filtered_logits[probs < threshold] = float('-inf')
    # 重新计算概率并采样
    filtered_probs = torch.softmax(filtered_logits, dim=-1)
    token_id = torch.multinomial(filtered_probs, num_samples=1).item()
    return token_id

12.3 参数调优

  • min_p 设为 0.1 时,若最大概率为 0.8,则保留概率 ≥ 0.08 的所有 token;若 max_prob 为 0.3,则保留 ≥ 0.03 的 token。可见其动态调整候选数量。

  • 较小的 min_p(如 0.05)接近全量采样,较大的 min_p(如 0.2)则更保守,类似 Top‑p 但机制不同。

  • 可与温度缩放联用,先温度缩放再 min‑p 过滤。

12.4 与 Top‑p 的比较

  • Top‑p 总是确保保留累积概率为 p 的最小集合,可能导致在平坦分布时保留过多 token(因为需要累积到 p),而 min‑p 在平坦分布时阈值也较低,因此保留更多,但在尖锐分布时,min‑p 可能仅保留 1~2 个 token,更接近贪心。两者可结合使用。

实现典型采样(Typical Sampling),基于局部熵选择 token 集合。

13.1 原理

典型采样(Typical Sampling, Meister et al., 2022)不是基于绝对概率,而是基于信息论中的“典型集”概念。它计算每个 token 的信息量 I = -log P,并与分布的局部熵比较。只保留信息量接近局部熵(即典型概率附近)的 token。具体做法:计算分布 P 的局部熵 H = -∑ P log P;对于每个 token,计算其信息量 I_i = -log P_i;保留那些满足 |I_i - H| < τ * H 的 token(τ 为阈值因子,如 0.2)。然后重新归一化采样。

此方法能筛选出具有“典型”概率值的 token,避免长尾噪声,同时在高概率 token 过多时限制候选集,生成更符合人类习惯的文本。

13.2 实现

def typical_sampling(logits, tau=0.2, temperature=1.0):
    """
    Args:
        logits: (vocab_size,) 张量。
        tau: 阈值因子,通常 0.2~0.5。
        temperature: 温度。
    Returns:
        sampled_token_id: int
    """
    if temperature != 1.0:
        logits = logits / temperature
    probs = torch.softmax(logits, dim=-1)
    # 计算局部熵 H
    log_probs = torch.log(probs)
    entropy = -torch.sum(probs * log_probs)  # 标量
    # 计算每个 token 的信息量
    I = -log_probs  # shape (vocab_size,)
    # 保留接近熵的 token
    mask = torch.abs(I - entropy) < tau * entropy
    # 将被过滤的 logits 置为 -inf
    filtered_logits = logits.clone()
    filtered_logits[~mask] = float('-inf')
    # 重新 softmax 并采样
    filtered_probs = torch.softmax(filtered_logits, dim=-1)
    token_id = torch.multinomial(filtered_probs, num_samples=1).item()
    return token_id

13.3 参数说明

  • tau:较小的 tau 使得候选集更小、更保守;较大的 tau 放宽约束。实践中 0.2 是一个较好的起点。

  • 局部熵 H 反映了分布的不确定性。当分布平坦时 H 大,阈值放宽,保留更多 token;当分布尖锐时 H 小,阈值收紧,只保留少数高概率 token。

13.4 与其他采样策略的差异

  • Top‑k/Top‑p 依赖于绝对概率排序;Min‑p 依赖于最大概率比例;Typical Sampling 依赖信息量与局部熵的偏离,理论基础更深厚,在某些场景下能产生更自然连贯的文本。

在束搜索的每步扩展中用采样替代固定 top‑k 选择,实现采样式束搜索。

14.1 原理

传统束搜索每步扩展时选择 top-k 个最高得分的 token,这可能导致束多样性不足。采样式束搜索(Stochastic Beam Search)在扩展时采用随机采样(如 top-k 采样、温度采样)从概率分布中抽取下一个 token,而不是确定性地选择 top-k。这样可以在保持束搜索框架的同时增加生成多样性,避免所有束收敛到相同序列。

注意:当引入采样后,束的得分不再是纯对数概率和,还需要考虑采样带来的随机性,通常仍使用累积对数概率作为得分。

14.2 实现

def stochastic_beam_search(model, input_ids, beam_width=3, max_length=50,
                           temperature=1.0, top_k=10, eos_token_id=None):
    """
    每步从 top-k 中按概率采样一个 token 来扩展束。
    """
    device = input_ids.device
    prompt = input_ids[0].tolist()
    prompt_len = len(prompt)
    beams = [(prompt, 0.0)]  # (tokens, score)
    finished = []

    for step in range(max_length - prompt_len):
        if not beams:
            break
        next_candidates = []
        for tokens, score in beams:
            if eos_token_id is not None and tokens[-1] == eos_token_id:
                finished.append((tokens, score))
                continue
            input_t = torch.tensor([tokens], device=device)
            outputs = model(input_t)
            next_logits = outputs[0, -1, :]
            # 采样式选择下一个 token
            next_token = top_k_sampling(next_logits, k=top_k, temperature=temperature)
            new_tokens = tokens + [next_token]
            # 计算对数概率(用于排序得分)
            log_probs = torch.log_softmax(next_logits, dim=-1)
            new_score = score + log_probs[next_token].item()
            if eos_token_id is not None and next_token == eos_token_id:
                finished.append((new_tokens, new_score))
            else:
                next_candidates.append((new_tokens, new_score))

        # 按得分排序保留 beam_width 个最佳
        next_candidates.sort(key=lambda x: x[1], reverse=True)
        beams = next_candidates[:beam_width]

    # 选取最佳序列
    final = finished + beams
    final.sort(key=lambda x: x[1], reverse=True)
    return final[0][0]

14.3 讨论

  • 因为每步是采样而非选 top-k,束中可能出现重复序列,可以通过维护已生成序列的哈希集合来去重。

  • 采样参数(temperature, top_k)强烈影响多样性和质量,需要调优。

  • 采样式束搜索常用于需要多样性的任务,如图像描述生成多个候选。


构建灵活的 LogitsProcessor 管道,依次施加温度、Top‑k、Top‑p、重复惩罚等。

15.1 设计思路

借鉴 transformers 库的 LogitsProcessorList,我们可以构建一个可组合的处理器管道。每个处理器是一个函数,接受 (logits, context) 返回修改后的 logits。管道按顺序调用。

15.2 实现

class LogitsProcessor:
    def __call__(self, logits, generated_ids):
        return logits

class TemperatureProcessor(LogitsProcessor):
    def __init__(self, temperature):
        self.temperature = temperature
    def __call__(self, logits, generated_ids):
        if self.temperature != 1.0:
            logits = logits / self.temperature
        return logits

class TopKProcessor(LogitsProcessor):
    def __init__(self, k):
        self.k = k
    def __call__(self, logits, generated_ids):
        topk_logits, topk_indices = torch.topk(logits, self.k, dim=-1)
        mask = torch.full_like(logits, float('-inf'))
        mask.scatter_(0, topk_indices, topk_logits)
        return mask

class TopPProcessor(LogitsProcessor):
    def __init__(self, p):
        self.p = p
    def __call__(self, logits, generated_ids):
        sorted_logits, sorted_indices = torch.sort(logits, descending=True)
        sorted_probs = torch.softmax(sorted_logits, dim=-1)
        cumsum = sorted_probs.cumsum(dim=-1)
        mask = cumsum > self.p
        mask[1:] = mask[:-1].clone()
        mask[0] = False
        sorted_logits[mask] = float('-inf')
        result = torch.full_like(logits, float('-inf'))
        result.scatter_(0, sorted_indices, sorted_logits)
        return result

class RepetitionPenaltyProcessor(LogitsProcessor):
    def __init__(self, penalty, mode='frequency'):
        self.penalty = penalty
        self.mode = mode
    def __call__(self, logits, generated_ids):
        if not generated_ids or self.penalty == 0.0:
            return logits
        if self.mode == 'frequency':
            counts = torch.bincount(torch.tensor(generated_ids, device=logits.device),
                                    minlength=logits.shape[0])
            return logits - counts.float() * self.penalty
        elif self.mode == 'presence':
            unique_ids = set(generated_ids)
            mask = torch.zeros(logits.shape[0], dtype=torch.bool, device=logits.device)
            mask[list(unique_ids)] = True
            return logits - mask.float() * self.penalty
        return logits

class LogitsProcessorPipeline:
    def __init__(self, processors):
        self.processors = processors
    def __call__(self, logits, generated_ids):
        for proc in self.processors:
            logits = proc(logits, generated_ids)
        return logits

15.3 使用示例

pipeline = LogitsProcessorPipeline([
    TemperatureProcessor(0.8),
    RepetitionPenaltyProcessor(0.5, mode='frequency'),
    TopKProcessor(50),
    TopPProcessor(0.9),
])
# 在每步解码时
logits = pipeline(logits, generated_ids)
probs = torch.softmax(logits, dim=-1)
next_token = torch.multinomial(probs, 1).item()

15.4 扩展性

  • 可以方便地添加其他约束(如白名单、黑名单、长度惩罚、EOS 禁止等)。

  • 注意处理顺序:温度应最先,因为缩放后影响后续概率;重复惩罚也应在截断前;Top‑k/Top‑p 通常在最后。


实现流式生成器,每次 yield 一个 token,支持实时调整采样参数。

16.1 设计

生成器是一个可迭代对象,每一步生成一个 token 并通过 yield 返回。为了支持动态调整参数,生成器可以接受一个外部队列或回调函数,在每步生成前检查是否有新参数传入。简化实现:在 yield 返回 token 的同时,外部可以通过 send 方法传入新的参数(如 temperature),利用 Python 生成器的双向通信。

16.2 实现

def stream_generate(model, input_ids, max_length=50, eos_token_id=None, **initial_params):
    """
    流式生成器,每次 yield 一个 token。
    可以通过 generator.send(params) 实时调整参数。
    """
    # 初始参数
    params = initial_params.copy()
    generated = input_ids[0].tolist()
    model.eval()
    with torch.no_grad():
        while len(generated) < max_length:
            inputs = torch.tensor([generated], device=input_ids.device)
            outputs = model(inputs)
            logits = outputs[0, -1, :]

            # 应用当前参数(如温度、top_k、top_p 等)
            temperature = params.get('temperature', 1.0)
            top_k = params.get('top_k', 0)
            top_p = params.get('top_p', 1.0)
            if temperature != 1.0:
                logits = logits / temperature
            probs = torch.softmax(logits, dim=-1)
            if top_k > 0:
                topk_probs, topk_indices = torch.topk(probs, top_k)
                filtered_probs = torch.zeros_like(probs)
                filtered_probs.scatter_(0, topk_indices, topk_probs)
                probs = filtered_probs / filtered_probs.sum()
            if top_p < 1.0:
                sorted_probs, sorted_indices = torch.sort(probs, descending=True)
                cumsum = sorted_probs.cumsum(dim=-1)
                mask = cumsum > top_p
                mask[1:] = mask[:-1].clone()
                mask[0] = False
                sorted_probs[mask] = 0.0
                probs = torch.zeros_like(probs)
                probs.scatter_(0, sorted_indices, sorted_probs)
                probs = probs / probs.sum()

            # 采样
            next_token = torch.multinomial(probs, 1).item()
            generated.append(next_token)
            # yield 返回 token,并接收可能的新参数
            new_params = yield next_token
            if new_params is not None:
                params.update(new_params)

            if eos_token_id is not None and next_token == eos_token_id:
                break

16.3 使用示例

gen = stream_generate(model, input_ids, max_length=50, temperature=0.9, top_p=0.95)
first_token = next(gen)  # 启动生成器
# 后续可以动态调整
gen.send({'temperature': 1.2})  # 在下一个 token 生成时应用新温度

16.4 注意

  • 生成器的双向通信需要先调用 next(gen)gen.send(None) 启动。

  • 状态更新应确保参数安全,例如检查参数合法性。

  • 对于需要更多复杂处理(如重复惩罚)的参数调整,可扩展参数字典。


实现带最小生成长度约束的束搜索,未达长度前禁止输出 EOS。

17.1 原理

在一些任务(如强制生成摘要)中,我们希望模型生成长度至少达到 min_length。实现方式:在束搜索的扩展阶段,如果当前序列长度(生成部分)小于 min_length,则禁止选择 EOS token——将其 logits 置为 -inf。但注意 EOS 的禁止可能导致序列无法终止,因此当所有束长度都超过 min_length 后,正常允许 EOS。

17.2 实现

在束搜索扩展循环中,对于每条序列,判断其生成部分的长度(len(tokens) - prompt_len),如果小于 min_length,则在获取 logits 后强制将 EOS token 的 logits 置为 -inf

def beam_search_min_length(model, input_ids, beam_width=3, max_length=50,
                           min_length=10, eos_token_id=None):
    device = input_ids.device
    prompt = input_ids[0].tolist()
    prompt_len = len(prompt)
    beams = [(prompt, 0.0)]
    finished = []

    for step in range(max_length - prompt_len):
        if not beams:
            break
        next_candidates = []
        for tokens, score in beams:
            # 若当前序列已生成部分 >= min_length 且遇到 EOS,则移入finished
            gen_len = len(tokens) - prompt_len
            if eos_token_id is not None and tokens[-1] == eos_token_id and gen_len >= min_length:
                finished.append((tokens, score))
                continue
            # 否则继续扩展
            input_t = torch.tensor([tokens], device=device)
            outputs = model(input_t)
            next_logits = outputs[0, -1, :]
            # 强制禁止 EOS 如果 gen_len < min_length
            if gen_len < min_length and eos_token_id is not None:
                next_logits[eos_token_id] = float('-inf')
            log_probs = torch.log_softmax(next_logits, dim=-1)
            topk_log_probs, topk_indices = torch.topk(log_probs, beam_width)
            for i in range(beam_width):
                next_token = topk_indices[i].item()
                new_score = score + topk_log_probs[i].item()
                new_tokens = tokens + [next_token]
                next_candidates.append((new_tokens, new_score))

        if not next_candidates:
            break
        next_candidates.sort(key=lambda x: x[1], reverse=True)
        beams = next_candidates[:beam_width]

    final = finished + beams
    final.sort(key=lambda x: x[1], reverse=True)
    return final[0][0]

17.3 注意事项

  • min_length 过大,可能导致束搜索无法在 max_length 内终止,最后可能选择未完成序列。

  • 更精细的实现:对于已完成但长度不足的序列,不放入 finished,而是强制继续扩展(禁止 EOS)。当所有束都超过 min_length 后,允许正常结束。


在束搜索中加入去重机制,检测并排除重复序列路径。

18.1 原理

普通束搜索可能产生内容完全相同的多个束,浪费计算且降低多样性。去重机制通过维护一个哈希集合记录已经出现过的序列(或最近 n 步的关键特征),当某条新候选序列与已存在的序列(通常指已完成的束或同一束内其他分支)重复时,降低其得分或直接移除。

18.2 实现

可以在每步扩展后,对候选列表进行去重:如果两条候选序列的 token 列表完全相同,则只保留得分最高的那条(或直接丢弃重复)。也可以基于 n-gram 哈希进行去重。

def beam_search_with_dedup(model, input_ids, beam_width=3, max_length=50, eos_token_id=None):
    device = input_ids.device
    prompt = input_ids[0].tolist()
    prompt_len = len(prompt)
    beams = [(prompt, 0.0)]
    finished = []

    for step in range(max_length - prompt_len):
        if not beams:
            break
        next_candidates = []
        for tokens, score in beams:
            if eos_token_id is not None and tokens[-1] == eos_token_id:
                finished.append((tokens, score))
                continue
            input_t = torch.tensor([tokens], device=device)
            outputs = model(input_t)
            next_logits = outputs[0, -1, :]
            log_probs = torch.log_softmax(next_logits, dim=-1)
            topk_log_probs, topk_indices = torch.topk(log_probs, beam_width)
            for i in range(beam_width):
                next_token = topk_indices[i].item()
                new_score = score + topk_log_probs[i].item()
                new_tokens = tokens + [next_token]
                next_candidates.append((new_tokens, new_score))

        # 去重:相同 token 序列只保留得分最高的
        dedup = {}
        for tokens, score in next_candidates:
            key = tuple(tokens)  # 以完整序列作为键
            if key not in dedup or score > dedup[key][1]:
                dedup[key] = (tokens, score)
        unique_candidates = list(dedup.values())

        unique_candidates.sort(key=lambda x: x[1], reverse=True)
        beams = unique_candidates[:beam_width]

    final = finished + beams
    final.sort(key=lambda x: x[1], reverse=True)
    return final[0][0]

18.3 优化

  • 完全基于完整序列去重在长序列下开销较大,可改用滑动窗口的 n-gram 哈希集合来近似去重。

  • 去重可能减少束的有效数量,如果剩余束不足 beam_width,可允许保留较少束。

  • 在多样束搜索中,去重尤为重要。


实现对比束搜索,结合参考模型惩罚项 score = log P_model - λ * log P_ref。

19.1 原理

对比解码(Contrastive Decoding)旨在通过对比一个较弱参考模型(如小模型或 amatuer model)来增强主模型的生成质量。束搜索中,可以修改得分函数:新得分 = log P_main(token) - λ * log P_ref(token)。当主模型和参考模型对某 token 的预测一致时,参考项会抵消一部分主模型的得分,抑制高频但无意义的词;当主模型明显偏好而参考模型不喜欢时,得分被增强,从而鼓励主模型生成更自信、更有信息量的内容。

19.2 实现

def contrastive_beam_search(main_model, ref_model, input_ids, beam_width=3, max_length=50,
                            lambda_val=0.5, eos_token_id=None):
    """
    main_model: 主模型
    ref_model: 参考模型(较弱)
    """
    device = input_ids.device
    prompt = input_ids[0].tolist()
    prompt_len = len(prompt)
    beams = [(prompt, 0.0)]
    finished = []

    for step in range(max_length - prompt_len):
        if not beams:
            break
        next_candidates = []
        for tokens, score in beams:
            if eos_token_id is not None and tokens[-1] == eos_token_id:
                finished.append((tokens, score))
                continue
            input_t = torch.tensor([tokens], device=device)
            # 主模型 logits
            main_logits = main_model(input_t)[0, -1, :]
            # 参考模型 logits
            ref_logits = ref_model(input_t)[0, -1, :]
            # 计算对比得分
            main_log_probs = torch.log_softmax(main_logits, dim=-1)
            ref_log_probs = torch.log_softmax(ref_logits, dim=-1)
            scores = main_log_probs - lambda_val * ref_log_probs
            # 选择 topk 进行扩展
            topk_scores, topk_indices = torch.topk(scores, beam_width, dim=-1)
            for i in range(beam_width):
                next_token = topk_indices[i].item()
                new_score = score + topk_scores[i].item()
                new_tokens = tokens + [next_token]
                next_candidates.append((new_tokens, new_score))

        next_candidates.sort(key=lambda x: x[1], reverse=True)
        beams = next_candidates[:beam_width]

    final = finished + beams
    final.sort(key=lambda x: x[1], reverse=True)
    return final[0][0]

19.3 注意事项

  • lambda_val 控制参考模型的惩罚力度,值越大,越抑制常见词,鼓励主模型与众不同的输出。

  • 参考模型通常应较弱,否则会削弱主模型的表现。

  • 该方法可有效减少文本退化,提升信息量。


实现基于 token 白名单的约束解码,将不允许的 token 置为 -inf。

20.1 原理

在某些应用(如受控词汇生成、特定领域的输出格式化)中,我们只想从一组预先定义好的 token 集合(白名单)中选择下一个 token。实现很简单:在每一步解码前,将不在白名单中的 token 的 logits 强制设为 -inf,然后 softmax 时这些 token 的概率即为 0,保证采样或贪心时不会被选中。

20.2 实现

def whitelist_constraint(logits, allowed_token_ids):
    """
    allowed_token_ids: list or tensor of token IDs that are allowed.
    其他 token 置为 -inf。
    """
    mask = torch.full_like(logits, float('-inf'))
    allowed = torch.tensor(allowed_token_ids, device=logits.device, dtype=torch.long)
    mask[allowed] = logits[allowed]
    return mask

在解码循环中:

logits = model(...)
logits = whitelist_constraint(logits, allowed_ids)
# 后续进行 softmax 和采样或贪心

20.3 扩展

  • 可以结合黑名单(将特定 token 排除)。

  • 白名单通常用于:

  • 强制输出 JSON 键或特定标点。
  • 限制回答只在特定选项内。
  • 受控词汇(如医疗术语)生成。

实现批量束搜索,支持 batch 内多样本的并行束搜索。

21.1 原理

在实际应用中,我们通常需要同时为 batch 中的多个样本进行束搜索。可以通过将 batch 内所有样本的束展开为一个大的 batch 维度,利用 GPU 并行计算。实现时,维护一个形状为 (batch_size * beam_width, cur_len) 的输入张量,每个样本的束占据连续的 beam_width 行。扩展时,从该大 batch 计算 logits,然后为每个样本独立执行 top‑k 选择、得分更新、束重组。注意需要使用索引偏移来区分不同样本的束。

21.2 实现(简化版)

假设模型接受 (batch_size, seq_len) 并返回 (batch_size, seq_len, vocab_size) 的 logits。

def batched_beam_search(model, input_ids, beam_width=3, max_length=50, eos_token_id=None):
    """
    input_ids: (batch_size, prompt_len) 初始 prompt 批量。
    返回: (batch_size, seq_len) 每个样本的最佳序列。
    """
    batch_size, prompt_len = input_ids.shape
    device = input_ids.device
    # 初始束:将每个样本的 prompt 重复 beam_width 次
    # 形状: (batch_size * beam_width, prompt_len)
    expanded_input = input_ids.unsqueeze(1).expand(-1, beam_width, -1).reshape(batch_size * beam_width, prompt_len)
    # 每个束的得分,初始 0
    beam_scores = torch.zeros(batch_size * beam_width, device=device)
    # 标记束是否已完成
    done = torch.zeros(batch_size * beam_width, dtype=torch.bool, device=device)
    # 存储最佳序列(每个样本独立)
    # 此处简化:仅返回最后得分最高的束对应的序列。实际还需考虑提前终止。

    for step in range(max_length - prompt_len):
        if done.all():
            break
        # 前向
        outputs = model(expanded_input)
        next_logits = outputs[:, -1, :]  # (B*BW, vocab)
        log_probs = torch.log_softmax(next_logits, dim=-1)

        # 计算所有可能扩展的得分:每个束扩展 beam_width 个候选,共 (B*BW * beam_width) 个
        # 为简便,我们取每个束的 top-beams 个 token
        topk_log_probs, topk_indices = torch.topk(log_probs, beam_width, dim=-1)  # (B*BW, beam_width)
        # 计算新得分
        new_scores = beam_scores.unsqueeze(1) + topk_log_probs  # (B*BW, BW)

        # 重塑为 (B, BW*BW) 以便为每个样本独立选取 top BW
        new_scores = new_scores.view(batch_size, beam_width * beam_width)
        topk_scores, topk_pos = torch.topk(new_scores, beam_width, dim=-1)  # (B, BW)
        # topk_pos 指示在 (BW*BW) 中的索引,需要还原为原束索引和 token 索引
        beam_indices = topk_pos // beam_width  # (B, BW) 原束索引
        token_indices = topk_pos % beam_width  # (B, BW) token 索引(在 topk_indices 中的位置)
        # 获取实际 token ID
        next_tokens = torch.gather(topk_indices, 1, token_indices + beam_indices * beam_width) # 需正确索引,简化可改用循环。
        # 为简洁,此处略去详细索引实现,实际可参考 HuggingFace 源码。

        # 更新 expanded_input 和 beam_scores
        # ...

    # 最终从每个样本的束中选择最佳序列(略)
    return best_sequences

由于完整实现较复杂,面试中可描述原理并给出核心索引逻辑,强调 batch 内并行带来的效率提升。

21.3 关键点

  • 需要正确维护 beam_scoresdone 标志,以及处理 EOS 提前终止。

  • 通过将 beam_width 维度融入 batch 维度,充分利用 GPU 并行。

  • 最终选择每个样本的最高分序列作为输出。