Python实现K-means聚类:调整后的中心节点变量为何未更新?
K-means聚类中心节点更新问题解决方法
问题场景
你编写了K-means聚类的Python代码,期望通过center_readjustment函数更新聚类中心,并用while循环重复执行主逻辑。当前center_readjustment内部能正确打印更新后的中心,但main0函数中打印的random_node_1仍是初始值,无法完成中心的迭代更新。
问题根源
你调用center_readjustment函数时,仅执行了函数逻辑,但没有将函数返回的新中心赋值回原变量。Python中,即便列表是可变对象,但你在函数内部将random_node重新赋值为元组(random_node = ((new_node_x, new_node_y))),会断开和原列表的引用,必须通过返回值重新绑定原变量才能完成更新。
解决步骤
- 接收函数返回值并更新变量:在
main0中调用center_readjustment时,将返回的新中心重新赋值给对应的节点变量 - 完善函数返回逻辑:当中心未发生变化时,返回原节点,避免出现
None值 - 让
main0返回更新后的中心:方便while循环获取新的中心,继续下一轮迭代 - 添加while循环控制迭代:直到所有中心不再变化或达到最大迭代次数,停止循环
修改后完整代码
import numpy as np import random import matplotlib.pyplot as plt data = np.loadtxt("k-mean-input.dat") numlist_x = list(range(-8, 15, 1)) numlist_y = list(range(-13, 12, 1)) # 初始化中心节点 random_node_1 = [random.choice(numlist_x), random.choice(numlist_y)] random_node_2 = [random.choice(numlist_x), random.choice(numlist_y)] random_node_3 = [random.choice(numlist_x), random.choice(numlist_y)] def distance(a, b): result = np.sqrt((a[0] - b[0])**2 + (a[1] - b[1])**2) return result def center_readjustment(matrix, current_node): total_x = 0 total_y = 0 for point in matrix: total_x += point[0] total_y += point[1] new_node_x = total_x / len(matrix) new_node_y = total_y / len(matrix) new_node = (new_node_x, new_node_y) # 对比新旧中心,返回新中心或原中心 if new_node != tuple(current_node): print(f"更新后的中心: {new_node}") return new_node else: return current_node def main0(data, node1, node2, node3): print(f"当前中心节点: {node1}, {node2}, {node3}") plt.figure() plt.scatter(node1[0], node1[1], color='red', marker='x', s=100) plt.scatter(node2[0], node2[1], color='green', marker='x', s=100) plt.scatter(node3[0], node3[1], color='blue', marker='x', s=100) first = [] second = [] third = [] for point in data: d1 = distance(node1, point) d2 = distance(node2, point) d3 = distance(node3, point) min_dist = min(d1, d2, d3) if min_dist == d1: first.append(point) elif min_dist == d2: second.append(point) else: third.append(point) # 绘制聚类点 for p in first: plt.scatter(p[0], p[1], color='red', s=10) for p in second: plt.scatter(p[0], p[1], color='green', s=10) for p in third: plt.scatter(p[0], p[1], color='blue', s=10) plt.show() # 接收返回值,更新中心节点 new_node1 = center_readjustment(first, node1) new_node2 = center_readjustment(second, node2) new_node3 = center_readjustment(third, node3) return new_node1, new_node2, new_node3 # 主循环:迭代直到中心不再变化 end = False iteration = 0 max_iterations = 100 # 设置最大迭代次数,防止死循环 while not end and iteration < max_iterations: iteration += 1 print(f"===== 第 {iteration} 次迭代 =====") new_node1, new_node2, new_node3 = main0(data, random_node_1, random_node_2, random_node_3) # 判断所有中心是否都未变化 if (new_node1 == tuple(random_node_1) and new_node2 == tuple(random_node_2) and new_node3 == tuple(random_node_3)): end = True print("中心节点不再变化,迭代结束") else: # 更新全局中心变量,用于下一轮迭代 random_node_1, random_node_2, random_node_3 = new_node1, new_node2, new_node3 print(f"最终聚类中心: {random_node_1}, {random_node_2}, {random_node_3}")
关键修改点说明
- 赋值返回值:在
main0中,将center_readjustment的返回值赋值给新变量,再返回给主循环更新全局中心 - 统一数据类型:将中心节点统一为可对比的类型,避免类型不一致导致的对比错误
- 添加迭代终止条件:通过对比新旧中心是否一致,或者设置最大迭代次数,防止无限循环
- 优化距离计算:提前计算每个点到三个中心的距离,减少重复计算,提升效率
内容的提问来源于stack exchange,提问作者스코티쉬폴드
相关产品推荐
相关产品推荐

