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

修复Python中K近邻算法绘图函数的边绘制错误问题

K近邻邻接图绘制错误排查与修复

问题根源分析

你的代码存在几个关键错误,导致边绘制异常:

  • 全局变量依赖:函数nearest_neighbor_graph内部直接使用全局的xy变量,完全忽略了传入的x、y参数,这是核心问题
  • 数组维度不匹配:生成的xy = np.random.rand(2, 40)是(特征数, 样本数)的形状,而距离计算逻辑需要(样本数, 特征数)(即(40,2)),维度错误会导致距离矩阵计算混乱
  • 散点图参数顺序错误:plt.scatter(x, y, psize)中第三个参数默认是颜色参数c,而非点大小s,会导致点大小设置失效甚至颜色异常
  • 冗余边绘制:使用argsort会保留所有近邻的排序结果,包括点自身和双向连接,导致同一条边被重复绘制两次

修正后的代码

import matplotlib.pyplot as plt
import numpy as np

def nearest_neighbor_graph(x, y, k, pcolor='blue', psize=20, ecolor='black', figsize=(6,6)):
    plt.figure(figsize=figsize)
    # 将传入的x、y组合成(样本数, 2)的坐标数组,替代全局变量
    xy = np.column_stack((x, y))
    # 显式指定s参数设置点大小
    plt.scatter(x, y, s=psize, color=pcolor)
    # 计算所有点对的平方距离
    dist_sq = np.sum((xy[:, np.newaxis, :] - xy[np.newaxis, :, :]) ** 2, axis=-1)
    # 使用argpartition高效获取k个近邻(排除自身)
    nearest_partition = np.argpartition(dist_sq, k+1, axis=1)
    for i in range(xy.shape[0]):
        # 跳过第0个元素(点自身,距离为0),取前k+1个里的后k个有效近邻
        for j in nearest_partition[i, 1:k+1]:
            plt.plot(*zip(xy[i], xy[j]), color=ecolor)

np.random.seed(10)
# 生成(40,2)的坐标数组,再拆分为x和y传入函数
xy = np.random.rand(40, 2)
nearest_neighbor_graph(xy[:,0], xy[:,1], 1, ecolor='red', psize=50)
nearest_neighbor_graph(xy[:,0], xy[:,1], 2, pcolor='green', psize=100, ecolor='black')
plt.show()

关键修改说明

  1. 移除全局变量依赖:在函数内部用np.column_stack((x, y))将传入的x、y组合成正确形状的坐标数组,确保函数使用当前传入的参数
  2. 修正数组维度:生成xy时使用np.random.rand(40, 2),符合(样本数, 特征数)的格式,距离计算逻辑正常工作
  3. 修复散点图参数:将plt.scatter(x, y, psize)改为plt.scatter(x, y, s=psize, color=pcolor),显式指定点大小参数
  4. 优化近邻选择:恢复使用argpartition(比argsort更高效,适合只取前k个元素的场景),并且跳过每个点的第0个近邻(即点自身),避免绘制无意义的自环边,同时减少重复边的绘制

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 14:05:24