
1. 項目概述K最近鄰K-Nearest Neighbors簡稱KNN算法是機器學習領域最基礎也最經典的算法之一。作為監督學習中的分類算法KNN以其簡單直觀、無需訓練過程的特性成為無數機器學習初學者的第一個實戰案例。今天我們就用兩個經典數據集——鳶尾花分類和手寫數字識別帶大家真正動手實現KNN算法。提示本文假設讀者已經了解Python基礎語法和機器學習基本概念。如果尚未安裝Python環境建議先配置AnacondaJupyter Notebook的開發環境。KNN算法的核心思想可以用一句俗語概括近朱者赤近墨者黑。當我們需要對一個新樣本進行分類時只需找到訓練集中與它最接近的K個鄰居根據這些鄰居的類別投票決定新樣本的類別。這種懶惰學習Lazy Learning的特性使得KNN算法實現簡單但計算復雜度會隨著數據規模增大而顯著增加。2. 核心原理與數學基礎2.1 KNN算法三要素KNN算法的實現主要依賴三個關鍵要素距離度量常用歐氏距離Euclidean Distance對于二維空間中的兩點(x1,y1)和(x2,y2)其距離計算公式為distance sqrt((x2-x1)^2 (y2-y1)^2)對于更高維度的數據公式可自然擴展。在文本分類等場景中也常使用曼哈頓距離或余弦相似度。K值選擇K是算法中的超參數表示考慮最近鄰的數量。K值過小容易過擬合對噪聲敏感K值過大會使分類邊界模糊。通常通過交叉驗證確定最佳K值。分類決策規則一般采用多數表決法即K個鄰居中出現次數最多的類別作為預測結果。也可以根據距離加權投票近距離的鄰居擁有更大權重。2.2 算法流程分解一個完整的KNN分類流程包括以下步驟數據準備加載數據集劃分訓練集和測試集特征標準化對數據進行歸一化處理重要距離計算測試樣本與所有訓練樣本的距離排序找鄰居按距離升序排列選取前K個投票決策統計K個鄰居的類別分布結果輸出將得票最多的類別作為預測結果性能評估計算準確率等指標3. 鳶尾花分類實戰3.1 數據集介紹鳶尾花數據集Iris是機器學習領域的Hello World包含150個樣本每個樣本有4個特征花萼長度sepal length花萼寬度sepal width花瓣長度petal length花瓣寬度petal width目標變量是鳶尾花的三個品種Iris SetosaIris VersicolourIris Virginica3.2 代碼實現步驟# 導入必要庫 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 accuracy_score # 加載數據 iris load_iris() X, y iris.data, iris.target # 數據分割 X_train, X_test, y_train, y_test train_test_split(X, y, test_size0.3, random_state42) # 特征標準化 scaler StandardScaler() X_train scaler.fit_transform(X_train) X_test scaler.transform(X_test) # 注意使用訓練集的參數轉換測試集 # 創建KNN模型 knn KNeighborsClassifier(n_neighbors5) # 訓練模型 knn.fit(X_train, y_train) # 預測測試集 y_pred knn.predict(X_test) # 評估準確率 accuracy accuracy_score(y_test, y_pred) print(f模型準確率: {accuracy:.2f})3.3 關鍵問題與調優特征標準化的重要性不同特征的量綱差異會導致距離計算偏向大數值特征標準化使所有特征具有相同的重要性常用方法Z-score標準化StandardScaler或MinMax縮放K值選擇實驗 通過交叉驗證尋找最佳K值from sklearn.model_selection import cross_val_score k_range range(1, 31) k_scores [] for k in k_range: knn KNeighborsClassifier(n_neighborsk) scores cross_val_score(knn, X_train, y_train, cv5, scoringaccuracy) k_scores.append(scores.mean()) # 繪制K值與準確率關系圖 import matplotlib.pyplot as plt plt.plot(k_range, k_scores) plt.xlabel(K值) plt.ylabel(交叉驗證準確率) plt.show()可視化決策邊界 由于鳶尾花有4個特征我們可以選擇兩個主要特征進行降維可視化from matplotlib.colors import ListedColormap # 選擇前兩個特征 X_2d X_train[:, :2] # 創建網格點 h 0.02 # 步長 x_min, x_max X_2d[:, 0].min() - 1, X_2d[:, 0].max() 1 y_min, y_max X_2d[:, 1].min() - 1, X_2d[:, 1].max() 1 xx, yy np.meshgrid(np.arange(x_min, x_max, h), np.arange(y_min, y_max, h)) # 訓練僅使用兩個特征的KNN knn_2d KNeighborsClassifier(n_neighbors5) knn_2d.fit(X_2d, y_train) # 預測網格點 Z knn_2d.predict(np.c_[xx.ravel(), yy.ravel()]) Z Z.reshape(xx.shape) # 繪制決策邊界 cmap_light ListedColormap([#FFAAAA, #AAFFAA, #AAAAFF]) plt.contourf(xx, yy, Z, cmapcmap_light, alpha0.8) # 繪制訓練點 plt.scatter(X_2d[:, 0], X_2d[:, 1], cy_train, edgecolork, s20) plt.xlabel(標準化花萼長度) plt.ylabel(標準化花萼寬度) plt.title(KNN決策邊界(K5)) plt.show()4. 手寫數字識別實戰4.1 MNIST數據集簡介MNIST數據集包含70,000張手寫數字(0-9)的28x28像素灰度圖像是計算機視覺領域的經典入門數據集。每個像素點的值范圍是0-255表示灰度強度。4.2 完整實現代碼from sklearn.datasets import fetch_openml from sklearn.model_selection import train_test_split from sklearn.preprocessing import MinMaxScaler from sklearn.neighbors import KNeighborsClassifier from sklearn.metrics import accuracy_score, confusion_matrix import matplotlib.pyplot as plt import numpy as np # 加載數據 mnist fetch_openml(mnist_784, version1) X, y mnist.data, mnist.target # 數據預覽 plt.figure(figsize(10,5)) for i in range(20): plt.subplot(2,10,i1) plt.imshow(X.iloc[i].values.reshape(28,28), cmapgray) plt.title(fLabel: {y[i]}) plt.axis(off) plt.show() # 數據分割 - 使用前10000個樣本加速演示 X X[:10000] y y[:10000] X_train, X_test, y_train, y_test train_test_split(X, y, test_size0.2, random_state42) # 特征縮放 scaler MinMaxScaler() X_train scaler.fit_transform(X_train) X_test scaler.transform(X_test) # 創建KNN模型 knn KNeighborsClassifier(n_neighbors5, n_jobs-1) # n_jobs-1使用所有CPU核心 # 訓練模型 knn.fit(X_train, y_train) # 預測測試集 y_pred knn.predict(X_test) # 評估模型 accuracy accuracy_score(y_test, y_pred) print(f模型準確率: {accuracy:.2f}) # 混淆矩陣 cm confusion_matrix(y_test, y_pred) plt.figure(figsize(10,8)) plt.imshow(cm, cmapBlues) plt.colorbar() plt.xlabel(預測標簽) plt.ylabel(真實標簽) plt.title(混淆矩陣) plt.show()4.3 性能優化技巧降維處理原始784維特征(28x28)計算距離耗時嚴重使用PCA降維保留95%方差from sklearn.decomposition import PCA pca PCA(n_components0.95) # 保留95%方差 X_train_pca pca.fit_transform(X_train) X_test_pca pca.transform(X_test) print(f原始維度: {X_train.shape[1]}) print(f降維后維度: {X_train_pca.shape[1]})近似最近鄰算法當數據量很大時使用近似算法加速from sklearn.neighbors import NearestNeighbors # 使用BallTree算法 nn NearestNeighbors(n_neighbors5, algorithmball_tree) nn.fit(X_train_pca) # 查詢測試樣本的鄰居 distances, indices nn.kneighbors(X_test_pca) # 手動實現投票 from collections import Counter y_pred [] for idx in indices: votes y_train.iloc[idx] pred Counter(votes).most_common(1)[0][0] y_pred.append(pred) accuracy accuracy_score(y_test, y_pred) print(f近似KNN準確率: {accuracy:.2f})距離加權投票近距離的鄰居應該有更大的投票權重# 自定義權重函數 def inverse_distance(weights): return 1 / (weights 1e-6) # 避免除以零 weighted_knn KNeighborsClassifier(n_neighbors5, weightsinverse_distance) weighted_knn.fit(X_train_pca, y_train) y_pred_weighted weighted_knn.predict(X_test_pca) print(f加權KNN準確率: {accuracy_score(y_test, y_pred_weighted):.2f})5. 常見問題與解決方案5.1 計算效率問題問題表現數據集較大時預測速度很慢內存消耗高解決方案使用KD樹或BallTree數據結構加速鄰居搜索knn KNeighborsClassifier(algorithmkd_tree) # 或ball_tree對大數據集使用近似最近鄰算法如LSH降維處理減少特征數量考慮使用GPU加速庫如cuML5.2 類別不平衡問題問題表現某些類別樣本數遠多于其他類別多數表決法會偏向多數類解決方案使用距離加權投票對多數類進行欠采樣或對少數類過采樣調整類別權重參數knn KNeighborsClassifier(weightsdistance)5.3 高維災難問題問題表現特征維度很高時所有樣本的距離趨于相似分類性能下降解決方案特征選擇去除無關特征使用PCA等降維方法考慮使用更適合高維數據的算法如SVM5.4 參數調優技巧K值選擇從Ksqrt(N)開始嘗試N為訓練樣本數使用網格搜索交叉驗證from sklearn.model_selection import GridSearchCV param_grid {n_neighbors: range(1, 20)} grid GridSearchCV(KNeighborsClassifier(), param_grid, cv5) grid.fit(X_train, y_train) print(f最佳K值: {grid.best_params_[n_neighbors]})距離度量選擇歐氏距離默認適用于連續特征曼哈頓距離對異常值更魯棒余弦相似度適用于文本數據6. 項目擴展與進階方向6.1 自定義距離度量在某些特定場景可能需要自定義距離函數。例如對于圖像數據可以嘗試以下距離# 自定義距離函數示例直方圖相交距離 def histogram_intersection(a, b): return np.minimum(a, b).sum() # 使用自定義距離的KNN custom_knn KNeighborsClassifier(n_neighbors5, metrichistogram_intersection) custom_knn.fit(X_train, y_train)6.2 多輸出KNNKNN也可以用于多輸出任務每個樣本有多個目標變量from sklearn.datasets import make_regression from sklearn.neighbors import KNeighborsRegressor # 生成多輸出回歸數據 X, y make_regression(n_samples1000, n_features10, n_targets2) # 多輸出KNN回歸 knn_reg KNeighborsRegressor(n_neighbors5) knn_reg.fit(X, y)6.3 在線學習實現標準KNN不支持增量學習但可以通過以下方式實現class OnlineKNN: def __init__(self, k5): self.k k self.X None self.y None def partial_fit(self, X_new, y_new): if self.X is None: self.X X_new self.y y_new else: self.X np.vstack([self.X, X_new]) self.y np.concatenate([self.y, y_new]) def predict(self, X_test): from sklearn.neighbors import NearestNeighbors nn NearestNeighbors(n_neighborsself.k) nn.fit(self.X) distances, indices nn.kneighbors(X_test) predictions [] for idx in indices: votes self.y[idx] pred Counter(votes).most_common(1)[0][0] predictions.append(pred) return np.array(predictions)6.4 與其他算法結合KNN可以與其他算法結合構建更強大的模型KNN特征工程使用KNN提取樣本鄰居的統計特征作為新特征例如計算每個樣本的K個最近鄰的類別分布集成學習方法構建多個不同參數的KNN模型進行投票例如使用不同的K值和距離度量from sklearn.ensemble import VotingClassifier knn1 KNeighborsClassifier(n_neighbors5) knn2 KNeighborsClassifier(n_neighbors10, weightsdistance) knn3 KNeighborsClassifier(n_neighbors7, metricmanhattan) ensemble VotingClassifier( estimators[(knn5, knn1), (knn10, knn2), (knn7, knn3)], votinghard) ensemble.fit(X_train, y_train)在實際項目中KNN雖然簡單但在特征工程良好、數據規模適中的情況下往往能取得出人意料的好效果。特別是在需要快速驗證想法或建立基線模型的場景中KNN因其實現簡單、無需復雜調參的優勢仍然是機器學習工具箱中的重要成員。