25,KNN算法实现鸢尾花分类,分类结果和准确率随着k值变化过程如下图所示。

from sklearn import datasets
iris = datasets.load_iris()
X, y = iris.data, iris.target

这数据集长得就像个乖巧的表格,4列特征分别是花萼长宽、花瓣长宽,目标值对应三种鸢尾花。顺手把数据拆成训练集和测试集:

from sklearn.model_selection import train_test_split
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42)

这里有个小细节,random_state设成42纯粹是程序员的恶趣味,保证每次切分结果一致。接下来搞个骚操作——数据标准化。KNN这算法对尺度敏感得像踩了尾巴的猫,不处理的话距离计算要出幺蛾子:

from sklearn.preprocessing import StandardScaler
scaler = StandardScaler().fit(X_train)
X_train = scaler.transform(X_train)
X_test = scaler.transform(X_test)

标准化前后数据分布变化肉眼可见,原本花瓣长度动不动就五六厘米,现在都缩放到-1到1之间。上主菜KNN模型,先整个k=5试试水:

from sklearn.neighbors import KNeighborsClassifier
knn = KNeighborsClassifier(n_neighbors=5)
knn.fit(X_train, y_train)
print(f"准确率:{knn.score(X_test, y_test):.2%}")  # 输出:准确率:97.78%

嚯,这准确率看着挺唬人,但别急着开香槟。咱们把k值从1到30撸一遍看看效果:

import matplotlib.pyplot as plt
accuracy = []
for k in range(1, 31):
    knn = KNeighborsClassifier(n_neighbors=k).fit(X_train, y_train)
    accuracy.append(knn.score(X_test, y_test))

plt.plot(range(1,31), accuracy, marker='o')
plt.xlabel('k值'), plt.ylabel('准确率')
plt.show()

跑出来的曲线跟过山车似的——k=1时准确率直接掉到92%,k=7冲到100%,接着又慢慢下滑。这说明选k值跟找对象似的,太小容易过拟合(跟邻居太亲密),太大又容易欠拟合(跟远房亲戚搞暧昧)。实际项目中建议用交叉验证来找最佳k值,不过咱们这个简单案例直接看图说话更直观。最后提醒下,KNN在特征工程到位的小数据集上能打,但数据量上百万级别的话...建议换个算法保命。

25,KNN算法实现鸢尾花分类,分类结果和准确率随着k值变化过程如下图所示。

更多推荐