现代 LLM 采样参数完全指南
深度讲解 temperature、top-p 等采样策略原理与使用场景。对构建 LLM 应用的程序员实用价值高。
深度讲解 temperature、top-p 等采样策略原理与使用场景。对构建 LLM 应用的程序员实用价值高。
大语言模型(LLM)通过获取一段文本(例如用户提示词)并计算下一个单词来工作。在更技术的术语中,称为 token。LLM 拥有一个词汇表(或有效 token 的字典),在训练和推理(文本生成过程)中会引用这些 token。更多内容见下文。首先,你需要理解为什么我们使用 token(子词)而不是单词或字母。但首先,让我们了解一些在下面章节中没有深入解释的技术术语的简明词汇表:
Logits:模型对其词汇表中每个 token 输出的原始、未归一化的分数。较高的 logits 表示模型认为更可能出现的 token。 Softmax:一个数学函数,将 logits 转换为适当的概率分布——介于 0 和 1 之间的值,总和为 1。 Entropy:概率分布中的不确定性或随机性的度量。更高的 entropy 意味着模型对应该选择哪个 token 不太确定。 Perplexity:与 entropy 相关,perplexity 衡量模型对文本的"惊讶"程度。更低的 perplexity 表示更高的置信度。 n-gram:n 个 token 的连续序列。例如,"once upon a" 是一个 3-gram。 Context window(或序列长度):LLM 一次能处理的最大 token 数,包括提示词和生成的输出。 Probability distribution:为所有可能的结果(token)分配概率的函数,使其总和为 1。可以把它想象成百分比:如果 1% 是 0.01,50% 是 0.5,100% 是 1.0。
你的直觉可能是使用单词或字母词汇表来实现 LLM。但我们使用的是子词:一些常见的单词在词汇表中被完整保留(例如 the 或 apple 可能是单个 token,因为它们在英语中非常常见),但其他单词被分段为常见的子词(例如 bi-fur-cat-ion)。为什么这样做?有好几个非常充分的理由:
有很多理由。举几个例子:LLM 的 context window(一次能处理的 token 数量)是有限的。使用字符级 tokenization,即使是中等数量的文本也会导致序列长度爆炸(过少的文本却产生过多的 token)。单词 tokenization 可能需要 12 个 token,而在子词系统中只需要 2 或 3 个。更长的序列还需要进行更多的自注意力计算。但更重要的是,模型需要学习跨越多个位置的高级模式。理解 t-h-e 代表单个概念需要在三个位置之间连接信息,而不是一个位置。这也可能导致有意义的关系距离更远。相关概念可能相隔数十或数百个位置。
纯单词级 tokenization 要求我们创建覆盖整个英语词汇表的词汇表,如果处理多种语言,规模会大很多倍。这会使嵌入矩阵变得不合理地庞大和昂贵。它对新词或稀有词的处理也很困难。当模型遇到不在其词汇表中的单词时,通常会用"未知"token 替换它,几乎完全丧失语义信息。子词 tokenization 可以通过组合现有的子词 token 来表示新词。例如,如果我们创造一个叫 "grompuficious" 的新词,子词 tokenizer 可能会将其表示为 g-romp-u-ficious,这取决于具体的 tokenizer。另一个值得一提的方面是形态学感知:许多语言通过组合语素(前缀、词根、后缀)来创建单词。例如,如前所述,unhappiness 可以分解为 un-happi-ness;子词 tokenization 可以自然地捕捉这些关系。它还允许我们执行跨语言迁移。对于具有复杂形态或复合结构的语言(例如德语或芬兰语,其中单词可能是极长的组合),这特别有用。
如果语言模型使用新的 tokenizer,开发团队可能会从训练数据中取一个具有代表性的样本,并训练一个 tokenizer 来找到数据集中最常见的子词。他们会预先设定一个词汇表大小,然后 tokenizer 会尽力找到足够的子词来填充该列表。
在训练期间,模型会处理许多 TB 级别的文本,并为 token 构建内部概率映射。例如,在看到 "How are" 之后通常跟随 "you?" 的 token 时,它会学到这是最可能的下一个 token 序列。一旦这个映射被构建到令人满意的程度,训练就会停止,一个检查点被发布到公众(或保持私密并通过 API 提供,例如 OpenAI)。在推理期间,用户会向 LLM 提供文本,LLM 会根据通过训练学到的概率来决定下一个 token 是什么。但它不会只决定一个 token:它会考虑其词汇表中存在的每一个可能 token,为每个分配一个概率分数,并且(根据你的采样器)只输出最可能的 token,即分数最高的那个。这会产生相当无聊的输出(除非你需要确定性),所以这就是采样(Sampling)发挥作用的地方。
现在我们理解了 LLM 如何使用 token 来分解和表示文本,让我们探索它们实际如何生成内容。LLM 中的文本生成过程涉及两个关键步骤:
预测:对于每个位置,模型计算其词汇表中所有可能下一个 token 的概率分布。
选择:模型必须从这个分布中选择一个 token 来添加到不断增长的文本中。
第一步是固定的——由模型在训练后的参数决定。但是第二步——token 选择——是采样发生的地方。虽然我们可以简单地总是选择最可能的 token(称为"贪心"采样),但这往往会产生重复、确定性的文本。采样引入了受控的随机性,使输出更加多样化。
如上所述,LLM 会选择概率最高的 token 进行生成。采样是一种引入受控随机性的做法。如果使用纯粹的“贪心”采样,它每次都会选择排名第一的选项,但那也太无聊了!我们使用温度、惩罚或截断等采样方法,允许生成结果出现一些创造性的变化。本文将介绍所有主流的采样方法,并从通俗易懂和技术两个角度解释它们的工作原理。
本文中的算法均以伪代码形式呈现,将数学符号与编程概念结合在一起。以下是一些有助于理解这些算法的说明:
L:logits 张量(模型输出的原始分数)
P:概率(对 logits 应用 softmax 后得到)
←:赋值操作(相当于编程中的 =)
|x|:根据上下文表示 x 的绝对值或长度/大小
x[i]:访问 x 的第 i 个元素
∨:逻辑 OR 操作
¬:逻辑 NOT 操作
∞:无穷大(通常通过将 logits 设置为负无穷来屏蔽 token)
argmax(x):返回 x 中最大值的索引
∈:“属于”或“是……的元素”(例如,x ∈ X 表示 x 是集合 X 中的元素)
这里提供的算法以清晰易懂为目标,而非追求最优性能。生产环境中的实现通常会:
尽可能对操作进行向量化,以提升效率
处理边界情况和数值稳定性问题(不过,下文的算法偶尔也会特别标出需要注意这些问题的部分)
如果框架需要,则为多个序列加入批处理
在有益的情况下缓存中间结果
可以把它想象成 LLM 的“创造力旋钮”。在较低的温度下(接近 0),模型会变得非常谨慎且可预测——它几乎总是选择概率最高的下一个词。这就像你每次去最喜欢的餐厅都点同一道菜,因为你知道自己会喜欢(又或者,你只是并不知道还有其他选择)。在较高的温度下(例如 0.7~1.0),模型会变得很有创造力,也更愿意冒险。它可能会选择概率排名第三或第四的词,而不是永远选择排名第一的词。这会让文本更加多样、更有趣,但也会增加出错的概率。非常高的温度(高于 1.0)会让模型变得狂野且不可预测,除非将其与其他采样方法(例如 min-p)结合使用,以约束模型。
技术原理:温度通过直接操纵整个词表上的概率分布来发挥作用。模型会为词表中的每个 token 生成 logits(未归一化分数),然后将其除以温度值。当温度小于 1 时,高 logits 会相对变得更高,低 logits 会相对变得更低,从而形成更加尖锐的分布,使得分数最高的 token 更有可能被选中。当温度大于 1 时,分布会变得更加平坦,高分与低分 token 之间的概率差距缩小,从而增加随机性。应用温度之后,修改过的 logits 会被转换为概率分布(使用 softmax),随后从该分布中随机采样一个 token。从数学效果来看,温度 T 对概率的变换,本质上相当于先将每个概率提升至 1/T 次幂,然后重新归一化。
1
2
3
4
5
6
7
8
Algorithm 1 Temperature Sampling
Required: Logits tensor L, temperature parameter T
Output: Modified logits with adjusted probability distribution
1: if T < 0.1 then
2: L ← L - max(L) + 1 // Shift range to [-inf, 1] for numerical stability
3: end if
4: L ← L / T // Apply temperature scaling
5: return L
这种方法会抑制模型重复任何此前出现过的 token,无论该 token 已经使用过多少次。可以把它想象成一位希望确保每个人都有机会发言的派对主持人。如果 Tim 已经发过一次言,那么他再次发言的意愿就会受到些许抑制,不管他之前是说过一次还是十次。通常不推荐这种方法,因为存在更好的惩罚策略(参见:DRY)。
技术原理:存在惩罚会对生成文本中出现过的任何 token 施加固定惩罚。首先,我们使用输出掩码识别此前已经使用过的 token(对于至少出现过一次的 token,该掩码值为 True)。然后,从这些 token 的 logits 中减去存在惩罚值。这样一来,此前使用过的 token 再次被选中的可能性就会降低,无论它们出现得有多频繁。任何至少使用过一次的 token 都会受到相同程度的惩罚。
1
2
3
4
5
6
7
8
Algorithm 2 Presence Penalty
Required: Logits tensor L, output tokens O, penalty weight λp
Output: Modified logits with penalty applied for token presence
1: Vsize ← |L[0]| // Vocabulary size from logits dimension
2: Moutput ← BinaryMask(O, Vsize) // Create binary mask where token has appeared at least once
3: P ← λp · Moutput // Calculate penalty matrix
4: L ← L - P // Apply presence penalty to logits
5: return L
这种方法会根据 token 已经使用过的次数对其进行抑制。简单来说,它就是存在惩罚,但还会将出现次数纳入考虑。一个词出现得越频繁,再次出现的可能性就越低。
技术原理:频率惩罚会将每个 token 此前出现的次数乘以惩罚值,再从该 token 的 logit 分数中减去结果。我们会追踪每个 token 在生成输出中出现过多少次;如果某个 token 已经出现三次,其 logit 就会减少 3 x (frequency penalty)。这样便形成了一种渐进式惩罚:每当 token 被重复使用一次,惩罚都会随之增加。
1
2
3
4
5
6
7
8
Algorithm 3 Frequency Penalty
Required: Logits tensor L, output tokens O, penalty weight λf
Output: Modified logits with penalty applied proportional to token frequency
1: Vsize ← |L[0]| // Vocabulary size from logits dimension
2: Coutput ← TokenCounts(O, Vsize) // Count occurrences of each token in output
3: P ← λf · Coutput // Calculate penalty matrix proportional to counts
4: L ← L - P // Apply frequency penalty to logits
5: return L
重复惩罚的工作方式与前两种方法略有不同:它会同时惩罚提示词和生成输出中的 token,并以不同方式处理正 logits 和负 logits。对于正分数,它会将分数除以惩罚系数(使其变小);对于负分数,它会将分数乘以惩罚系数(使其变得更负)。这种方法有助于跳出循环,但当取值较为激进时,会以连贯性为代价。
技术原理:该惩罚会应用于提示词或当前生成文本中已经出现过的 token。我们会创建一个合并提示词掩码与输出掩码的掩码。对于这个合并掩码中的 token,我们会根据 logits 是正数还是负数,通过除法或乘法来施加惩罚。这种方法可以避免对原本就不太可能出现的 token 施加过度惩罚,同时有效降低高分 token 的概率。它最显著的特征是对正 logits 和负 logits 采取不同处理方式,这有助于维持更加均衡的概率分布。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
Algorithm 4 Repetition Penalty
Required: Logits tensor L, prompt tokens P, output tokens O, penalty factor λr
Output: Modified logits with asymmetric penalty for repeated tokens
1: Vsize ← |L[0]| // Vocabulary size from logits dimension
2: Mprompt ← BinaryMask(P, Vsize) // Create binary mask for tokens in prompt
3: Moutput ← BinaryMask(O, Vsize) // Create binary mask for tokens in output
4: M ← Mprompt ∨ Moutput // Combined mask for tokens in either prompt or output
5: R ← matrix of size |L| × Vsize filled with λr
6: R[¬M] ← 1.0 // No penalty for tokens not previously seen
7: for each position (i,j) in L do
8: if L[i,j] > 0 then
9: L[i,j] ← L[i,j] / R[i,j] // Divide positive logits
10: else
11: L[i,j] ← L[i,j] * R[i,j] // Multiply negative logits
12: end if
13: end for
14: return L
DRY 采样就像安排了一位编辑,专门检查你写作中的重复模式。假设你正在写一个故事,并且已经使用过短语“once upon a time”——DRY 会抑制你再次使用完全相同的短语。但它比简单地防止词语重复(例如上面的三种惩罚)聪明得多。
DRY 会在文本中寻找重复模式(称为 n-gram)。如果它发现你之前写过类似“the cat sat on the”的内容,而你即将重复这一模式,它就会降低对下一个会延续该重复模式的词的偏好。重复模式越长,抑制力度就越强。这可以防止文本陷入循环或反复使用相同短语,使输出保持新鲜和多样。DRY 的特殊之处在于,它会考虑已经存在的模式。它尤其适合创意写作,因为重复的文本读起来会很不自然。
技术细节:DRY 采样通过检测 n-gram 重复,并惩罚那些会延续这些模式的 token 来工作。该算法检查目前为止生成的 token 序列,并识别以最近生成的 token 结尾的重复模式。我们会追踪最后一个 token 在文本其他位置出现的地方,然后比较这些位置之前的上下文。当它找到匹配的上下文(重复的 n-gram)时,就会识别此前跟在这些模式之后的 token,并在当前 logits 分布中对它们施加惩罚。关键参数包括 multiplier(惩罚强度)、base(n-gram 越长时惩罚增强的幅度)、需要考虑的最小 n-gram 长度,以及需要检查的最大 n-gram 长度。该算法还会识别“序列断点”(例如标点符号),它们会重置模式匹配;同时还支持范围限制,只考虑近期文本以提高效率。
惩罚会根据匹配模式的长度以指数形式施加,匹配越长,惩罚越强。这形成了一个动态系统,既能防止重复,又允许文本自然流动;通过避免语言模型经常出现的重复模式,使输出文本更加多样、更像人类写作。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
Algorithm 5 DRY (Don't Repeat Yourself) Sampling
Required: Logits tensor L, input tokens I, output tokens O, multiplier λ, base b,
minimum n-gram length Nmin, sequence breaker tokens B, range limit r,
maximum n-gram length Nmax, maximum occurrences M, early exit threshold E
Output: Modified logits with penalties for repeating patterns
1: for each sequence s where λ > 0 do
2: // Prepare token sequence by concatenating prompt and generated tokens
3: p ← Length of valid tokens in I[s]
4: q ← Length of valid tokens in O[s]
5: T ← Concatenate(I[s][1:p], O[s][1:q])
6:
7: if r > 0 then
8: T ← T[-r:] // Consider only the last r tokens if range limit is set
9: end if
10:
11: if |T| < 2 then continue end if // Skip if sequence too short
12:
13: last ← T[-1] // Last token in sequence
14: if last ∈ B then continue end if // Skip if last token is a sequence breaker
15:
16: // Create mask for sequence breaker positions
17: breakMask ← ZeroVector(|T|)
18: for each token id b ∈ B do
19: breakMask ← breakMask ∨ (T = b)
20: end for
21:
22: // Find maximum allowed n-gram length before hitting a sequence breaker
23: maxN ← 0
24: for i from 1 to min(|breakMask|, Nmax) do
25: if breakMask[-i] then break end if
26: maxN ← i
27: end for
28:
29: if maxN ≤ Nmin then continue end if // Skip if maximum n-gram too short
30:
31: // Initialize array to track longest matching n-gram for each token
32: ngramLengths ← ZeroVector(VocabularySize)
33:
34: // Find all positions where the last token appears
35: endpoints ← FindIndices(T = last)
36: if |endpoints| < 2 then continue end if
37:
38: // Remove the last occurrence (current position)
39: endpoints ← endpoints[:-1]
40:
41: // Limit number of previous occurrences to check
42: if |endpoints| > M then
43: endpoints ← endpoints[-M:]
44: end if
45:
46: // Check each previous occurrence of the last token for matching contexts
47: for each idx in Reverse(endpoints) do
48: if idx = |T| - 1 then continue end if
49:
50: matchLen ← 0
51: // Look backward to find matching context
52: for u from 1 to min(idx, maxN) do
53: if breakMask[idx - u] then break end if
54: if T[idx - u] ≠ T[-u - 1] then break end if
55: matchLen ← u
56: end for
57:
58: if matchLen > 0 then
59: nextToken ← T[idx + 1] // Token that followed this pattern before
60: newLen ← matchLen + 1
61: ngramLengths[nextToken] ← max(ngramLengths[nextToken], newLen)
62:
63: if newLen ≥ E then break end if // Early exit if match is long enough
64: end if
65: end for
66:
67: // Apply penalties to tokens that would continue repeating patterns
68: penaltyMask ← (ngramLengths > 0)
69: if any(penaltyMask) then
70: scales ← b ^ (ngramLengths[penaltyMask] - Nmin) // Exponential scaling by pattern length
71: L[s][penaltyMask] ← L[s][penaltyMask] - λ * scales
72: end if
73: end for
74: return L
模型不会考虑所有可能的下一个词(数量可能多达数万个),而是将范围缩小到最有可能的 K 个候选项。如果 K 为 40,模型就只会从概率最高的 40 个候选词中进行选择。这种方法既能防止模型选中概率极低的词,又能保留一定的随机性。
技术细节:该方法会对所有可能的下一个 token 的 logits 进行排序,只保留最高的 K 个值,并将其他值全部设为负无穷。首先,我们按升序排列 logits,并获取它们的索引。然后,找出第 K 大的值,为所有低于该阈值的值创建一个掩码。任何低于该阈值的 logit 都会被设为 -inf,确保这些 token 在应用 softmax 后的概率实际上为零。随后,使用保存的索引将过滤后的 logits 恢复为原始顺序。
1
2
3
4
5
6
7
8
9
Algorithm 6 Top-K Sampling
Required: Logits tensor L, parameter k
Output: Modified logits with all but top-k options filtered out
1: Lsorted, Lidx ← Sort(L, descending=False) // Sort logits in ascending order
2: kth ← Lsorted[|Lsorted| - k] // Find the kth largest logit value
3: mask ← Lsorted < kth // Create mask for values below the threshold
4: Lsorted[mask] ← -∞ // Filter out tokens below threshold
5: L ← Unsort(Lsorted, Lidx) // Restore original ordering
6: return L
Top-P 不像 Top-K 那样选择固定数量的候选项,而是选择一个最小的词集合,使其累积概率超过阈值 P。这就像是在说:“我只考虑占这家餐厅全部订单 90% 的那些菜品。”如果 P 为 0.9,模型就会纳入刚好足以