1、基本介绍

一、名字解析

  1. 温度调节(Temperature Scaling)

“温度” 借用物理学概念:温度越高,系统越混乱;温度越低,系统越有序。在语言模型生成中,温度用于控制输出分布的随机性程度——高温使概率分布更平坦,增加随机性;低温使分布更尖锐,趋向于选择高概率词。

  1. 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 iItop-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=jItop-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 ]

五、它是干什么的?(作用与优点)

  1. 避免不合理输出
    Top‑k 通过丢弃尾部大量低概率 token,有效防止采样出语法错误、语义不通或无意义的词。

  2. 引入可控的随机性
    温度参数让你可以精细调节“保守度”与“创造性”的平衡:

    • 翻译、摘要等要求高准确性的任务:低温(( T \approx 0.6 \sim 0.8 ))+ 较小 k(如 10)→ 结果更确定。
    • 对话生成、故事创作等需要多样性的任务:高温(( T \approx 1.0 \sim 1.2 ))+ 较大 k(如 50)→ 输出更丰富。
  3. 缓解曝光偏差
    在训练中使用 Free Running 时,若采样策略与推理一致,能让模型更好地适应自己生成时的输入分布,减少训练‑推理差异。

  4. 实现简单,兼容性强
    只需在 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 维度。

七、注意事项

  1. k 值选择
    • k 过小 → 候选集太窄,可能错过合理但概率略低的词。
    • k 过大 → 引入过多低质量候选,随机性失控。
      通常通过实验确定,常见取值范围 10~100。
  2. 与 Free Running 的一致性
    若在训练的计划采样阶段也采用相同的采样策略(温度与 k 值可略低于推理,但保留随机性),有助于减小曝光偏差。
  3. 与其他方法的组合
    Top‑k 与 Top‑p 可结合:先取 top‑k 进一步过滤尾部,再在剩余候选中做 top‑p 动态筛选。两者结合可得到更稳定、高质量的生成。
  4. 适用性
    虽然束搜索在机器翻译等确定性任务中表现更好,但 Top‑k + 温度调节在无法使用束搜索时是最佳的替代方案,尤其适合从零实现模型、训练阶段引入随机性的场景。

参数设置建议(以汉译英为例)

任务类型TemperatureTop‑k
机器翻译(推理)0.7 ~ 1.010 ~ 20
摘要生成0.8 ~ 1.020 ~ 50
创意写作1.0 ~ 1.250 ~ 100
对话系统0.9 ~ 1.130 ~ 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 个面的骰子(每个面权重不同),一次试验会落到哪个面。


