KNN算法详解

    简单的说,K近邻算法是采用不同特征值之间的距离方法进行分类。

   该方法优点:精确值高、对异常值不敏感、无数据输入假定

   缺点:计算复杂度高、空间复杂度高

   适用范围:数据型和标称型

   现在我们来讲KNN算法的工作原理:存在一个样本数据集,也称作训练样本集,并且样本中每条数据都存在标签。将新输入的没有标签的数据与训练样本数据集中每条数据进行距离计算,选择前K个最小距离。并统计出现次数最多的分类。将该分类作为新数据的标签。

  例如:使用k-近邻算法分类爱情片和动作片。

            训练样本数据集如下:

            

           此时给定一部电影的统计数据:打斗镜头:18,接吻镜头:90,判断该电影属于哪种类型?

          分别计算该条数据与样本集中各条数据的距离:

          

         假设K=3,将距离从小到大排列,选择前三条数据,并统计出该三条数据中所属类别最多的类别。并将该类别赋值给新输入数据。如上所属,该输入数据的类别为爱情片。

        K-近邻算法的一般流程:

       (1)收集数据

       (2)准备数据:距离计算所需的数值,最好是结构欧化数据格式

       (3)分析数据

       (4)训练数据:此步骤不适用与K-近邻算法

       (5)测试数据:计算错误率

       (6)使用数据:首先需要输入样本数据和结构化数据的输出结果,然后运用K近邻算法判断输入数据分别属于哪一类,最后应用对计算出的分类执行后续处理

        实现代码如下:

1import numpy as np 2import operator 3from os import listdir 4def createDateSet(): 5 group=np.array([[1,1.1],[1,1],[0,0],[0,0.1]]) 6 labels=['A','A','B','B'] 7 return group,labels 8#分类 9def classify(inX,dataSet,labels,k): 10 dataSetSize=dataSet.shape[0] 11 diffMat=np.tile(inX,(dataSetSize,1))-dataSet 12 sqDiffMat=diffMat**2 13 sqDistances=sqDiffMat.sum(axis=1) 14 distances=sqDistances**0.5 15 ##argsort()根据元素的值从大到小对元素进行排序,返回下标 16 sortedDistance=distances.argsort() 17 classCount={} 18 for i in range(k): 19 voteIlabel=labels[sortedDistance[i]] 20 classCount[voteIlabel]=classCount.get(voteIlabel,0)+1 21 sortedClassCount=np.sort(classCount.items(),key=operator.itemgetter(1),reverse=True) 22 return sortedClassCount[0][0] 23#获取数据 24def filematrix(filename): 25 fr=open(filename) 26 arrayLine=fr.readlines() 27 numberOfLines=len(arrayLine) 28 returnMat=np.zeros((numberOfLines,3)) 29 classLabelVector=[] 30 index=0 31 for line in arrayLine: 32 line=line.strip().split('\t') 33 returnMat[index,:]=line[0:3] 34 classLabelVector.append(int(line[-1])) 35 index+=1 36 return returnMat,classLabelVector 37#归一化数据 38def autoNorm(dataSet): 39 minVals=dataSet.min(0) #0表示列 40 maxVals=dataSet.max(0) 41 ranges=maxVals-minVals 42 normDataSet=np.zeros(np.shape(dataSet)) 43 m=dataSet.shape[0] 44 normDataSet=dataSet-np.tile(minVals,(m,1)) 45 normDataSet=normDataSet/np.tile(ranges,(m,1)) 46 return normDataSet,ranges,minVals 47 48def datingClassTest(): 49 hoRatio=0.10 50 datingDataMat,datingLabels=filematrix('../data/testSet.txt') 51 normMat,rangs,minVals=autoNorm(datingDataMat) 52 m=normMat.shape[0] 53 numTestVecs=int(m*hoRatio) 54 errorCount=0.0 55 for i in range(numTestVecs): 56 classifierResult=classify(normMat[i,:],normMat[numTestVecs:m,:],datingLabels[numTestVecs:m],3) 57 print("预测类别为:",classifierResult,"实际类别:",datingLabels[i]) 58 if(classifierResult!=datingLabels[i]): 59 errorCount+=1 60 print("总的错误率:",(errorCount/float(numTestVecs))) 61 62def img2vector(filename): 63 returnVect=np.zeros((1,1024)) 64 fr=open(filename) 65 for i in range(32): 66 lineStr=fr.readlines() 67 for j in range(32): 68 returnVect[0,32*i+j]=int(lineStr[j]) 69 return returnVect 70 71def handwritingClassTest(): 72 hwLabels = [] 73 trainingFileList = listdir('trainingDigits') #load the training set 74 m = len(trainingFileList) 75 trainingMat = np.zeros((m,1024)) 76 for i in range(m): 77 fileNameStr = trainingFileList[i] 78 fileStr = fileNameStr.split('.')[0] #take off .txt 79 classNumStr = int(fileStr.split('_')[0]) 80 hwLabels.append(classNumStr) 81 trainingMat[i,:] = img2vector('trainingDigits/%s' % fileNameStr) 82 testFileList = listdir('testDigits') #iterate through the test set 83 errorCount = 0.0 84 mTest = len(testFileList) 85 for i in range(mTest): 86 fileNameStr = testFileList[i] 87 fileStr = fileNameStr.split('.')[0] #take off .txt 88 classNumStr = int(fileStr.split('_')[0]) 89 vectorUnderTest = img2vector('testDigits/%s' % fileNameStr) 90 classifierResult = classify(vectorUnderTest, trainingMat, hwLabels, 3) 91 print ("the classifier came back with: %d, the real answer is: %d" % (classifierResult, classNumStr)) 92 if (classifierResult != classNumStr): errorCount += 1.0 93 print ("\nthe total number of errors is: %d" % errorCount) 94 print ("\nthe total error rate is: %f" % (errorCount/float(mTest)))
点赞
收藏

