机器学习K近邻算法的实现主要是参考《机器学习实战》这本书。

一、K近邻(KNN)算法

K最近邻(k-Nearest Neighbour,KNN)分类算法,理解的思路是:如果一个样本在特征空间中的K个最相似(即特征空间最邻近)的样本中的大多数属于某一个类别,则该样本也属于这个类别。

我们采用一个图来进行说明(如下):

图中的蓝色小正方形和红色的小正方形属于两类不同的样本数据,图正中间的绿色的圆代表的是待分类的数据。现在我们可以根据K最近邻算法来判断绿色的圆属于哪一类数据?

  • 如果K=3,绿色圆点的最近的3个邻居是2个红色小三角形和1个蓝色小正方形,由于红色三角形所占比例为2/3,选择概率大的一类,我们可以认为绿色的这个待分类点属于红色的三角形一类。
  • 如果K=5,绿色圆点的最近的5个邻居是2个红色三角形和3个蓝色的正方形,由于蓝牙正方形所占比例为3/5,选择概率大的一类,我们可以认为绿色的这个待分类点属于蓝色的正方形一类。

以上就是对K最近邻算法的理解。

K-近邻算法的特点是:

优点:精度高、对异常值不敏感、无数据输入假定。
缺点:计算复杂度高、空间复杂度高。
适用数据范围:数值型和标称型。

二、算法实践

现在我们对《机器学习实战》中K最近邻的代码进行实现

首先要导入样本数据,我们建立KNN.py的模块,用来生成数据库,代码如下:

 from numpy import *
import operator def createDataSet():
group = array([[1.0,1.1],[1.0,1.0],[0,0],[0,0.1]])
labels = ['A','A','B','B']
return group, labels

在代码中,我们导入了Python的两个模块:科学计算包Numpy和运算符模块。Python中导入Nump模块可以参考之前的Python安装reque模块文章。

然后实现K-近邻算法,

K-近邻算法的具体思想如下:

(1)计算已知类别数据集中的点与当前点之间的距离

(2)按照距离递增次序排序

(3)选取与当前点距离最小的k个点

(4)确定前k个点所在类别的出现频率

(5)返回前k个点中出现频率最高的类别作为当前点的预测分类

Python代码如下:

 # coding : utf-8

 from numpy import *
import operator
import kNN group, labels = kNN.createDataSet() def classify(inX, dataSet, labels, k):
dataSetSize = dataSet.shape[0] #得到数组的行数,即知道有几个数据集
diffMat = tile(inX, (dataSetSize,1)) - dataSet #样本与训练集的差值矩阵
sqDiffMat = diffMat**2 #各元素分别平方
sqDistances = sqDiffMat.sum(axis=1) #每一行元素求和
distances = sqDistances**0.5 #开方,得到距离
sortedDistances = distances.argsort() #按照distances中元素进行升序排列后得到相应下表的列表
classCount = {} #选择距离最小的k个点
for i in range(k):
numOflabel = labels[sortedDistances[i]]
classCount[numOflabel] = classCount.get(numOflabel,0) + 1
sortedClassCount = sorted(classCount.iteritems(), key=operator.itemgetter(1),reverse=True) #按classCount字典的value(类别出现次数)降序排列
return sortedClassCount[0][0] test = classify([0,0], group, labels, 3)
print test

运算结果为:

B

输出结果是B:说明我们新的数据([0,0])是属于B类。

代码解析:

classify函数的参数为:

  • inX:用于分类的输入向量
  • dataSet:训练样本集合
  • labels:标签向量
  • k:K-近邻算法中的k

shape函数:是numpy.core.fromnumeric中的函数,它的功能是读取矩阵的长度,比如shape[0]就是读取矩阵第一维度的长度。它的输入参数可以是一个整数表示维                       度,也可以是一个矩阵。如下所示:

>>> from numpy import *
>>> shape([1])
(1L,)
>>> shape([[1],[2]])
(2L, 1L)
>>> shape(3)
()
>>>
>>> e=eye(3)
>>> e
array([[ 1., 0., 0.],
[ 0., 1., 0.],
[ 0., 0., 1.]])
>>> e.shape
(3L, 3L)
>>> e.shape[0]
3L
>>>

一维矩阵,shape([1]),结果为(1L,);

    二维矩阵,shape([1],[2]), 结果为(2L,1L);

  单独的数字,返回值为空;

  创建一个单位矩阵e,可以判断e的形状,e.shape的结果为(3L,3L); 我们只取第一维度长度(即行数)e.shape[0]的值为3L。