二、数学原理

  1. 多项分布的概率质量函数(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!p1n1pVnV
其中 ( \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,或者最多做少量有放回/无放回的抽样,并不需要那个复杂的组合公式。你完全可以忽略它,把它当作“按概率抽一个”的工具。


三、所以,你只需要知道三件事

  1. torch.multinomial 按给定权重随机选一个(或多个)索引
  2. 权重不需要归一化,函数内部会自动处理。
  3. 它常用于文本生成,在 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 采样了。

  1. 单次抽样(( N=1 ))

当 ( N = 1 ) 时,多项分布退化为 类别分布(Categorical Distribution),即一次试验中每个类别被抽中的概率就是 ( p_i )。torch.multinomial 主要处理这种单次抽样(num_samples=1),或者有放回/无放回的多次抽样(num_samples>1)。

你可以把复杂的数学定义抛开,直接这样理解:

probs=[0.2, 0.5, 0.3] 意味着:

  • 索引 020% 的机会被抽中。
  • 索引 150% 的机会被抽中。
  • 索引 230% 的机会被抽中。

💡 为什么这么简单,还要叫“多项分布”?

其实,“多项分布” 这个名字主要是为了强调**“多次抽样”时的统计规律,但在实际写代码(比如 torch.multinomial)时,我们往往只关心单次**结果。

我们可以分两个层面来看:

  1. 单次层面(你现在的理解)

这就是一个简单的“抽奖”动作。

  • 就像你手里有一个不均匀的骰子,扔一次,看它停在哪一面。
  • 代码表现:torch.multinomial(probs, 1) 返回一个数字(比如 1)。
  • 结论:概率就是 probs[i]
  1. 多次层面(数学定义的“多项分布”)

这是指“扔很多次”后的统计结果。

  • 如果你扔 1000 次,数学上会预测:索引 0 大约出现 200 次,索引 1 大约出现 500 次,索引 2 大约出现 300 次。
  • 代码表现:torch.multinomial(probs, 1000) 返回 1000 个数字。
  • 结论:虽然名字叫“多项分布”,但 torch.multinomial 这个函数本质上就是帮你执行一次次独立的“单次抽奖”(在有放回模式下)。

📌 总结

在写代码(特别是做 AI 推理)时,你就把它当成一个“加权随机数生成器”

  • 给一堆权重(概率)。
  • 让它吐出一个索引。
  • 权重越大,被吐出来的概率越大。

就这么简单!

  1. 采样过程

给定权重向量 ( \mathbf{w} = [w_1, w_2, \dots, w_V] )(不必归一化),torch.multinomial 内部:

  • 有放回采样:每次独立地根据归一化后的概率 ( p_i = w_i / \sum w_j ) 抽取一个索引。
  • 无放回采样:依次抽取,每次抽取后将被抽中的索引从候选集中移除,剩余概率重新归一化,再继续下一次抽取。

实际实现通常使用 逆变换采样:(具体实现按照这个理解就完全够了)

  1. 计算累积概率(或累积权重)( c_i = \sum_{j=1}^i w_j )。
  2. 生成均匀随机数 ( u \in [0, \text{总和}) )。
  3. 找到最小的 ( 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) 内部步骤:

  1. 计算累积概率:[0.2, 0.7, 1.0]
  2. 生成均匀随机数 u ∈ [0, 1),可不服从正太分布哈,例如 0.65)。
  3. 找到第一个 i 使得累积概率 ≥ u:i=1(因为 0.7 ≥ 0.65)。
  4. 返回 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


五、代码示例

  1. 单次抽样(无放回,单次即一次)
import torch

probs = torch.tensor([0.2, 0.5, 0.3])
sample = torch.multinomial(probs, 1)
print(sample)    # 可能输出 tensor([1])
  1. 有放回采样 5 次
samples = torch.multinomial(probs, 5, replacement=True)
print(samples)  # 可能输出 tensor([1, 1, 2, 0, 1]),允许重复
  1. 无放回采样(需 num_samples ≤ 类别数
samples = torch.multinomial(probs, 2, replacement=False)
print(samples)  # 输出两个不同的索引,如 tensor([1, 2])
  1. 批量处理(输入形状 [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]])
  1. 正确使用 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,需特别注意。


七、注意事项

  1. 输入不需要严格归一化
    torch.multinomial 会自动将输入视为权重,通过除以总和进行归一化。但为清晰起见,通常传入 softmax 后的概率或未归一化的正数权重。

  2. 无放回时 num_samples 不能超过类别总数
    replacement=False,则采样数量必须 ≤ 最后一维大小,否则会报错。

  3. 数据类型
    输入应为浮点型(float32/float64),返回值为 torch.long

  4. 随机性控制
    可通过 torch.manual_seed(seed)generator 参数固定随机性,便于复现。

  5. torch.distributions.Categorical 的关系
    Categorical 是更高级的分布封装,提供了 sample() 等方法,但 torch.multinomial 更底层、更轻量。

  6. 输入值建议非负
    尽管函数内部会处理,但为了数值稳定性,建议输入非负权重。若包含负值,可能导致意外行为。


八、在文本生成中的应用

在 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 详解

  1. 函数签名
torch.multinomial(
    input,                # 【必填】输入张量(Tensor),包含每个类别被选中的概率(权重)
    num_samples,          # 【必填】整数(int),表示要抽取多少个样本
    replacement=False,    # 【可选】布尔值,默认是 False(不放回采样)。
                          #        设为 True 表示允许同一个类别被重复抽到(放回采样)
    *,                    # (* 后面是关键字参数,调用时必须写成 key=value 的形式)
    generator=None,       # 【可选】默认是 None。用于控制随机性的生成器,设了它可以让结果可复现
    out=None              # 【可选】默认是 None。指定输出结果的存储位置(一般不用管)
)
  1. 参数详解
  • 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 wi0)。
      • 不需要归一化:你不需要手动把它变成概率(和为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)

    • 含义: 用于控制随机性的生成器。如果你需要复现结果,可以通过它设置随机种子。
  1. 返回值
  • 类型: torch.LongTensor
  • 形状:
    • 如果 input 是 1 维,返回形状为 [num_samples]
    • 如果 input 是 2 维,返回形状为 [batch_size, num_samples]
  • 内容: 采样得到的索引 (0 到 N − 1 N-1 N1)。

  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]) -> 索引互不相同

  1. 进阶:配合 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
  1. 常见坑与注意事项
问题说明解决方案
负数 Logits直接传入包含负数的 Logits必须先经过 Softmax 转为概率,否则结果错误。
数据类型错误传入 int 类型的张量使用 .float() 转换类型。
无放回超限num_samples > 类别数量,且 replacement=False减小 num_samples 或改为 replacement=True
全零权重输入张量某一行全是 0检查数据源,确保至少有一个正权重。
Device 不一致input 在 GPU 上,但 generator 在 CPU 上确保所有张量和生成器都在同一个设备上。
  1. 总结

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 TensorNumPy 数组


一句话总结

... 是一个“智能占位符”,它的意思是:“这里省略了一堆冒号 :,请自动帮我把剩下的维度都填满。”

