不能直接用scipy.spatial.distance.cdist,因其在小规模数据上开销大、不便于向量化更新,且对一维centroids报错;而np.linalg.norm配合广播更鲁棒高效,能自动处理k=1场景并保持数值稳定与dtype一致。

为什么不能直接用 scipy.spatial.distance.cdist 就完事?
因为 K-means 每轮迭代都要计算所有样本到所有聚类中心的欧氏距离,生成一个 (n_samples, n_clusters) 的距离矩阵。虽然 cdist 能做,但它在小规模数据上开销大、不便于后续向量化更新;更关键的是——它默认把输入当二维数组处理,若你传入的 centroids 是一维(比如只有一簇),cdist 会报 ValueError: XA must be a 2-dimensional array,而 NumPy 原生广播机制反而更鲁棒。
用 np.linalg.norm + 广播构造距离矩阵
核心思路是把样本点 X(shape (n, d))和中心点 centroids(shape (k, d))通过广播对齐成 (n, k, d),再逐点求 L2 范数。实际不用显式扩展维度,靠 np.newaxis 或 None 插入轴即可:
-
X[:, None, :]把X变成(n, 1, d) -
centroids[None, :, :]把centroids变成(1, k, d) - 二者相减后 shape 为
(n, k, d),再对最后一维求np.linalg.norm(..., axis=2)
完整一行写法:
distances = np.linalg.norm(X[:, None, :] - centroids[None, :, :], axis=2)
避免 np.sqrt(np.sum(...)) 手动展开的坑
有人会写 np.sqrt(np.sum((X[:, None, :] - centroids[None, :, :])**2, axis=2)),逻辑没错,但有三个实际问题:
快速生成专业的 Python 脚本和应用代码。一键创建完整项目结构,支持CLI、API、爬虫、Bot、Django等多种项目类型,包含完整的项目结构、配置文件、依赖管理、测试、README和文档。
- **数值不稳定**:平方后可能溢出,尤其当特征值很大时;
np.linalg.norm内部做了缩放处理 - **性能差**:显式平方+求和比
norm多一次内存写入,实测慢 10%–20% - **dtype 不一致**:若
X是float32,手动写法可能因中间计算升为float64,而norm默认保持输入精度
当 centroids 只有一个时怎么不出错?
单中心场景(比如初始化第一轮或 k=1)最容易暴露维度 bug。错误写法:X - centroids 会触发广播失败((n,d) − (d,) 成功,但 (n,d) − (1,d) 更安全)。正确做法始终统一用 None 显式扩维:
# 即使 centroids.shape == (1, d),也这么写<br>distances = np.linalg.norm(X[:, None, :] - centroids[None, :, :], axis=2)
这样无论
k==1 还是 k==100,形状都自动对齐,无需分支判断。真正容易被忽略的是:距离矩阵必须是浮点型(哪怕输入是整数),否则后续 argmin 索引虽能运行,但若中间出现 nan 或截断误差,会导致簇分配跳变——务必确认 X 和 centroids 至少一个是 float 类型。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!










