You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.28 18:59:56