AI核采样:Top-k不够用了之后

本地跑 7B 模型写技术摘要,Temperature 设 0.7 时句子顺得可疑,同一段落里「此外」「另外」来回出现;调到 1.2,又开始编造不存在的 API 参数。Top-k 能挡一部分长尾噪声,但 k 设小了太死板,设大了又和没限制差不多。

核采样(Nucleus / Top-p)按累积概率动态截断候选集,是我后来真正稳住输出质量的那招。这篇记实现细节和可视化对比,不是推公式。

为什么会有这篇文章?

最近在折腾本地大模型的时候,我遇到了一个很经典的问题:模型生成的文本要么"太顺",像复读机一样重复相同的句子;要么"太疯",Temperature 一调高,就开始胡言乱语,生成一堆语法正确但逻辑不通的废话。

为了解决这个问题,我深入研究了生成策略。大家都知道 Temperature(温度)控制随机性,但仅仅靠 Temperature 往往不够。于是我把目光投向了 Nucleus Sampling(核采样,也叫 Top-p 采样)

市面上的教程大多是数学公式堆砌,看完还是不知道怎么在代码里落地。既然我是来"解决实际问题"的,那我就用最直观的方式,把这事儿讲清楚,并附上可以直接跑的 Python 代码和可视化方案。

这篇文章要解决的核心问题是:如何通过核采样,在保持文本创造力和连贯性之间找到一个完美的平衡点?

背景知识:概率分布的"长尾"

在开始写代码之前,先用人话解释下大模型是怎么生成下一个字的。

模型预测下一个字时,会给词表里的每个字(或者 Token)算一个概率。比如预测"我吃…“的下一个字:

  • “饭”:概率 40%
  • “苹果”:概率 20%
  • “了”:概率 15%
  • “西瓜”:概率 10%
  • …(剩下几千个字瓜分剩下的 15%)

这就是所谓的"概率分布”。这里有个长尾效应:头部几个词占据了大部分概率,尾部一大堆词的概率都很小。

贪婪解码就是无脑选概率最大的"饭"。 Top-k 采样是从概率最大的 K 个词里随机选一个。比如 K=3,就在"饭"、“苹果”、“了"里选。

但这就有个坑:如果 K=3,但"饭"的概率高达 90%,另外两个才 5%,这时候强行从中选一个,选到"饭"的概率其实已经非常高了,随机性很弱;反过来,如果概率分布很扁平,头部三个词各占 30%,K=3 又可能漏掉后面概率 29% 的好词。

核采样就是为了解决这个"一刀切"的问题而生的。

实现需求:动态截断而非固定数量

我的需求很简单:

  1. 不要像贪婪解码那样死板。
  2. 不要像 Top-k 那样设定死板的数量 K。
  3. 要根据概率分布的实际情况,动态决定保留哪些词。
  4. 保留的词加起来的概率要达到一个阈值 P(比如 0.9)。

这就是 Top-p 采样的核心逻辑:从高概率到低概率累加,直到累加和超过 P,剩下的直接丢弃。

实现过程:手把手写一个核采样函数

为了彻底搞懂,我用 PyTorch 手写了一个核采样函数。这里没有黑魔法,全是基础的张量操作。

第一步:Softmax 与 Temperature

首先,我们需要把模型的输出 Logits 转换成概率分布,并应用 Temperature。Temperature 越高,分布越平缓;越低,分布越陡峭。

import torch
import torch.nn.functional as F

def apply_temperature(logits, temperature):
    """
    应用温度参数
    logits: [vocab_size] 或 [batch_size, vocab_size]
    temperature: float
    """
    # 防止除以0
    if temperature == 0:
        temperature = 1e-5
    return logits / temperature

第二步:Top-p 采样的核心逻辑

这是最关键的一步。我们需要:

  1. 把概率按从大到小排序。
  2. 计算累积概率。
  3. 找到累积概率刚刚超过 p 的那个位置。
  4. 把后面的概率全部置 0。
  5. 重新归一化。
