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

EM算法拟合高斯混合模型聚类异常:三高斯收敛为同一模型求助

高斯混合模型EM算法聚类异常排查:三个高斯参数收敛为相同值问题

我有一个含150个样本的4维多元数据集,需要用EM算法将其划分为3个聚类,目标是求解三个聚类对应的高斯分布参数。我参考《Bishop-Pattern-Recognition-and-Machine-Learning-2006》中的高斯混合模型聚类EM算法实现了代码,但运行后发现三个高斯的均值、协方差矩阵参数最终收敛为完全相同的结果,而数据集实际应该对应三个不同聚类。请帮忙排查代码中的错误。

原始代码

import numpy as np, math, pandas as pd
from scipy.stats import multivariate_normal
from decimal import Decimal


def clustering(points):
    points = list(set(points)) #To eliminate redundant points

    '''
    Initialize random gaussians!
    '''
    mean = [np.array([5,4,1.5,0.4]), np.array([4.5,5,0.5,0.5]), np.array([6,3,2,1])]
    covariance1 = np.array([[16,0,0,0], [0,4,0,0], [0,0,9,0], [0,0,0,1]])
    covariance2 = np.array([[4,0,0,0], [0,16,0,0], [0,0,1,0], [0,0,0,9]])
    covariance3 = np.array([[16,0,0,0], [0,1,0,0], [0,0,9,0], [0,0,0,4]])
    covariance = [covariance1, covariance2, covariance3]
    weights = np.array([0.3, 0.3, 0.4])

    init_log_likelihood = None #Log likelihood to test convergence, initially None
    threshold = 0.0000000000000000000001 #Threshold difference for convergence

    iterations = 0 #for counting iterations

    while True:
        iterations +=1 

        prob_dict = {point:[] for point in points}
        new_means = [np.array([0,0,0,0]) for _ in range(3)]
        new_covariance = [np.array([[0,0,0,0], [0,0,0,0], [0,0,0,0], [0,0,0,0]]) for _ in range(3)]

        log_likelihood = Decimal(0) #Using decimal to increase the precision of floating point numbers

        for point in points:

            prob_sum = 0 #Prob sum for normalization

            for _ in range(3):
                prob =  multivariate_normal.pdf(point, mean[_], covariance[_], allow_singular=True) #Finding gaussian probability
                prob = np.dot(weights[_], prob)+1 #Multiply gaussian probability with weight
                prob_dict[point].append(prob)
                prob_sum += prob
            
            prob_sum = Decimal(prob_sum)

            log_likelihood += prob_sum.ln() #comprehensive computation of log likelihood

            for _ in range(3):
                prob_dict[point][_] = prob_dict[point][_]/float(prob_sum) #Normalizing the probability
                new_means[_] = new_means[_]+np.dot(prob_dict[point][_], list(point)) #New means for the next iteration

        for _ in range(3):
            weights[_] = sum([prob_item[1][_] for prob_item in prob_dict.items()]) #New weights
            new_means[_] = np.divide(new_means[_], weights[_])
            for point in points:
              #New covariance
              new_covariance[_] = new_covariance[_]+np.dot(np.dot(prob_dict[point][_], (np.array(point)-new_means[_])),  np.transpose(np.array(point)-new_means[_]))
        
        for _ in range(3):
            new_covariance[_] = np.divide(new_covariance[_],weights[_])
            weights[_] = weights[_]/len(points)        
        
        mean = new_means 
        covariance = new_covariance

        #check for convergence
        if init_log_likelihood and abs(init_log_likelihood-log_likelihood)<=threshold:
            break

        init_log_likelihood = log_likelihood

    print("Means: ", mean)
    print("Covariance: ", covariance)
    print("Log-likehilhood: ", log_likelihood)
    print("Iterations: ", iterations)

