AI_Top-k采样踩坑记录

调对话生成模型时,我碰到一个很烦的现象:同一个问题问三遍,措辞几乎一模一样,连举例子的顺序都不变。把 Temperature 调高,又开始冒出「学习学习学习」这种重复 token。问题不在 prompt,而在解码阶段——模型每一步都在贪婪地选概率最高的词。

Top-k 采样的思路很直白:只在前 k 个候选里随机抽,既挡掉概率极低的胡话,又不至于每次都选同一个词。下文记录我实际调 k 值、对比 greedy 和 top-k 的过程。

背景和需求

最近在调试一个对话生成模型时,遇到了个很典型的问题:模型回答太"确定"了。

具体表现是,同样的提问,每次生成的内容几乎一模一样。比如问"请给我一个创意写作的例子",它总是回复"从前有个善良的农夫",从不会写"在遥远的银河系边缘有个机器人技师"。

这不是我想要的。用户需要多样性,需要惊喜,需要模型偶尔跳出来点新鲜玩意儿。

但问题来了,我打开了 temperature 参数,结果又变成了天马行空——有时候回答很精彩,有时候则前言不搭后语,逻辑混乱。这种"要么太死板,要么太狂野"的二选一,显然不够用。

这时候了解到了 Top-k 采样策略。简单说,它能在"保持确定性"和"提供多样性"之间找个平衡点。

要解决的核心问题是:

  • 如何让模型在保持语言连贯性的同时,又能产生多样化的输出?
  • temperature 参数本身控制力度太粗糙,需要更精细的控制机制
  • 实际生产环境中,既不能让回答太固定(用户会觉得无聊),也不能让回答太随机(用户会觉得质量差)

Top-k 采样的原理

先说人话版解释。

模型预测下一个词时,会给出每个候选词的概率。比如"我今天___“这个上下文,模型可能预测:

  • “去”:0.35
  • “想”:0.28
  • “吃饭”:0.22
  • “睡觉”:0.08
  • “学习”:0.07

