预测样本x的标记。 获取方式通过找到与样本x最近的topK个点,并查看它们的标签。 查找里面占某类标签最多的那类标签 (书中3.1 3.2节) :param trainDataMat:训练集数据集 :param trainLabelMat:训练集标签集 :param x:要预测的样本x :param topK:选择参考最邻近样本的数目(样本数目的选择关系到正确率,详看3.2.3 K值的选择) :return:预测的标记
(trainDataMat, trainLabelMat, x, topK)
| 64 | |
| 65 | |
| 66 | def getClosest(trainDataMat, trainLabelMat, x, topK): |
| 67 | ''' |
| 68 | 预测样本x的标记。 |
| 69 | 获取方式通过找到与样本x最近的topK个点,并查看它们的标签。 |
| 70 | 查找里面占某类标签最多的那类标签 |
| 71 | (书中3.1 3.2节) |
| 72 | :param trainDataMat:训练集数据集 |
| 73 | :param trainLabelMat:训练集标签集 |
| 74 | :param x:要预测的样本x |
| 75 | :param topK:选择参考最邻近样本的数目(样本数目的选择关系到正确率,详看3.2.3 K值的选择) |
| 76 | :return:预测的标记 |
| 77 | ''' |
| 78 | #建立一个存放向量x与每个训练集中样本距离的列表 |
| 79 | #列表的长度为训练集的长度,distList[i]表示x与训练集中第 |
| 80 | ## i个样本的距离 |
| 81 | distList = [0] * len(trainLabelMat) |
| 82 | #遍历训练集中所有的样本点,计算与x的距离 |
| 83 | for i in range(len(trainDataMat)): |
| 84 | #获取训练集中当前样本的向量 |
| 85 | x1 = trainDataMat[i] |
| 86 | #计算向量x与训练集样本x的距离 |
| 87 | curDist = calcDist(x1, x) |
| 88 | #将距离放入对应的列表位置中 |
| 89 | distList[i] = curDist |
| 90 | |
| 91 | #对距离列表进行排序 |
| 92 | #argsort:函数将数组的值从小到大排序后,并按照其相对应的索引值输出 |
| 93 | #例如: |
| 94 | # >>> x = np.array([3, 1, 2]) |
| 95 | # >>> np.argsort(x) |
| 96 | # array([1, 2, 0]) |
| 97 | #返回的是列表中从小到大的元素索引值,对于我们这种需要查找最小距离的情况来说很合适 |
| 98 | #array返回的是整个索引值列表,我们通过[:topK]取列表中前topL个放入list中。 |
| 99 | #----------------优化点------------------- |
| 100 | #由于我们只取topK小的元素索引值,所以其实不需要对整个列表进行排序,而argsort是对整个 |
| 101 | #列表进行排序的,存在时间上的浪费。字典有现成的方法可以只排序top大或top小,可以自行查阅 |
| 102 | #对代码进行稍稍修改即可 |
| 103 | #这里没有对其进行优化主要原因是KNN的时间耗费大头在计算向量与向量之间的距离上,由于向量高维 |
| 104 | #所以计算时间需要很长,所以如果要提升时间,在这里优化的意义不大。(当然不是说就可以不优化了, |
| 105 | #主要是我太懒了) |
| 106 | topKList = np.argsort(np.array(distList))[:topK] #升序排序 |
| 107 | #建立一个长度时的列表,用于选择数量最多的标记 |
| 108 | #3.2.4提到了分类决策使用的是投票表决,topK个标记每人有一票,在数组中每个标记代表的位置中投入 |
| 109 | #自己对应的地方,随后进行唱票选择最高票的标记 |
| 110 | labelList = [0] * 10 |
| 111 | #对topK个索引进行遍历 |
| 112 | for index in topKList: |
| 113 | #trainLabelMat[index]:在训练集标签中寻找topK元素索引对应的标记 |
| 114 | #int(trainLabelMat[index]):将标记转换为int(实际上已经是int了,但是不int的话,报错) |
| 115 | #labelList[int(trainLabelMat[index])]:找到标记在labelList中对应的位置 |
| 116 | #最后加1,表示投了一票 |
| 117 | labelList[int(trainLabelMat[index])] += 1 |
| 118 | #max(labelList):找到选票箱中票数最多的票数值 |
| 119 | #labelList.index(max(labelList)):再根据最大值在列表中找到该值对应的索引,等同于预测的标记 |
| 120 | return labelList.index(max(labelList)) |
| 121 | |
| 122 | |
| 123 | def model_test(trainDataArr, trainLabelArr, testDataArr, testLabelArr, topK): |
no test coverage detected