如何展示图像聚类任务中高斯混合模型(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
相关产品推荐
相关产品推荐