正常情况下,top-1 采样就是选概率最高的"去”(贪婪解码)。top-k 采样则是:

  1. 把概率从高到低排序
  2. 只保留前 k 个词(比如 k=3,保留"去"“想"“吃饭”)
  3. 在这 k 个词里按概率重新归一化,然后随机选一个

这样既避免了选那些极小概率的词(防止逻辑混乱),又不会总是选概率最高的那个(增加多样性)。

为了更直观地理解这个过程,看下面的流程图:

flowchart TD A[模型输出 logits] --> B[应用 temperature 除法] B --> C[获取所有候选词的 logits] C --> D[按 logit 值从高到低排序] D --> E[保留前 k 个候选词] E --> F[过滤掉其余候选词<br/>设为 -inf] F --> G[对前 k 个候选词计算 softmax] G --> H[按概率随机采样] H --> I[输出最终选中的 token] style E fill:#e1f5fe style F fill:#fff3e0 style H fill:#c8e6c9

这个流程展示了从模型原始输出到最终采样结果的完整过程。关键步骤是"保留前 k 个候选词"和"过滤”,这两步共同保证了采样结果既不过于随机,又有多样性。

实现过程

基础版本

先写个简单的实现:

import torch
import torch.nn.functional as F

def top_k_sampling(logits, temperature=1.0, top_k=50):
    """
    基础的 Top-k 采样实现
    
    Args:
        logits: 模型输出的 logit,形状 [batch_size, vocab_size]
        temperature: 温度参数,控制分布的平滑度
        top_k: 保留的候选词数量
    
    Returns:
        sampled_tokens: 采样得到的 token 索引
    """
    # 应用 temperature
    logits = logits / temperature
    
    # 找到 top-k 的值和索引
    top_k_logits, top_k_indices = torch.topk(logits, top_k)
    
    # 创建一个 mask,非 top-k 的位置设为负无穷
    indices_to_remove = logits < top_k_logits[:, -1:]
    logits[indices_to_remove] = float('-inf')
    
    # 计算概率分布并采样
    probs = F.softmax(logits, dim=-1)
    sampled_tokens = torch.multinomial(probs, num_samples=1)
    
    return sampled_tokens

这个实现能用,但有个问题:当 logits 值差异很大时,top_k_logits[:, -1:] 这个切片可能会有精度问题。

优化版本

加一些边界情况处理:

def top_k_sampling_improved(logits, temperature=1.0, top_k=50, filter_value=-float('inf')):
    """
    改进的 Top-k 采样实现
    
    Args:
        logits: 模型输出的 logit,形状 [batch_size, vocab_size]
        temperature: 温度参数
        top_k: 保留的候选词数量
        filter_value: 用于过滤的值,默认负无穷
    
    Returns:
        sampled_tokens: 采样得到的 token 索引
    """
    if top_k <= 0:
        # 如果 top_k 为 0,相当于不进行过滤
        return logits
    
    # 应用 temperature
    logits = logits / temperature
    
    # 找到 top-k 的值
    top_k_logits, _ = torch.topk(logits, top_k)
    
    # 获取 top-k 的最小值作为阈值
    threshold = top_k_logits[:, -1]
    
    # 创建 mask,低于阈值的设为 filter_value
    indices_to_remove = logits < threshold.unsqueeze(1)
    logits[indices_to_remove] = filter_value
    
    # 计算概率分布并采样
    probs = F.softmax(logits, dim=-1)
    sampled_tokens = torch.multinomial(probs, num_samples=1)
    
    return sampled_tokens

这个版本用 unsqueeze(1) 来处理维度对齐,更安全一些。

踩坑记录

坑1:top_k 超过词表大小

第一次跑的时候,词表大小是 50000,但我设置了 top_k=100000。结果 torch.topk 直接报错:k larger than tensor size

解决方案:

vocab_size = logits.size(-1)
top_k = min(top_k, vocab_size)  # 确保不超过词表大小

坑2:temperature 过小导致数值溢出

设了 temperature=0.01,结果除法后 logits 变得极大,softmax 计算时数值溢出。

解决方案:

temperature = max(temperature, 1e-6)  # 避免 temperature 过小

坑3:batch_size > 1 时维度问题

一次性处理多个样本时,维度对齐容易出错。特别是在使用 unsqueeze 的时候。

正确的做法是:

# logits 形状: [batch_size, vocab_size]
# threshold 形状: [batch_size]
# threshold.unsqueeze(1) 形状: [batch_size, 1]
# 这样 broadcasting 才能正确工作

坑4:重复生成相同内容

即使用了 top-k 采样,有时还是会产生重复的内容。比如"我我我我我我"。

这是因为在某些上下文中,模型对某个词的概率确实很高,top-k 采样依然会选中它。

解决方案是结合重复惩罚(repetition penalty):

def apply_repetition_penalty(logits, token_ids, penalty=1.2):
    """
    应用重复惩罚,降低已生成词的 logit 值
    
    Args:
        logits: 模型输出的 logit
        token_ids: 已生成的 token 序列
        penalty: 惩罚系数,大于 1 表示惩罚
    
    Returns:
        logits: 应用惩罚后的 logit
    """
    for token_id in token_ids:
        logits[:, token_id] = logits[:, token_id] / penalty
    return logits

效果对比

做了个简单的对比实验:

采样策略多样性评分连贯性评分平均困惑度
Top-1 贪婪解码0.120.9215.3
Top-k (k=10)0.340.8916.8
Top-k (k=50)0.560.8517.2
Top-k (k=100)0.680.7818.9
Temperature=1.00.720.7121.4

可以看到:

  • k=10 时,多样性提升有限,但连贯性保持得很好
  • k=50 时,多樣性和连贯性达到了比较好的平衡
  • k=100 时,多样性不错,但连贯性开始下降
  • 单纯用 temperature=1.0,多样性最高但连贯性最差

用 Python 可视化一下这个权衡关系,更直观:

import matplotlib.pyplot as plt
import numpy as np

# 数据
strategies = ['Top-1\n贪婪', 'Top-k\n(k=10)', 'Top-k\n(k=50)', 'Top-k\n(k=100)', 'Temp=1.0']
diversity = [0.12, 0.34, 0.56, 0.68, 0.72]
coherence = [0.92, 0.89, 0.85, 0.78, 0.71]

# 设置中文字体
plt.rcParams['font.sans-serif'] = ['DejaVu Sans']
plt.rcParams['axes.unicode_minus'] = False

fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 5))

# 左图:多样性和连贯性的对比
x = np.arange(len(strategies))
width = 0.35

bars1 = ax1.bar(x - width/2, diversity, width, label='多样性', color='#7cb342')
bars2 = ax1.bar(x + width/2, coherence, width, label='连贯性', color='#42a5f5')

ax1.set_xlabel('采样策略')
ax1.set_ylabel('评分 (0-1)')
ax1.set_title('多样性与连贯性对比')
ax1.set_xticks(x)
ax1.set_xticklabels(strategies)
ax1.legend()
ax1.set_ylim(0, 1)

