k-th nearest neighbor (k-NN)
k-NN 是監督式學習 (Supervised learning) 的一種,名稱非常簡明扼要,就是尋找「K 個最相近的鄰居」。
這個演算法在實作時,會找到附近 K 個最近的點,根據鄰居的類別來判斷自己要歸在哪一類。雖然它是監督式學習,但其實並不需要訓練模型參數,而是將所有訓練資料儲存起來進行即時對比。
圖:KnnClassification.svg,Antti Ajanki (AnAj),CC BY-SA 3.0。
我們可以藉由調整 K 的數值來增加演算法的 Noise Margin。然而,此演算法存在著儲存空間需求大(空間複雜度高)的問題,且容易受到數據不平衡的影響。
在實作上,核心在於計算點與點之間的距離。我使用了 Scipy 的函數來實作,為了方便觀察,先取 K=1,並將結果與 sklearn 的 KNN 進行比較。
實作思路是利用 for 迴圈計算每個測試資料與所有訓練資料的距離,並取最近者的類別作為預測結果。
完整的程式如下:
from scipy.spatial import distance
from sklearn import datasets
from sklearn.cross_validation import train_test_split
from sklearn.neighbors import KNeighborsClassifier
from sklearn.metrics import accuracy_score
iris = datasets.load_iris()
X = iris.data
y_ = iris.target
X_train, X_test, y_train, y_test = train_test_split(X,y_,test_size = 0.5)
class easyknn():
def fit(self, X_train, y_train):
self.X_train = X_train
self.y_train = y_train
def predict(self, X_test):
predictions = []
for i in X_test:
label = self.closest(i)
predictions.append(label)
return predictions
def closest(self, row):
min_dist = distance.euclidean(row,self.X_train[0])
min_index = 0
for i in range(1,len(self.X_train)):
dist = distance.euclidean(row,self.X_train[i])
if dist < min_dist:
min_dist = dist
min_index = i
return self.y_train[min_index]
# sklearn or scipy
# knn = KNeighborsClassifier()
knn = easyknn()
knn.fit(X_train,y_train)
predictions = knn.predict(X_test)
print(accuracy_score(y_test,predictions))
準確率比較:
sklearn knn: 0.9733- 手刻
knn: 0.9467
這支程式在做的事
上面讀進來的是 scikit-learn 內建的 iris,150 筆鳶尾花、每筆四個測量值,test_size = 0.5 把它對半切開,訓練集 75 筆,測試集也是 75 筆。手刻的那個類別只有三個方法,而且 fit() 沒有在訓練:它把訓練資料原封不動存起來就結束了,真正的工作全部發生在預測的時候,由 closest() 逐筆比對距離。
兩個準確率不是同一次執行的成績。GitHub 上那支檔案裡,# knn = KNeighborsClassifier() 是註解掉的,底下 knn = easyknn() 才是活的,跑一次只會印出一個數字。train_test_split 又沒有給 random_state,每執行一次就重切一次,所以 0.9733 和 0.9467 各自來自不同的切分、不同的測試集。
兩邊的 K 也不一樣。手刻的版本只看最近的那一個鄰居,KNeighborsClassifier() 沒帶參數就是問五個鄰居再投票。不過在 iris 上,K 換來的差距不大。固定 test_size=0.5、跑 200 組隨機切分,1NN 的平均準確率是 0.9525、5NN 是 0.9581,差 0.55 個百分點;逐組去比,5NN 比較高的佔 44.0%、1NN 比較高的佔 25.5%,其餘打平。(這 200 組裡,手刻的那個類別跟 KNeighborsClassifier(n_neighbors=1) 有 199 組給出一模一樣的預測。)同一個 1NN 光是換切法,最低 0.8800、最高 1.0000,0.9733 和 0.9467 都在這個範圍裡面。
這支程式現在直接跑會停在第三行。sklearn.cross_validation 這個模組已經不在了,scikit-learn 1.8.0 會回 ModuleNotFoundError,train_test_split 現在在 sklearn.model_selection 底下,那一行換掉就能跑。
暴力搜尋
closest() 每被呼叫一次,就把訓練集從頭到尾走一遍。75 筆測試資料,每一筆都跟 75 筆訓練資料各算一次距離,跑完一次 predict(),distance.euclidean 一共被呼叫 5,625 次,其中真正留下來的只有 75 個結果,其餘的算完就丟掉。這個量在筆電上跑起來不會慢。
scikit-learn 的文件把這種做法叫 brute force,也就是把資料集裡每一對點的距離都算出來。文件上寫這種做法在小樣本上很有競爭力,樣本數一大就很快變得不可行,而 iris 的 75 筆屬於前者。
k-NN 沒有訓練出來的參數可以查,資料本身就是模型,每來一筆新的查詢,手上有的就是全部的訓練資料。訓練集從 75 筆變成 75 萬筆,照 closest() 這個寫法,一次預測就是 75 萬次距離計算。空間複雜度高也是同一個來源:資料得全部留著,才有東西可以比。
樹狀索引
scikit-learn 除了 brute force,還有兩種樹狀結構:KD tree 和 ball tree。它們做的是同一件事,把樣本之間的距離資訊先聚合編碼起來,用來減少實際要算的距離次數。點 A 離點 B 很遠、點 B 離點 C 很近,那 A 跟 C 就不必真的算一次距離。
這就是先算好放著、下次再用的部分,而且算出來的鄰居跟暴力搜尋一樣。拿 75 萬筆、四個維度的隨機資料,一次送 100 筆查詢、每筆取 5 個鄰居實測(scikit-learn 1.8.0、NumPy 2.4.2、Python 3.14.3):algorithm='brute' 沒有索引要建,100 次查詢 0.061 秒;algorithm='kd_tree' 建樹花 1.77 秒,之後 100 次查詢 0.0039 秒。兩邊回傳的鄰居索引完全相同,換成 ball tree 也相同。(秒數跟機器有關,這裡看的是同一台機器上的比例。)
KNeighborsClassifier 預設是讓它自己挑用哪一種。文件上的簽名是 KNeighborsClassifier(n_neighbors=5, *, weights='uniform', algorithm='auto', ...),'auto' 的說明是依照傳進 fit() 的資料決定最合適的演算法。把原文那組訓練資料(75 筆、四個維度)餵進去,fit 完之後看 _fit_method,得到的是 kd_tree。
向量檢索
同樣是找最近的鄰居,現在的語意搜尋和 RAG 也在做這件事。文字會先換算成一串座標(這段換算寫在 Embedding 是什麼?AI 怎麼知道兩句話意思一樣),查詢也一樣,接下來就是挑出離它最近的那幾筆,跟 closest() 想做的是同一件事。差別在座標有多長:iris 一筆是四個測量值,sentence-transformers/all-mpnet-base-v2 這種句子模型,一句話換算出來是 768 個數字。
維度一上去,樹狀索引就不划算了。scikit-learn 文件上寫著,維度 D 大起來之後,成本會逼近 O[DN],而樹狀結構本身的額外開銷會讓查詢比暴力搜尋還慢。把筆數固定在 2 萬筆、只換維度,四個維度的時候,kd_tree 的 100 次查詢是 0.0021 秒、brute 是 0.0030 秒;換成 768 個維度,kd_tree 變成 2.90 秒、brute 是 0.036 秒。兩邊回傳的鄰居還是一樣,快慢對調了。同一批 768 維的資料交給 KNeighborsClassifier(),_fit_method 挑的是 brute。
到了這個維度,精確搜尋能做的就剩下把全部比過一遍,一次查詢的時間跟筆數乘上維度成正比。向量資料庫管理的是大批 embedding 向量,AI 應用成長得快,需要存下來、建索引的向量數量也跟著增加,那一遍就掃不完了。2016 年的 HNSW 論文提的是近似最近鄰搜尋(approximate nearest neighbor search),它把向量接成一組分層的鄰近圖,搜尋從上層開始往下找,成本以對數的方式成長。
pgvector 不下索引的時候,Postgres 做的是全部比過一遍的精確搜尋,recall 是滿的,也就是那個 for 迴圈的資料庫版本;下了索引才換成近似搜尋,速度是拿 recall 換來的。加了近似索引之後,同一句查詢、同一批資料,撈回來的東西可以不一樣。


