加载Mnist数据集 :param fileName:要加载的数据集路径 :return: list形式的数据集及标记
(fileName)
| 17 | import time |
| 18 | |
| 19 | def loadData(fileName): |
| 20 | ''' |
| 21 | 加载Mnist数据集 |
| 22 | :param fileName:要加载的数据集路径 |
| 23 | :return: list形式的数据集及标记 |
| 24 | ''' |
| 25 | print('start to read data') |
| 26 | # 存放数据及标记的list |
| 27 | dataArr = []; labelArr = [] |
| 28 | # 打开文件 |
| 29 | fr = open(fileName, 'r') |
| 30 | # 将文件按行读取 |
| 31 | for line in fr.readlines(): |
| 32 | # 对每一行数据按切割福','进行切割,返回字段列表 |
| 33 | curLine = line.strip().split(',') |
| 34 | |
| 35 | # Mnsit有0-9是个标记,由于是二分类任务,所以将>=5的作为1,<5为-1 |
| 36 | if int(curLine[0]) >= 5: |
| 37 | labelArr.append(1) |
| 38 | else: |
| 39 | labelArr.append(-1) |
| 40 | #存放标记 |
| 41 | #[int(num) for num in curLine[1:]] -> 遍历每一行中除了以第一哥元素(标记)外将所有元素转换成int类型 |
| 42 | #[int(num)/255 for num in curLine[1:]] -> 将所有数据除255归一化(非必须步骤,可以不归一化) |
| 43 | dataArr.append([int(num)/255 for num in curLine[1:]]) |
| 44 | |
| 45 | #返回data和label |
| 46 | return dataArr, labelArr |
| 47 | |
| 48 | def perceptron(dataArr, labelArr, iter=50): |
| 49 | ''' |
no outgoing calls
no test coverage detected