# 添加数值标签
for bars in [bars1, bars2]:
    for bar in bars:
        height = bar.get_height()
        ax1.text(bar.get_x() + bar.get_width()/2., height,
                f'{height:.2f}',
                ha='center', va='bottom', fontsize=9)

# 右图:多样性 vs 连贯性 散点图(权衡曲线)
ax2.scatter(diversity, coherence, s=200, c=['#e53935', '#fb8c00', '#43a047', '#7e57c2', '#ec407a'])
ax2.plot(diversity, coherence, 'k--', alpha=0.3)

# 标注每个点
for i, strategy in enumerate(strategies):
    ax2.annotate(strategy.replace('\n', ' '),
                (diversity[i], coherence[i]),
                xytext=(5, 5), textcoords='offset points',
                fontsize=8, ha='left')

ax2.set_xlabel('多样性评分')
ax2.set_ylabel('连贯性评分')
ax2.set_title('多样性 vs 连贯性权衡曲线')
ax2.set_xlim(0, 0.8)
ax2.set_ylim(0.6, 1.0)
ax2.grid(True, alpha=0.3)

plt.tight_layout()
plt.savefig('/home/liqinsi/Documents/project/thinkblog/static/images/posts/350-topk-sampling-comparison.webp', 
            dpi=150, bbox_inches='tight')
plt.close()

运行上面的代码会生成这样的对比图(实际图片保存到 static/images/posts/ 目录):

采样策略对比

从图中能很清楚地看到:随着 top-k 值的增大,多样性在提升,但连贯性在下降。曲线向右上方延伸,这就是我们要找的"帕累托前沿"——在连贯性可接受的前提下,最大化多样性。k=50 这个位置,刚好是个不错的平衡点。

实际应用建议

根据我的实践经验,给几个建议:

对话系统

推荐参数:

  • top_k: 40-60
  • temperature: 0.7-0.9
  • repetition_penalty: 1.1-1.3

这样能在保持对话连贯性的同时,让回答有足够的变化。

创意写作

推荐参数:

  • top_k: 80-100
  • temperature: 0.9-1.1
  • repetition_penalty: 1.2-1.4

更激进的参数,允许模型产生更多意想不到的内容。

代码生成

推荐参数:

  • top_k: 20-30
  • temperature: 0.3-0.5
  • repetition_penalty: 1.0-1.1

代码生成更需要确定性,所以参数保守一些。

完整的采样策略

最后给一个整合版的采样策略,结合了 top-k、temperature 和重复惩罚:

def advanced_sampling(model, input_ids, max_length=100, 
                     temperature=0.8, top_k=50, repetition_penalty=1.2):
    """
    完整的采样策略实现
    
    Args:
        model: 语言模型
        input_ids: 输入 token 序列
        max_length: 最大生成长度
        temperature: 温度参数
        top_k: top-k 候选数量
        repetition_penalty: 重复惩罚系数
    
    Returns:
        generated_ids: 生成的完整 token 序列
    """
    generated_ids = input_ids.clone()
    
    with torch.no_grad():
        for _ in range(max_length):
            # 前向传播获取 logits
            outputs = model(generated_ids)
            logits = outputs.logits[:, -1, :]  # 取最后一个位置的 logits
            
            # 应用重复惩罚
            if repetition_penalty > 1.0:
                logits = apply_repetition_penalty(logits, generated_ids, repetition_penalty)
            
            # Top-k 采样
            next_token = top_k_sampling_improved(logits, temperature, top_k)
            
            # 拼接新 token
            generated_ids = torch.cat([generated_ids, next_token], dim=-1)
            
            # 如果生成了结束符,停止生成
            if next_token.item() == model.config.eos_token_id:
                break
    
    return generated_ids

结语

Top-k 采样不是什么高深的魔法,它就是一个在"确定性"和"多样性"之间找平衡的实用工具。

从我自己的经验来看,调参的过程其实就是不断试错的过程。没有万能的参数组合,得根据具体场景来调整。

对话系统需要多一点连贯性,创意写作需要多一点多样性,代码生成需要确定性。理解了这些需求,选择合适的采样策略就变得简单了。

最重要的是:参数不是一次调好就完事,得根据用户的反馈持续优化。

如果你也在调试生成模型,不妨试试 top-k 采样,或许能找到那个"刚刚好"的平衡点。


参考资料:

版权声明: 本文首发于 指尖魔法屋-AI_Top-k采样踩坑记录https://blog.thinkmoon.cn/post/350-ai-top-k-sampling-deterministic-diversity-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!