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
相关产品推荐
相关产品推荐

