diff --git a/docs/ml.md b/docs/ml.md index 97bc22d..81938f0 100644 --- a/docs/ml.md +++ b/docs/ml.md @@ -234,4 +234,35 @@ 特征选择: 在特征数量非常庞大的情况下,随机森林能够通过其特征重要性评估功能,自动筛选出对预测有重要贡献的特征。 非线性关系: - 当数据中的特征和目标变量之间存在复杂的非线性关系时,单棵线性模型(如线性回归)可能表现不佳,但随机森林可以通过多棵树的集成来捕捉这些复杂的非线性关系。 \ No newline at end of file + 当数据中的特征和目标变量之间存在复杂的非线性关系时,单棵线性模型(如线性回归)可能表现不佳,但随机森林可以通过多棵树的集成来捕捉这些复杂的非线性关系。 + +K近邻(KNN):基于实例的学习方法,通过测量样本之间的距离来进行分类或预测 + 原理: + 距离度量: + KNN通过计算样本之间的距离来判断样本的相似性. + 欧式距离: + 曼哈顿距离: + 分类: + 对于一个待分类的样本,KNN算法会找到与该样本最相似(距离最近)的K个样本。 + 然后,根据这K个样本中多数的类别来决定待分类样本的类别。投票机制:选择出现次数最多的类别。 + 回归: + 对于回归任务,KNN则会根据这K个最近邻样本的输出值进行平均,得到预测值。 + K的选择: + K值(邻居的个数)是KNN算法的一个重要超参数。如果K值过小,可能会导致模型过拟合;如果K值过大,可能会导致模型欠拟合。 + 通常K值的选择依赖于交叉验证。 + 优缺点: + 优点: + · 简单易懂:KNN算法实现非常简单,易于理解和解释。 + · 无模型假设:KNN是非参数化算法,不需要做出关于数据分布的假设。 + · 适用范围广泛:可以用于分类和回归任务,适用性广。 + · 良好的性能:对于小数据集,KNN可以达到很好的分类效果。 + 缺点: + · 计算量大:KNN的预测阶段需要计算待预测样本与所有训练样本的距离,因此计算复杂度较高,尤其在数据量很大的时候,时间开销大。 + · 存储需求高:KNN需要存储所有的训练数据,这对于数据量大的情况可能会成为瓶颈。 + · 对异常值敏感:KNN算法对噪声和离群点非常敏感,因为它只是通过邻居的类别来决定预测值,容易受异常值影响。 + · 特征维数 curse of dimensionality(维度灾难):随着特征维度的增加,计算距离变得不那么准确,KNN算法的表现会显著下降。高维数据使得样本之间的距离差异缩小,影响模型的判断。 + 适用场景: + 图像分类:例如手写数字识别、物体分类等任务。 + 推荐系统:通过用户相似性或物品相似性来推荐内容。 + 语音识别:通过语音特征和其他样本之间的相似性来分类不同的语音命令。 + 医疗诊断:例如根据症状数据预测患者是否患有某种疾病。 \ No newline at end of file diff --git a/knn_classification.py b/knn_classification.py new file mode 100644 index 0000000..b77a80c --- /dev/null +++ b/knn_classification.py @@ -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") diff --git a/output/knn_c.png b/output/knn_c.png new file mode 100644 index 0000000..d83b8f6 Binary files /dev/null and b/output/knn_c.png differ