它的学名叫 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 torchimport numpy as np,并且数据必须是 Tensor 或 Array。

总结

看到 ...,你就把它当成**“等等等等”或者“剩下的全都要”**。

  • 它不是三个点,它是一个智能填充工具
  • 它专门用来拯救那些维度太多、写冒号写到吐的代码。

5、[-1] vs […, -1] 详解

这两个操作在一维列表(或一维张量)中效果是一模一样的,但在多维数据中,区别就非常大了。

简单来说:... 是一个“偷懒”的符号,意思是“这里省略了一堆维度”。

我们可以用**“俄罗斯套娃”或者“书架”**来打比方。

  1. [-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]
  • ...:代表前面所有的批次维度(如果有的话),我们都要照顾到。
  • -1:取最后一个元素。
  • 作用:因为上一步的结果是从大到小排列的,所以最后一个元素就是这 top_k 个分数里最小的那个。这个分数就是我们的“淘汰线”。
  • 接上例:从 [5.0, 4.0, 3.0] 中取出最后一个,也就是 3.0。任何低于 3.0 的分数都将被淘汰。
  1. 增加一个维度 [..., 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]]

  1. 当 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[[[?, ?]], [[?, ?]]]

  1. index[0, 0, 0] 位置

    • 值是 0
    • 意思是:去 第0层 拿。
    • 去哪拿?保持 index 当前位置的其他坐标不变(第0排,第0个)。
    • 结果:去 input[0, 0, 0] 拿了 000
  2. index[0, 0, 1] 位置

    • 值是 1
    • 意思是:去 第1层 拿。
    • 去哪拿?保持 index 当前位置的其他坐标不变(第0排,第1个)。
    • 结果:去 input[1, 0, 1] 拿了 101

结论dim=0 时,index 的值控制的跳转。


  1. 当 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排的数据
    

取值过程解析

  1. index[0, 0, :] 位置(第0层,第0行输出):

    • 值全是 0
    • 意思是:去 第0排 拿。
    • 去哪拿?保持层是0,保持列位置不变。
    • 结果:把 input[0, 0, :] 的数据搬过来。即 000, 001, 002, 003
  2. index[0, 1, :] 位置(第0层,第1行输出):

    • 值全是 2
    • 意思是:去 第2排 拿。
    • 结果:把 input[0, 2, :] 的数据搬过来。即 020, 021, 022, 023

结论dim=1 时,index 的值控制**排(行)**的跳转。


  1. 当 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]]])
    

取值过程解析

  1. index[0, 0, 0] 位置

    • 值是 3
    • 意思是:去 第3列 拿。
    • 保持层0、排0不变。
    • 结果:去 input[0, 0, 3] 拿了 003
  2. 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维 (列)列号去第几个找?

💡 避坑指南(重要!)

  1. 形状匹配规则
    index 的形状不需要input 一样,但输出的形状会和 index 完全一样。

    • 铁律:除了操作的那个维度 dim 之外,indexinput其他所有维度的大小必须一致(或者支持广播)。
    • 例子:如果 input(2, 3, 4),在 dim=1 操作时,index 的第0维(层)必须是2,第2维(列)必须是4。
  2. 索引不能越界
    index 里的数字,绝对不能超过 inputdim 维度上的长度。

    • 例子:如果 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=0dim=1dim=2torch.gather 的收集规则。

假设输入 input 形状为 (D0, D1, D2),即三个维度的大小分别为 D0D1D2
索引张量 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)
  1. 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] = 1
  • output[0][0][1] = input[1][0][1] = 14
  • output[1][0][0] = input[1][0][0] = 13
output_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]]]
  1. 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] = 1
  • output[0][0][1] = input[0][1][1] = 6
  • output[0][1][0] = input[0][2][0] = 9
output_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]]]
  1. 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] = 1
  • output[0][0][1] = input[0][0][1] = 2
  • output[0][1][0] = input[0][1][3] = 8
output_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]]])

⚠️ 避坑指南(重要!)

  1. 形状匹配规则(最重要!)

    • index 的形状不需要input 完全一样,输出的形状会和 index 完全一样。
    • 铁律:除了操作的那个维度 dim 之外,indexinput其他所有维度的大小必须一致(或者支持广播)。
    • 例子:如果 input(2, 5),在 dim=1 操作时,index 的行数(第0维)必须是 2。
  2. 索引不能越界
    index 里的数字,绝对不能超过 inputdim 维度上的长度。

    • 比如 input(2, 5),在 dim=1 时,index 里的数字只能是 0, 1, 2, 3, 4。
  3. 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 采样 + 温度调节

这段代码做的事情是:

  1. 输入模型输出的 logits(未归一化的分数,形状可能是 [batch, vocab][vocab])。
  2. 用温度调节 logits。
  3. 只保留概率最高的 k 个 token(Top‑k)。
  4. 在这 k 个 token 中按概率随机采样一个。
  5. 返回这个 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


Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