一文搞懂大模型生成时的随机性控制:Top‑k + Temperature 核心逻辑与源码级拆解
文章目录
1、基本介绍
一、名字解析
- 温度调节(Temperature Scaling)
“温度” 借用物理学概念:温度越高,系统越混乱;温度越低,系统越有序。在语言模型生成中,温度用于控制输出分布的随机性程度——高温使概率分布更平坦,增加随机性;低温使分布更尖锐,趋向于选择高概率词。
- Top‑k 采样(Top‑k Sampling)
“Top‑k” 指只保留模型预测概率最高的 (k) 个 token,丢弃其余低概率 token,然后仅从这 (k) 个 token 中随机采样。这样做既避免了采样到极不合理的词,又保留了适度的随机性。
因为 Top‑k 在采样之前,先截断了概率最低的那部分词(只保留概率最高的 k 个),然后仅从这 k 个词中随机选取。这样,那些概率极低(如 10−510−5 以下)的“不合理词”根本不会出现在候选池里,自然就不会被抽到。
需要强调的是,这个优势是相对于“纯随机采样”而言的——纯随机采样会从整个词表中按原始概率抽取,有可能抽到尾部的低质量词;而 Argmax 虽然也不会选低概率词,但它完全没有随机性,容易导致输出单调、重复。Top‑k 则在保留随机性的同时,用截断保证了候选词的基本质量。
两者结合,就是:先通过温度调节软化/锐化概率分布,再从中筛选出概率最高的 k 个 token,最后按重新归一化的概率进行随机采样。
二、背景:为什么需要它?
语言模型的自回归解码中,最朴素的方法是 Argmax(贪婪解码)——每步都选概率最大的 token。其缺点包括:
- 缺乏多样性,容易产生重复或僵化的输出;
- 一旦选错,错误会持续累积,无法修正;
- 生成的句子往往过于“安全”而显得生硬。
为了引入可控的随机性,同时避免采样到完全离谱的词,便有了“Top‑k 采样 + 温度调节”这一经典解码策略。
在自回归语言模型(如 Transformer)中,每一步都需要根据当前已生成的 token 决定下一个 token。最朴素的方法有两种:
| 方法 | 做法 | 主要问题 |
|---|---|---|
| Argmax(贪婪解码) | 每次选择概率最大的 token | 缺乏多样性,易重复、僵化;一旦选错,错误会持续累积 |
| 纯随机采样 | 按原始 softmax 概率随机采样 | 可能采样到极低概率的无意义词,生成质量不稳定 |
为了在 质量(避免胡言乱语)与 多样性(避免重复、生硬)之间取得平衡,同时保留 可控的随机性,便诞生了“Top‑k 采样 + 温度调节”这一经典解码策略。
三、数学公式详解
步骤1:模型输出 logits
设词汇表大小为 ( V ),当前时刻模型输出的 logits 向量为:
z
=
[
z
1
,
z
2
,
…
,
z
V
]
\mathbf{z} = [z_1, z_2, \dots, z_V]
z=[z1,z2,…,zV]
这些值表示模型对每个 token 的原始打分。
步骤2:温度缩放
引入温度参数 ( T > 0 ),对 logits 进行缩放:
z
′
=
z
T
\mathbf{z}' = \frac{\mathbf{z}}{T}
z′=Tz
- 当 ( T > 1 ) 时,logits 绝对值变小,后续 softmax 后的概率分布更均匀(高概率与低概率词的差距缩小),随机性增强。
- 当 ( 0 < T < 1 ) 时,logits 绝对值变大,概率分布更尖锐,高概率词的权重更大,随机性降低。
- 当 ( T = 1 ) 时,保持原始分布。
温度越高,高分和低分的差距拉得越小;温度越低,高分和低分的差距拉得越大
步骤3:计算概率分布(全词表,可选,)
对缩放后的 logits 应用 softmax,得到温度调节后的全词表概率分布:
p
i
=
exp
(
z
i
/
T
)
∑
j
=
1
V
exp
(
z
j
/
T
)
p_i = \frac{\exp(z_i / T)}{\sum_{j=1}^{V} \exp(z_j / T)}
pi=∑j=1Vexp(zj/T)exp(zi/T)
这一步在实际高效实现中通常被省略,改为对筛选后的候选集直接做 softmax(下文也有)。
步骤4:Top‑k 筛选
设定参数 ( k )(通常为 10~100)。从分布 ( \mathbf{p} ) 中找出概率最大的 ( k ) 个 token,记它们的索引集合为 ( \mathcal{I}_{\text{top-k}} )。对于不在集合内的 token,将其概率置为 0:
p
^
i
=
{
p
i
if
i
∈
I
top-k
0
otherwise
\hat{p}_i = \begin{cases} p_i & \text{if } i \in \mathcal{I}_{\text{top-k}} \\ 0 & \text{otherwise} \end{cases}
p^i={pi0if i∈Itop-kotherwise
步骤5:重新归一化
由于 Top‑k 操作后概率和不再为 1,需重新归一化,得到最终采样用的分布:
q
i
=
p
^
i
∑
j
∈
I
top-k
p
^
j
q_i = \frac{\hat{p}_i}{\sum_{j \in \mathcal{I}_{\text{top-k}}} \hat{p}_j}
qi=∑j∈Itop-kp^jp^i
步骤6:随机采样
根据分布 ( \mathbf{q} ) 进行多项式采样(multinomial sampling),即按概率 ( q_i ) 随机抽取一个 token。这一步通常用 PyTorch 的 torch.multinomial 实现。
这一步是随机的,因此每次生成结果可能不同(除非固定随机种子)。
四、流程图解
原始 logits z
│
▼
[ Temperature 缩放: z / T ]
│
▼
[ 取 top-k 个最大 logits(可直接在 logits 上取) ]
│
▼
[ 对 top-k logits 做 softmax → 概率分布 ]
│
▼
[ 重新归一化(若候选集概率和≠1,此步已由 softmax 自动完成)]
│
▼
[ Multinomial 采样 → 下一个 token ]
五、它是干什么的?(作用与优点)
-
避免不合理输出
Top‑k 通过丢弃尾部大量低概率 token,有效防止采样出语法错误、语义不通或无意义的词。 -
引入可控的随机性
温度参数让你可以精细调节“保守度”与“创造性”的平衡:- 翻译、摘要等要求高准确性的任务:低温(( T \approx 0.6 \sim 0.8 ))+ 较小 k(如 10)→ 结果更确定。
- 对话生成、故事创作等需要多样性的任务:高温(( T \approx 1.0 \sim 1.2 ))+ 较大 k(如 50)→ 输出更丰富。
-
缓解曝光偏差
在训练中使用 Free Running 时,若采样策略与推理一致,能让模型更好地适应自己生成时的输入分布,减少训练‑推理差异。 -
实现简单,兼容性强
只需在 softmax 之前做温度缩放,再加上一个 Top‑k 筛选,即可用纯 PyTorch 实现,不依赖任何高级库。
六、与相关方法的对比
| 方法 | 特点 | 适用场景 |
|---|---|---|
| Argmax | 每次选最高概率,确定性 | 极小模型、快速测试、必须确定输出的场景 |
| 随机采样(无过滤) | 从全词汇表采样 | 多样性极高,但易产生无意义词 |
| Top‑k 采样 | 从概率最高的 k 个 token 中采样 | 平衡质量与多样性 |
| Top‑p(核采样) | 从累积概率超过 p 的最小 token 集合中采样 | 动态调整候选集大小,更灵活 |
| 束搜索 | 保留多条候选路径,取整体最优 | 翻译、摘要等追求高质量的任务 |
Top‑k + 温度 常与 Top‑p 结合使用(例如先 Top‑k 再 Top‑p),但单独使用已能大幅提升生成质量。
七、实际代码示例(PyTorch)
体现原理的版本:
import torch
import torch.nn.functional as F
def top_k_sampling(logits, k=50, temperature=1.0):
# logits: [vocab_size] 或 [batch, vocab_size]
# 温度缩放
logits = logits / temperature
# softmax 得到概率
probs = F.softmax(logits, dim=-1)
# Top‑k 筛选
top_k_probs, top_k_indices = torch.topk(probs, k)
# 重新归一化
top_k_probs = top_k_probs / top_k_probs.sum(dim=-1, keepdim=True) # 这里又变成了概率, 和为1
# 采样
sampled_idx = torch.multinomial(top_k_probs, num_samples=1)
# 映射回原始 token id
token_id = top_k_indices.gather(dim=-1, index=sampled_idx)
return token_id
高效版本:先取 top-k 的 logits,再 softmax
该代码的详细解释在后面有详情
import torch
import torch.nn.functional as F
def top_k_sampling_with_temperature(logits, k=10, temperature=1.0):
"""
logits: [batch_size, vocab_size] 或 [vocab_size]
k: 保留的候选 token 数量
temperature: 温度参数 (>0)
返回: 下一个 token 的索引,维度与输入 batch 维度一致(若输入为 1D,返回 Python int)
"""
# 统一处理 batch 维度
was_1d = (logits.dim() == 1)
if was_1d:
logits = logits.unsqueeze(0) # [1, vocab_size]
# 1. 温度缩放
logits = logits / temperature
# 2. 在 logits 上直接取 top-k 【高效实现:先取 top-k 的 logits,再 softmax(避免计算全词表)】
# topk: 默认是降序排列(从大到小排列)
top_k_logits, top_k_indices = torch.topk(logits, k, dim=-1)
# 3. 对 top-k logits 做 softmax(得到归一化后的概率)
top_k_probs = F.softmax(top_k_logits, dim=-1)
# 4. 从 top-k 中采样(torch.multinomial 后面有详情)
sampled_idx_in_topk = torch.multinomial(top_k_probs, num_samples=1) # [batch, 1]
# 5. 映射回原始词表索引(torch.gather 后面有详情)
# .squeeze(-1) 这个是降维,不是升维
next_token = torch.gather(top_k_indices, -1, sampled_idx_in_topk).squeeze(-1) # [batch]
# 恢复原始维度
if was_1d:
next_token = next_token.item()
return next_token
说明:
- 上述代码采用了“先取 top‑k 的 logits,再 softmax”的高效方式,与数学公式中的“全 softmax 再截断”在数学上等价(因为 softmax 的分母只依赖于候选集内部)。
- 若需批处理,函数已支持 batch 维度。
七、注意事项
- k 值选择
- k 过小 → 候选集太窄,可能错过合理但概率略低的词。
- k 过大 → 引入过多低质量候选,随机性失控。
通常通过实验确定,常见取值范围 10~100。
- 与 Free Running 的一致性
若在训练的计划采样阶段也采用相同的采样策略(温度与 k 值可略低于推理,但保留随机性),有助于减小曝光偏差。 - 与其他方法的组合
Top‑k 与 Top‑p 可结合:先取 top‑k 进一步过滤尾部,再在剩余候选中做 top‑p 动态筛选。两者结合可得到更稳定、高质量的生成。 - 适用性
虽然束搜索在机器翻译等确定性任务中表现更好,但 Top‑k + 温度调节在无法使用束搜索时是最佳的替代方案,尤其适合从零实现模型、训练阶段引入随机性的场景。
参数设置建议(以汉译英为例)
| 任务类型 | Temperature | Top‑k |
|---|---|---|
| 机器翻译(推理) | 0.7 ~ 1.0 | 10 ~ 20 |
| 摘要生成 | 0.8 ~ 1.0 | 20 ~ 50 |
| 创意写作 | 1.0 ~ 1.2 | 50 ~ 100 |
| 对话系统 | 0.9 ~ 1.1 | 30 ~ 60 |
针对汉译英任务:
- 训练中的 Free Running(计划采样):建议 T = 0.9,k = 10~15。保留一定探索性,让模型习惯自己生成的分布。
- 推理(暂不用束搜索时):建议 T = 0.7,k = 10。偏向确定性,提高翻译准确率。
十、总结
Top‑k 采样 + 温度调节 是一种简单、高效、可控的文本生成策略,在现代语言模型中被广泛使用:
- 温度 控制整体随机性强度;
- Top‑k 负责过滤低质量候选;
- 二者结合,既避免了 argmax 的僵硬,又防止了无约束采样的荒谬。
它不仅是推理阶段的有效解码方法,也是训练中计划采样(Free Running)的理想选择,能显著提升模型对自身生成数据的适应能力。掌握这一技术,将为你在自定义 Transformer 项目中实现高质量翻译生成打下坚实基础。
2、多项式采样(multinomial sampling)
以下是对上一份“多项式采样(multinomial sampling)”回答的修订版,修正了代码示例中关于 logits 的错误表述,并补充了相关说明,确保内容准确、严谨。
一、名字由来
multinomial 是 “多项分布” 的英文。
- “multi” = 多个
- “nominal” = 名义的、类别的
多项分布(Multinomial Distribution) 是二项分布的推广,用于描述在 多次独立试验 中,每个可能结果出现的次数 的概率分布。在 torch.multinomial 中,它被简化为 单次试验的抽样 或 有放回/无放回的多次抽样。
简单理解:多项分布 = 投掷一个有 V 个面的骰子(每个面权重不同),一次试验会落到哪个面。
二、数学原理
- 多项分布的概率质量函数(PMF)(看不懂就直接跳过,不重要)
假设试验有 ( V ) 种可能结果,每种结果的概率为 ( p_1, p_2, \dots, p_V ),且 ( \sum_{i=1}^V p_i = 1 )。在 (N) 次独立试验中,各结果出现次数 ( n_1, n_2, \dots, n_V ) 的概率为:
P
(
n
1
,
…
,
n
V
)
=
N
!
n
1
!
⋯
n
V
!
p
1
n
1
⋯
p
V
n
V
P(n_1, \dots, n_V) = \frac{N!}{n_1! \cdots n_V!} p_1^{n_1} \cdots p_V^{n_V}
P(n1,…,nV)=n1!⋯nV!N!p1n1⋯pVnV
其中 ( \sum n_i = N )。
你感觉复杂很正常,因为那个公式确实比较抽象。我来用更直白的方式解释清楚,让你知道
torch.multinomial到底在算什么,不需要死记硬背这个公式。
一、核心概念:从“掷骰子”理解
想象你有一个 不均匀的骰子:
- 面 1 出现的概率是 0.2
- 面 2 出现的概率是 0.5
- 面 3 出现的概率是 0.3
一次试验:投一次这个骰子,结果会是 1、2 或 3 中的一个,概率就是上面这些。
torch.multinomial做的就是这件事:给定一组概率(权重),它帮你“投一次骰子”(或多次),返回投出来的面编号。这其实就是 类别分布(Categorical Distribution),多项分布在 ( N=1 ) 时的特例。
二、那复杂的公式是干什么的?
那个概率质量函数(PMF)描述的是 连续投 N 次骰子后,每个面出现的次数恰好是某个组合的概率。
例如:投 10 次,面 1 出现 2 次、面 2 出现 5 次、面 3 出现 3 次的概率是多少?那个公式就是算这个的。
但在
torch.multinomial的常规用法中,我们 几乎只用它做单次抽样(num_samples=1),或者最多做少量有放回/无放回的抽样,并不需要那个复杂的组合公式。你完全可以忽略它,把它当作“按概率抽一个”的工具。
三、所以,你只需要知道三件事
torch.multinomial按给定权重随机选一个(或多个)索引。- 权重不需要归一化,函数内部会自动处理。
- 它常用于文本生成,在 softmax 之后从中采样下一个 token,而不是每次都选最大的。
四、更直观的代码对照
import torch # 骰子每个面的概率 probs = torch.tensor([0.2, 0.5, 0.3]) # 投一次骰子 sample = torch.multinomial(probs, 1) print(sample) # 可能是 tensor([1]) 对应面 2每次运行可能得到不同结果,但面 2 被抽中的概率最高。
五、总结:别被公式吓跑
那个复杂公式是为“多次试验次数分布”准备的,而你在文本生成中只用到了它最简单的功能——按概率随机选一个。把这个核心理解透,就足够你实现 Top‑k 采样了。
- 单次抽样(( N=1 ))
当 ( N = 1 ) 时,多项分布退化为 类别分布(Categorical Distribution),即一次试验中每个类别被抽中的概率就是 ( p_i )。torch.multinomial 主要处理这种单次抽样(num_samples=1),或者有放回/无放回的多次抽样(num_samples>1)。
你可以把复杂的数学定义抛开,直接这样理解:
probs=[0.2, 0.5, 0.3]意味着:
- 索引 0 有 20% 的机会被抽中。
- 索引 1 有 50% 的机会被抽中。
- 索引 2 有 30% 的机会被抽中。
💡 为什么这么简单,还要叫“多项分布”?
其实,“多项分布” 这个名字主要是为了强调**“多次抽样”时的统计规律,但在实际写代码(比如
torch.multinomial)时,我们往往只关心单次**结果。我们可以分两个层面来看:
- 单次层面(你现在的理解)
这就是一个简单的“抽奖”动作。
- 就像你手里有一个不均匀的骰子,扔一次,看它停在哪一面。
- 代码表现:
torch.multinomial(probs, 1)返回一个数字(比如1)。- 结论:概率就是
probs[i]。
- 多次层面(数学定义的“多项分布”)
这是指“扔很多次”后的统计结果。
- 如果你扔 1000 次,数学上会预测:索引 0 大约出现 200 次,索引 1 大约出现 500 次,索引 2 大约出现 300 次。
- 代码表现:
torch.multinomial(probs, 1000)返回 1000 个数字。- 结论:虽然名字叫“多项分布”,但
torch.multinomial这个函数本质上就是帮你执行一次次独立的“单次抽奖”(在有放回模式下)。📌 总结
在写代码(特别是做 AI 推理)时,你就把它当成一个“加权随机数生成器”:
- 给一堆权重(概率)。
- 让它吐出一个索引。
- 权重越大,被吐出来的概率越大。
就这么简单!
- 采样过程
给定权重向量 ( \mathbf{w} = [w_1, w_2, \dots, w_V] )(不必归一化),torch.multinomial 内部:
- 有放回采样:每次独立地根据归一化后的概率 ( p_i = w_i / \sum w_j ) 抽取一个索引。
- 无放回采样:依次抽取,每次抽取后将被抽中的索引从候选集中移除,剩余概率重新归一化,再继续下一次抽取。
实际实现通常使用 逆变换采样:(具体实现按照这个理解就完全够了)
- 计算累积概率(或累积权重)( c_i = \sum_{j=1}^i w_j )。
- 生成均匀随机数 ( u \in [0, \text{总和}) )。
- 找到最小的 ( i ) 使得 ( c_i \geq u ),则结果索引为 ( i )。
对于无放回,重复上述过程但每次排除已选索引。
🌰 示例
工作原理示例(单次抽样,我们 几乎只用它做单次抽样(num_samples=1))
假设概率向量:
probs = torch.tensor([0.2, 0.5, 0.3]) # 三个类别,概率分别为 0.2, 0.5, 0.3
torch.multinomial(probs, 1) 内部步骤:
- 计算累积概率:
[0.2, 0.7, 1.0]。 - 生成均匀随机数 u ∈ [0, 1),可不服从正太分布哈,例如 0.65)。
- 找到第一个
i使得累积概率 ≥ u:i=1(因为 0.7 ≥ 0.65)。 - 返回
1。
解释1:为什么最后一个值能被抽到?
在逆变换采样中:
- 累积概率序列:
[0.2, 0.7, 1.0]- 随机数 u∈[0,1) 均匀分布
关键点:虽然 u 不能等于 1,但它可以无限接近 1,例如 0.95。
此时,第一个满足“累积概率 ≥ u”的是最后一个累积概率1.0,因为它 ≥ 0.95。
所以索引 2 仍然有概率被抽中,其概率正好等于第三个区间的长度 1.0−0.7=0.31.0−0.7=0.3。总结:
- 索引 0 对应区间 [0, 0.2)
- 索引 1 对应区间 [0.2, 0.7)
- 索引 2 对应区间 [0.7, 1.0)
每个区间的长度等于对应概率,所以最后一个区间能正常覆盖。
解释2:累积概率不会让后面的更难选,它只是把概率值转换成了区间长度。
- 每个索引 被选中的概率 = 它对应的 区间长度
- 区间长度 = 它的原始概率
举例:
概率: [0.2, 0.5, 0.3] 累积: [0.2, 0.7, 1.0] 区间: [0,0.2) [0.2,0.7) [0.7,1.0) 长度: 0.2 0.5 0.3
- 索引 0 的区间长度 0.2 → 被选概率 20%
- 索引 1 的区间长度 0.5 → 被选概率 50%
- 索引 2 的区间长度 0.3 → 被选概率 30%
后面的区间并不因为“累计”而变小,它的长度就是它自己的概率。
“后面部分不容易选到”只发生在它的原始概率本身就很小时,这正是我们想要的。你担心的“前部分的容易选到”是因为前面索引的概率(0.2、0.5)加起来已经 0.7,随机数落在前面的概率自然大,但这完全由原始概率决定,不是累积方法造成的。
多次运行则会:
- 以 0.2 的概率分别返回 0
- 以 0.5 的概率分别返回 1
- 以 0.3 的概率分别返回 2
无放回采样的内部机制(补充)
初始权重: [2, 5, 3], 总和=10
第1次: 抽中索引1(权重5)
剩余: [2, 3] (移除索引1)
重新归一化: 概率变为 [2/5, 3/5] = [0.4, 0.6]
第2次: 从 [0, 2] 中按 [0.4, 0.6] 抽取
...
注意:无放回采样不是简单地把概率置零再归一化,而是物理移除已选索引,保证不会重复。
三、torch.multinomial 是干什么的?
功能:从给定的概率分布(或权重)中进行随机采样,返回采样的索引。
典型场景:
- 文本生成:从 softmax 输出的概率分布中采样下一个 token,替代 argmax 以引入随机性。
- 强化学习:从动作概率分布中采样动作,用于探索。
- 重采样:如粒子滤波中根据权重抽取样本。
四、函数签名与参数(后面有详情)
torch.multinomial(input, num_samples, replacement=False, *, generator=None, out=None)
input(Tensor):输入张量,形状(..., V),最后一维表示每个类别的权重(不需要归一化,函数会自动处理)。通常传入概率或未归一化的分数。num_samples(int):采样的次数(即每个分布抽取几个索引)。replacement(bool):是否允许重复采样(有放回)。True表示有放回,False表示无放回(此时num_samples必须 ≤ 最后一维的大小)。generator:可选,随机数生成器。out:可选,输出张量。
返回值:形状为 (..., num_samples) 的张量,元素是采样的索引(0 到 V-1),数据类型为 torch.long。
五、代码示例
- 单次抽样(无放回,单次即一次)
import torch
probs = torch.tensor([0.2, 0.5, 0.3])
sample = torch.multinomial(probs, 1)
print(sample) # 可能输出 tensor([1])
- 有放回采样 5 次
samples = torch.multinomial(probs, 5, replacement=True)
print(samples) # 可能输出 tensor([1, 1, 2, 0, 1]),允许重复
- 无放回采样(需
num_samples ≤ 类别数)
samples = torch.multinomial(probs, 2, replacement=False)
print(samples) # 输出两个不同的索引,如 tensor([1, 2])
- 批量处理(输入形状 [batch, V])
batch_probs = torch.tensor([
[0.1, 0.9],
[0.5, 0.5],
[0.8, 0.2]
])
samples = torch.multinomial(batch_probs, 1) # 每行独立采样一个
print(samples) # 形状 [3,1],如 tensor([[1], [0], [0]])
- 正确使用 logits 进行采样
torch.multinomial 将输入直接视为权重,进行线性归一化。若想基于 softmax 概率采样,应先用 softmax 处理:
logits = torch.tensor([1.0, 2.0, 3.0]) # 未归一化分数
probs = torch.softmax(logits, dim=-1) # 转为概率
samples = torch.multinomial(probs, 1) # 基于 softmax 概率采样
如果直接传入 logits(如 torch.multinomial(logits, 1)),则采样基于线性权重 [1,2,3],等价于概率 [1/6, 2/6, 3/6],不等价于 softmax,需特别注意。
七、注意事项
-
输入不需要严格归一化
torch.multinomial会自动将输入视为权重,通过除以总和进行归一化。但为清晰起见,通常传入 softmax 后的概率或未归一化的正数权重。 -
无放回时
num_samples不能超过类别总数
若replacement=False,则采样数量必须 ≤ 最后一维大小,否则会报错。 -
数据类型
输入应为浮点型(float32/float64),返回值为torch.long。 -
随机性控制
可通过torch.manual_seed(seed)或generator参数固定随机性,便于复现。 -
与
torch.distributions.Categorical的关系
Categorical是更高级的分布封装,提供了sample()等方法,但torch.multinomial更底层、更轻量。 -
输入值建议非负
尽管函数内部会处理,但为了数值稳定性,建议输入非负权重。若包含负值,可能导致意外行为。
八、在文本生成中的应用
在 Transformer 解码过程中,你会先获得 logits,然后:
# 假设 decoder_output 形状 [batch, vocab]
logits = decoder_output[:, -1, :] # 取最后一个时间步
# 应用温度调节
logits = logits / temperature
# 可选:Top‑k 或 Top‑p 过滤
top_k_logits, top_k_indices = torch.topk(logits, k)
# 采样
probs = torch.softmax(top_k_logits, dim=-1)
# 单次抽样
# probs = [0.2, 0.5, 0.3] 表示抽中第 0 个的概率是 0.2,第 1 个的概率是 0.5,第 2 个的概率是 0.3。torch.multinomial(probs, 1) 就是按照这个概率进行一次随机抽取。
sampled_idx = torch.multinomial(probs, 1) #
# torch.gather 后面有详情
next_token = torch.gather(top_k_indices, -1, sampled_idx)
这里的 torch.multinomial 就是从候选集中随机选择一个 token,实现“随机采样”而非贪心选择。
九、总结
| 要点 | 说明 |
|---|---|
| 名称 | multinomial = 多项分布 |
| 功能 | 根据给定的概率/权重进行随机抽样 |
| 参数 | input(权重), num_samples(抽样次数), replacement(是否放回) |
| 返回值 | 采样的索引(long 型) |
| 核心原理 | 基于累积概率分布和均匀随机数实现类别采样;无放回时逐步移除已选索引并重新归一化 |
| 应用 | 文本生成、强化学习、重采样等需要随机选择的场景 |
掌握了 torch.multinomial,你就拥有了在生成过程中引入可控随机性的基本工具,这也是实现 Top‑k 采样、Top‑p 采样等策略的核心依赖。
3、多项式采样 - API
📘 torch.multinomial API 详解
- 函数签名
torch.multinomial(
input, # 【必填】输入张量(Tensor),包含每个类别被选中的概率(权重)
num_samples, # 【必填】整数(int),表示要抽取多少个样本
replacement=False, # 【可选】布尔值,默认是 False(不放回采样)。
# 设为 True 表示允许同一个类别被重复抽到(放回采样)
*, # (* 后面是关键字参数,调用时必须写成 key=value 的形式)
generator=None, # 【可选】默认是 None。用于控制随机性的生成器,设了它可以让结果可复现
out=None # 【可选】默认是 None。指定输出结果的存储位置(一般不用管)
)
- 参数详解
-
input(Tensor)-
含义: 包含权重的张量。
-
形状: 可以是 1维 (单个分布,形状
[num_categories]) 或 2维 (批量分布,形状[batch_size, num_categories])。 不能是其它维度,
must be 1 or 2 dim -
数据类型: 必须是浮点型 (
torch.float16,torch.float32,torch.float64)。如果是整数类型会报错。 -
数值要求:
- 必须是非负的 ( w i ≥ 0 w_i \ge 0 wi≥0)。
- 不需要归一化:你不需要手动把它变成概率(和为1),函数内部会自动对最后一维进行归一化(除以总和)。
- 全零错误:如果某一行的所有权重都为 0,会报错(因为无法计算概率)。
- 关于 Logits:虽然可以直接传入 Logits,但必须确保 Logits 是非负的(例如经过 ReLU 处理)。如果 Logits 包含负数(这是常态),不能直接传给
multinomial,必须先经过Softmax转为概率。
-
-
num_samples(int)- 含义: 你要从分布中抽取多少个样本。
- 注意: 如果
input是 2 维的,这个参数表示每一行都要抽取这么多样本。
-
replacement(bool)- 默认值:
False False(无放回): 抽到的元素不会放回,同一个索引在一次采样操作中不会重复出现。- 限制: 此时
num_samples必须 ≤ \le ≤input的最后一维大小(类别总数)。
- 限制: 此时
True(有放回): 抽到的元素会放回,同一个索引可以重复出现。- 优势:
num_samples可以大于类别总数。
- 优势:
- 默认值:
-
generator(torch.Generator)- 含义: 用于控制随机性的生成器。如果你需要复现结果,可以通过它设置随机种子。
- 返回值
- 类型:
torch.LongTensor - 形状:
- 如果
input是 1 维,返回形状为[num_samples]。 - 如果
input是 2 维,返回形状为[batch_size, num_samples]。
- 如果
- 内容: 采样得到的索引 (0 到 N − 1 N-1 N−1)。
- 常用操作与代码示例
基础用法:单次采样 (最常用)
这是 LLM 生成文本时最常用的模式,每次生成一个 token。
import torch
# 假设这是模型输出的概率(已经归一化,且非负)
probs = torch.tensor([0.1, 0.2, 0.7])
# 抽取 1 个样本
# 结果大概率是 2 (因为 0.7 最大)
result = torch.multinomial(probs, num_samples=1)
print(result) # 输出示例: tensor([2])
批量采样 (Batch Processing)
在训练或并行生成时,我们经常需要同时处理多个序列。torch.multinomial 支持直接传入 2D 张量。
# 2个序列,每个序列对应 3 个候选词的概率
batch_probs = torch.tensor([
[0.1, 0.1, 0.8],
[0.8, 0.1, 0.1]
])
# 每个序列各抽取 1 个样本
result = torch.multinomial(batch_probs, num_samples=1)
print(result)
# 输出示例:
# tensor([[2],
# [0]])
有放回采样
如果你想一次性生成多个 token(例如在束搜索或并行解码中),可以设置 replacement=True。
weights = torch.tensor([0.5, 0.5])
# 抽取 5 次,允许重复
result = torch.multinomial(weights, num_samples=5, replacement=True)
print(result)
# 输出示例: tensor([0, 1, 1, 0, 1])
无放回采样 (去重)
如果你需要从一组选项中选出 k k k 个不重复的元素。
weights = torch.tensor([1.0, 1.0, 1.0, 1.0, 1.0]) # 5个选项,权重相等
# 抽取 3 个不重复的索引
result = torch.multinomial(weights, num_samples=3, replacement=False)
print(result)
# 输出示例: tensor([0, 4, 2]) -> 索引互不相同
- 进阶:配合 Temperature 和 Top-K
在实际的大模型推理中,我们通常处理的是原始 Logits(包含负数)。因此,必须先进行 Softmax 归一化才能传给 multinomial。
这是一个标准的带温度调节的采样流程:
import torch
import torch.nn.functional as F
# 1. 假设这是模型输出的原始 logits (包含负数,未归一化)
logits = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0])
# 2. 温度调节 (Temperature)
temperature = 0.8
scaled_logits = logits / temperature
# 3. (可选) Top-K 过滤
# 只保留概率最大的 k 个,其他的置为负无穷
top_k = 3
# 注意:这里操作的是 scaled_logits,而不是原始的 logits
# 关于 【torch.topk(scaled_logits, top_k)[0][..., -1, None]】 后面会有详解
a = torch.topk(scaled_logits, top_k)[0][..., -1, None]
print(a) # tensor([3.7500]) , 至于为什么是这样, 后面有详解
indices_to_remove = scaled_logits < a
scaled_logits[indices_to_remove] = float('-inf')
# 4. 计算概率 (Softmax)
# 这一步至关重要:将 Logits 转换为非负的概率分布
probs = F.softmax(scaled_logits, dim=-1)
# 5. 多项式采样
next_token_idx = torch.multinomial(probs, num_samples=1)
print(f"选中的索引: {next_token_idx.item()}") # 比如:4
- 常见坑与注意事项
| 问题 | 说明 | 解决方案 |
|---|---|---|
| 负数 Logits | 直接传入包含负数的 Logits | 必须先经过 Softmax 转为概率,否则结果错误。 |
| 数据类型错误 | 传入 int 类型的张量 | 使用 .float() 转换类型。 |
| 无放回超限 | num_samples > 类别数量,且 replacement=False | 减小 num_samples 或改为 replacement=True。 |
| 全零权重 | 输入张量某一行全是 0 | 检查数据源,确保至少有一个正权重。 |
| Device 不一致 | input 在 GPU 上,但 generator 在 CPU 上 | 确保所有张量和生成器都在同一个设备上。 |
- 总结
torch.multinomial 的核心就一句话:给它一堆非负权重,它根据权重比例给你返回索引。
- 推理时:通常
num_samples=1,配合 Temperature 和 Softmax 使用。 - 训练/探索时:可能用到
replacement=True来生成多样化的序列。 - 输入:可以是概率,也可以是非负权重。如果是原始 Logits,务必先 Softmax。
4、[…]:占位符
核心前提:这可不是普通列表能玩的
首先,也是最重要的一点:Python 原生的列表是不支持 ... 语法的!
如果你直接写 a = [[1, 2, 3]] 然后尝试 a[..., -1],Python 会直接报错 TypeError。因为原生列表只认识整数索引或切片,看不懂 Ellipsis 对象。
... 是 NumPy 和 PyTorch 等科学计算库的“特权”。 所以,接下来的所有例子,我们都默认 a 是一个 PyTorch Tensor 或 NumPy 数组。
一句话总结
... 是一个“智能占位符”,它的意思是:“这里省略了一堆冒号 :,请自动帮我把剩下的维度都填满。”
它的学名叫 Ellipsis(省略号),是写高维张量代码时的“偷懒”神器。
举个最直观的例子
假设你有一个 5维 的张量(比如视频数据):
[批次, 时间, 颜色, 高度, 宽度]
形状是:(2, 10, 3, 32, 32)
如果你想取第一个视频的所有画面数据:
-
写法 A(不用省略号):
你需要手动写满剩下的冒号,非常手酸且容易数错。data[0, :, :, :, :] -
写法 B(使用省略号):
data[0, ...]
发生了什么?
0:锁定了第 1 个维度(批次)。...:PyTorch/NumPy 自动帮你补全了后面剩下的 4 个冒号:, :, :, :。- 结果:完全一样,但代码清爽了无数倍。
核心规则与玩法
... 会代表**“剩下的所有维度”**。它会根据你写的位置,自动膨胀来填补空缺。
假设有一个 3维 张量 x,形状 (2, 3, 4)。
-
放在后面:
x[0, ...]- 含义:我要第 0 个元素,后面剩下的维度我全都要。
- 等同于:
x[0, :, :] - 结果形状:
(3, 4)
-
放在前面:
x[..., 0]- 含义:前面的维度我全都要(遍历所有),只要每个里面的第 0 个元素。
- 等同于:
x[:, :, 0] - 结果形状:
(2, 3)
-
夹在中间:
x[0, ..., 1]- 含义:我要第 0 个块,中间剩下的维度全要,但最后只要索引为 1 的元素。
- 等同于:
x[0, :, 1] - 结果形状:
(3,)
为什么它这么好用?
-
拯救“冒号密集恐惧症”
维度越高,冒号越多。用...可以让代码瞬间清爽。 -
让代码“不挑数据”(通用性)
这是它最强大的地方。- 如果你写死
x[0, :, :],万一明天数据变成了 4 维或 10 维,代码就报错了。 - 但如果你写
x[0, ...],不管数据是 3 维、4 维还是 100 维,它都能完美运行——它会自动适配剩下的所有维度。
- 如果你写死
避坑指南
-
一个切片操作里,只能有一个
...。x[0, ...](正确)x[..., 0](正确)x[..., 0, ...](报错! Python 会困惑你到底想省略哪一部分)
-
必须导入库
别忘了import torch或import numpy as np,并且数据必须是 Tensor 或 Array。
总结
看到 ...,你就把它当成**“等等等等”或者“剩下的全都要”**。
- 它不是三个点,它是一个智能填充工具。
- 它专门用来拯救那些维度太多、写冒号写到吐的代码。
5、[-1] vs […, -1] 详解
这两个操作在一维列表(或一维张量)中效果是一模一样的,但在多维数据中,区别就非常大了。
简单来说:... 是一个“偷懒”的符号,意思是“这里省略了一堆维度”。
我们可以用**“俄罗斯套娃”或者“书架”**来打比方。
[-1]:只针对最外面的一层
- 含义:我要取最外层容器的最后一个元素。
- 比喻:你有一个书架,
[-1]就是拿走书架最右边的那一整层(不管这一层里有多少书)。
[..., -1]:穿透所有中间层,直达核心
- 含义:
...代表“中间的任意层”,-1代表“每一层的最后一个”。它的意思是:不管套了多少层娃,我要取最里面的那个娃娃的最后一个。 - 比喻:你要打开书架上的每一个抽屉,从每一个抽屉里都拿出最右边的那本书。
🌰 举个具体的例子(二维数据)
假设我们有一个二维列表(就像一个表格):
data = [
[1, 2, 3], # 第0行
[4, 5, 6] # 第1行
]
操作 A:data[-1]
- 动作:取列表的最后一个元素。
- 结果:
[4, 5, 6] - 解释:它把
[4, 5, 6]当作一个整体拿走了。
操作 B:data[..., -1]
- 动作:
...说:“我要遍历前面的所有维度(也就是每一行)”。-1说:“在每一行里,我要最后一个元素”。
- 结果:
[3, 6] - 解释:它穿透到了内部,分别取了第0行的最后一个(3)和第1行的最后一个(6)。
🌰 再上一个 3维 的例子,看看这两个操作的区别有多大。
假设我们有一个 3维张量(形状是 2x2x3),你可以把它想象成 2个班级,每个班级有 2个小组,每个小组有 3个学生。
数据如下:
import torch
# 形状: (2个班级, 2个小组, 3个学生)
tensor = torch.tensor([
[ [1, 2, 3], [4, 5, 6] ], # 班级 0
[ [7, 8, 9], [10, 11, 12] ] # 班级 1
])
a = tensor[..., -1]
print(a)
# tensor([[ 3, 6],
# [ 9, 12]])
操作一:tensor[-1]
含义:取最外层(第0维)的最后一个元素。
-
动作:不管里面有多复杂,我只取“班级”维度的最后一个,也就是**“班级 1”的完整数据**。
-
结果:
tensor([ [ 7, 8, 9], [10, 11, 12] ]) -
形状变化:从
(2, 2, 3)变成了(2, 3)。维度降低了,因为你切走了一层皮。
操作二:tensor[..., -1]
含义:... 代表“前面的所有维度保持不变”,-1 代表“取最里面(最后一维)的最后一个元素”。
-
动作:
- 保留“班级”维度。
- 保留“小组”维度。
- 在“学生”维度上,只取最后一个(也就是每个小组的第3个学生)。
-
结果:
tensor([ [ 3, 6], # 班级 0 的各组最后一名 [ 9, 12] # 班级 1 的各组最后一名 ]) -
形状变化:从
(2, 2, 3)变成了(2, 2)。维度也降低了,但它是把最里面的维度“压扁”提取出来了。
📌 直观对比图
-
tensor[-1]:就像切蛋糕,横着切一刀,把最下面那一整块拿走了。
-
tensor[..., -1]:就像用吸管插进蛋糕,竖着插到底,把每一块蛋糕的最右边那一角都吸出来了。
🧠 记忆口诀
[-1]:“我要最后那一大块。”(针对最外层)[..., -1]:“我要每一小块里的最后一个。”(针对最内层)
💡 为什么要用 ...?
在 PyTorch 或 NumPy 中,数据经常是很多维的(比如 [Batch大小, 句子长度, 词向量维度])。
如果你想取每一个句子的最后一个词,你不需要写死维度(比如 [:, -1, :]),你可以直接写 [..., -1]。
- 好处:不管你的数据是 2 维的、3 维的还是 10 维的,
[..., -1]永远能精准地帮你取到最内层的最后一个数据,代码写起来更通用、更简洁。
📌 总结
[-1]:切走最后一片(不管这一片有多厚)。[..., -1]:在所有片里,都只取最后那个芯。
6、[None]:增加一个维度
已经把 ... 这个“占位符”搞明白了,那现在我们可以毫无障碍地来拆解 [None] 了。
在 PyTorch 或 NumPy 的索引操作里,None 的作用非常单一且强大:它是一个“维度扩充器”。
📌 一句话总结
None 的作用就是:在它出现的那个位置,强行插入一个长度为 1 的新维度。
它不会改变数据里的数值,只会改变数据的形状(Shape),把数据“撑”起来。
🌰 最直观的例子:从“线”变“板”
想象你有一个一维的列表(像一条线):
import torch
x = torch.tensor([1, 2, 3])
print(x.shape) # 输出: torch.Size([3])
它只有 3 个数字,是一维的。
如果你加上 None
y = x[:, None]
print(y)
# 输出:
# tensor([[1],
# [2],
# [3]])
print(y.shape) # 输出: torch.Size([3, 1])
发生了什么?
:表示“把原来的数据都保留”。None表示“在这里加一个新的维度”。- 结果:原本趴着的一维数组,被
None拉了起来,变成了一个 3行1列 的二维矩阵(列向量)。
🧩 常见的三种“变身”玩法
假设我们有一个二维张量 A,形状是 (3, 4)(3行4列)。
玩法 1:A[:, None, :](在中间加一层)
- 含义:
::保留第1维(3)。None:插入一个新维度(变成1)。::保留第2维(4)。
- 结果形状:
(3, 1, 4) - 形象理解:把原本扁平的矩阵,变成了“三明治”,中间夹了一层厚度为1的维度。
玩法 2:A[None, :, :](在最前面加一层)
- 含义:
None:插入一个新维度(变成1)。::保留原来的所有维度(3, 4)。
- 结果形状:
(1, 3, 4) - 形象理解:把原本的一个矩阵,变成了一个“只有1张图”的批次(Batch)。这在深度学习输入数据时非常常用。
玩法 3:A[..., None](在最后加一层)
- 含义:
...:代表前面所有的维度(3, 4)。None:在最后插入一个新维度(变成1)。
- 结果形状:
(3, 4, 1) - 形象理解:把每个数字都包进了一个小盒子里。
🤔 为什么要这么麻烦?(核心用途)
你可能会问:“把 [1, 2, 3] 变成 [[1], [2], [3]] 有什么意义?”
意义在于“广播机制”(Broadcasting)。
当你想要做矩阵运算(比如乘法、加法)时,PyTorch 要求两个张量的形状必须匹配。
- 如果一个是
(3, 4),另一个是(4,),它们可能无法直接按你想要的方式运算。 - 但如果你用
None把(4,)变成(1, 4)或者(4, 1),PyTorch 就能瞬间明白:“哦,你是想把这个向量应用到每一行(或每一列)上!”
📌 总结
看到索引里的 None,你就把它翻译成:“在这里切一刀,增加一个厚度为1的维度”。
- 它是维度变换的神器。
- 它等价于函数
torch.unsqueeze(x, dim=...),但在代码里写[:, None]更简洁、更 Pythonic。
7、torch.topk(scaled_logits, top_k)[0][…, -1, None]
🔍 代码拆解:torch.topk(scaled_logits, top_k)[0][..., -1, None]
这行代码看起来像天书,但其实它只是一个三步走的流水线。
它的终极目标:找到 Top-K 里最小的那个分数(也就是“门槛分数”),并且把它调整成正确的形状,以便后续把低于这个门槛的分数全部过滤掉。
🥇 第一步:选出优胜者 torch.topk(...)
torch.topk(scaled_logits, top_k)
- 作用:从所有的分数中,找出分数最高的
top_k个。 - 返回值:一个元组
(values, indices)。values:这top_k个优胜者的具体分数。indices:这些分数在原始列表中的位置(索引)。
🎯 第二步:只要分数,不要索引 [0]
torch.topk(scaled_logits, top_k)[0]
- 作用:我们只关心分数是多少,不关心它们原来在哪。所以用
[0]取出元组里的第一个元素,也就是values。 - 结果:一个只包含
top_k个最高分数的张量。 - 关键点:
torch.topk默认是从大到小排列的。 - 举例:假设
top_k=3,原始分数是[1.0, 5.0, 2.0, 4.0, 3.0]。- 这一步的结果就是
[5.0, 4.0, 3.0]。
- 这一步的结果就是
🚪 第三步:找出“门槛”并调整形状 [..., -1, None]
这是最关键的一步,它包含两个动作:
- 找到“门槛”分数
[..., -1]
...:代表前面所有的批次维度(如果有的话),我们都要照顾到。-1:取最后一个元素。- 作用:因为上一步的结果是从大到小排列的,所以最后一个元素就是这
top_k个分数里最小的那个。这个分数就是我们的“淘汰线”。 - 接上例:从
[5.0, 4.0, 3.0]中取出最后一个,也就是3.0。任何低于3.0的分数都将被淘汰。
- 增加一个维度
[..., None]
None:在当前位置插入一个长度为 1 的新维度。- 作用:这是为了广播机制。我们得到的“门槛”分数(如
3.0)需要和原始的scaled_logits进行<比较。为了让 PyTorch 能够自动、正确地将这个“门槛”分数应用到原始 logits 的每一个对应位置上,我们必须给它增加一个维度,把它变成一个“列向量”。
📌 形状变化追踪表(假设输入是 2个句子,每个5个词)
假设 scaled_logits 的形状是 (2, 5),top_k=3。
| 步骤 | 代码片段 | 结果形状 | 说明 |
|---|---|---|---|
| 1. 原始数据 | scaled_logits | (2, 5) | 2个句子,每句5个词 |
| 2. 选出TopK | .topk(...)[0] | (2, 3) | 每句只留分数最高的3个 |
| 3. 找门槛 | [..., -1] | (2,) | 取出每句的第3高分(最低入围分) |
| 4. 升维 | [..., None] | (2, 1) | 关键! 变成列向量,准备广播 |
最终结果:一个形状为 (2, 1) 的张量,里面装着每个句子的“淘汰门槛分数”,准备用于下一步的过滤操作。
没问题,那我们就把维度再加一层,看看当第 0 维变成批量大小 B 时,整个流程会有什么不同。
假设现在的输入是一个 3维 张量,形状为 (B, S, V),比如 (2, 3, 5)。
- B=2:批次大小,代表 2 个样本。
- S=3:序列长度,代表每个样本有 3 个句子。
- V=5:词表大小,代表每个句子有 5 个词的分数。
torch.topk 默认是在最后一个维度(也就是词表维度)上进行操作的。
📊 3维数据下的形状变化追踪表
| 步骤 | 代码片段 | 结果形状 | 详细说明 |
|---|---|---|---|
| 1. 原始数据 | scaled_logits | (2, 3, 5) | 2个批次,每个3句,每句5词 |
| 2. 选出TopK | .topk(...)[0] | (2, 3, 3) | 最后一维从 5 变成了 k=3。保留了前两个维度。 |
| 3. 找门槛 | [..., -1] | (2, 3) | 取最后一维的最后一个数。也就是每个句子的“门槛分”。 |
| 4. 升维 | [..., None] | (2, 3, 1) | **关键!**在最后强行加一个维度,准备广播。 |
🧠 深度解析
操作对象
虽然数据变成了 3 维,但 torch.topk(x, k) 依然只关心最后一维。它相当于在每一个“句子”上独立地做了一次 Top-K 筛选。
省略号的作用
在步骤 3 和 4 中,... 完美地代表了前面的 (2, 3) 这两个维度。
[..., -1]的意思是:不管前面是 2 维还是 10 维,我只在乎最后一个维度的最后一个数。[..., None]的意思是:不管前面是什么形状,我只在最后面加一个维度。
广播的用途
最终得到的形状是 (2, 3, 1)。
这个张量通常会被用来和原始的 scaled_logits(形状 (2, 3, 5))进行比较。
- PyTorch 会自动把
(2, 3, 1)在最后一维复制 5 次,变成(2, 3, 5)。 - 这样就可以实现:用每个句子的门槛分,去过滤该句子原本的 5 个词。
🌰 代码演示
import torch
# 为了让结果一眼能看懂,我们这里用整数,不用随机数
# 假设数据是:[[10, 20, 30, 40, 50], [5, 4, 3, 2, 1]]
# 也就是 2个句子,每个5个词
logits = torch.tensor([[10, 20, 30, 40, 50], [5, 4, 3, 2, 1]])
top_k = 3
print(f"原始数据:\n{logits}")
print(f"形状: {logits.shape}\n")
# 1. 选出 Top-K
# 结果应该是每行最大的3个数
topk_values = torch.topk(logits, top_k)[0]
print(f"1. TopK 选出的分数 (每行最大的 {top_k} 个):\n{topk_values}")
print(f" 形状变化: {logits.shape} -> {topk_values.shape}")
# 2. 取出“门槛”分数 [..., -1]
# 取每行最后一个(也就是TopK里最小的那个)
threshold = topk_values[..., -1]
print(f"\n2. 取出的门槛分数 (每行的第 {top_k} 大分):\n{threshold}")
print(f" 形状变化: {topk_values.shape} -> {threshold.shape}")
# 3. 增加维度 [..., None]
# 强行把 (2,) 变成 (2, 1)
threshold_expanded = topk_values[..., -1, None]
print(f"\n3. 增加维度后的门槛 (准备广播):\n{threshold_expanded}")
print(f" 形状变化: {threshold.shape} -> {threshold_expanded.shape}")
# 输出:
原始数据:
tensor([[10, 20, 30, 40, 50],
[ 5, 4, 3, 2, 1]])
形状: torch.Size([2, 5])
1. TopK 选出的分数 (每行最大的 3 个):
tensor([[50, 40, 30],
[ 5, 4, 3]])
形状变化: torch.Size([2, 5]) -> torch.Size([2, 3])
2. 取出的门槛分数 (每行的第 3 大分):
tensor([30, 3])
形状变化: torch.Size([2, 3]) -> torch.Size([2])
3. 增加维度后的门槛 (准备广播):
tensor([[30],
[ 3]])
形状变化: torch.Size([2]) -> torch.Size([2, 1])
8、torch.gather(小难)
三维确实是理解 torch.gather 的分水岭。但只要掌握了**“对号入座”**的规律,其实比二维更直观。
我们把三维张量想象成一个**“多层货架”**:
dim=0(层):代表第几层货架。dim=1(行):代表货架上的第几排。dim=2(列):代表排里的第几个位置。
torch.gather 的核心逻辑永远是:index 里的数字,就是用来替换 dim 对应的那个坐标的。
下面我们用同一个“货架”数据,分别演示 dim=0, 1, 2 是怎么取的。
通用公式
对于输出张量中任意一个位置 (i, j, k)(下标按维度顺序),它的值由以下规则确定:
- 当
dim=0时:
output[i][j][k] = input[ index[i][j][k] ] [j] [k]
即第一维的索引用index中的值代替,其他维索引保持不变。 - 当
dim=1时:
output[i][j][k] = input[i] [ index[i][j][k] ] [k]
即第二维的索引用index中的值代替,其他维索引保持不变。 - 当
dim=2时:
output[i][j][k] = input[i] [j] [ index[i][j][k] ]
即第三维的索引用index中的值代替,其他维索引保持不变。
📦 准备数据:一个 (2层, 3排, 4个) 的货架
假设我们的 input 形状是 (2, 3, 4),数据如下(为了方便看,我用坐标值来命名数据,比如 012 代表第0层第1排第2个):
import torch
# 形状: (2, 3, 4) -> (层, 排, 个)
# 数据内容模拟坐标:
# 第0层: [[000, 001, 002, 003],
# [010, 011, 012, 013],
# [020, 021, 022, 023]]
#
# 第1层: [[100, 101, 102, 103],
# [110, 111, 112, 113],
# [120, 121, 122, 123]]
- 当 dim=0 时:跨层取货 (换层)
含义:index 里的数字代表**“去第几层拿”。
规则:index 的位置决定了我们在哪一排、哪一个,而 index 的值决定了去哪个层**。
-
input:(2, 3, 4) -
index: 假设我们只想取第0层和第1层的特定数据,形状设为(2, 1, 2)。# index 形状 (2, 1, 2) # 这里的数字代表“层号” index = torch.tensor([[[0, 1]], # 里面数字是几,就取第几层。 [[1, 0]]]) # 里面数字是几,就取第几层。
取值过程解析:
我们要填充 output 的 [[[?, ?]], [[?, ?]]]。
-
看
index的[0, 0, 0]位置:- 值是
0。 - 意思是:去 第0层 拿。
- 去哪拿?保持
index当前位置的其他坐标不变(第0排,第0个)。 - 结果:去
input[0, 0, 0]拿了000。
- 值是
-
看
index的[0, 0, 1]位置:- 值是
1。 - 意思是:去 第1层 拿。
- 去哪拿?保持
index当前位置的其他坐标不变(第0排,第1个)。 - 结果:去
input[1, 0, 1]拿了101。
- 值是
结论:dim=0 时,index 的值控制层的跳转。
- 当 dim=1 时:跨排取货 (换排)
含义:index 里的数字代表**“去第几排拿”。
规则:index 的值决定了去哪个排**,其他坐标(层、个)保持不变。
-
input:(2, 3, 4) -
index: 假设形状为(2, 2, 4),意思是每层取2排。# index 形状 (2, 2, 4) # 这里的数字代表“排号” index = torch.tensor([[[0, 0, 0, 0], # 第0层,取第0排的数据 [2, 2, 2, 2]], # 第0层,取第2排的数据 [[1, 1, 1, 1], # 第1层,取第1排的数据 [0, 0, 0, 0]]]) # 第1层,取第0排的数据
取值过程解析:
-
看
index的[0, 0, :]位置(第0层,第0行输出):- 值全是
0。 - 意思是:去 第0排 拿。
- 去哪拿?保持层是0,保持列位置不变。
- 结果:把
input[0, 0, :]的数据搬过来。即000, 001, 002, 003。
- 值全是
-
看
index的[0, 1, :]位置(第0层,第1行输出):- 值全是
2。 - 意思是:去 第2排 拿。
- 结果:把
input[0, 2, :]的数据搬过来。即020, 021, 022, 023。
- 值全是
结论:dim=1 时,index 的值控制**排(行)**的跳转。
- 当 dim=2 时:跨列取货 (换位置)
含义:index 里的数字代表**“去第几个位置拿”。
规则:index 的值决定了去哪个列**,其他坐标(层、排)保持不变。这是最像二维 dim=1 的情况。
-
input:(2, 3, 4) -
index: 假设形状为(2, 3, 2),意思是每排只取2个数据。# index 形状 (2, 3, 2) # 这里的数字代表“列号” index = torch.tensor([[[3, 2], # 第0层第0排,取第3个和第2个 [1, 0], # 第0层第1排,取第1个和第0个 [0, 1]], # 第0层第2排,取第0个和第1个 [[0, 1], # 第1层... [2, 3], [3, 3]]])
取值过程解析:
-
看
index的[0, 0, 0]位置:- 值是
3。 - 意思是:去 第3列 拿。
- 保持层0、排0不变。
- 结果:去
input[0, 0, 3]拿了003。
- 值是
-
看
index的[0, 0, 1]位置:- 值是
2。 - 意思是:去 第2列 拿。
- 结果:去
input[0, 0, 2]拿了002。
- 值是
结论:dim=2 时,index 的值控制**列(具体元素)**的跳转。
📌 终极总结表
对于形状为 (D0, D1, D2) 的输入:
| dim 设置 | 操作对象 | index 里的数字代表什么? | 也就是… |
|---|---|---|---|
| dim=0 | 第0维 (层) | 层号 | 去第几层找? |
| dim=1 | 第1维 (排) | 排号 | 去第几排找? |
| dim=2 | 第2维 (列) | 列号 | 去第几个找? |
💡 避坑指南(重要!)
-
形状匹配规则:
index的形状不需要和input一样,但输出的形状会和index完全一样。- 铁律:除了操作的那个维度
dim之外,index和input的其他所有维度的大小必须一致(或者支持广播)。 - 例子:如果
input是(2, 3, 4),在dim=1操作时,index的第0维(层)必须是2,第2维(列)必须是4。
- 铁律:除了操作的那个维度
-
索引不能越界:
index里的数字,绝对不能超过input在dim维度上的长度。- 例子:如果
input是(2, 3, 4),在dim=1(排)时,index里的数字只能是 0, 1, 2。
- 例子:如果
🧮 通用公式
理解这个公式,你就能掌握任意维度的 gather:
o u t [ i ] [ j ] [ k ] = i n p u t [ i ] [ index [ i ] [ j ] [ k ] ] [ k ] ( 当 dim=1 时 ) out[i][j][k] = input[i][ \text{index}[i][j][k] ][k] \quad (\text{当 dim=1 时}) out[i][j][k]=input[i][index[i][j][k]][k](当 dim=1 时)
通俗解释:
输出张量里的每一个位置,都去 index 里看那个位置写的数字是多少,然后拿着这个数字去 input 里对应的 dim 维度上取值。
再来看一遍:
我们以三维张量为例,用最直观的方式解释
dim=0、dim=1、dim=2时torch.gather的收集规则。假设输入
input形状为(D0, D1, D2),即三个维度的大小分别为D0、D1、D2。
索引张量index必须和input有相同的维度数,且除了dim维度外,其他维度的大小必须与input一致。
输出output的形状与index完全相同。
通用公式
对于输出张量中任意一个位置
(i, j, k)(下标按维度顺序),它的值由以下规则确定:
当
dim=0时:
output[i][j][k] = input[ index[i][j][k] ][j][k]
即第一维的索引用index中的值代替,其他维索引保持不变。当
dim=1时:
output[i][j][k] = input[i][ index[i][j][k] ][k]
即第二维的索引用index中的值代替,其他维索引保持不变。当
dim=2时:
output[i][j][k] = input[i][j][ index[i][j][k] ]
即第三维的索引用index中的值代替,其他维索引保持不变。
具体示例
我们用形状
(2, 3, 4)的输入,手动演示三种情况。import torch input = torch.tensor([ [ # 第0组 (D0=0) [1, 2, 3, 4], # 第0行 (D1=0) [5, 6, 7, 8], # 第1行 [9,10,11,12] # 第2行 ], [ # 第1组 (D0=1) [13,14,15,16], # 第0行 [17,18,19,20], # 第1行 [21,22,23,24] # 第2行 ] ]) # 形状 (2, 3, 4)
dim=0(在第一个维度上收集)我们构造一个
index,形状为(2, 3, 4)(为了演示,让索引值在 0~1 之间)。index_dim0 = torch.tensor([ [[0,1,0,1], [1,0,1,0], [0,1,0,1]], [[1,0,1,0], [0,1,0,1], [1,0,1,0]] ])根据公式
output[i][j][k] = input[ index[i][j][k] ][j][k],例如:
output[0][0][0] = input[ index[0][0][0] ][0][0] = input[0][0][0] = 1output[0][0][1] = input[1][0][1] = 14output[1][0][0] = input[1][0][0] = 13output_dim0 = torch.gather(input, dim=0, index=index_dim0) print(output_dim0) # 结果(手动验证部分): # [[[ 1,14, 3,16], # [17, 6,19, 8], # [ 9,22,11,24]], # [[13, 2,15, 4], # [ 5,18, 7,20], # [21,10,23,12]]]
dim=1(在第二个维度上收集)构造
index,形状(2, 2, 4)(D1 维度大小可以随意,但 D0 和 D2 必须与 input 一致)。这里我们让每个组取 2 行(因为 index 的 D1=2)。index_dim1 = torch.tensor([ [[0,1,2,0], [2,0,1,1]], [[1,0,2,2], [0,2,1,0]] ]) # 形状 (2, 2, 4)公式
output[i][j][k] = input[i][ index[i][j][k] ][k],例如:
output[0][0][0] = input[0][ index[0][0][0] ][0] = input[0][0][0] = 1output[0][0][1] = input[0][1][1] = 6output[0][1][0] = input[0][2][0] = 9output_dim1 = torch.gather(input, dim=1, index=index_dim1) print(output_dim1) # 结果(部分): # [[[ 1, 6,11, 4], # [ 9, 2, 7, 8]], # [[14,13,23,16], # [13,22,19,16]]]
dim=2(在第三个维度上收集)构造
index,形状(2, 3, 3)(D2 维度大小变为 3,其他维度不变)。index_dim2 = torch.tensor([ [[0,1,2], [3,0,1], [2,3,0]], [[1,2,3], [0,2,1], [3,1,0]] ]) # 形状 (2, 3, 3)公式
output[i][j][k] = input[i][j][ index[i][j][k] ],例如:
output[0][0][0] = input[0][0][0] = 1output[0][0][1] = input[0][0][1] = 2output[0][1][0] = input[0][1][3] = 8output_dim2 = torch.gather(input, dim=2, index=index_dim2) print(output_dim2) # 结果: # [[[ 1, 2, 3], # [ 8, 5, 6], # [11,12, 9]], # [[14,15,16], # [17,19,18], # [24,22,21]]]
总结
dim替换的维度 公式(对于输出位置 (i,j,k))0 第1维 input[ index[i][j][k] ][j][k]1 第2维 input[i][ index[i][j][k] ][k]2 第3维 input[i][j][ index[i][j][k] ]核心:
gather让你可以在某一维上自由选择索引,其他维保持不变。index的形状决定了输出在该维度上的大小,index里的值告诉你去取input的哪一位置(在该维度上)。
9、torch.gather - API
📝 完整的函数签名
torch.gather(
input, # 【必填】源张量(Tensor),也就是你的“数据库”,我们要从这里取数据
dim, # 【必填】维度轴(int),指定沿着哪个维度去“抓”数据
index, # 【必填】索引张量(Tensor),这是“寻宝图”,指定要取的数据在 dim 维度上的下标
*, # (* 后面是关键字参数,调用时必须写成 key=value 的形式)
sparse_grad=False, # 【可选】默认是 False。用于反向传播时是否返回稀疏梯度(一般不用管)
out=None # 【可选】默认是 None。指定输出结果存放的张量(一般不用管)
) -> Tensor
📊 参数详解
| 参数 | 必填/可选 | 说明 |
|---|---|---|
input | 【必填】 | 源数据。形状可以是任意的(比如 (2, 3, 4))。 |
dim | 【必填】 | 操作轴。指定沿着哪个维度去“抓”数据。• dim=0:沿着行抓(跨层/跨行)。• dim=1:沿着列抓(跨排/跨列)。• dim=-1:沿着最后一个维度抓(最常用)。 |
index | 【必填】 | 索引模具。这是最关键的部分。• 它里面的数字:代表在 dim 维度上的下标。• 它的形状:决定了输出结果的形状。 |
sparse_grad | 【可选】 | 默认 False。用于反向传播时是否返回稀疏梯度(一般不用管)。 |
out | 【可选】 | 默认 None。指定输出结果存放的张量。 |
🧠 核心逻辑
一句话口诀:
dim 定轴,index 定值。
index 里的数字,就是用来替换 dim 那个轴坐标的。
通用数学公式:
假设 index 的形状和 input 完全一致(或者可以通过广播对齐),那么输出张量中的每一个元素遵循以下规则:
o u t [ i ] [ j ] [ k ] . . . = i n p u t [ i ] [ j ] [ k ] . . . 但是在 d i m 维度上的索引被替换为 i n d e x [ i ] [ j ] [ k ] . . . out[i][j][k]... = input[i][j][k]... \text{ 但是在 } dim \text{ 维度上的索引被替换为 } index[i][j][k]... out[i][j][k]...=input[i][j][k]... 但是在 dim 维度上的索引被替换为 index[i][j][k]...
通俗解释:
输出张量里的每一个位置,都去 index 里看那个位置写的数字是多少,然后拿着这个数字去 input 里对应的 dim 维度上取值。
🛠️ 常用操作与场景演示
场景一:2D 数据,按行取数 (dim=1)
这是最常见的场景,比如**“根据预测的类别ID,取出对应的概率值”**。
假设 input 是模型输出的概率,index 是我们要取的类别。
import torch
# 1. 源数据 (2行3列)
# 含义:样本1的三个类别概率 [0.1, 0.5, 0.9], 样本2的概率 [0.3, 0.2, 0.8]
input = torch.tensor([[0.1, 0.5, 0.9],
[0.3, 0.2, 0.8]])
# 2. 索引 (2行2列)
# 含义:样本1我想取第2列(下标2)和第0列(下标0)的值;样本2我想取第1列(下标1)的值...
# 注意:index 的值必须小于 input 对应维度的长度(这里是3)
index = torch.tensor([[2, 0],
[1, 2]])
# 3. 执行 Gather (dim=1 表示沿着列的方向去取)
# 逻辑:
# output[0,0] -> input[0, index[0,0]] -> input[0, 2] -> 0.9
# output[0,1] -> input[0, index[0,1]] -> input[0, 0] -> 0.1
# output[1,0] -> input[1, index[1,0]] -> input[1, 1] -> 0.2
# output[1,1] -> input[1, index[1,1]] -> input[1, 2] -> 0.8
output = torch.gather(input, dim=1, index=index)
print(output)
# 结果:
# tensor([[0.9000, 0.1000],
# [0.2000, 0.8000]])
场景二:3D 数据,跨层取货 (dim=0)
把三维张量想象成一个**“多层货架”:(层, 排, 个)。
dim=0 意味着 index 里的数字代表“去第几层拿”**。
# 1. 源数据 (2层, 3排, 4个)
# 为了方便看,我用坐标值来命名数据,比如 012 代表第0层第1排第2个
input = torch.tensor([
[[0, 1, 2, 3], # 第0层
[4, 5, 6, 7]],
[[10, 11, 12, 13], # 第1层
[14, 15, 16, 17]]
])
# 2. 索引 (1层, 1排, 2个)
# 这里的数字代表“层号”
index = torch.tensor([
[[0, 1]]
])
# 3. 执行 Gather (dim=0 表示沿着层的方向去取)
# 逻辑:
# output[0, 0, 0] -> input[index[0,0,0], 0, 0] -> input[0, 0, 0] -> 0
# output[0, 0, 1] -> input[index[0,0,1], 0, 1] -> input[1, 0, 1] -> 11
output = torch.gather(input, dim=0, index=index)
print(output)
# 结果:
# tensor([[[ 0, 11]]])
⚠️ 避坑指南(重要!)
-
形状匹配规则(最重要!):
index的形状不需要和input完全一样,输出的形状会和index完全一样。- 铁律:除了操作的那个维度
dim之外,index和input的其他所有维度的大小必须一致(或者支持广播)。 - 例子:如果
input是(2, 5),在dim=1操作时,index的行数(第0维)必须是 2。
-
索引不能越界:
index里的数字,绝对不能超过input在dim维度上的长度。- 比如
input是(2, 5),在dim=1时,index里的数字只能是 0, 1, 2, 3, 4。
- 比如
-
与
index_select的区别:torch.index_select(input, dim, index):index是一个一维向量,取出来的数据会拼在一起,输出形状会改变。torch.gather(input, dim, index):index是多维的,它像是一个“模具”,输出形状严格跟随index。
10、Top-k 采样 + Temperature 调节 代码详解
import torch
import torch.nn.functional as F
def top_k_sampling_with_temperature(logits, k=10, temperature=1.0):
"""
logits: [batch_size, vocab_size] 或 [vocab_size]
k: 保留的候选 token 数量
temperature: 温度参数 (>0)
返回: 下一个 token 的索引,维度与输入 batch 维度一致(若输入为 1D,返回 Python int)
"""
# 统一处理 batch 维度
was_1d = (logits.dim() == 1)
if was_1d:
logits = logits.unsqueeze(0) # [1, vocab_size]
# 1. 温度缩放
logits = logits / temperature
# 2. 在 logits 上直接取 top-k 【高效实现:先取 top-k 的 logits,再 softmax(避免计算全词表)】
# topk: 默认是降序排列(从大到小排列)
top_k_logits, top_k_indices = torch.topk(logits, k, dim=-1)
# 3. 对 top-k logits 做 softmax(得到归一化后的概率)
top_k_probs = F.softmax(top_k_logits, dim=-1)
# 4. 从 top-k 中采样(torch.multinomial 后面有详情)
sampled_idx_in_topk = torch.multinomial(top_k_probs, num_samples=1) # [batch, 1]
# 5. 映射回原始词表索引(torch.gather 后面有详情)
# .squeeze(-1) 这个是降维,不是升维
next_token = torch.gather(top_k_indices, -1, sampled_idx_in_topk).squeeze(-1) # [batch]
# 恢复原始维度
if was_1d:
next_token = next_token.item()
return next_token
先明确任务:Top‑k 采样 + 温度调节
这段代码做的事情是:
- 输入模型输出的
logits(未归一化的分数,形状可能是[batch, vocab]或[vocab])。 - 用温度调节 logits。
- 只保留概率最高的
k个 token(Top‑k)。 - 在这
k个 token 中按概率随机采样一个。 - 返回这个 token 在原始词表中的索引。
关键难点:第 4 步采样的结果是 在 top‑k 候选中的相对位置(比如 0 表示候选列表中的第一个),需要把它映射回原始词表的绝对索引。torch.gather 就是用来做这个映射的。
一、准备一个具体例子
假设:
batch_size = 2(两个句子同时生成)vocab_size = 5(词表只有 5 个词,索引 0~4)k = 3(只保留概率最高的 3 个)
输入 logits 的形状 (2, 5),内容假设如下:
logits = torch.tensor([
[0.5, 2.1, 1.2, 0.8, 3.0], # 第 1 个样本
[1.0, 0.3, 2.5, 0.9, 1.8] # 第 2 个样本
])
为了简化演示,我们暂不使用温度调节(即 temperature=1.0),直接做 softmax 看概率:
probs = F.softmax(logits, dim=-1)
# 结果(四舍五入):
# [[0.06, 0.35, 0.12, 0.08, 0.39], # 第1个样本
# [0.10, 0.05, 0.45, 0.09, 0.31]] # 第2个样本
显然,每个样本概率最高的 3 个 token 是:
- 样本0:索引 4(0.39)、索引 1(0.35)、索引 2(0.12)
- 样本1:索引 2(0.45)、索引 4(0.31)、索引 0(0.10)
二、逐步执行代码
第 1 步:统一 batch 维度
was_1d = (logits.dim() == 1)
if was_1d:
logits = logits.unsqueeze(0)
如果输入是一维(单个样本),就变成 (1, vocab)。这里输入是二维,不变。
第 2 步:温度缩放
logits = logits / temperature # 本例 temperature=1.0,所以 logits 不变
实际使用中 temperature 可以调节(如 0.7 或 1.2),这里为 1.0 仅作演示。
第 3 步:取 top‑k 的 logits 和 indices
top_k_logits, top_k_indices = torch.topk(logits, k, dim=-1)
-
dim=-1表示在最后一维(词表维度)上取最大的 k 个。 -
top_k_logits形状(2, 3),内容是每行降序排列的 logits 值:[[3.0, 2.1, 1.2], [2.5, 1.8, 1.0]] -
top_k_indices形状(2, 3),内容是这些值对应的原始词表索引:[[4, 1, 2], [2, 4, 0]]解释:第 1 个样本中,最大的是索引 4(值 3.0),第二是索引 1(2.1),第三是索引 2(1.2)。
第 2 个样本中,最大的是索引 2(2.5),第二是索引 4(1.8),第三是索引 0(1.0)。
第 4 步:对 top‑k logits 做 softmax,得到概率分布
top_k_probs = F.softmax(top_k_logits, dim=-1)
top_k_probs形状(2, 3),每一行是那 3 个候选 token 的概率(归一化后)。
以样本0为例:logits[3.0, 2.1, 1.2],softmax 后概率约为[0.64, 0.26, 0.10](精确值:0.636, 0.259, 0.105)。
样本1类似,但具体数值不影响理解。
第 5 步:从 top‑k 中采样
sampled_idx_in_topk = torch.multinomial(top_k_probs, num_samples=1)
-
torch.multinomial根据每一行的概率分布,随机抽取一个索引(相对位置,即 0, 1, 2)。 -
sampled_idx_in_topk形状(2, 1),值可能是:[[1], # 第1个样本抽中了相对位置 1(即 top‑k 列表中的第2个候选) [2]] # 第2个样本抽中了相对位置 2(即 top‑k 列表中的第3个候选)注意:这里的 1 和 2 是相对位置,不是原始词表索引。
第 6 步:映射回原始词表索引(重点:torch.gather)
问题:
我们有一个 top_k_indices(装的是原始词表索引) 张量(形状 (2, 3)):
[[4, 1, 2],
[2, 4, 0]]
还有一个 sampled_idx_in_topk 张量(形状 (2, 1)):
[[1],
[2]]
对于第 1 个样本,我们需要取 top_k_indices[0][1] = 1(原始词表索引)。
对于第 2 个样本,我们需要取 top_k_indices[1][2] = 0(原始词表索引)。
torch.gather 就是用来做这个的。
next_token = torch.gather(top_k_indices, -1, sampled_idx_in_topk).squeeze(-1)
详细解释 torch.gather 在这里的工作方式
input = top_k_indices,形状(2, 3)。dim = -1(等价于dim=1),表示在最后一个维度(即列方向)上进行收集。index = sampled_idx_in_topk,形状(2, 1)。
torch.gather 的通用规则:
输出在位置 (i, j, k, ...) 的值,等于 input 在相同位置但将 dim 维的索引替换为 index[i, j, k, ...] 后的值。
对于本例(二维,dim=1),可以简化为:
output[i][j] = input[i][ index[i][j] ]。
因为 index 的形状是 (2, 1),所以输出形状也是 (2, 1)。
具体计算:
- 输出
[0, 0]=input[0][ index[0][0] ]=input[0][1]= 1 - 输出
[1, 0]=input[1][ index[1][0] ]=input[1][2]= 0
所以 torch.gather 的结果是:
[[1],
[0]]
形状 (2, 1)。
然后 .squeeze(-1) 去掉最后一维,得到形状 (2,) 的张量:
[1, 0]
这就是最终采样到的原始词表索引。
第 7 步:恢复原始维度
if was_1d:
next_token = next_token.item()
如果原始输入是一维的,那么经过上述处理后 next_token 是一个形状 (1,) 的张量,调用 .item() 将其转换为 Python 整数标量。
三、为什么不能直接用 top_k_indices[range(batch), sampled_idx_in_topk.squeeze(-1)]?
也可以,但 gather 更通用,可以处理任意维度的张量,并且是 PyTorch 提供的标准索引收集函数,语义清晰,支持自动微分。对于本例两者等价,但理解 gather 对后续处理更高维数据(如 (batch, seq_len, vocab))很有帮助。
四、torch.gather 的直观理解(以本例为例)
torch.gather 就像是拿着一个“索引地图”去原张量里“捡”元素。
- 原张量
top_k_indices是一个二维表格。 - 索引张量
sampled_idx_in_topk告诉你在每一行,要取第几列。 - 结果就是每行取出的元素拼在一起。
你可以这样记忆(针对二维且 dim=1 的情况):
输出[i][j] = 输入[i][ 索引[i][j] ] # 当 dim=1 时(列方向)
输出[i][j] = 输入[ 索引[i][j] ][j] # 当 dim=0 时(行方向)
在我们的代码中 dim=-1(即最后一维,也就是列方向),所以是第一种。
五、完整代码加注释(重点标注 gather)
import torch
import torch.nn.functional as F
def top_k_sampling_with_temperature(logits, k=10, temperature=1.0):
was_1d = (logits.dim() == 1)
if was_1d:
logits = logits.unsqueeze(0) # [1, vocab_size]
logits = logits / temperature
top_k_logits, top_k_indices = torch.topk(logits, k, dim=-1)
top_k_probs = F.softmax(top_k_logits, dim=-1)
sampled_idx_in_topk = torch.multinomial(top_k_probs, num_samples=1) # [batch, 1]
# 关键步骤:用 sampled_idx_in_topk 作为列索引,从 top_k_indices 中取出对应的原始 token id
# top_k_indices: [batch, k] sampled_idx_in_topk: [batch, 1]
# gather 沿着 dim=-1(列)操作,对于每行 i,取第 sampled_idx_in_topk[i,0] 列的值
next_token = torch.gather(top_k_indices, -1, sampled_idx_in_topk).squeeze(-1) # [batch]
if was_1d:
next_token = next_token.item()
return next_token
六、总结
torch.gather在这里的作用是:把采样得到的“候选列表中的相对位置”翻译成“原始词表中的绝对索引”。- 它避免了手动写循环,直接利用张量操作完成映射。
- 理解
gather的关键是记住:当dim=1时,输出在(i, j)的值 = 输入在(i, index[i][j])的值。
11、最终代码
import torch
import torch.nn.functional as F
def top_k_sampling_with_temperature(logits, k=10, temperature=1.0):
"""
对 logits 进行 Top-k 采样,并结合温度参数调节分布。
参数:
logits: [batch_size, vocab_size] 或 [vocab_size],模型输出的原始分数
k: 保留的候选 token 数量
temperature: 温度参数 (>0)。
<1.0 使分布更尖锐(更确定),
>1.0 使分布更平滑(更随机)。
返回:
下一个 token 的索引 (int 或 Tensor)
"""
# --- 1. 预处理:统一维度 ---
# 记录输入是不是 1D 的,方便最后恢复
was_1d = (logits.dim() == 1)
if was_1d:
logits = logits.unsqueeze(0) # 变成 [1, vocab_size],方便统一处理
# --- 2. 温度调节 ---
# 温度越低,高分和低分的差距拉得越大(高分更高)
logits = logits / temperature
# --- 3. Top-k 筛选 ---
# 取出分数最高的 k 个 logits 和它们对应的索引
# top_k_logits: [batch, k]
# top_k_indices: [batch, k] (存的是原始词表里的 ID)
top_k_logits, top_k_indices = torch.topk(logits, k, dim=-1)
# --- 4. 计算概率 ---
# 只在这 k 个候选词上计算 Softmax,把它们变成概率分布
# 这一步非常关键,因为我们要在“小圈子”里采样,而不是全词表
probs = F.softmax(top_k_logits, dim=-1)
# --- 5. 采样 (核心补充部分) ---
# torch.multinomial 是真正的“掷骰子”环节
# num_samples=1 表示每个句子只抽 1 个词
# 返回的是:在 top-k 这个小圈子里的下标 (0 到 k-1)
sampled_idx_in_topk = torch.multinomial(probs, num_samples=1) # [batch, 1]
# --- 6. 映射回原始词表 (Gather) ---
# 拿着“小圈子里的下标”,去“原始词表 ID 列表”里查出真正的 ID
next_token = torch.gather(top_k_indices, -1, sampled_idx_in_topk) # [batch, 1]
# --- 7. 恢复维度 ---
next_token = next_token.squeeze(-1) # 变回 [batch]
if was_1d:
next_token = next_token.item() # 如果是 1D 输入,返回 Python 整数
return next_token
# --- 测试代码 ---
if __name__ == "__main__":
# 模拟一个 logits,假设词表大小为 100
# 我们故意让第 5 号和第 10 号位置的分数很高
dummy_logits = torch.randn(100)
# 运行函数
# k=5 表示只在前 5 名里选
# temperature=0.8 稍微增加一点确定性
result = top_k_sampling_with_temperature(dummy_logits, k=5, temperature=0.8)
print(f"采样到的 Token ID: {result}") # 比如输出 24
更多推荐
所有评论(0)