tile函数:tile函数是模板numpy.lib.shape_base中的函数。函数的形式是tile(A,reps)。
             A的类型众多,几乎所有类型都可以:array, list, tuple, dict, matrix以及基本数据类型int, string, float以及bool类型。
             reps的类型也很多,可以是tuple,list, dict, array, int,bool.但不可以是float, string, matrix类型。

    tile函数是为了扩充数组元素,看一下举例:

>>> from numpy import *
>>> a=array([1,2])
>>> tile(a,(3,2)) #构造3X2个copy
array([[1, 2, 1, 2],
[1, 2, 1, 2],
[1, 2, 1, 2]])
>>> tile(7,(3,2))
array([[7, 7],
[7, 7],
[7, 7]])
>>>

  简单理解tile(a,(3,2))就是对a进行3X2的复制。具体的理解可以参考这篇文章:http://blog.sina.com.cn/s/blog_6bd0612b0101cr3u.html

在算法程序中tile(inX, (dataSetSize,1))就是对inX进行了复制了dataSetSize行,1表示列的倍数。

diffMat = tile(inX, (dataSetSize,1)) - dataSet 这一行代码表示前一个二维数组矩阵的每一个元素减去后一个数组对应的元素值,实现了矩阵之间的减法。

sqDistances = sqDiffMat.sum(axis=1)这一行代码中,axis=1:参数等于1的时候,表示矩阵中行之间的数的求和,等于0的时候表示列之间数的求和。

classCount.get(numOflabel,0) + 1:这一行代码中使用了字典的get方法,在这里需要理解,get方法可以访问字典中键对应的值。key存在则返回对应value,键不存在返回None,也可以设定自己的返回参数,设置的方法是:变量名.get(键名,参数)。在这里我们设置的参数为0,说明在访问不到键对应值时,返回的参数为0,然后对其加1,如此循环,便可以得到key出现的频率。

sorted函数:

为Python内置的排序函数,可以对list或者iterator进行排序,官网文档见:http://docs.python.org/2/library/functions.html?highlight=sorted#sorted,该函数原型为:

sorted(iterable[, cmp[, key[, reverse]]])

参数解释:

(1)iterable指定要排序的list或者iterable;

(2)cmp为函数,指定排序时进行比较的函数,可以指定一个函数或者lambda函数,如:

students为类对象的list,每个成员有三个域,用sorted进行比较时可以自己定cmp函数,例如这里要通过比较第三个数据成员来排序,代码可以这样写:
      students = [('john', 'A', 15), ('jane', 'B', 12), ('dave', 'B', 10)]
      sorted(students, key=lambda student : student[2])
(3)key为函数,指定取待排序元素的哪一项进行排序,函数用上面的例子来说明,代码如下:
      sorted(students, key=lambda student : student[2])

key指定的lambda函数功能是去元素student的第三个域(即:student[2]),因此sorted排序时,会以students所有元素的第三个域来进行排序。

  也可以用operator.itemgetter函数来实现,例如要通过student的第三个域排序,可以这么写:
  sorted(students, key=operator.itemgetter(2)) 
  sorted函数也可以进行多级排序,例如要根据第二个域和第三个域进行排序,可以这么写:
  sorted(students, key=operator.itemgetter(1,2))  即先跟句第二个域排序,再根据第三个域排序。

(4)reverse参数就不用多说了,是一个bool变量,表示升序还是降序排列,默认为false(升序排列),定义为True时将按降序排列。

在算法程序中,我们也可以用另一种方法来获得最大概率的分类,代码修改如下:

     classCount = {}
for i in range(k):
numOflabel = labels[sortedDistances[i]]
classCount[numOflabel] = classCount.get(numOflabel,0) + 1 maxCount = 0
for key, value in classCount.items():
if value > maxCount:
maxCount = value
maxIndex = key
return maxIndex

同样可以实现。

