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

PyTorch中Log sum likelihood的实现方法咨询

实现高斯似然和的对数(Log of Sum Gaussian Likelihood)

核心逻辑

假设你提到的第二项是对多个高斯分布的似然值求和后取自然对数,数学表达式为:

$L_2 = \log\left(\sum_{k=1}^N \mathcal{N}(y; \mu_k, \sigma_k^2)\right)$

直接硬算「先exp再sum再log」很容易出现数值溢出,必须用log-sum-exp技巧,PyTorch自带的torch.logsumexp函数能帮你解决数值稳定问题。

一步步实现

1. 计算单个高斯的对数似然

每个高斯由均值$\mu_k$和方差$\sigma_k^2$定义,目标值$y$的对数似然直接套用高斯分布公式即可:

def gaussian_log_likelihood(y, mu, var):
    # var是方差,必须保证大于0
    log_p = -0.5 * (torch.log(2 * torch.tensor(math.pi)) + torch.log(var) + (y - mu)**2 / var)
    return log_p

2. 计算多个高斯似然的对数和

批量算出所有高斯的对数似然后,用torch.logsumexp直接计算log(sum(exp(log_ps)))——这和直接计算log(sum(似然值))等价,但完全避免了数值溢出:

import torch
import math

def log_sum_gaussian_likelihood(y, mus, vars):
    """
    参数说明:
        y: 目标值,形状可为 [batch_size] 或 [batch_size, 1]
        mus: 多个高斯的均值,形状 [batch_size, num_gaussians]
        vars: 多个高斯的方差,形状 [batch_size, num_gaussians](每个元素必须>0)
    返回:
        每个样本的log(sum(高斯似然)),形状 [batch_size]
    """
    # 扩展y的维度,方便和mus、vars做广播计算
    y = y.unsqueeze(-1)
    # 批量计算所有高斯的对数似然
    log_ps = -0.5 * (torch.log(2 * torch.tensor(math.pi)) + torch.log(vars) + (y - mus)**2 / vars)
    # 在高斯数量的维度上执行logsumexp
    log_sum_likelihood = torch.logsumexp(log_ps, dim=-1)
    return log_sum_likelihood

3. 整合到完整损失函数

你已经明确第一项用nn.GaussianNLLLoss,把第二项整合进去即可——注意损失函数是要最小化的,如果原公式中第二项是被减去的项,记得转成负值作为损失项:

import torch.nn as nn

class CustomLoss(nn.Module):
    def __init__(self):
        super().__init__()
        self.gaussian_nll = nn.GaussianNLLLoss()
    
    def forward(self, y_true, mu1, var1, mus2, vars2, lambda_weight=1.0):
        # 第一项:单个高斯的负对数似然损失
        loss1 = self.gaussian_nll(y_true, mu1, var1)
        # 第二项:先计算log(sum(高斯似然)),若原损失是L = loss1 - L2,则取负转为损失
        log_sum_likelihood = log_sum_gaussian_likelihood(y_true, mus2, vars2)
        loss2 = -log_sum_likelihood
        # 总损失,对loss2取均值适配批量训练
        total_loss = loss1 + lambda_weight * loss2.mean()
        return total_loss

关键注意事项

  • 方差必须为正:模型输出方差时,不要用原始输出,建议用torch.exp()或nn.Softplus()激活,还可加极小值避免零方差:
    # 模型输出log_var,转成方差:
    var = torch.exp(log_var)
    # 用Softplus更稳定:
    var = nn.Softplus()(log_var) + 1e-6
    
  • 禁止手动计算log(sum(exp(...))):一定要用torch.logsumexp,否则当对数似然数值较大时,exp会溢出变成inf,导致计算失效。

内容的提问来源于stack exchange,提问作者Ahmad Halimi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 23:46:31