评论区

加载中...

相关推荐

MySQL:[Err] 1292 - Incorrect datetime value: ‘0000-00-00 00:00:00‘ for column ‘CREATE_TIME‘ at row 1

文章目录问题用navicat导入数据时,报错:原因这是因为当前的MySQL不支持datetime为0的情况。解决修改sql\mode:sql\mode:SQLMode定义了MySQL应支持的SQL语法、数据校验等,这样可以更容易地在不同的环境中使用MySQL。全局s

Oracle 分组与拼接字符串同时使用

SELECTT.,ROWNUMIDFROM(SELECTT.EMPLID,T.NAME,T.BU,T.REALDEPART,T.FORMATDATE,SUM(T.S0)S0,MAX(UPDATETIME)CREATETIME,LISTAGG(TOCHAR(

皕杰报表之UUID

​在我们用皕杰报表工具设计填报报表时,如何在新增行里自动增加id呢?能新增整数排序id吗?目前可以在新增行里自动增加id,但只能用uuid函数增加UUID编码,不能新增整数排序id。uuid函数说明:获取一个UUID,可以在填报表中用来创建数据ID语法:uuid()或uuid(sep)参数说明:sep布尔值,生成的uuid中是否包含分隔符'',缺省为

手写Java HashMap源码

HashMap的使用教程HashMap的使用教程HashMap的使用教程HashMap的使用教程HashMap的使用教程22

皕杰报表(关于日期时间时分秒显示不出来)

在使用皕杰报表设计器时,数据据里面是日期型,但当你web预览时候,发现有日期时间类型的数据时分秒显示不出来,只有年月日能显示出来,时分秒显示为0:00:00。1.可以使用tochar解决,数据集用selecttochar(flowdate,"yyyyMMddHH:mm:ss")fromtablename2.也可以把数据库日期类型date改成timestamp

2020年前端实用代码段,为你的工作保驾护航

有空的时候,自己总结了几个代码段,在开发中也经常使用,谢谢。1、使用解构获取json数据let jsonData  id: 1,status: "OK",data: 'a', 'b';let  id, status, data: number   jsonData;console.log(id, status, number )