Sklearn KNN模型fit函数报错:n_neighbors需传入整数而非浮点数
解决sklearn KNeighborsClassifier的n_neighbors浮点数错误
问题描述
使用sklearn的KNN模型进行二分类(类别Y取值为1或2),特征为X1、X2、X3,运行模型训练代码时触发错误:
"n_neighbors does not take <class 'float'> value, enter integer value"
即使将数据集改为全整数类型,错误依然存在,代码如下:
#Importing necessary libraries import pandas as pd import numpy as np #Imports for KNN models from sklearn.model_selection import train_test_split from sklearn.neighbors import KNeighborsClassifier #Imports for testing the model from sklearn.metrics import confusion_matrix from sklearn.metrics import f1_score from sklearn.metrics import accuracy_score #Import the data file data = pd.read_csv("/content/drive/MyDrive/Python/Colab Notebooks/Onlyinttest.csv") #Split data X = data.loc[:,['X1','X2','X3']] Y = data.loc[:,'Y'] X_train, X_test, Y_train, Y_test = train_test_split(X,Y, random_state=0, test_size=0.2) #Determine k by using sqrt import math k = math.sqrt(len(Y_test)) print(k) #Make k uneven k = k-1 #KNN Model classifer = KNeighborsClassifier(n_neighbors=k, p=2,metric='euclidean') classifer.fit(X_train,Y_train)
问题原因
错误和数据集的数值类型完全无关,核心问题出在k的计算逻辑:
math.sqrt()返回的结果是浮点数类型,即使执行k = k-1,最终结果依然是浮点数KNeighborsClassifier的n_neighbors参数要求必须是正整数,不接受任何浮点数输入
解决方案
将计算得到的k转换为整数,同时保证k为正奇数(符合你想要避免投票平局的需求),可以采用以下两种方式:
方式1:强制转为整数
直接用int()截断浮点数的小数部分,快速得到整数k:
#Determine k by using sqrt import math k = math.sqrt(len(Y_test)) print(k) #Make k uneven and convert to integer k = int(k - 1) # 额外添加判断,避免极端情况(比如测试集样本数过少导致k为0或负数) if k <= 0: k = 1
方式2:向下取整(语义更清晰)
用math.floor()明确对数值向下取整,逻辑更直观:
#Determine k by using sqrt import math k = math.sqrt(len(Y_test)) print(k) #Make k uneven and floor to integer k = math.floor(k - 1) # 确保k为正整数 if k <= 0: k = 1
修改后的完整代码
#Importing necessary libraries import pandas as pd import numpy as np #Imports for KNN models from sklearn.model_selection import train_test_split from sklearn.neighbors import KNeighborsClassifier #Imports for testing the model from sklearn.metrics import confusion_matrix from sklearn.metrics import f1_score from sklearn.metrics import accuracy_score #Import the data file data = pd.read_csv("/content/drive/MyDrive/Python/Colab Notebooks/Onlyinttest.csv") #Split data X = data.loc[:,['X1','X2','X3']] Y = data.loc[:,'Y'] X_train, X_test, Y_train, Y_test = train_test_split(X,Y, random_state=0, test_size=0.2) #Determine k by using sqrt import math k = math.sqrt(len(Y_test)) print(k) #Make k uneven and convert to integer k = int(k - 1) # 确保k是正整数,防止测试集样本数过少导致k无效 if k <= 0: k = 1 #KNN Model classifer = KNeighborsClassifier(n_neighbors=k, p=2,metric='euclidean') classifer.fit(X_train,Y_train)
内容的提问来源于stack exchange,提问作者Fabian
相关产品推荐
相关产品推荐

