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% 的好词。
核采样就是为了解决这个"一刀切"的问题而生的。
实现需求:动态截断而非固定数量
我的需求很简单:
- 不要像贪婪解码那样死板。
- 不要像 Top-k 那样设定死板的数量 K。
- 要根据概率分布的实际情况,动态决定保留哪些词。
- 保留的词加起来的概率要达到一个阈值 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 采样的核心逻辑
这是最关键的一步。我们需要:
- 把概率按从大到小排序。
- 计算累积概率。
- 找到累积概率刚刚超过 p 的那个位置。
- 把后面的概率全部置 0。
- 重新归一化。
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()
流程图
为了更直观地理解这个过程,我画了个流程图:
踩坑记录:Temperature 与 Top-p 的相爱相杀
实现了代码之后,我并没有马上得到完美的结果。在调参过程中,我踩了几个非常典型的坑,值得记录一下。
坑一:Top-p 太低,变成了伪贪婪
刚开始我设 top_p=0.5,想着保留一半的概率肯定够多样化了。结果生成的文本极其生硬,几乎和贪婪解码没区别。
原因:在大多数情况下,前几个高概率词的概率之和就已经超过 0.5 了。比如"我吃…“的例子,“饭”(0.4)+“苹果”(0.2)=0.6,已经截断了,后续的词完全没机会被选中。
解决:Top-p 通常需要设得比较高,比如 0.9 或 0.95,才能留出足够的随机空间。
坑二:Temperature 太高配合高 Top-p,直接崩盘
为了追求"更有创意”,我同时设了 temperature=1.5 和 top_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()
观察结果
运行这段代码后,你会看到一个很有意思的现象(虽然我这里不能展示图片,但我可以描述):
- 头部保留:前几个高概率的 Token(索引 0, 1, 2, 3)的相对位置没有变,依然占据主要地位。
- 尾部截断:在对数坐标下,你会发现原本那条长长的尾巴(很小的概率值)被"切断"了,变成了 0。
- 重新分配:虽然绝对概率值变了,但头部 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/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。