def nucleus_sampling(logits, top_p=0.9, temperature=1.0):
    """
    实现 Top-p (Nucleus) 采样
    logits: 模型输出的原始 logits [vocab_size]
    top_p: 核采样的阈值,通常 0.9 或 0.95
    temperature: 温度参数
    """
    # 1. 应用 Temperature
    scaled_logits = apply_temperature(logits, temperature)

    # 2. 计算 Softmax 得到概率分布
    probs = F.softmax(scaled_logits, dim=-1)

    # 3. 按概率从大到小排序
    sorted_probs, sorted_indices = torch.sort(probs, descending=True)

    # 4. 计算累积概率
    cumulative_probs = torch.cumsum(sorted_probs, dim=-1)

    # 5. 移除累积概率超过 top_p 的 Token
    # 这里的减法是为了确保在临界点时保留该 Token (例如 0.9 - 0.9 = 0)
    sorted_indices_to_remove = cumulative_probs > top_p
    # 第一个超过 top_p 的 Token 也要保留
    sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
    sorted_indices_to_remove[..., 0] = 0

    # 6. 将被移除的 Token 概率置为 0
    sorted_probs[sorted_indices_to_remove] = 0.0

    # 7. 重新归一化
    sorted_probs /= sorted_probs.sum(dim=-1, keepdim=True)

    # 8. 根据排序后的索引,把概率放回原来的位置(这一步为了后续从原始采样)
    # 但实际上采样可以直接在 sorted_probs 上进行,然后映射回原始索引
    # 这里为了简化展示,直接采样
    next_token_idx = torch.multinomial(sorted_probs, num_samples=1)

    # 9. 将排序后的索引映射回原始词汇表索引
    original_next_token_idx = sorted_indices.gather(dim=-1, index=next_token_idx)

    return original_next_token_idx.item()

流程图

为了更直观地理解这个过程,我画了个流程图:

graph TD A[输入 Logits] --> B[应用 Temperature] B --> C[Softmax 转概率] C --> D[概率降序排序] D --> E[计算累积概率] E --> F{累积 > Top-p?} F -- 是 --> G[丢弃后续 Token] F -- 否 --> H[保留 Token] G --> I[概率置零] H --> I I --> J[重新归一化] J --> K[多项式采样] K --> L[输出 Token]

踩坑记录:Temperature 与 Top-p 的相爱相杀

实现了代码之后,我并没有马上得到完美的结果。在调参过程中,我踩了几个非常典型的坑,值得记录一下。

坑一:Top-p 太低,变成了伪贪婪

刚开始我设 top_p=0.5,想着保留一半的概率肯定够多样化了。结果生成的文本极其生硬,几乎和贪婪解码没区别。

原因:在大多数情况下,前几个高概率词的概率之和就已经超过 0.5 了。比如"我吃…“的例子,“饭”(0.4)+“苹果”(0.2)=0.6,已经截断了,后续的词完全没机会被选中。

解决:Top-p 通常需要设得比较高,比如 0.90.95,才能留出足够的随机空间。

坑二:Temperature 太高配合高 Top-p,直接崩盘

为了追求"更有创意”,我同时设了 temperature=1.5top_p=0.99。结果模型开始输出完全不相关的词,逻辑崩塌。

原因:Temperature 把概率分布拉得很平,Top-p 又几乎保留了所有词,这就相当于随机乱选。

解决:这两个参数是联动的。

  • 如果想稳重一点:temperature=0.7, top_p=0.9
  • 如果想创意一点:temperature=1.0, top_p=0.95
  • 切忌两个同时拉满

坑三:边界条件的死循环

在实现 sorted_indices_to_remove 时,我最开始写成了:

sorted_indices_to_remove = cumulative_probs > top_p

这会导致当某个 Token 的概率恰好让累积概率刚刚超过 top_p 时,它被移除,但累积概率又掉回 top_p 以下,下一个 Token 又被保留,逻辑非常混乱。

解决:必须处理偏移量,也就是代码中那行看起来很奇怪的:

sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()

意思是把移除信号向后移一位,确保第一个导致超标的 Token 被保留,而它之后的才被移除。

