PyTorch Autograd在含采样场景下的工作机制及计算图疑问
关于Categorical采样与Autograd的问题解答
Autograd能正常工作吗?
可以,但直接对Categorical分布做离散采样的操作本身是不可导的——采样得到的是离散类别索引,从概率到索引的映射是阶跃式的,没有连续梯度,直接这么做的话,损失到网络参数的梯度会断裂,Autograd没法完成梯度传递。
要让梯度正常传递,常用两种方案:
- REINFORCE(策略梯度):把采样的离散变量当作“常量”,通过「损失 × 采样动作的对数概率」的方式,让梯度借助对数概率的分支传回网络参数(对数概率是网络输出的连续函数,对参数可导)。
- Gumbel-Softmax重参数化:用连续的Gumbel噪声+Softmax来近似离散采样,让整个过程变成连续可导的,Autograd能直接沿着计算图传递梯度。
采样时的计算图结构
1. 直接离散采样(梯度断裂情况)
计算链条:网络参数 → 输出类别概率 → Categorical分布 → 离散类别索引 → 计算损失
这里「分布→离散索引」这一步是梯度断点,因为离散采样没有梯度,损失的梯度没法传回网络参数。
2. REINFORCE方法的计算图
调整后的链条会拆成两个分支:
网络参数 → 输出类别概率 → 分支1:Categorical分布→离散采样→计算损失 → 分支2:计算采样索引对应的对数概率 最终:损失 × 对数概率 → 反向传播
通过把损失和对数概率相乘,相当于用对数概率作为“梯度传递的桥梁”,让损失的梯度能通过对数概率的分支传回网络参数——因为类别概率是网络输出的连续函数,对数概率对参数的梯度是可计算的。
3. Gumbel-Softmax方法的计算图
用连续近似替代离散采样,整个链条全可导:网络参数 → 输出类别概率 → 加入Gumbel噪声 → Softmax得到连续近似概率 → 计算损失
Gumbel噪声的采样用重参数化技巧(-ln(-ln(U)),U是均匀分布采样),让噪声采样过程也能传递梯度,所以Autograd可以直接沿着整个链条把损失的梯度传回网络参数。
代码示例
REINFORCE实现(PyTorch)
import torch import torch.nn as nn import torch.distributions as dist class SimpleNet(nn.Module): def __init__(self, input_dim, num_classes): super().__init__() self.fc = nn.Linear(input_dim, num_classes) def forward(self, x): logits = self.fc(x) return dist.Categorical(logits=logits) net = SimpleNet(10, 3) optimizer = torch.optim.SGD(net.parameters(), lr=0.01) # 训练循环 x = torch.randn(1, 10) for _ in range(10): optimizer.zero_grad() cat_dist = net(x) sample = cat_dist.sample() # 离散采样 # 示例损失:惩罚采样到的类别 loss = -sample.float() # REINFORCE核心:损失乘以采样动作的对数概率 loss = loss * cat_dist.log_prob(sample) loss.backward() optimizer.step()
Gumbel-Softmax实现片段
def gumbel_softmax(logits, tau=1.0): # 生成Gumbel噪声,加小epsilon避免log(0) gumbel_noise = -torch.log(-torch.log(torch.rand_like(logits) + 1e-20) + 1e-20) y = logits + gumbel_noise return torch.softmax(y / tau, dim=-1) logits = net(x) approx_sample = gumbel_softmax(logits) # 基于连续近似样本计算损失 loss = -approx_sample[:, 1] # 示例:惩罚第二个类别的概率 loss.backward()
内容的提问来源于stack exchange,提问作者peter
相关产品推荐
相关产品推荐

