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

Python中KL散度与JS散度函数输出异常原因排查

KL/JS散度计算函数输出异常的原因

问题重现

你定义的代码如下:

import numpy as np

p = np.array([1., 0., 0.])
q = np.array([0., 0., 1.])
m = np.array([0.5, 0. , 0.5])

def kl(p, q):
    return np.sum(np.log2(np.divide(p, q, where=q!=0), where=p!=0))

def js(p, q):
    m = 1/2 * (p + q)
    return 1/2 * (kl(p, m) + kl(q, m))

实际调用输出:

>>> js(p, q)
3.0
>>> kl(p, m)
2.0
>>> kl(p, m)
1.0

预期输出:

>>> js(p, q)
1.0
>>> kl(p, m)
1.0
>>> kl(q, m)
1.0

核心问题分析

1. KL散度公式实现错误

你编写的kl函数完全不符合KL散度的数学定义:

  • 正确的KL散度公式为:$D_{KL}(p||q) = \sum_{i} p_i \log_2\left(\frac{p_i}{q_i}\right)$
  • 你的函数漏掉了关键的p_i乘积项,直接对$\log_2(p_i/q_i)$求和,这会导致计算结果完全偏离预期。

2. numpy where参数使用不当导致结果不稳定

在调用np.divide和np.log2时使用where参数,但未指定out数组,也未处理where为False的位置:

  • numpy的ufunc(如divide、log2)在使用where且未指定out时,会创建一个未完全初始化的数组,where为False的位置会残留内存中的垃圾值。
  • 求和时这些垃圾值会被计入结果,导致每次调用的输出可能不同;添加print语句时内存布局变化,垃圾值恰好变为0,因此得到正确结果。

修正后的代码

import numpy as np

p = np.array([1., 0., 0.])
q = np.array([0., 0., 1.])
m = np.array([0.5, 0. , 0.5])

def kl(p, q):
    # 过滤掉p或q为0的项,避免除以零和对数无意义的情况
    valid_mask = (p > 0) & (q > 0)
    # 按照KL散度公式计算有效项的和
    return np.sum(p[valid_mask] * np.log2(p[valid_mask] / q[valid_mask]))

def js(p, q):
    m = 0.5 * (p + q)
    return 0.5 * (kl(p, m) + kl(q, m))

测试修正后的代码,输出与预期完全一致:

>>> js(p, q)
1.0
>>> kl(p, m)
1.0
>>> kl(q, m)
1.0

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 23:30:44