自定义Kmeans算法处理2维数据正常,784维数据拼接报错如何解决?
错误原因
- 硬编码特征维度:初始化每个聚类的存储数组时,代码写死了
reshape(2, 0),仅适配2维输入数据。当输入变为784维数据时,初始化的数组维度和实际数据维度不匹配,拼接时报错。 - 变量名未统一:函数入参为
data,但拼接样本时用了未定义的变量X[i],2维测试场景下若全局存在X变量会隐式正常运行,切换输入后直接触发异常。
修复后完整代码
import numpy as np # parameter c is for how many cluster do you want def kmeans(data, c, iter, state): np.random.seed(state) data = np.array(data) m = data.shape[0] # 样本数量 n = data.shape[1] # 特征维度 Centroids = np.array([]).reshape(n, 0) for i in range(c): rand = np.random.randint(0, m-1) Centroids = np.c_[Centroids, data[rand]] result = {} for _ in range(iter): distance = np.array([]).reshape(m, 0) # 欧氏距离计算 for k in range(c): tempDist = np.sum((data - Centroids[:,k])**2, axis=1) distance = np.c_[distance, tempDist] C = np.argmin(distance, axis=1) + 1 Y = {} # 修复点1:把固定的2替换为动态获取的特征维度n for k in range(c): Y[k+1] = np.array([]).reshape(n, 0) # 修复点2:把X[i]替换为函数入参data[i] for i in range(m): Y[C[i]] = np.c_[Y[C[i]], data[i]] for k in range(c): Y[k+1] = Y[k+1].T for k in range(c): Centroids[:,k] = np.mean(Y[k+1], axis=0) result = Y return result
验证说明
修复后既兼容原有2维测试用例,也支持784维等高维数据的聚类计算,无需再根据输入维度修改代码。
内容的提问来源于stack exchange,提问作者Thomas Iskandar
相关产品推荐
相关产品推荐

