【问题标题】:Kullback-Leibler divergence from Gaussian pm,pv to Gaussian qm,qv从高斯 pm,pv 到高斯 qm,qv 的 Kullback-Leibler 散度
【发布时间】:2022-01-04 21:26:59
【问题描述】:

我正在尝试计算从 Gaussian#1 到 Gaussian#2 的 Kullback-Leibler 散度 我有两个高斯的平均值和标准差 我从http://www.cs.cmu.edu/~chanwook/MySoftware/rm1_Spk-by-Spk_MLLR/rm1_PNCC_MLLR_1/rm1/python/sphinx/divergence.py尝试了这段代码

def gau_kl(pm, pv, qm, qv):
    """
    Kullback-Leibler divergence from Gaussian pm,pv to Gaussian qm,qv.
    Also computes KL divergence from a single Gaussian pm,pv to a set
    of Gaussians qm,qv.
    Diagonal covariances are assumed.  Divergence is expressed in nats.
    """
    if (len(qm.shape) == 2):
        axis = 1
    else:
        axis = 0
    # Determinants of diagonal covariances pv, qv
    dpv = pv.prod()
    dqv = qv.prod(axis)
    # Inverse of diagonal covariance qv
    iqv = 1./qv
    # Difference between means pm, qm
    diff = qm - pm
    return (0.5 *
            (numpy.log(dqv / dpv)            # log |\Sigma_q| / |\Sigma_p|
             + (iqv * pv).sum(axis)          # + tr(\Sigma_q^{-1} * \Sigma_p)
             + (diff * iqv * diff).sum(axis) # + (\mu_q-\mu_p)^T\Sigma_q^{-1}(\mu_q-\mu_p)
             - len(pm)))                     # - N

我使用均值和标准差作为输入,但代码的最后一行(len(pm)) 导致错误,因为均值是一个数字,我这里的 len 函数我看不懂。

注意。两组(即高斯)不相等,这就是为什么我不能使用 scipy.stats.entropy

【问题讨论】:

    标签: python gaussian


    【解决方案1】:

    以下函数计算任意两个多元正态分布之间的 KL-Divergence(协方差矩阵不需要是对角线)(其中 numpy 被导入为 np)

    def kl_mvn(m0, S0, m1, S1):
        """
        Kullback-Liebler divergence from Gaussian pm,pv to Gaussian qm,qv.
        Also computes KL divergence from a single Gaussian pm,pv to a set
        of Gaussians qm,qv.
        
    
        From wikipedia
        KL( (m0, S0) || (m1, S1))
             = .5 * ( tr(S1^{-1} S0) + log |S1|/|S0| + 
                      (m1 - m0)^T S1^{-1} (m1 - m0) - N )
        """
        # store inv diag covariance of S1 and diff between means
        N = m0.shape[0]
        iS1 = np.linalg.inv(S1)
        diff = m1 - m0
    
        # kl is made of three terms
        tr_term   = np.trace(iS1 @ S0)
        det_term  = np.log(np.linalg.det(S1)/np.linalg.det(S0)) #np.sum(np.log(S1)) - np.sum(np.log(S0))
        quad_term = diff.T @ np.linalg.inv(S1) @ diff #np.sum( (diff*diff) * iS1, axis=1)
        #print(tr_term,det_term,quad_term)
        return .5 * (tr_term + det_term + quad_term - N) 
    

    【讨论】:

    • 您的帖子说“协方差矩阵不需要是对角线”,但是您的函数的文档字符串与此相矛盾。
    • 是的,是哪个?
    • @Peter:已修复。它也可以接受非对角协方差。
    • noiceeeeeeeeeeeee
    【解决方案2】:

    如果你还有兴趣...

    该函数需要多元高斯协方差矩阵的对角线条目,而不是您提到的标准偏差。如果您的输入是单变量高斯,那么 pvqv 都是对应高斯方差的长度为 1 的向量。

    此外,len(pm) 对应于均值向量的维度。在多元正态分布部分here中确实是k。对于单变量高斯,k 为 1,对于二元高斯,k 为 2,以此类推。

    【讨论】:

      猜你喜欢
      • 2011-06-19
      • 1970-01-01
      • 2017-12-18
      • 2014-09-14
      • 1970-01-01
      • 1970-01-01
      • 2014-12-24
      • 2012-04-17
      • 1970-01-01
      相关资源
      最近更新 更多