首词元截断 LLM 采样
从理论角度来说,LLM 给出的最佳回复应当是始终选择概率最高的那个 token。这种做法会让 LLM 的输出完全确定,而对许多应用来说这并不是一个好特性。因此,为了在保持 LLM 创造性的同时维持对上下文的贴合,近年来人们提出了各种不同的采样算法。
如今最常用的算法之一——差不多算是默认选择——叫做 top-p:它是核采样(nucleus sampling)的一种形式,把得分最高的 token 收集起来直到累计概率达到“p”,然后进行加权随机采样。
在这篇博客文章中,我将探讨为什么我认为核采样可能不是最好的方法,并展示一种简单易懂的替代方案,以避免核采样的问题。该算法目前仍在完善中,但现在发布出来,希望能引发一些讨论和尝试。
logits 中藏着金子
尽管 LLM 的 logits 是其内部运作中少数完全可以理解的部分之一,我却很少看到有人有兴趣研究它们的特征、探索更高级的采样方法、检测并向用户提示不确定性和可能的幻觉。将连续各步 token 的概率分布可视化,是一个简单而实用的练习,可以带来不少洞见:

在图中我们可以看到按概率着色的前 32 个候选 token(白色 = 0,蓝色 = 1)、被选中的 token 以及被选中 token 的排名(最高概率 = 0,次高 = 1,依此类推)。
在上面的例子中,Mistral 基础模型知道 Umberto Eco(翁贝托·埃科)的出生和去世日期,所以它自信地把大部分总概率分配给了最可能的 token。其他时候模型会更加困惑,原因要么是有多种方式可以表达文本的延续,要么是它对某些事实并不确定。询问一个训练中没学过生日的名字的日期,就会产生不同的 token 概率分布。
不过总的来说,我们要避免的是选出次优的 token,让 LLM 的生成偏离低困惑度、少幻觉的路径。核采样在这方面会失败,因为累积到“p”为止所包含的 token,取决于分布情况,可能远弱于第一选择。比如考虑上图中的“writer”这个 token。Umberto Eco 确实是一位作家。与 writer 相关联的 token 虽然得分不是很高,但仍比第二选择高出许多。然而 p 为 0.3 时可能会把一个低值如 0.01 的第二个 token 也累积进来,并以个位数的概率输出它,从而有把生成引上错误路径的风险。
相反,对于下一个 token“who”,存在多个得分相对接近的选择,它们都可以用于不同的生成结果。
利用雪崩效应
这里的一个关键观察是:为了多样性而选择过弱的 token 并不划算:即使我们只在有几个好候选的时刻利用生成来使输出多样化,雪崩效应也会帮我们:输入上下文会发生改变,从而 LLM 的输出会被扰动,使其更有可能产生文本的某种替代版本。
First Token Cutoff(FTC)算法
免责声明:这里描述的算法尚未经过任何科学审查。我花了几天时间试验了几种具有“限制最差 token”特性的采样算法,而这个算法在适用性、效果和可理解性之间看起来达到了最佳平衡。
这里描述的算法可以非正式地表述如下:
- 当 LLM 对某个候选强烈偏向时,直接选它。
- 当存在多个可行候选时,产生不同选项。
- 最差可能 token 的选取应被限制在给定范围内。
以往的工作,如 Tail Free Sampling,也指出应当在 LLM 输出的一小组高质量 token 中进行选择。但在 TFS 中,这个集合是通过求导数来确定的,即选出一条曲线之后没有出现陡峭下降的 token 簇,在那条曲线之后 token 质量急剧下降。
而在本文提出的算法中,我们希望选择遵循一个相对于最高得分 token T0 更有界、更易理解的截断点,赋予 T0 相对于所有其他 token 的特殊含义:LLM 在发出该 token 时的确定程度(基本上就是困惑度的一个代理指标)。因此,该算法拒绝所有与 T0 相比差于给定百分比的 token。
这个截断百分比取值范围为 0 到 1,称为“co”。一个示例且可行的“co”可以是 0.5。
算法的工作方式如下:
- 计算 logits 的 softmax()。
- 按概率对 token 排序。
- 设 T0 为最佳 token 的概率,计算所有其他 token 的比率:r = 1 - (T[i] / T0)
- 只选择 r <= co 的 token。
- 在选出的 token 中进行加权随机抽取。
注意,这样一来,无论 token 的数值是否平滑单调递减,我们能纳入候选集合的 token 都有一个硬性上限。而其他试图识别高分簇的方法则不具备这一特性。
实际例子
核采样 top-p 在实践中之所以看起来不会灾难性地失败,一个原因是第一个 token 的概率往往非常高,因此在困惑度低的时候,以及碰巧没有收集到再抽到低质量 token 的时候,生成都会沿着合理的路径继续。
当连续几个 token 的概率形如:
0.25, 0.14, 0.01
时,问题就更突出了。取 p=0.4 时,我们可能收集到第三个低质量 token,并以约 3% 的概率输出它。
现在考虑 co 值为 0.5 的 First Token Cutoff(token 可以比第一个最多差 50%):
第二个 token 的 r 值为:
r[t1] = 1-(0.14/0.25) = 0.44 # 0.44 <= 0.5,接受此 token
r[t2] = 1-(0.01/0.25) = 0.96 # 0.96 > 0.5,拒绝此 token
输出示例
Mistral 基础模型(非 instruct 版本)在 co=0.7 下针对提示语“Sorted sets are”的输出。以下三个输出都是成功的输出,并非为质量而精心挑选的。
- Sorted sets are a powerful data structure in Redis. They can be used to store sorted lists, to store unique values, to store scores for ranking, and to store a sorted list of sorted sets.
- Sorted sets are a very powerful data structure. They allow you to store data in a way that makes it easy to find the highest or lowest values in the set, and they also allow you to sort the data. This can be useful for many different tasks, such as ranking users by their score, or finding the most popular items in a database.
- Sorted sets are a very powerful data structure that can be used to solve many different problems. The most common use case is to store a list of unique elements, each of which has an associated value. For example, you could use a sorted set to store the names of all the people in your family, with their ages as the associated values.
为什么可理解性很重要
采样参数是 LLM 的最终用户或 API 用户必须调优的少数几样东西之一。试错往往是必需的,但拥有一个单一的可调参数,并且它能对应一个直接的现实世界描述和直觉,会有很大帮助。此外,“co”是一个线性参数,所以相比 temperature,甚至 top_p 的“p”(后者虽然是线性的,却强烈依赖于 logits 分布的形状),推理起来要简单得多。
后续工作
我的研究才刚刚开始,所以我会进一步研究和评估这个算法。最重要的是,我希望看到更多人对采样算法产生兴趣,并推动大家超越类似 top-p 的方法。
从 logits 分布中或许还能收集到有趣的信息。例如,线性探针很可能能够学会判断 LLM 的隐藏层何时在处理某些事实信息。这一点结合 token 困惑度,可以用来向 LLM 用户提示输出的某些部分很可能是错误的。总体而言,将 token 概率分布可视化信息量很大,并且在某种程度上让人能够亲手触摸到 LLM 是如何工作的,以及每一步有哪些候选。
参考实现
logits = mx.softmax(logits)
np_logits = np.array(logits) # MX -> NumPy
np_logits = np_logits.flatten()
sorted_indices = np.argsort(np_logits)
sorted_indices = sorted_indices[::-1]
co = 0.7
j = 1
t0 = np_logits[sorted_indices[0]]
while 1 - (np_logits[sorted_indices[j]] / t0) < co and j < len(np_logits):
j += 1
accepted_logits = []
for i in range(0,j):
accepted_logits.append(float(np_logits[sorted_indices[j]]))
accepted_logits = mx.array(accepted_logits)
idx = mx.random.categorical(accepted_logits)
idx = int(np.array(idx)) # Convert zero-dim array to scalar
token_id = sorted_indices[idx]致谢
多亏了 Apple 的 MLX 库以及为之不懈努力的聪明又出色的开发者们,这些实验做起来非常简单。MLX 极易上手,正如它理应如此:毕竟 LLM 本身深不可测,但推理本身是一个简单的过程。
随机一篇博客