代码错误分析

  • 错误1:错误移除重复样本
    points = list(set(points))存在两处问题:一是numpy数组无法直接作为集合元素,执行时会被转成元组破坏结构;二是移除重复样本会严重扭曲原始数据集的分布,EM算法依赖样本数量和分布特征区分聚类,丢失样本会导致算法无法捕捉真实聚类结构。

  • 错误2:加权概率计算错误添加+1
    E步中prob = np.dot(weights[_], prob)+1完全违反EM逻辑。正确操作应为计算权重×高斯概率密度,添加+1会无意义抬高所有分量的概率,导致后验概率归一化结果趋向一致,最终三个高斯参数收敛为相同值。

  • 错误3:对数似然计算逻辑错误
    当前代码对Σ(权重×高斯概率+1)取对数后累加,正确的对数似然应为每个样本的log(Σ(权重×高斯概率))之和。该错误会导致收敛判断失效,算法无法识别正确收敛点,甚至陷入错误局部极值。

  • 错误4:协方差计算使用更新后的均值
    M步中先更新均值new_means[_] = np.divide(new_means[_], weights[_]),再用新均值计算协方差。正确做法是使用E步时的旧均值计算协方差,否则会引入计算偏差,影响参数更新的正确性。

修正后的代码

import numpy as np
from scipy.stats import multivariate_normal
from decimal import Decimal


def clustering(points):
    # 保留原始样本,移除错误去重操作
    points = np.array(points)
    n_samples, n_dims = points.shape

    '''
    初始化高斯分量参数
    '''
    mean = [np.array([5,4,1.5,0.4]), np.array([4.5,5,0.5,0.5]), np.array([6,3,2,1])]
    covariance1 = np.array([[16,0,0,0], [0,4,0,0], [0,0,9,0], [0,0,0,1]])
    covariance2 = np.array([[4,0,0,0], [0,16,0,0], [0,0,1,0], [0,0,0,9]])
    covariance3 = np.array([[16,0,0,0], [0,1,0,0], [0,0,9,0], [0,0,0,4]])
    covariance = [covariance1, covariance2, covariance3]
    weights = np.array([0.3, 0.3, 0.4])

    init_log_likelihood = None
    threshold = 1e-6  # 调整为合理收敛阈值
    iterations = 0

    while True:
        iterations +=1 

        # 存储每个样本属于各分量的后验概率
        gamma = np.zeros((n_samples, 3))
        new_means = [np.zeros(n_dims) for _ in range(3)]
        new_covariance = [np.zeros((n_dims, n_dims)) for _ in range(3)]
        log_likelihood = Decimal(0)

        # E步:计算后验概率和对数似然
        for i in range(n_samples):
            point = points[i]
            # 计算每个分量的加权概率
            weighted_probs = [weights[k] * multivariate_normal.pdf(point, mean[k], covariance[k], allow_singular=True) 
                              for k in range(3)]
            total_prob = sum(weighted_probs)
            # 计算对数似然
            log_likelihood += Decimal(total_prob).ln()
            # 归一化得到后验概率
            gamma[i] = np.array(weighted_probs) / total_prob

        # M步:更新参数
        for k in range(3):
            # 计算分量的有效样本数
            N_k = np.sum(gamma[:, k])
            # 更新权重
            weights[k] = N_k / n_samples
            # 更新均值
            new_means[k] = np.sum(gamma[:, k, np.newaxis] * points, axis=0) / N_k
            # 更新协方差:使用旧均值计算
            diff = points - mean[k]
            new_covariance[k] = np.dot(gamma[:, k] * diff.T, diff) / N_k

        # 更新参数
        mean = new_means
        covariance = new_covariance

        # 检查收敛
        if init_log_likelihood is not None and abs(init_log_likelihood - log_likelihood) <= threshold:
            break

        init_log_likelihood = log_likelihood

    print("Means: ", mean)
    print("Covariance: ", covariance)
    print("Log-likehilhood: ", log_likelihood)
    print("Iterations: ", iterations)

内容的提问来源于stack exchange,提问作者Sasi Bhushan V Saladi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 22:05:22