基于NumPy的聚类数据均值计算:如何避免显式循环?
无显式循环的NumPy聚类均值计算方法
给你两种不用写循环的高效NumPy实现方法,比原有的Python循环效率高很多——毕竟NumPy底层是C实现的:
方法一:用np.bincount做加权累加
bincount可以直接按类别统计样本数,还能配合权重参数实现按类别累加每个维度的数值,最后做除法得到均值:
import numpy as np np.random.seed(42) k = 3 x = np.random.rand(100, 2) c = np.random.randint(0, k, size=x.shape[0]) # 统计每个聚类的样本数量 cluster_counts = np.bincount(c) # 按聚类ID累加每个维度的数值,再转置成(k, 2)的形状 cluster_sums = np.array([np.bincount(c, weights=x[:, dim]) for dim in range(x.shape[1])]).T # 计算均值 mu = cluster_sums / cluster_counts[:, np.newaxis]
方法二:用np.add.at原地累加
np.add.at可以按指定索引原地累加数值,一步完成所有类别的求和,再统计样本数后计算均值:
import numpy as np np.random.seed(42) k = 3 x = np.random.rand(100, 2) c = np.random.randint(0, k, size=x.shape[0]) mu = np.zeros((k, x.shape[1])) cluster_counts = np.zeros(k, dtype=int) # 按聚类ID把x的数值累加到mu对应位置 np.add.at(mu, c, x) # 统计每个聚类的样本数 np.add.at(cluster_counts, c, 1) # 计算均值 mu /= cluster_counts[:, np.newaxis]
这两种方法都完全避免了Python级别的显式循环,在样本量很大的时候,速度会比原循环方法快不少,且计算结果和原方法完全一致。
内容的提问来源于stack exchange,提问作者andywiecko
相关产品推荐
相关产品推荐

