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

sklearn AgglomerativeClustering的connectivity参数未按预期工作

问题描述

使用sklearn的AgglomerativeClustering基于给定距离矩阵聚类,要求部分对象绝对不能归为同一簇,且已设置合理聚类数。尝试过以下操作但未解决问题:

  • 将互斥对象的距离设为1,仍有互斥对象被分到同一簇;
  • 传入connectivity连通性矩阵,结果不符合预期(如对象2和3仍同簇);
  • 更换linkage类型、用kneighbors_graph生成连通矩阵,均无改善。

疑问:是否正确使用了connectivity参数?如何实现绝对互斥的聚类需求?

代码示例:

import numpy as np
from sklearn.cluster import AgglomerativeClustering

n_clusters = 2

distance_matrix = np.array([
[0.,1.,0.,0.,0.],
[1.,0.,0.,0.,0.], 
[0.,0.,0.,1.,0.], 
[0.,0.,1.,0.,0.], 
[0.,0.,0.,0.,0.]
])

connectivity_matrix = np.array([
[True, False, True, True, True],
[False, True, True, True, True],
[True, True, True, False, True],
[True, True, False, True, True],
[True, True, True, True, True]
])

agglomerativeClustering = AgglomerativeClustering(n_clusters=n_clusters, metric='precomputed', 
                                                  linkage='average')
cluster_labels_wo_connectivity = agglomerativeClustering.fit_predict(distance_matrix)

agglomerativeClustering = AgglomerativeClustering(n_clusters=n_clusters, metric='precomputed', 
                                                  linkage='average', connectivity=connectivity_matrix)
cluster_labels_with_connectivity = agglomerativeClustering.fit_predict(distance_matrix)

print("无连通矩阵的聚类结果: ", cluster_labels_wo_connectivity)
print("带连通矩阵的聚类结果: ", cluster_labels_with_connectivity)

解决方案

1. connectivity参数的局限性

connectivity参数仅能限制样本之间的直接合并路径:矩阵中False对应的样本对不能被直接合并,但允许通过其他样本间接合并(比如对象2和3不能直接合并,但可以通过对象4间接合并到同一簇)。因此它无法实现"绝对不能同簇"的硬约束,仅能影响合并顺序。

2. 实现绝对互斥的正确方法

要让部分对象绝对不被分到同一簇,最直接有效的方式是将互斥对象对的距离设为无穷大(np.inf):

  • 聚类算法会优先选择距离最小的簇对合并;
  • 互斥对的距离设为无穷大后,任何包含这两个对象的簇合并代价都会是无穷大,算法会完全避开这类合并,从而保证它们永远不在同一簇。

修改后的代码示例:

import numpy as np
from sklearn.cluster import AgglomerativeClustering

n_clusters = 2

# 修改距离矩阵:将互斥对(0,1)、(2,3)的距离设为无穷大
distance_matrix = np.array([
[0., np.inf, 0., 0., 0.],
[np.inf, 0., 0., 0., 0.], 
[0., 0., 0., np.inf, 0.], 
[0., 0., np.inf, 0., 0.], 
[0., 0., 0., 0., 0.]
])

agglomerativeClustering = AgglomerativeClustering(n_clusters=n_clusters, metric='precomputed', 
                                                  linkage='average')
cluster_labels = agglomerativeClustering.fit_predict(distance_matrix)

print("带硬互斥约束的聚类结果: ", cluster_labels)

3. 额外说明

  • 如果同时需要限制其他合并路径,可以结合connectivity参数,但硬互斥必须通过修改距离矩阵实现;
  • 确保聚类数n_clusters是合理的,即互斥约束下存在可行的聚类划分(比如不能要求2个簇,但有3组完全互斥的对象)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 20:50:11