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

如何为字典各键对应列表分别设置不同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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 23:54:17