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

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))),会断开和原列表的引用,必须通过返回值重新绑定原变量才能完成更新。

解决步骤

  1. 接收函数返回值并更新变量:在main0中调用center_readjustment时,将返回的新中心重新赋值给对应的节点变量
  2. 完善函数返回逻辑:当中心未发生变化时,返回原节点,避免出现None值
  3. 让main0返回更新后的中心:方便while循环获取新的中心,继续下一轮迭代
  4. 添加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,提问作者스코티쉬폴드

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 17:41:06