Python KNN多标签散点图绘制遇InvalidIndexError,求解决方法
解决Pandas DataFrame索引错误并绘制多标签散点图
问题背景
拥有CSV数据,想要用Matplotlib绘制以纬度(latitude)和经度(longitude)为X、Y轴的多标签散点图,现有代码如下:
import pandas as pd import numpy as np import matplotlib.pyplot as plt from sklearn import datasets from sklearn.model_selection import train_test_split from matplotlib.colors import ListedColormap cmap = ListedColormap(['#FF0000','#00FF00','#0000FF']) df = pd.DataFrame(data) X = df.drop(['Wisata','Link Gambar','Nama','Region'],axis = 1) # 仅保留纬度和经度 y = df['Wisata'] # 标签列 X_train,X_test,y_train,y_test = train_test_split(X,y,test_size = 0.2,random_state=1)
报错情况
执行以下绘制代码时出现错误:
plt.figure() plt.scatter(X[:,0],X[:,1],c=y,cmap=cmap,edgecolors='k',s=20) plt.show()
报错信息:
--------------------------------------------------------------------------- TypeError Traceback (most recent call last) D:\Anaconda\lib\site-packages\pandas\core\indexes\base.py in get_loc(self, key, method, tolerance) 3628 try: -> 3629 return self._engine.get_loc(casted_key) 3630 except KeyError as err: D:\Anaconda\lib\site-packages\pandas\_libs\index.pyx in pandas._libs.index.IndexEngine.get_loc() D:\Anaconda\lib\site-packages\pandas\_libs\index.pyx in pandas._libs.index.IndexEngine.get_loc() TypeError: '(slice(None, None, None), 0)' is an invalid key During handling of the above exception, another exception occurred: InvalidIndexError Traceback (most recent call last) ~\AppData\Local\Temp\ipykernel_12876\2344342603.py in <module> 1 plt.figure() ----> 2 plt.scatter(X[:,0],X[:,1],c=y,cmap=cmap,edgecolors='k',s=20) 3 plt.show() D:\Anaconda\lib\site-packages\pandas\core\frame.py in __getitem__(self, key) 3503 if self.columns.nlevels > 1: 3504 return self._getitem_multilevel(key) -> 3505 indexer = self.columns.get_loc(key) 3506 if is_integer(indexer): 3507 indexer = [indexer] D:\Anaconda\lib\site-packages\pandas\core\indexes\base.py in get_loc(self, key, method, tolerance) 3634 # InvalidIndexError. Otherwise we fall through and re-raise 3635 # the TypeError. -> 3636 self._check_indexing_error(key) 3637 raise 3638 D:\Anaconda\lib\site-packages\pandas\core\indexes\base.py in _check_indexing_error(self, key) 5649 # if key is not a scalar, directly raise an error (the code below 5650 # would convert to numpy arrays and raise later any way) - GH29926 -> 5651 raise InvalidIndexError(key) 5652 5653 @cache_readonly InvalidIndexError: (slice(None, None, None), 0) <Figure size 640x480 with 0 Axes>
错误原因
X是Pandas的DataFrame类型,不是numpy数组,不能使用[:,0]这种numpy数组的二维切片语法,这会触发Pandas的索引解析错误。
解决方法
方法1:将DataFrame转为numpy数组
把X转换成numpy数组后,即可使用切片索引:
plt.figure() # 转换为numpy数组 X_np = X.to_numpy() plt.scatter(X_np[:,0], X_np[:,1], c=y, cmap=cmap, edgecolors='k', s=20) plt.xlabel('经度') plt.ylabel('纬度') plt.title('多标签散点图') plt.show()
方法2:直接通过列名索引
更直观的方式是直接用列名获取经纬度数据,无需依赖列的位置(需根据实际CSV列名调整):
plt.figure() # 假设CSV中经度列名为longitude,纬度列名为latitude plt.scatter(X['longitude'], X['latitude'], c=y, cmap=cmap, edgecolors='k', s=20) plt.xlabel('经度') plt.ylabel('纬度') plt.title('多标签散点图') plt.show()
补充说明
- 如果
y是字符串类型的标签,Matplotlib的scatter函数要求c参数为数值类型,需要先将标签转为数值编码:
# 将字符串标签转为数值 y_num = pd.factorize(y)[0] plt.scatter(X['longitude'], X['latitude'], c=y_num, cmap=cmap, edgecolors='k', s=20)
- 若要绘制训练集/测试集的散点图,只需将
X替换为X_train/X_test,y替换为y_train/y_test即可。
内容的提问来源于stack exchange,提问作者Jessen Jie
相关产品推荐
相关产品推荐

