k-近鄰演算法(K-Nearest Neighbor),k-k-nearest
一、概述
k-近鄰演算法採用測量不同特徵值之間的距離方法進行分類
1、工作原理:
存在一個樣本資料集合,也稱作訓練樣本集,並且樣本集中每個資料都存在標籤,即我們知道樣本集中每一資料與所屬分類的對應關係。輸入沒有標籤的新資料後,將新資料的每個特徵與樣本集中資料對應的特徵進行比較,然後演算法提取樣本集中特徵最相似資料(最近鄰)的分類標籤。最後,選擇k個最相似資料中出現次數最多的分類,作為新資料的分類。
通常k取不大於20的整數,一般為了方便利用少數服從多數的投票法則(Majority-voting),k取質數。
2、舉例分析:電影分類
首先我們從動作片和愛情片中提取出兩個特徵--打鬥和接吻。並對已知類型的6部電影和未知類型電影的兩個特徵進行統計如下:
圖 1 打鬥與接吻特徵統計
這樣我們可以將7部電影抽象為二維座標系中的7個點,將兩個特徵分別抽象為對應點的X座標值和Y座標值,如:
圖 2:抽象後的特徵資料
然後就可以根據抽象得到的資料用散佈圖來表示:
圖 3:電影分類散佈圖
這時我們需要計算不同特徵值之間的距離,即圖3 中黃色點與其它各點之間的距離。這裡我們使用比較常用的歐氏距離公式(Euclidean Distance)
(關於距離的計算,還可以使用其它演算法。)
通過計算我們得到如下資料:
表1:已知電影與未知電影的距離 |
電影名稱 |
電影類型 |
與未知電影的距離 |
california Man |
Romance |
20.5 |
He's Not Really into Dudes |
Romance |
18.7 |
Beautiful Woman |
Romance |
19.2 |
Kevin Longblade |
Action |
115.3 |
Robo Slayer 3000 |
Action |
117.4 |
Amped II |
Action |
118.9 |
若k=3,則我們取距離值最小的3個點。在這3個點中Romance類型有3個,Action類型有0個,所以Romance類型出現頻率最高。因此我們判定未知類電影屬於Romance類型。
3、KNN分類演算法虛擬碼:
對未知類別屬性的資料集中的每個點依次執行以下操作:
(1)計算已知類別資料集中的點與當前點之間的距離;
(2)按照距離遞增次序排序;
(3)選取與當前點距離最小的k個點;
(4)確定前k個點所在類別的出現頻率;
(5)返回前k個點出現頻率最高的類別作為當前點的預測分類。
4、演算法優缺點
優點:
演算法簡單,容易實現;對異常值不敏感。
缺點:
空間複雜度高
需要大量空間儲存所有已知執行個體
計算複雜度高
需要比較所有已知執行個體與要分類的執行個體
二、執行個體:手寫辨識系統
程式運行在python3.6
1 #-*- coding:utf-8 -*- 2 3 from numpy import * 4 import operator 5 from os import listdir 6 7 def classify(inX, dataSet, labels, k): 8 """ 9 :param inX: 樣本資料10 :param dataSet: 已知資料11 :param labels: 已知資料的分類標籤12 :param k:選取的k值13 :return: 返回樣本資料的分類標籤14 """15 dataSetSize = dataSet.shape[0] #擷取矩陣行數16 17 #計算歐氏距離18 diffMat = tile(inX, (dataSetSize, 1)) - dataSet19 sqDiffMat = diffMat**220 sqDistances = sqDiffMat.sum(axis=1)21 distances = sqDistances**0.522 23 sortedDistIndicies = distances.argsort() #對索引進行排序(從小到大)24 classCount={}25 26 #選出距離最小的k個點27 for i in range(k):28 voteIlabel = labels[sortedDistIndicies[i]]29 classCount[voteIlabel] = classCount.get(voteIlabel,0) + 130 31 sortedClassCount = sorted(classCount.items(),32 key=operator.itemgetter(1),reverse=True)33 34 return sortedClassCount[0][0]35 36 37 def img2vector(filename):38 """39 :param filename: 輸入檔案名稱,用於擷取文本資料40 :return: 將文本資料以數組形式返回41 """42 returnVect = zeros((1, 1024))43 fr = open(filename)44 for i in range(32):45 lineStr = fr.readline()46 for j in range(32):47 returnVect[0,32*i+j] = int(lineStr[j])48 return returnVect49 50 def handwritingClassTest():51 hwLabels = []52 trainingFileList = listdir('trainingDigits') #擷取目錄下的內容(檔案名稱)53 m = len(trainingFileList)54 trainingMat = zeros((m, 1024))55 for i in range(m):56 fileNameStr = trainingFileList[i]57 fileStr = fileNameStr.split('.')[0]58 classNumStr = int(fileStr.split('_')[0])59 hwLabels.append(classNumStr)60 trainingMat[i,:] = img2vector('trainingDigits/%s' % fileNameStr)61 testFileList = listdir('testDigits')62 errorCount = 0.063 mTest = len(testFileList)64 for i in range(mTest):65 '''66 對檔案名稱字進行解析67 此程式中使用的檔案名稱字格式為:68 正確數字_編號.txt69 '''70 fileNameStr = testFileList[i]71 fileStr = fileNameStr.split('.')[0]72 classNumStr = int (fileStr.split('_')[0])73 74 vectorUnderTest = img2vector('testDigits/%s' % fileNameStr) #錄入測試資料75 #對測試資料進行分類76 classifierResult = classify(vectorUnderTest,77 trainingMat, hwLabels, 3)78 print("the classifier came back with: %d, the real answer is : %d" \79 % (classifierResult, classNumStr))80 if (classifierResult != classNumStr): errorCount += 1.081 print("the total number of errors is : %d" % errorCount)82 print("the total error rate is : %f" % (errorCount/float(mTest)))
運行結果:
我們可以看出k-近鄰演算法識別手寫數字程式,錯誤率為1.4%。
三、總結
kNN演算法是機器學習中分類演算法的一種,屬於監督學習。是分類資料時最簡單最有效演算法。但是執行效率低,運行非常耗時。
參考資料:
《機器學習實戰》