如何复现PyTorch中Categorical分布的Logits计算逻辑?
PyTorch中Categorical分布的Logits计算问题
关于Logit的基础理解
什么是Logit?Logit函数又称对数几率函数,可将0到1的概率值映射到负无穷到正无穷的范围。
问题场景
我尝试理解PyTorch中如何从概率密度函数计算Logits,但无法复现Categorical分布的输出结果。
测试代码与输出
初始化Categorical分布并获取logits的代码:
import torch from torch.distributions import Categorical probs = torch.tensor([0.1, 0.2, 0.5, 0.2]) categorical_distribution = Categorical(probs=probs) print(categorical_distribution.logits)
输出结果:
tensor([-2.3026, -1.6094, -0.6931, -1.6094])
我的尝试与偏差
我使用二分类场景下的Logit公式($\text{Logit}(p) = \ln\left(\frac{p}{1-p}\right)$)计算,代码如下:
import numpy as np print(np.log(0.1/(1-0.1)), np.log(0.2/(1-0.2)), np.log(0.5/(1-0.5)), np.log(0.2/(1-0.2)))
得到的结果却和PyTorch输出不符:
-2.19722457734 -1.38629436112 0.0 -1.38629436112
问题原因与复现方法
你混淆了二分类和多分类场景下的Logits定义:
- 你使用的是**伯努利分布(二分类)**的Logit公式,对应sigmoid函数的逆运算;
- 而PyTorch的
Categorical是多分类分布,其logits是softmax函数的输入,与概率的关系为:$\text{probs} = \text{softmax(logits)}$。反过来推导时,logits可以表示为$\text{logits} = \ln(\text{probs}) + C$($C$为任意常数,因为softmax对logits的平移具有不变性)。PyTorch中直接取$\ln(\text{probs})$作为logits值。
复现PyTorch结果的代码:
import torch probs = torch.tensor([0.1, 0.2, 0.5, 0.2]) logits = torch.log(probs) print(logits)
输出与Categorical分布的logits完全一致:
tensor([-2.3026, -1.6094, -0.6931, -1.6094])
内容的提问来源于stack exchange,提问作者Alexandre Tavares
相关产品推荐
相关产品推荐