机器学习 Python实践-K近邻算法的更多相关文章

  1. 机器学习03:K近邻算法

    本文来自同步博客. P.S. 不知道怎么显示数学公式以及排版文章.所以如果觉得文章下面格式乱的话请自行跳转到上述链接.后续我将不再对数学公式进行截图,毕竟行内公式截图的话排版会很乱.看原博客地址会有更 ...

  2. 02机器学习实战之K近邻算法

    第2章 k-近邻算法 KNN 概述 k-近邻(kNN, k-NearestNeighbor)算法是一种基本分类与回归方法,我们这里只讨论分类问题中的 k-近邻算法. 一句话总结:近朱者赤近墨者黑! k ...

  3. 机器学习实战笔记--k近邻算法

    #encoding:utf-8 from numpy import * import operator import matplotlib import matplotlib.pyplot as pl ...

  4. 机器学习随笔01 - k近邻算法

    算法名称: k近邻算法 (kNN: k-Nearest Neighbor) 问题提出: 根据已有对象的归类数据,给新对象(事物)归类. 核心思想: 将对象分解为特征,因为对象的特征决定了事对象的分类. ...

  5. 《机器学习实战》-k近邻算法

    目录 K-近邻算法 k-近邻算法概述 解析和导入数据 使用 Python 导入数据 实施 kNN 分类算法 测试分类器 使用 k-近邻算法改进约会网站的配对效果 收集数据 准备数据:使用 Python ...

  6. 用python实现k近邻算法

    用python写程序真的好舒服. code: import numpy as np def read_data(filename): '''读取文本数据,格式:特征1 特征2 -- 类别''' f=o ...

  7. 机器学习:1.K近邻算法

    1.简单案例:预测男女,根据身高,体重,鞋码 import numpy as np import matplotlib import sklearn from skleran.neighbors im ...

  8. 《机器学习实战》——K近邻算法

    三要素:距离度量.k值选择.分类决策 原理: (1) 输入点A,输入已知分类的数据集data (2) 求A与数据集中每个点的距离,归一化,并排序,选择距离最近的前K个点 (3) K个点进行投票,票数最 ...

  9. GridSearchCV网格搜索得到最佳超参数, 在K近邻算法中的应用

    最近在学习机器学习中的K近邻算法, KNeighborsClassifier 看似简单实则里面有很多的参数配置, 这些参数直接影响到预测的准确率. 很自然的问题就是如何找到最优参数配置? 这就需要用到 ...

随机推荐

  1. Express入门( node.js Web应用框架 )

    运用Express框架构建简单的NodeJS应用 Start  确认安装了NodeJS之后(最新的Node安装好后NPM也会自带安装了),npm可理解为nodejs的一个工具包.可通过查看版本来检测是 ...

  2. IEEE 754浮点数表示标准

    二进制数的科学计数法 C++中使用的浮点数包括采用的是IEEE标准下的浮点数表示方法.我们知道在数学中可以将任何十进制的数写成以10为底的科学计数法的形式,如下 其中显而易见,因为如果a比10大或者比 ...

  3. python 类与对象解析

    类成员:    # 字段        - 普通字段,保存在对象中,执行只能通过对象访问        - 静态字段,保存在类中,  执行 可以通过对象访问 也可以通过类访问            # ...

  4. 针对《面试心得与总结—BAT、网易、蘑菇街》一文中出现的技术问题的收集与整理

    最近,我在ImportNew网站上,看到了这篇文章,觉得总结的非常好,就默默的收藏起来了,觉得日后一定要好好整理学习一下,昨天突然发现在脉脉的行业头条中,居然也推送了这篇文章,更加坚定了我整理的信心. ...

  5. 如何在Mongodb中实现数据超时自动删除功能?

    在工作过程中,我们难免会遇到这样的问题,我们想保存一些数据,但是我们对这些数据的要求并不高,有时候往往只是想要某个时间范围内的数据,比如我们如果永远只关心从当前时间往前推半年内的数据特性,那么我们就不 ...

  6. Eclipse中遇到main方法不能运行 的情况

    java.lang.UnsupportedClassVersionError: Bad version number in .class file 造成这种过错是ni的支撑Tomcat运行的JDK版本 ...

  7. ZeroMQ API(五) 传输模式

    1.使用TCP的单播传输:zmq_tcp(7) 1.1 名称 zmq_tcp - 使用TCP的ZMQ单播传输 1.2 概要 TCP是一种无处不在,可靠的单播传输.当通过具有ZMQ的网络连接分布式应用程 ...

  8. 安装python包时遇到"error: Microsoft Visual C++ 9.0 is required"的简答

    简答 在Windows下用pip安装Scrapy报如下错误, error: Microsoft Visual C++ 9.0 is required (Unable to find vcvarsall ...

  9. Header File Dependencies

    [Header File Dependencies] 什么时候可以用前置声明替代include? 1.当 declare/define pointer&reference 时. 2.当 dec ...

  10. soj1010. Zipper

    1010. Zipper Constraints Time Limit: 1 secs, Memory Limit: 32 MB Description Given three strings, yo ...