结果验证:用数据说话

为了验证我的核采样实现是否有效,我写了一段可视化代码,对比了不同策略下 Token 的概率分布变化。

可视化脚本

import matplotlib.pyplot as plt
import numpy as np

# 模拟一个词表大小为 1000 的 Logits 分布
# 假设前几个词概率很高,后面是长尾
vocab_size = 1000
logits = torch.randn(vocab_size)
# 手动放大前几个词的 Logits,模拟真实的高概率预测
logits[0] += 5
logits[1] += 3
logits[2] += 2
logits[3] += 1

# 设置参数
temperature = 0.8
top_p = 0.9

# 1. 原始 Softmax
probs_orig = F.softmax(logits, dim=-1).numpy()

# 2. Top-p 采样后的分布
# 这里复用上面的 nucleus_sampling 逻辑来获取处理后的概率矩阵用于绘图
def get_filtered_probs(logits, top_p, temperature):
    scaled_logits = logits / temperature
    probs = F.softmax(scaled_logits, dim=-1)
    sorted_probs, sorted_indices = torch.sort(probs, descending=True)
    cumulative_probs = torch.cumsum(sorted_probs, dim=-1)
    sorted_indices_to_remove = cumulative_probs > top_p
    sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
    sorted_indices_to_remove[..., 0] = 0
    sorted_probs[sorted_indices_to_remove] = 0.0
    sorted_probs /= sorted_probs.sum()
    # 恢复原始顺序
    original_probs = torch.zeros_like(probs)
    original_probs.scatter_(0, sorted_indices, sorted_probs)
    return original_probs.numpy()

probs_filtered = get_filtered_probs(logits, top_p, temperature)

# 绘图
plt.figure(figsize=(12, 6))
indices = np.arange(vocab_size)

plt.bar(indices, probs_orig, alpha=0.5, label='Original Probabilities', width=1.0)
plt.bar(indices, probs_filtered, alpha=0.8, label=f'Nucleus Sampled (Top-p={top_p})', width=1.0)

plt.yscale('log') # 使用对数坐标以便观察长尾
plt.xlabel('Token Index')
plt.ylabel('Probability (Log Scale)')
plt.title('Impact of Nucleus Sampling on Probability Distribution')
plt.legend()
plt.grid(True, which="both", ls="-", alpha=0.2)
plt.show()

观察结果

运行这段代码后,你会看到一个很有意思的现象(虽然我这里不能展示图片,但我可以描述):

  1. 头部保留:前几个高概率的 Token(索引 0, 1, 2, 3)的相对位置没有变,依然占据主要地位。
  2. 尾部截断:在对数坐标下,你会发现原本那条长长的尾巴(很小的概率值)被"切断"了,变成了 0。
  3. 重新分配:虽然绝对概率值变了,但头部 Token 之间的相对关系保持了稳定。

这正是我们想要的:保留了主要语义的可能性,切断了产生"胡说八道"的长尾噪声。

结语

折腾完这一圈,我对"生成策略"的理解从"调参玄学"变成了"数学直觉”。

核采样之所以比 Top-k 更先进,本质上是因为它自适应。它不关心你保留了几个词,只关心你保留了多少"确定性"。这种从"数量"到"质量"的思维转变,不仅适用于大模型,其实在很多工程场景都能看到影子。

最后,分享一个我个人常用的配置组合,供大家直接抄作业:

  • 创作类任务(写小说、头脑风暴)temperature=1.0, top_p=0.95
  • 总结类任务(摘要、提炼)temperature=0.7, top_p=0.9
  • 代码类任务temperature=0.2, top_p=0.95 (代码对确定性要求高,但也需要一点随机性来避免陷入死循环)

希望这篇文章能帮你把核采样这个概念从"听说过"变成"能上手"。下次调参时,别光盯着 Temperature 看了,试着动一动 top_p,或许会有惊喜。

版权声明: 本文首发于 指尖魔法屋-AI核采样:Top-k不够用了之后https://blog.thinkmoon.cn/post/351-ai-nucleus-sampling-top-k-p-sampling-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!