k-NN 没有特别的训练过程,给定训练集,标签,k,计算待预测特征到训练集的所有距离,选取前k个距离最小的训练集,k个中标签最多的为预测标签

约会类型分类、手写数字识别分类

  1. 计算输入数据到每一个训练数据的距离
  2. 选择前k个,判断其中类别最多的类作为预测类
import numpy as np
import operator
import matplotlib
import matplotlib.pyplot as plt # inX: test data, N features (1xN)
# dataSet: M samples, N features (MxN)
# label: for M samples (1xM)
# k: k-Nearest Neighbor
def classify0(inX, dataSet, labels, k):
dataSetSize = dataSet.shape[0]
diffMat = np.tile(inX, (dataSetSize, 1)) - dataSet
distances = np.sum(diffMat**2, axis=1)**0.5
sortDistances = distances.argsort() # 计算距离
classCount = {}
for i in range(k):
voteLable = labels[sortDistances[i]]
classCount[voteLable] = classCount.get(voteLable, 0) + 1
sortedClassCount = sorted(classCount.items(), key=operator.itemgetter(1), reverse=True) # 找出最多投票的类
result = sortedClassCount[0][0]
# print("Predict: ", result)
return result # 将一个文件写入矩阵,文件有4列,最后一列为labels,以\t间隔
def file2matrix(filename):
with open(filename) as f:
arrayLines = f.readlines()
# print(arrayLines) # 有\n
numberOfLines = len(arrayLines) # 将txt文件按行读入为一个list,一行为一个元素
returnMat = np.zeros((numberOfLines, 3))
classLabelVector = []
index = 0
for line in arrayLines:
line = line.strip()
listFromLine = line.split('\t')
returnMat[index,:] = listFromLine[0:3]
classLabelVector.append(int(listFromLine[-1]))
index += 1
return returnMat, classLabelVector # 画一些图
def ex3():
datingDateMat, datingLables = file2matrix("datingTestSet2.txt")
fig = plt.figure()
ax = fig.add_subplot(1,2,1)
ax.scatter(datingDateMat[:,1], datingDateMat[:,2], s=15.0*np.array(datingLables), c=15.0*np.array(datingLables))
ax2 = fig.add_subplot(1,2,2)
ax2.scatter(datingDateMat[:,0], datingDateMat[:,1], s=15.0*np.array(datingLables), c=15.0*np.array(datingLables))
plt.show() # 将数据集归一化[0 1]之间 (value - min)/(max - min)
def autoNorm(dataSet):
minVals = dataSet.min(axis=0)
maxVals = dataSet.max(axis=0)
ranges = maxVals - minVals
m = dataSet.shape[0]
normDataSet = dataSet - np.tile(minVals, (m,1))
normDataSet = normDataSet/np.tile(ranges, (m,1))
return normDataSet, ranges, minVals # 分类器,输入数据集,归一化参数,labels,70%作为训练集,30%测试集
def datingClassTest(normDataSet, ranges, minVals, labels):
m = normDataSet.shape[0]
numOfTrain = int(m*0.7)
trainIndex = np.arange(m)
np.random.shuffle(trainIndex)
dataSet = normDataSet[trainIndex[0:numOfTrain],:]
testSet = normDataSet[trainIndex[numOfTrain:],:]
labels = np.array(labels)
dataSetLabels = labels[trainIndex[0:numOfTrain]]
testSetLabels = labels[trainIndex[numOfTrain:]] k = int(input("Input k: "))
results = []
for inX in testSet:
result = classify0(inX, dataSet, dataSetLabels, k)
results.append(result)
compResultsAndLable = np.argwhere(results==testSetLabels)
acc = len(compResultsAndLable)/len(testSetLabels)
print("Accuracy: {:.2f}".format(acc))
print("Error: {:.2f}".format(1-acc)) classList = ['not at all', 'in small doses', 'in large doses']
inX1 = float(input("1: percentage of time spent playing video games? "))
inX2 = float(input("2: frequent flier miles earned per year? "))
inX3 = float(input("3: liters of ice cream consumed per year? "))
inXUser = [inX1,inX2,inX3]
inXUser = (inXUser - minVals)/ranges
result = classify0(inXUser, dataSet, dataSetLabels, k)
print("Predict: ", classList[result]) if __name__ == '__main__':
# # -- ex1 --
# inX = [1, 1]
# dataSet = np.array([[1.0, 1.1], [1.0, 1.0], [0, 0], [0, 0.1]])
# labels = ['A', 'A', 'B', 'B']
# k = 3
# classify0(inX, dataSet, labels, k) # # -- ex2 --
datingDateMat, datingLables = file2matrix("datingTestSet2.txt") # # -- ex3 --
# ex3() # #-- ex4 --
# normDataSet, ranges, minVals = autoNorm(datingDateMat) # # -- ex5 --
# datingClassTest(normDataSet, ranges, minVals, datingLables)
import numpy as np
import matplotlib
import matplotlib.pyplot as plt
import os
import operator def img2vector(filename):
with open(filename) as f:
lines = f.readlines()
return_vector = []
for line in lines:
line = line.strip()
for j in line:
return_vector.append(int(j))
return return_vector # inX: test data, N features (1xN)
# dataSet: M samples, N features (MxN)
# label: for M samples (1xM)
# k: k-Nearest Neighbor
def classify0(inX, dataSet, labels, k):
dataSetSize = dataSet.shape[0]
diffMat = np.tile(inX, (dataSetSize, 1)) - dataSet
distances = np.sum(diffMat**2, axis=1)**0.5
sortDistances = distances.argsort() # 计算距离
classCount = {}
for i in range(k):
voteLable = labels[sortDistances[i]]
classCount[voteLable] = classCount.get(voteLable, 0) + 1
sortedClassCount = sorted(classCount.items(), key=operator.itemgetter(1), reverse=True) # 找出最多投票的类
result = sortedClassCount[0][0]
# print("Predict: ", result)
return result def handwriting_class_test(data_set, training_labels, test_set, test_labels, k):
results = []
for i in range(len(test_set)):
result = classify0(test_set[i], data_set, training_labels, k)
results.append(result)
# print('predict: ', result, 'answer: ', test_labels[i])
compare_results = np.argwhere(results==test_labels)
acc = len(compare_results)/len(test_labels)
print("Accuracy: {:.5f}".format(acc))
print("Error: {:.5f}".format(1-acc)) if __name__ == '__main__':
dir_path = r'H:\ML\MachineLearninginAction\02kNN\digits'
training_path = os.path.join(dir_path, r'trainingDigits')
test_path = os.path.join(dir_path, r'testDigits') training_files_list = os.listdir(training_path)
test_files_list = os.listdir(test_path) # 计算训练集矩阵与labels
m = len(training_files_list)
# m = 5
data_set = np.zeros((m, 1024))
training_labels = np.zeros(m)
for i in range(m):
data_set[i] = img2vector(os.path.join(training_path, training_files_list[i]))
training_labels[i] = training_files_list[i].split('_')[0]
# 测试集矩阵与labels
mt = len(test_files_list)
test_set = np.zeros((mt,1024))
test_labels = np.zeros(mt)
for i in range(mt):
test_set[i] = img2vector(os.path.join(test_path, test_files_list[i]))
test_labels[i] = test_files_list[i].split('_')[0]
k = 3
handwriting_class_test(data_set, training_labels, test_set, test_labels, k)

