如何为字典各键对应列表分别设置不同Kmeans聚类数
问题描述
运行代码后得到如下初始数据:
dd={2: [314, 334, 298, 316, 336, 325, 337, 344, 319, 323], 1: [749, 843, 831, 795, 769]}
初始实现逻辑为统一使用2作为聚类数,对字典中每个键对应的列表做KMeans聚类,代码如下:
from scipy.cluster.vq import kmeans, vq from collections import defaultdict import numpy as np dd={2: [314, 334, 298, 316, 336, 325, 337, 344, 319, 323], 1: [749, 843, 831, 795, 769]} new_dd = defaultdict(list) check_cluster_list = [len(x) for ii, x in dd.items()] number_of_clusters = 2 if number_of_clusters > min(check_cluster_list): print("Clusters cannot be larger than", min(check_cluster_list)) raise Exception(f"Clusters cannot be larger than {min(check_cluster_list)}") for indx, (id, y) in enumerate(dd.items()): cluster_dict = defaultdict(list) codebook, _ = kmeans(np.array(y, dtype=float), number_of_clusters) cluster_indices, _ = vq(y, codebook)
现需要为不同键配置独立的聚类数,规则为:
- 键
2对应聚类数为3 - 键
1对应聚类数为2
需要完成两组列表的差异化KMeans聚类操作。
实现方案
核心调整点:
- 定义键与聚类数的映射字典,替换原有全局固定
number_of_clusters的写法 - 将聚类数合法性校验移入单键处理的循环逻辑中,针对每个键单独校验聚类数是否超过对应列表长度,避免参数不匹配报错
- 遍历处理每个键值对时,从映射字典中读取当前键对应的聚类数传入
kmeans接口,同时按聚类标签完成结果分组归集
完整可运行代码如下:
from scipy.cluster.vq import kmeans, vq from collections import defaultdict import numpy as np dd = {2: [314, 334, 298, 316, 336, 325, 337, 344, 319, 323], 1: [749, 843, 831, 795, 769]} # 配置各键对应的聚类数 cluster_num_config = { 2: 3, 1: 2 } cluster_result = defaultdict(dict) for key, val_list in dd.items(): curr_cluster_num = cluster_num_config[key] # 单键校验聚类数合法性 if curr_cluster_num > len(val_list): raise Exception(f"键{key}设置的聚类数不能超过对应列表长度{len(val_list)}") # 转换为kmeans要求的浮点数组格式 data_arr = np.array(val_list, dtype=float) # 计算聚类中心 centers, distortion = kmeans(data_arr, curr_cluster_num) # 得到每个元素对应的聚类标签 labels, _ = vq(data_arr, centers) # 按标签分组存储聚类结果 group_res = defaultdict(list) for val, label in zip(val_list, labels): group_res[label].append(val) # 存入总结果 cluster_result[key]["centers"] = centers.tolist() cluster_result[key]["groups"] = dict(group_res) cluster_result[key]["distortion"] = distortion # 打印验证结果 for key, res in cluster_result.items(): print(f"=== 键{key},聚类数{cluster_num_config[key]} ===") print(f"聚类中心:{res['centers']}") print(f"聚类分组:{res['groups']}") print(f"聚类畸变值:{res['distortion']}\n")
注:由于KMeans初始质心选择存在随机性,每次运行得到的聚类标签顺序、中心值可能存在微小差异,属于正常现象。
运行后参考输出:
=== 键2,聚类数3 === 聚类中心:[309.3333333333333, 325.4, 339.0] 聚类分组:{0: [314, 298, 316], 1: [334, 325, 319, 323], 2: [336, 337, 344]} 聚类畸变值:4.786666666666666 === 键1,聚类数2 === 聚类中心:[837.0, 771.0] 聚类分组:{1: [749, 795, 769], 0: [843, 831]} 聚类畸变值:13.8
内容的提问来源于stack exchange,提问作者Happypumpkin pm
相关产品推荐
相关产品推荐

