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

Python实现广义狄利克雷分布KL散度的代码调试求助

广义狄利克雷分布KL散度实现问题

我正在研究一篇论文,需要实现其中公式11里两个广义狄利克雷分布的KL散度,相关公式截图如下:
公式截图1
公式截图2

参数说明

  • 解码器网络输出的参数(Alpha₁、Beta₁)示例张量:
decoder_alpha = torch.tensor([[0.2, 0.3, 0.3, 0.4],[0.2, 0.3, 0.3, 0.4],[0.2, 0.3, 0.3, 0.4], [0.2, 0.3, 0.3, 0.4], [0.2, 0.3, 0.3, 0.4], [0.2, 0.3, 0.3, 0.4]])

decoder_beta = torch.tensor([[0.3, 0.6, 0.4, 0.8],[0.2, 0.3, 0.3, 0.4],[0.2, 0.3, 0.3, 0.4], [0.2, 0.3, 0.3, 0.4], [0.2, 0.3, 0.3, 0.4], [0.2, 0.3, 0.3, 0.4]])
  • 预设先验参数(Alpha₂、Beta₂)示例张量:
prior_alpha = torch.tensor([[0.1, 0.1, 0.4, 0.1],[0.8, 0.7, 0.1, 0.4],[0.2, 0.8, 0.9, 0.1], [0.1, 0.5, 0.2, 0.4], [0.1, 0.2, 0.1, 0.4], [0.2, 0.1, 0.3, 0.3]])

prior_beta = torch.tensor([[0.7, 0.6, 0.1, 0.2],[0.5, 0.8, 0.1, 0.2],[0.2, 0.8, 0.5, 0.4], [0.2, 0.6, 0.1, 0.4], [0.6, 0.8, 0.3, 0.2], [0.2, 0.6, 0.3, 0.9]])

我的PyTorch实现代码

decoderParamSum = decoder_alpha + decoder_beta
priorParamSum = prior_alpha + prior_beta
alphaParamsDiff = decoder_alpha - prior_alpha
numerator = torch.lgamma(decoderParamSum) + torch.lgamma(prior_alpha) + torch.lgamma(prior_beta)
denomirator = torch.lgamma(decoder_alpha) + torch.lgamma(decoder_beta) + torch.lgamma(priorParamSum)
firstTerm = torch.sum((numerator - denomirator),dim=1)
secondTerm = torch.sum((torch.digamma(decoderParamSum)-torch.digamma(decoder_beta)), dim=1)
secondTerm = torch.reshape(secondTerm, (input.shape[0], 1))
secondTerm = torch.digamma(decoder_alpha) - torch.digamma(decoder_beta) - secondTerm
secondTerm = torch.sum(torch.multiply(alphaParamsDiff, secondTerm), dim=1)
thirdTerm = torch.cumsum((torch.digamma(decoderParamSum)-torch.digamma(decoder_beta)), dim=1)
thirdTerm = torch.reshape(thirdTerm,(input.shape[0], 1))
v1 = torch.cat([decoder_beta[:,:-1] -decoder_alpha[:, 1:] - decoder_beta[:, 1:], decoder_beta[:, -1:] - 1], dim=-1)
v2 = torch.cat([prior_beta[:,:-1] -prior_alpha[:, 1:] - prior_beta[:, 1:], prior_beta[:, -1:] - 1], dim=-1)

thirdTerm = torch.sum((torch.multiply((v1-v2), thirdTerm)), dim=1)

KLD = firstTerm - secondTerm + thirdTerm

问题

运行上述代码后得到的损失值远不符合预期,出现较大负值且不稳定,推测代码实现存在问题。恳请帮忙检查该KL散度的实现代码,或告知是否已有现成的Python实现方案(我已检索网络但未找到)。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 01:06:18