k-NN——算法实现的更多相关文章

  1. kaggle赛题Digit Recognizer:利用TensorFlow搭建神经网络(附上K邻近算法模型预测)

    一.前言 kaggle上有传统的手写数字识别mnist的赛题,通过分类算法,将图片数据进行识别.mnist数据集里面,包含了42000张手写数字0到9的图片,每张图片为28*28=784的像素,所以整 ...

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

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

  3. 《机器学习实战》学习笔记一K邻近算法

     一. K邻近算法思想:存在一个样本数据集合,称为训练样本集,并且每个数据都存在标签,即我们知道样本集中每一数据(这里的数据是一组数据,可以是n维向量)与所属分类的对应关系.输入没有标签的新数据后,将 ...

  4. [Machine-Learning] K临近算法-简单例子

    k-临近算法 算法步骤 k 临近算法的伪代码,对位置类别属性的数据集中的每个点依次执行以下操作: 计算已知类别数据集中的每个点与当前点之间的距离: 按照距离递增次序排序: 选取与当前点距离最小的k个点 ...

  5. k近邻算法的Java实现

    k近邻算法是机器学习算法中最简单的算法之一,工作原理是:存在一个样本数据集合,即训练样本集,并且样本集中的每个数据都存在标签,即我们知道样本集中每一数据和所属分类的对应关系.输入没有标签的新数据之后, ...

  6. 基本分类方法——KNN(K近邻)算法

    在这篇文章 http://www.cnblogs.com/charlesblc/p/6193867.html 讲SVM的过程中,提到了KNN算法.有点熟悉,上网一查,居然就是K近邻算法,机器学习的入门 ...

  7. 聚类算法:K-means 算法(k均值算法)

    k-means算法:      第一步:选$K$个初始聚类中心,$z_1(1),z_2(1),\cdots,z_k(1)$,其中括号内的序号为寻找聚类中心的迭代运算的次序号. 聚类中心的向量值可任意设 ...

  8. 从K近邻算法谈到KD树、SIFT+BBF算法

    转自 http://blog.csdn.net/v_july_v/article/details/8203674 ,感谢july的辛勤劳动 前言 前两日,在微博上说:“到今天为止,我至少亏欠了3篇文章 ...

  9. Python实现kNN(k邻近算法)

    Python实现kNN(k邻近算法) 运行环境 Pyhton3 numpy科学计算模块 计算过程 st=>start: 开始 op1=>operation: 读入数据 op2=>op ...

  10. 机器学习之K近邻算法(KNN)

    机器学习之K近邻算法(KNN) 标签: python 算法 KNN 机械学习 苛求真理的欲望让我想要了解算法的本质,于是我开始了机械学习的算法之旅 from numpy import * import ...

