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

如何展示图像聚类任务中高斯混合模型(GMM)的聚类分布

问题说明

我使用附带的代码完成了多幅图像的高斯混合模型(GMM)训练,目前已经实现了在图像直方图上展示GMM拟合结果的功能,但还需要进一步展示GMM的聚类分布效果。

原有实现代码:

# Code for GMM
import os
import matplotlib.pyplot as plt
import numpy as np
import cv2


img = cv2.imread("test.jpg")

#Convert MxNx3 image into Kx3 where K=MxN
img2 = img.reshape((-1,3))  #-1 reshape means, in this case MxN

from sklearn.mixture import GaussianMixture as GMM

#covariance choices, full, tied, diag, spherical
gmm_model = GMM(n_components=6, covariance_type='full').fit(img2)  #tied works better than full
gmm_labels = gmm_model.predict(img2)

#Put numbers back to original shape so we can reconstruct segmented image
original_shape = img.shape
segmented = gmm_labels.reshape(original_shape[0], original_shape[1])
cv2.imwrite("test_segmented.jpg", segmented)


gmm_model.means_

gmm_model.covariances_

gmm_model.weights_

print(gmm_model.means_, gmm_model.covariances_, gmm_model.weights_)

data = img2.ravel()
data = data[data != 0]
data = data[data != 1]  #Removes background pixels (intensities 0 and 1)
gmm = GMM(n_components = 6)
gmm = gmm.fit(X=np.expand_dims(data,1))
gmm_x = np.linspace(0,255,256)
gmm_y = np.exp(gmm.score_samples(gmm_x.reshape(-1,1)))


#Plot histograms and gaussian curves
fig, ax = plt.subplots()
ax.hist(img.ravel(),255,[2,256], density=True, stacked=True)
ax.plot(gmm_x, gmm_y, color="crimson", lw=2, label="GMM")

ax.set_ylabel("Frequency")
ax.set_xlabel("Pixel Intensity")

plt.legend()
plt. grid(False)

plt.show()
实现方法

现有代码已经完成GMM核心训练流程,补充两部分逻辑即可实现聚类分布效果展示:

  • 修正逻辑冗余问题:代码先后训练了两个独立GMM模型,第一个基于RGB三通道像素值训练,用于像素聚类分割;第二个基于单通道像素强度训练,用于直方图拟合。两个模型参数不互通,可视化时需对应场景使用。
  • 聚类空间分布展示:使用三通道GMM模型输出的各聚类RGB均值,替换对应标签位置的像素值,生成彩色分割图,可直观呈现不同聚类在图像上的区域划分。
  • 直方图片段分布展示:基于单通道GMM模型的参数,单独绘制每个高斯分量的分布曲线,可直观呈现每个聚类在像素强度维度的占比和分布范围。

修改后可直接运行的完整代码如下:

import matplotlib.pyplot as plt
import numpy as np
import cv2
from sklearn.mixture import GaussianMixture as GMM

# 读取图像
img = cv2.imread("test.jpg")
img2 = img.reshape((-1,3))  # 将MxNx3图像转换为Kx3格式,K=M*N

# 训练三通道GMM用于图像分割,固定随机种子保证结果可复现
gmm_model = GMM(n_components=6, covariance_type='full', random_state=42).fit(img2)
gmm_labels = gmm_model.predict(img2)

# ---------------------- 聚类空间分布可视化 ----------------------
original_shape = img.shape
# 生成灰度标签分割图,拉伸标签值到0-255范围避免图像过暗
segmented_gray = gmm_labels.reshape(original_shape[0], original_shape[1])
cv2.imwrite("test_segmented_gray.jpg", segmented_gray * (255//gmm_model.n_components))

# 生成彩色聚类分割图,每个聚类用对应分量的RGB均值填充
segmented_color = np.zeros_like(img2)
for cls_idx in range(gmm_model.n_components):
    segmented_color[gmm_labels == cls_idx] = gmm_model.means_[cls_idx]
segmented_color = segmented_color.reshape(img.shape)
cv2.imwrite("test_segmented_color.jpg", segmented_color)

# ---------------------- 直方图聚类分布可视化 ----------------------
# 训练单通道GMM用于直方图拟合
data = img2.ravel()
data = data[data != 0]
data = data[data != 1]  # 过滤强度为0、1的背景像素
gmm_hist = GMM(n_components=6, random_state=42)
gmm_hist.fit(X=np.expand_dims(data,1))
gmm_x = np.linspace(0,255,256)
gmm_y_total = np.exp(gmm_hist.score_samples(gmm_x.reshape(-1,1)))

# 绘制带单聚类分量的直方图拟合结果
fig, ax = plt.subplots(figsize=(10,6))
ax.hist(img.ravel(), 255, [2,256], density=True, stacked=True, alpha=0.5, label='像素分布')
ax.plot(gmm_x, gmm_y_total, color="crimson", lw=2, label="GMM整体拟合")
# 逐个绘制每个聚类对应的高斯分布曲线
for i in range(gmm_hist.n_components):
    weight = gmm_hist.weights_[i]
    mu = gmm_hist.means_[i][0]
    var = gmm_hist.covariances_[i][0][0]
    single_curve = weight * (1/(np.sqrt(2*np.pi*var))) * np.exp(-(gmm_x - mu)**2/(2*var))
    ax.plot(gmm_x, single_curve, lw=1.5, ls='--', label=f'聚类{i+1}')

ax.set_ylabel("频率")
ax.set_xlabel("像素强度")
plt.legend()
plt.grid(False)
plt.savefig("gmm_hist_with_clusters.jpg", dpi=150, bbox_inches='tight')
plt.show()

# 可选:并排展示原图、灰度分割图、彩色分割图做效果对比
fig, axes = plt.subplots(1,3, figsize=(18,6))
axes[0].imshow(cv2.cvtColor(img, cv2.COLOR_BGR2RGB))
axes[0].set_title("原图")
axes[0].axis('off')

axes[1].imshow(segmented_gray, cmap='gray')
axes[1].set_title("灰度标签聚类结果")
axes[1].axis('off')

axes[2].imshow(cv2.cvtColor(segmented_color, cv2.COLOR_BGR2RGB))
axes[2].set_title("彩色均值聚类结果")
axes[2].axis('off')

plt.savefig("gmm_segment_compare.jpg", dpi=150, bbox_inches='tight')
plt.show()

运行后将生成4份结果文件:

  • test_segmented_gray.jpg:灰度标签形式的聚类分割结果
  • test_segmented_color.jpg:聚类均值填充的彩色分割结果,可直接查看不同聚类的空间分布
  • gmm_hist_with_clusters.jpg:标注了每个单独聚类高斯曲线的直方图拟合图,可查看每个聚类的像素强度分布
  • gmm_segment_compare.jpg:原图、灰度分割、彩色分割的并排对比图

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 12:48:21