完成--KNN分类任务
This commit is contained in:
parent
f4b473e4ea
commit
c0a1931321
33
docs/ml.md
33
docs/ml.md
@ -234,4 +234,35 @@
|
||||
特征选择:
|
||||
在特征数量非常庞大的情况下,随机森林能够通过其特征重要性评估功能,自动筛选出对预测有重要贡献的特征。
|
||||
非线性关系:
|
||||
当数据中的特征和目标变量之间存在复杂的非线性关系时,单棵线性模型(如线性回归)可能表现不佳,但随机森林可以通过多棵树的集成来捕捉这些复杂的非线性关系。
|
||||
当数据中的特征和目标变量之间存在复杂的非线性关系时,单棵线性模型(如线性回归)可能表现不佳,但随机森林可以通过多棵树的集成来捕捉这些复杂的非线性关系。
|
||||
|
||||
K近邻(KNN):基于实例的学习方法,通过测量样本之间的距离来进行分类或预测
|
||||
原理:
|
||||
距离度量:
|
||||
KNN通过计算样本之间的距离来判断样本的相似性.
|
||||
欧式距离:
|
||||
曼哈顿距离:
|
||||
分类:
|
||||
对于一个待分类的样本,KNN算法会找到与该样本最相似(距离最近)的K个样本。
|
||||
然后,根据这K个样本中多数的类别来决定待分类样本的类别。投票机制:选择出现次数最多的类别。
|
||||
回归:
|
||||
对于回归任务,KNN则会根据这K个最近邻样本的输出值进行平均,得到预测值。
|
||||
K的选择:
|
||||
K值(邻居的个数)是KNN算法的一个重要超参数。如果K值过小,可能会导致模型过拟合;如果K值过大,可能会导致模型欠拟合。
|
||||
通常K值的选择依赖于交叉验证。
|
||||
优缺点:
|
||||
优点:
|
||||
· 简单易懂:KNN算法实现非常简单,易于理解和解释。
|
||||
· 无模型假设:KNN是非参数化算法,不需要做出关于数据分布的假设。
|
||||
· 适用范围广泛:可以用于分类和回归任务,适用性广。
|
||||
· 良好的性能:对于小数据集,KNN可以达到很好的分类效果。
|
||||
缺点:
|
||||
· 计算量大:KNN的预测阶段需要计算待预测样本与所有训练样本的距离,因此计算复杂度较高,尤其在数据量很大的时候,时间开销大。
|
||||
· 存储需求高:KNN需要存储所有的训练数据,这对于数据量大的情况可能会成为瓶颈。
|
||||
· 对异常值敏感:KNN算法对噪声和离群点非常敏感,因为它只是通过邻居的类别来决定预测值,容易受异常值影响。
|
||||
· 特征维数 curse of dimensionality(维度灾难):随着特征维度的增加,计算距离变得不那么准确,KNN算法的表现会显著下降。高维数据使得样本之间的距离差异缩小,影响模型的判断。
|
||||
适用场景:
|
||||
图像分类:例如手写数字识别、物体分类等任务。
|
||||
推荐系统:通过用户相似性或物品相似性来推荐内容。
|
||||
语音识别:通过语音特征和其他样本之间的相似性来分类不同的语音命令。
|
||||
医疗诊断:例如根据症状数据预测患者是否患有某种疾病。
|
||||
41
knn_classification.py
Normal file
41
knn_classification.py
Normal file
@ -0,0 +1,41 @@
|
||||
# 导入必要的库
|
||||
from sklearn.datasets import load_iris
|
||||
from sklearn.model_selection import train_test_split
|
||||
from sklearn.preprocessing import StandardScaler
|
||||
from sklearn.neighbors import KNeighborsClassifier
|
||||
from sklearn.metrics import classification_report, accuracy_score
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
# 1. 加载数据集
|
||||
iris = load_iris()
|
||||
X = iris.data # 特征数据
|
||||
y = iris.target # 目标标签
|
||||
|
||||
# 2. 划分数据集(80%训练,20%测试)
|
||||
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)
|
||||
|
||||
# 3. 特征标准化(很重要,KNN对特征尺度敏感)
|
||||
scaler = StandardScaler()
|
||||
X_train = scaler.fit_transform(X_train)
|
||||
X_test = scaler.transform(X_test)
|
||||
|
||||
# 4. 构建KNN分类模型(这里K=3)
|
||||
knn = KNeighborsClassifier(n_neighbors=3)
|
||||
|
||||
# 5. 训练模型
|
||||
knn.fit(X_train, y_train)
|
||||
|
||||
# 6. 预测
|
||||
y_pred = knn.predict(X_test)
|
||||
|
||||
# 7. 评估模型性能
|
||||
print("Accuracy:", accuracy_score(y_test, y_pred))
|
||||
print("Classification Report:")
|
||||
print(classification_report(y_test, y_pred))
|
||||
|
||||
# 可视化(仅用于二维数据的情况)
|
||||
plt.scatter(X_test[:, 0], X_test[:, 1], c=y_pred, cmap='viridis', marker='o', label='Predictions')
|
||||
plt.xlabel("Feature 1")
|
||||
plt.ylabel("Feature 2")
|
||||
plt.title("KNN Classification (Predictions)")
|
||||
plt.savefig("./output/knn_c.png")
|
||||
BIN
output/knn_c.png
Normal file
BIN
output/knn_c.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 24 KiB |
Loading…
Reference in New Issue
Block a user