如何以数值稳定方式计算对应概率均值的Logits(神经网络集成)
计算集成模型的平均Logit(数值稳定版)
要得到满足 torch.sigmoid(l_averaged) ≈ torch.mean(torch.sigmoid(l)) 的标量Logit l_averaged,且全程避免直接调用sigmoid(保证数值稳定),可以通过Logit逆变换+数值稳定的对数求和技巧实现,核心思路如下:
推导过程
已知sigmoid(x) = 1/(1+exp(-x)),我们需要找到l_averaged使得:
sigmoid(l_averaged) = (1/n) * Σsigmoid(l_i)
对两边取Logit逆变换(即logit(p) = ln(p/(1-p))),可得:
l_averaged = logit( (1/n)Σsigmoid(l_i) )
直接计算sigmoid(l_i)会在l_i极大/极小时出现数值下溢/溢出,因此我们用softplus和logsumexp改写计算逻辑:
sigmoid(l_i) = exp(l_i - softplus(l_i)),其中softplus(x) = ln(1+exp(x))(PyTorch内置数值稳定实现)1 - sigmoid(l_i) = exp(-softplus(l_i))
代入Logit公式后,最终可转化为:
l_averaged = logsumexp(l_i - softplus(l_i)) - logsumexp(-softplus(l_i))
实现代码
import torch import torch.nn.functional as F def compute_averaged_logit(l: torch.Tensor) -> torch.Tensor: """ 计算数值稳定的平均Logit,满足sigmoid(l_averaged) ≈ mean(sigmoid(l)) 参数: l: 形状为[n]的Logit张量 返回: 标量平均Logit """ # 计算每个Logit的softplus(数值稳定) softplus_l = F.softplus(l) # 计算分子的对数和 numerator_log = torch.logsumexp(l - softplus_l, dim=0) # 计算分母的对数和 denominator_log = torch.logsumexp(-softplus_l, dim=0) # 得到平均Logit return numerator_log - denominator_log
验证示例
# 测试用例1:对称Logit l = torch.tensor([0.0, 0.0]) l_avg = compute_averaged_logit(l) print(f"平均Logit: {l_avg.item()}") # 输出0.0 print(f"sigmoid(平均Logit): {torch.sigmoid(l_avg).item()}") # 输出0.5,与mean(sigmoid(l))一致 # 测试用例2:正负Logit l = torch.tensor([1.0, -1.0]) l_avg = compute_averaged_logit(l) print(f"平均Logit: {l_avg.item():.4f}") # 输出≈-0.0020 print(f"sigmoid(平均Logit): {torch.sigmoid(l_avg).item():.4f}") # 输出≈0.4995,与mean(sigmoid(l))一致 # 测试用例3:极端值Logit l = torch.tensor([100.0, -100.0]) l_avg = compute_averaged_logit(l) print(f"平均Logit: {l_avg.item()}") # 输出0.0 print(f"sigmoid(平均Logit): {torch.sigmoid(l_avg).item()}") # 输出0.5,避免了直接计算sigmoid的数值问题
优势
- 数值稳定:通过
softplus和logsumexp避免了大/小Logit下的溢出/下溢问题 - 无sigmoid调用:全程不需要计算
sigmoid(l),直接基于原始Logit计算 - 适配集成场景:完美适用于神经网络集成的预测结果平均需求
内容的提问来源于stack exchange,提问作者CrabMan
相关产品推荐
相关产品推荐

