首 Token 截断的 LLM 采样
原文由 Salvatore Sanfilippo 于 发布,订阅该博客
从理论上讲,LLM 给出的最佳回复是通过始终选择概率最高的 token 得到的。这种方式会让 LLM 的输出变得确定,而这对许多应用场景来说并不是一个理想的特性。因此,为了在保持贴合上下文的同时兼顾 LLM 的创造性,近年来提出了多种不同的采样算法。
如今最常用、几乎成为默认方案的一种叫做 top-p:它是一种核采样(nucleus sampling)的形式,会将得分最高的若干 token 累加,直到其概率总和达到“p”,然后在这些 token 中进行加权随机采样。
在这篇博文中,我将探讨为什么我认为核采样可能不是最佳方案,并展示一种简单易懂的替代方法来规避核采样的问题。该算法目前仍在完善中,但我希望通过现在发布它,能激发一些讨论和尝试。
Logits 中蕴藏的价值
尽管 LLM 的 logits 是其内部机制中少数几个完全可理解的部分,我却很少看到有人去研究它们的特性、探索更先进的采样方法,或是利用它们来检测和提示用户模型的不确定性与可能的幻觉。将连续 token 的概率分布可视化,是一个简单而实用的方法,有助于获得一些直观认识:

在上图中,可以看到按概率着色的前 32 个候选 token(白色 = 0,蓝色 = 1),以及被选中的 token 及其排名(概率最高为 0,其次为 1,依此类推)。
在上面的例子中,Mistral 基础模型知道 Umberto Eco 的生卒日期,因此它以极高的确定性将绝大部分概率赋给了最可能的 token。而在其他情况下,模型会显得更加困惑——要么是因为文本的后续有多种表达方式,要么是因为它对某些事实并不确定。比如,询问一个训练阶段未曾学过生日的人名日期时,token 的概率分布就会呈现出截然不同的形态。
不过,总的来说我们想要避免的是选中次优的 token,从而让 LLM 的生成偏离最低困惑度和最低幻觉的路径。而核采样在这方面做得并不好,因为按“p”累加 token 的方式,取决于具体的分布,有可能会纳入比首选 token 弱得多的选项。例如,上图中“writer”这个 token。Umberto Eco 确实是一位作家,与之对应的 token 虽然得分并非极高,但仍远高于第二名的候选。然而,当 p 取 0.3 时,仍有可能把概率仅为 0.01 的第二名 token 也累加进来,并以个位数的概率将其采样出来,从而有风险将生成引向错误的路径。
相比之下,对于后面的“who”这个 token,则存在多个得分相对接近的选项,它们都可以用来生成不同的文本变体。
利用雪崩效应
这里的一个关键观察是,为了增加多样性而去选择过弱的 token 并不划算:即便我们只在存在少数几个优质候选的生成时刻来制造多样性,雪崩效应也会帮到我们:输入的上下文会发生变化,从而扰动 LLM 的后续输出,使得模型更有可能产生出文本的另一种版本。
首 Token 截断(FTC)算法
免责声明:本文所述算法尚未经过任何科学验证。我花了几天时间试验了几种具有“最差 token 有界”特性的不同采样算法,这一种看起来在适用性、效果和可理解性之间取得了最佳平衡。
这里提出的算法可以非正式地概括如下:
- 当 LLM 强烈偏向某个候选时,就选择它。
- 当存在多个可行的候选时,则产生多样化的选择。
- 被选中的最差 token 应被限制在给定的范围内。
以往的研究,例如 Tail Free Sampling,也曾指出应在 LLM 生成的少量高质量 token 集合中进行选择。不过在 TFS 中,这个集合是通过求导来识别的,以选出对应于曲线未出现陡峭下降的那一组 token——在此之后 token 的质量会急剧降低。
而在本文提出的算法中,我们希望选择遵循一个相对于最高分 token T0 更为有界且易于理解的截断,将 T0 赋予相对于其他所有 token 的特殊意义:即 LLM 在生成该 token 时的确定性程度(基本上可视为困惑度的一个代理指标)。因此,该算法会拒绝所有相对于 T0 差于给定百分比的 token。
这个截断百分比取值范围为 0 到 1,被称为“co”。例如,0.5 就是一个可行且合理的 co 取值。
该算法的工作流程如下:
- 计算 logits 的 softmax()。
- 按概率对 token 进行排序。
- 给定最佳 token 的概率 T0,按 r = 1 - (T[i] / T0) 计算其他所有 token 的比值
- 仅保留满足 r <= co 的 token
- 在筛选出的 token 中进行加权随机采样。
需要注意的是,通过这种方式,无论 token 的概率是否呈现平滑单调递减,都对可纳入候选集合的 token 设下了硬性上限。而在其他试图识别高分簇的方法中,则并非如此。
实践示例
核采样 top-p 在实践中之所以看起来没有灾难性地失败,一个原因在于首个 token 的概率往往非常高,因此当困惑度较低时,再加上很多时候我们并不会偶然地收集并选中低质量的 token,生成过程就能沿着合理的路径继续下去。
而当出现如下这样的连续 token 概率时,问题就会变得比较棘手:
0.25, 0.14, 0.01
当 p=0.4 时,我们可能会把第三个低质量的 token 也纳入进来,并以约 3% 的概率将其采样出来。
现在考虑使用首 Token 截断,取 co 值为 0.5(即 token 最多可比首个 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 基础模型(非指令微调版)在 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 本身是难以捉摸的,但推理过程本身却是一个简单的过程。
随机一篇博客
评论
登录后参与讨论