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

如何正确访问Python函数内创建的costs数组并获取有效值?

解决K-Means聚类代码中main函数costs数组无法正确输出的问题

问题描述

编写了一段实现K-Means聚类的Python代码,尝试访问main()函数内创建的costs数组时,直接执行print(costs)无法得到正确数值,需要排查并修复问题。

问题分析

代码存在几个关键问题导致costs数组输出异常:

  • 未定义距离计算函数d:plot_k_means函数中调用了d(M[k], X[n])但未实现该函数,会直接抛出运行错误,后续代码无法执行。
  • 变量名冲突:plot_k_means内部循环使用了k作为循环变量,覆盖了函数参数K,导致聚类数量逻辑出错。
  • costs数组初始化不合理:用np.empty(10)创建数组并将costs[0]设为None,会导致数组类型变为object,混合数值与非数值类型,输出和绘图都会异常。
  • 重复计算cost:main函数中已经从plot_k_means返回了final_cost,却再次调用cost(X, R, M)计算,属于冗余操作,且可能因中间变量问题导致结果不一致。

修复步骤

  • 补充距离计算函数:实现计算欧氏距离的d函数,放在cost函数之前。
  • 修正变量名冲突:将plot_k_means内部循环的k改为kk,避免覆盖参数K。
  • 优化costs数组初始化:将costs初始化为长度为10的数组,costs[0]设为0或直接从索引1开始赋值,避免混入None。
  • 移除冗余计算:直接使用plot_k_means返回的final_cost赋值给costs[k]。

修正后的完整代码

from __future__ import print_function, division
from future.utils import iteritems
from builtins import range, input
# Note: you may need to update your version of future
# sudo pip install -U future

import numpy as np
import matplotlib.pyplot as plt

# 补充缺失的距离计算函数
def d(a, b):
    return np.linalg.norm(a - b)

def cost(X, R, M):
    cost = 0
    for k in range(len(M)):
        # method 2
        diff = X - M[k]
        sq_distances = (diff * diff).sum(axis=1)
        cost += (R[:,k] * sq_distances).sum()
    return cost

def plot_k_means(X, K, max_iter=20, beta=3.0, show_plots=False):
    N, D = X.shape
    exponents = np.empty((N, K))

    # initialize M to random
    initial_centers = np.random.choice(N, K, replace=False)
    M = X[initial_centers]

    costs = []
    # 修改循环变量名,避免覆盖K
    for i in range(max_iter):
        # step 1: determine assignments / resposibilities
        for kk in range(K):
            for n in range(N):
                exponents[n,kk] = np.exp(-beta*d(M[kk], X[n]))
        R = exponents / exponents.sum(axis=1, keepdims=True)

        # step 2: recalculate means
        M = R.T.dot(X) / R.sum(axis=0, keepdims=True).T

        c = cost(X, R, M)
        costs.append(c)
        if i > 0:
            if np.abs(costs[-1] - costs[-2]) < 1e-5:
                break

        if len(costs) > 1:
            if costs[-1] > costs[-2]:
                pass

    if show_plots:
        plt.plot(costs)
        plt.title("Costs")
        plt.show()

        random_colors = np.random.random((K, 3))
        colors = R.dot(random_colors)
        plt.scatter(X[:,0], X[:,1], c=colors)
        plt.show()

    final_cost = costs[-1]
    return M, R, final_cost


def get_simple_data():
    # assume 3 means
    D = 2 # so we can visualize it more easily
    s = 4 # separation so we can control how far apart the means are
    mu1 = np.array([0, 0])
    mu2 = np.array([s, s])
    mu3 = np.array([0, s])

    N = 900 # number of samples
    X = np.zeros((N, D))
    X[:300, :] = np.random.randn(300, D) + mu1
    X[300:600, :] = np.random.randn(300, D) + mu2
    X[600:, :] = np.random.randn(300, D) + mu3
    return X


def main():
    X = get_simple_data()

    plt.scatter(X[:,0], X[:,1])
    plt.show()

    # 初始化costs数组,避免混入None
    costs = np.zeros(10)
    # K=1时单独处理
    M, R, final_cost = plot_k_means(X, 1, show_plots=False)
    costs[1] = final_cost
    
    for k in range(2, 10):
        M, R, final_cost = plot_k_means(X, k, show_plots=False)
        costs[k] = final_cost
    
    plt.plot(range(10), costs)
    plt.title("Cost vs K")
    plt.xlabel("K")
    plt.ylabel("Cost")
    plt.show()

    print("Costs数组输出:")
    print(costs)
    
if __name__ == '__main__':
  main()

验证效果

运行修正后的代码,main函数中的costs数组会正确存储不同K值对应的聚类成本,执行print(costs)会输出正常的数值数组,同时"Cost vs K"的折线图也能正常显示。

内容的提问来源于stack exchange,提问作者Nishant

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 16:13:36