多项式回归绘图时plt.scatter报IndexError的原因及修复方法
错误成因
- 核心问题:
X_test_imputed是NumPy数组而非Pandas DataFrame/Series。NumPy数组不支持用字符串列名(比如'Level_of_Hemoglobin')索引,只能用整数、切片、布尔数组这类合法索引,所以触发了IndexError。 - 常见诱因:数据预处理阶段(比如用填充缺失值的
SimpleImputer),调用transform()方法后直接返回的是NumPy数组,原DataFrame的列名结构丢失了——比如你可能写了X_test_imputed = imputer.transform(X_test),这就把带列名的DataFrame转成了无列名的数组。
修复方案
有两种简单可行的解决方式:
方式1:把数组转回带列名的DataFrame
如果需要保留列名方便后续操作,将NumPy数组转回DataFrame即可,前提是你知道原特征的列名(可以从训练集的DataFrame里获取):
# 假设X_train是预处理前的训练集DataFrame,保留了列名 import pandas as pd X_test_imputed = pd.DataFrame(X_test_imputed, columns=X_train.columns) # 重新执行散点图代码就正常了 plt.scatter(X_test_imputed['Level_of_Hemoglobin'], y_test, color='blue', label='Actual')
方式2:直接用数组的整数索引
如果不需要DataFrame结构,直接用目标特征在数组中的位置索引(比如Level_of_Hemoglobin是第0列的话):
plt.scatter(X_test_imputed[:, 0], y_test, color='blue', label='Actual')
快速验证方法
先打印type(X_test_imputed)看数据类型:
- 如果输出是
<class 'numpy.ndarray'>,就确认是数组导致的问题; - 如果是
<class 'pandas.core.frame.DataFrame'>,再检查列名拼写(比如大小写、空格是否和原数据集一致)。
内容的提问来源于stack exchange,提问作者Nima_Ebr
相关产品推荐
相关产品推荐

