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

使用KMeans迭代构建层级聚类树时递归逻辑异常问题求助

问题修复方案

核心Bug原因

你遇到的层级节点子节点长度完全一致的问题,是Python可变默认参数陷阱导致的:
你在定义Node类的__init__方法时,将children的默认值设为了空列表[]。Python中列表、字典这类可变对象作为函数/方法默认参数时,只会在定义阶段创建一次,后续所有没有显式传入children参数的Node实例,都会共用同一个列表对象。你在递归过程中每次往节点的children追加元素时,本质都是修改同一个全局共享的列表,因此所有层级节点的子节点长度完全相同。

其他注意事项

  1. 你代码中的K变量未提前定义,运行会直接报错,建议将其设为cluster函数的入参,提升灵活性。
  2. 统计当前簇的标签计数时,无需遍历全局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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 18:15:06