使用KMeans迭代构建层级聚类树时递归逻辑异常问题求助
问题修复方案
核心Bug原因
你遇到的层级节点子节点长度完全一致的问题,是Python可变默认参数陷阱导致的:
你在定义Node类的__init__方法时,将children的默认值设为了空列表[]。Python中列表、字典这类可变对象作为函数/方法默认参数时,只会在定义阶段创建一次,后续所有没有显式传入children参数的Node实例,都会共用同一个列表对象。你在递归过程中每次往节点的children追加元素时,本质都是修改同一个全局共享的列表,因此所有层级节点的子节点长度完全相同。
其他注意事项
- 你代码中的
K变量未提前定义,运行会直接报错,建议将其设为cluster函数的入参,提升灵活性。 - 统计当前簇的标签计数时,无需遍历全局
labels数组,直接统计kmeans.labels_即可,性能更高。
修复后代码
修正Node类定义
from typing import Any, Optional, List import numpy as np from sklearn.cluster import MiniBatchKMeans from collections import Counter class Node: def __init__(self, name: str, mean: np.array, children: Optional[List[Any]] = None): self.name = name self.mean = mean # 每次实例化单独创建空列表,避免共享 self.children = children if children is not None else []
修正聚类递归函数
def cluster(idx: np.array, parent: Optional[Node]=None, parent_label: str="", K:int=10): kmeans = MiniBatchKMeans(n_clusters=K) kmeans.fit(data[idx]) # 生成当前层级标签 current_labels = [parent_label + ">" + str(label) for label in kmeans.labels_] labels[idx] = current_labels # 直接统计当前批次的标签计数,减少全局遍历 for label, count in Counter(kmeans.labels_).items(): full_label = parent_label + ">" + str(label) node = Node( full_label, kmeans.cluster_centers_[label] ) if count > 1000: cluster(labels == full_label, node, full_label, K=K) parent.children.append(node)
测试代码
# 测试数据 np.random.seed(42) data = np.random.randn(100000, 768) # 初始化根节点和标签数组 head_node = Node("head", 0) labels = np.array(["" for _ in range(len(data))], dtype=object) idx = labels == "" # 执行聚类,K可自行调整 cluster(idx, head_node, K=10) # 验证子节点长度 print(len(head_node.children), len(head_node.children[2].children)) # 输出为 10, 10(符合预期,第一层10个簇,每个簇1万条,再拆分10个刚好1000条停止)
内容的提问来源于stack exchange,提问作者sachinruk
相关产品推荐
相关产品推荐

