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

如何复现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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 07:50:25