Gumbel Softmax输入疑问:归一化对数概率vs非归一化概率
Gumbel Softmax输入要求的明确结论
原始论文定义与PyTorch实现的差异
原始Gumbel-Max技巧的输入要求
原始论文中Gumbel-Max的采样公式为:z = one_hot(argmax_i[g_i + log π_i])这里的输入必须是归一化类别概率的对数形式(log π_i),其中π_i是满足Σπ_i=1的合法类别概率。这是从类别分布采样的理论推导出发的设定——Gumbel噪声直接作用于对数概率,通过argmax得到离散采样结果。
PyTorch
torch.gumbel_softmax的输入要求
PyTorch官方文档明确规定,函数默认接受非归一化的类别概率(即模型输出的logits)。这是因为实现中做了封装:- 当
logits=True(默认参数)时,输入被当作未归一化的logits,内部会自动处理为对数概率形式后添加Gumbel噪声; - 当
logits=False时,需要输入已经归一化的类别概率,函数会先对其取对数再执行后续操作。
- 当
原理层面的合理性解释
从实践和数值稳定性角度,直接使用无约束的logits作为输入更合理:
- logits是模型的原始输出,无需提前做Softmax归一化,避免了Softmax在极端数值下的溢出或下溢问题;
- log π_i和logits之间仅差一个全局归一化常数:
log π_i = logits - log(sum(exp(logits))),而argmax操作对这个常数不敏感(常数不改变argmax的结果),因此用logits替代log π_i不会改变采样的分布特性,这也是PyTorch实现可以直接用logits的核心原因。
内容的提问来源于stack exchange,提问作者Sammy Cui
相关产品推荐
相关产品推荐

