在搭建神经网络时,我们经常会遇到“从一堆候选里挑出若干个”的需求,比如从特征图中选取关键区域、从词表中采样离散token、或者从一组专家模型里路由样本。最直觉的写法是tensor[index]或者torch.gather,但这种硬索引操作会把梯度挡在门外,被选中的元素即使参与了后续计算,也无法把损失信号回传到选择依据本身。换句话说,模型可以学会用好被选中的内容,却学不会“该选谁”。这篇文章就来系统梳理如何在PyTorch中实现可微分的张量选择,从原理到代码逐层展开。

一、为什么硬索引会阻断梯度
先看一个最小示例。假设有一个分数张量scores,我们希望取最大值对应的那个元素参与后续计算:
import torch scores = torch.randn(5, requires_grad=True) values = torch.randn(5) # 硬索引:取出最大值位置 idx = scores.argmax() picked = values[idx] picked.backward() # 对scores求导会得到全零 print(scores.grad) # tensor([0., 0., 0., 0., 0.])
运行后会发现scores.grad全是零。原因在于,argmax输出的是一个整数索引,整数张量是不可导的,反向传播到这一步就断了。索引操作本身在数学上是一个阶梯函数,输入连续变化时输出离散跳变,导数几乎处处为零。
更本质地说,可微分编程要求整个计算流程构成一个连续可导的函数链。而“选择”这个动作天然是离散的:要么选中,要么不选中,没有中间状态。要让它可导,思路只有两条:要么用连续的权重去“软选择”,让每个候选都按权重参与计算;要么用重参数化技巧,把随机性从计算图里挪出去。这两条思路分别对应softmax软选择和Gumbel-Softmax采样。
二、把索引改写成加权求和:one-hot 等价形式
理解软选择的关键,是先看清硬索引和矩阵乘法的等价关系。取索引为k的元素,等价于用one-hot向量右乘:
import torch import torch.nn.functional as F values = torch.randn(5) k = 2 # 写法一:硬索引 a = values[k] # 写法二:one-hot 点乘 onehot = F.one_hot(torch.tensor(k), 5).float() b = (onehot * values).sum() print(a, b) # 两者数值相同
这两种写法结果完全一样,但第二种形式透露了改进方向:one-hot向量如果是从一个可导的网络输出的,比如softmax的结果,那么整个“选择”过程就变成可导的了。这就是软选择的核心思想——不真正做选择,而是让每个候选元素都按权重参与加权求和。
基于softmax的软选择实现非常简洁:
scores = torch.randn(5, requires_grad=True) values = torch.randn(5) # 软选择:权重由softmax产生 weights = F.softmax(scores, dim=0) picked = (weights * values).sum() picked.backward() print(scores.grad) # 非零,梯度可以回传
这种写法其实就是注意力机制的最简版本。Transformer里的注意力本质上就是对value向量做一次权重由query和key决定的软选择。它的优点是完全可导、实现简单;缺点是“选择”不彻底,所有候选都贡献了梯度,当候选数量非常大时,软选择会退化为一种平滑的平均,失去“挑出少数几个”的语义。另外,softmax输出永远不可能是严格的one-hot,推理阶段如果需要离散决策,还得额外做argmax。
三、温度系数:在软与硬之间调节
softmax带温度的形式是softmax(logits / tau),其中tau是温度系数。温度越低,分布越尖锐,接近one-hot;温度越高,分布越平坦,接近均匀分布。这个参数在可微选择中至关重要,因为它直接控制了软选择和硬选择之间的权衡。
logits = torch.tensor([1.0, 2.0, 3.0, 5.0])
for tau in [10.0, 1.0, 0.1, 0.01]:
probs = F.softmax(logits / tau, dim=0)
print(f"tau={tau}: {probs.detach().numpy().round(4)}")
运行结果可以直观看到,tau等于10时四项权重都比较接近,tau等于0.01时几乎全部概率集中到了最大值上。实践中的常见做法是训练初期用较高的温度让梯度流动顺畅,随着训练推进逐渐退火到低温,让选择行为逐渐硬化。不过温度太低时softmax的梯度会变得极小,数值上容易出现下溢,所以退火策略不能太激进,一般从1.0逐步降到0.1或0.5左右比较稳妥。
需要注意的是,无论温度多低,softmax在数学上永远不会输出精确的zero和one。如果下游模块严格要求离散输入(比如把选择结果作为另一个查表操作的索引),纯softmax方案就不适用了,这时需要引入下一节的采样技巧。
四、Gumbel-Softmax:可导的离散采样
有一类场景不仅要软选择,还要真正采样出离散的类别,且采样过程保持可导,比如从分类分布里采样token、离散潜变量模型的训练。直接用torch.distributions.Categorical采样会得到整数,梯度同样断掉。Gumbel-Softmax通过重参数化技巧解决了这个问题:把随机噪声加到logits上再做softmax,采样就被改写成了确定性函数加外部噪声,梯度可以从softmax结果流回logits。
logits = torch.randn(5, requires_grad=True) # Gumbel-Softmax采样,hard=True时前向输出one-hot,反向仍按软分布传梯度 y_soft = F.gumbel_softmax(logits, tau=1.0, hard=False) y_hard = F.gumbel_softmax(logits, tau=1.0, hard=True) print(y_hard) # 精确的one-hot向量 print(y_hard.grad_fn) # 仍在计算图中,可反向传播
这里最值得注意的是hard=True的模式。它在前向传播时把输出强制转成one-hot,保证下游拿到的是真正的离散选择;在反向传播时利用直通估计器的技巧,把one-hot处的梯度按软分布的梯度来计算。这样一来,前向是硬选择,反向是软梯度,兼顾了离散语义和端到端训练,是从分类分布中做可微采样最实用的方案。
温度参数在Gumbel-Softmax中的作用同样是控制分布尖锐程度。训练时通常采用温度退火:初始tau取1到5之间,随训练逐步降到0.1附近。如果发现采样结果频繁翻转不稳定,往往是温度降得太快或学习率过大导致的。
五、top-k 选择的软替代与工程实践
除了选一个,还经常要选多个。硬top-k可以用torch.topk,但它和argmax一样不可导。一种常见的软替代是先过softmax或sigmoid得到权重,再乘上一个掩码,让非top-k位置的权重乘以可导的放松因子。更简洁的做法是直接用温度很低的softmax近似,或者采用如下的加权top-k松弛:
scores = torch.randn(10, requires_grad=True) values = torch.randn(10, 10) # 假设10个候选,每个候选是一个向量 k = 3 weights = F.softmax(scores / 0.5, dim=0) # 软top-k:权重最大的k项保留,其余置零后再归一化 topk_weights, topk_idx = torch.topk(weights, k) mask = torch.zeros_like(weights).scatter_(0, topk_idx, 1.0) soft_topk = weights * mask soft_topk = soft_topk / soft_topk.sum(dim=0, keepdim=True) picked = soft_topk @ values # 加权聚合被选中的k个候选
这种实现里,掩码本身由top-k产生,是离散的,但权重是可导的,梯度仍然能通过weights流回scores。它的效果是:哪些位置被选中取决于当前分数排序,而被选中位置贡献多少则完全可导。对于路由、稀疏激活这类场景已经够用。如果需要选中的集合本身也可导地学习,则要借助更复杂的方案,比如基于optimal transport的soft sorting,或者sparsemax这类输出稀疏概率分布的替代激活函数。
工程上还有几个细节值得留意。第一,soft掩码中直接置零的位置梯度也为零,如果希望未被选中的位置也能收到一定的探索信号,可以在掩码上叠加一个很小的均匀底噪。第二,批量训练时不同样本的top-k索引不同,用scatter生成掩码时注意维度对齐,避免维度不匹配的广播错误。第三,推理阶段可以直接把软权重换成硬选择,消除softmax平滑带来的数值偏差,这也是训练和推理行为解耦的常见做法。
总结一下方案选型:如果只是要从候选中做加权聚合,softmax软选择最简单可靠;如果需要采样离散类别且参与训练,用F.gumbel_softmax并开启hard模式;如果要做稀疏的多选,用掩码加归一化的软top-k;只有当推理阶段才需要真正离散决策时,把软权重argmax一下即可。理解了“用连续权重替代离散索引”这个核心思想,遇到具体业务时就能灵活组合出合适的可微选择方案。
PyTorch可微分张量选择gumbel softmax修改时间:2026-09-04 07:40:49