随机推荐

  1. gorm声明模型

    模型定义 模型是标准的结构体,由go的基本数据类型.实现了Scanner和Valuer接口的自定义类型及其指针或别名组成 例如: type User struct { ID uint Name str ...

  2. golang中使用switch语句根据年月计算天数

    package main import "fmt" func main() { days := CalcDaysFromYearMonth(2021, 9) fmt.Println ...

  3. Redis 源码简洁剖析 03 - Dict Hash 基础

    Redis Hash 源码 Redis Hash 数据结构 Redis rehash 原理 为什么要 rehash? Redis dict 数据结构 Redis rehash 过程 什么时候触发 re ...

  4. js-reduce方法源码

    // 数组中的reduce方法源码复写 //先说明一下reduce原理:总的一句,reduce方法主要是把数组遍历, //然后把数组的每个元素传入回调函数中,回调函数怎么处理,就会的到什么样的效果 A ...

  5. IDEA Debug常用快捷键

    快捷键 介绍 F7 步入:进入到方法内部执行.一般步入自定义的方法.区别于强行步入 F8 步过:不会进入到方法内部,直接执行. F9 恢复程序:下面有断点则运行到下一断点,否则结束程序. Shift+ ...

  6. npm 查看一个包的版本信息

    有了npm 我们能够简单的一段代码就下载我们需要的包,但是包是不断更新的, 所以我们要关注包的版本信息: 现在,假设我们需要 jquery ,但是jquery现在有很多版本,我们如何通过npm查看呢? ...

  7. node.js中的fs.appendFile方法使用说明

    方法说明: 该方法以异步的方式将 data 插入到文件里,如果文件不存在会自动创建.data可以是任意字符串或者缓存. 语法: 代码如下: fs.appendFile(filename, data, ...

  8. php7.3编译安装 支持微擎2.0

    再次整理   //一下配置在命令粘贴时注意句尾加 \ , 在 \ 后不能有空格,不然会自动执行,相当于回车./configure --prefix=/usr/local/php \ --with-co ...

  9. FileInputStream 类与 FileReader 类的区别

    FileInputStream 类与 FileReader 类的区别: 两个类的构造函数的形式和参数都是相同的,参数为 File 对象或者表示路径的 String ,它们到底有何区别呢? FileIn ...

  10. .NET 6全文检索引擎Lucene.NET 4.8简单封装

    前言 因为最近在做一个检索数据的工具.最开始用的Mysql8自带的全文检索功能.但是发现这货数据量超过百万之后,检索速度直线下降. 于是想到Lucene.net.花了一晚上时间做了简单的封装.可以直接 ...