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

