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

Python MyKMeans类fit方法运行速度远慢于同逻辑独立函数,如何优化?

性能差异的根本原因

你认为两段逻辑完全相同,但实际存在两个关键差异,根本不是类封装导致的性能损耗:

  1. 独立函数的return语句缩进错误,提前终止了迭代
    你写的独立函数中return语句的缩进级别与while循环内的if/else判断同级,也就是说第一次迭代完成后无论是否收敛,都会直接返回结果,根本没有完成后续的迭代流程。而类中的fit方法会一直循环到质心完全不再变化才停止,迭代次数相差数倍,自然运行时间差距极大。你可以将独立函数的return语句调整到与while循环同级的位置,再次测试会发现运行时间和类方法基本一致。
  2. 两段代码都存在质心计算位置错误的问题,额外增加了大量无效计算
    你的类方法和独立函数中,new_centroids的计算逻辑被缩进放在了遍历每个样本的for i, point循环内部,也就是说每给一个样本分配完簇,就会全量计算一次质心。而KMeans的正确逻辑是:所有样本都完成簇分配之后,再统一计算一次新质心,这个错误会导致你多执行了N次(N为样本数量)非常耗时的groupby操作,是类方法运行慢的核心原因之一。
保留类结构的优化方案

按照正确的KMeans逻辑修正代码后,类实现的性能和独立函数完全一致,还可以通过以下方式进一步提速:

  • 修正质心计算位置,把new_centroids的计算移到样本遍历循环之外,每个迭代轮次只计算一次质心
  • 替换pandas的groupby操作为numpy原生分组均值计算,避免DataFrame转换的额外开销
  • 启用你定义的max_iter参数,避免数据无法收敛时的无限循环
  • 用numpy广播优化欧氏距离计算,替换掉两层for循环,大幅减少循环开销

修正后的类实现示例

import numpy as np

def euclidean_dist(a, b):
    return np.sqrt(np.sum((a - b)**2))

class MyKMeans:
    def __init__(self, max_iter = 300):
        self.max_iter = max_iter
        self.centroids = None
        self.clusters = None

    def fit(self, X, k):
        n_samples = X.shape[0]
        self.clusters = np.zeros(n_samples, dtype=int)
        
        # 随机初始化质心
        random_idxs = np.random.choice(n_samples, size=k, replace=False)
        self.centroids = X[random_idxs, :].copy()

        for _ in range(self.max_iter):
            # 步骤1:全量样本分配簇
            for i, point in enumerate(X):
                min_d = float('inf')
                for idx, centroid in enumerate(self.centroids):
                    d = euclidean_dist(centroid, point)
                    if d < min_d:
                        min_d = d
                        self.clusters[i] = idx
            
            # 步骤2:统一计算新质心(每个迭代仅执行1次)
            new_centroids = np.zeros_like(self.centroids)
            for cluster_idx in range(k):
                cluster_points = X[self.clusters == cluster_idx]
                if len(cluster_points) > 0:
                    new_centroids[cluster_idx] = cluster_points.mean(axis=0)
                else:
                    # 空簇随机重置质心
                    new_centroids[cluster_idx] = X[np.random.choice(n_samples)]
            
            # 判断收敛
            if np.allclose(self.centroids, new_centroids):
                break
            self.centroids = new_centroids

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 07:36:07