Python实现bp神经网络识别MNIST数据集

前言

训练时读入的是.mat格式的训练集,测试正确率时用的是png格式的图片

代码

1#!/usr/bin/env python3 2# coding=utf-8 3import math 4import sys 5import os 6import numpy as np 7from PIL import Image 8import scipy.io as sio 9 10 11def sigmoid(x): 12 return np.array(list(map(lambda i: 1 / (1 + math.exp(-i)), x))) 13 14 15def get_train_pattern(): 16 # 返回训练集的特征和标签 17 # current_dir = os.getcwd() 18 current_dir = "/home/lxp/F/developing_folder/intelligence_system/bpneuralnet/" 19 train = sio.loadmat(current_dir + "mnist_train.mat")["mnist_train"] 20 train_label = sio.loadmat( 21 current_dir + "mnist_train_labels.mat")["mnist_train_labels"] 22 train = np.where(train > 180, 1, 0) # 二值化 23 return train, train_label 24 25 26def get_test_pattern(): 27 # 返回测试集 28 # base_url = os.getcwd() + "/test/" 29 base_url = "/home/lxp/F/developing_folder/intelligence_system/bpneuralnet/mnist_test/" 30 test_img_pattern = [] 31 for i in range(10): 32 img_url = os.listdir(base_url + str(i)) 33 t = [] 34 for url in img_url: 35 img = Image.open(base_url + str(i) + "/" + url) 36 img = img.convert('1') # 二值化 37 img_array = np.asarray(img, 'i') # 转化为int数组 38 img_vector = img_array.reshape( 39 img_array.shape[0] * img_array.shape[1]) # 展开成一维数组 40 t.append(img_vector) 41 test_img_pattern.append(t) 42 return test_img_pattern 43 44 45class BPNetwork: 46 # 神经网络类 47 def __init__(self, in_count, hiden_count, out_count, in_rate, hiden_rate): 48 """ 49 50 :param in_count: 输入层数 51 :param hiden_count: 隐藏层数 52 :param out_count: 输出层数 53 :param in_rate: 输入层学习率 54 :param hiden_rate: 隐藏层学习率 55 """ 56 # 各个层的节点数量 57 self.in_count = in_count 58 self.hiden_count = hiden_count 59 self.out_count = out_count 60 61 # 输入层到隐藏层连线的权重随机初始化 62 self.w1 = 0.2 * \ 63 np.random.random((self.in_count, self.hiden_count)) - 0.1 64 65 # 隐藏层到输出层连线的权重随机初始化 66 self.w2 = 0.2 * \ 67 np.random.random((self.hiden_count, self.out_count)) - 0.1 68 69 # 隐藏层偏置向量 70 self.hiden_offset = np.zeros(self.hiden_count) 71 # 输出层偏置向量 72 self.out_offset = np.zeros(self.out_count) 73 74 # 输入层学习率 75 self.in_rate = in_rate 76 # 隐藏层学习率 77 self.hiden_rate = hiden_rate 78 79 def train(self, train_img_pattern, train_label): 80 if self.in_count != len(train_img_pattern[0]): 81 sys.exit("输入层维数与样本维数不等") 82 # for num in range(10): 83 # for num in range(10): 84 for i in range(len(train_img_pattern)): 85 if i % 5000 == 0: 86 print(i) 87 # 生成目标向量 88 target = [0] * 10 89 target[train_label[i][0]] = 1 90 # for t in range(len(train_img_pattern[num])): 91 # 前向传播 92 # 隐藏层值等于输入层*w1+隐藏层偏置 93 hiden_value = np.dot( 94 train_img_pattern[i], self.w1) + self.hiden_offset 95 hiden_value = sigmoid(hiden_value) 96 97 # 计算输出层的输出 98 out_value = np.dot(hiden_value, self.w2) + self.out_offset 99 out_value = sigmoid(out_value) 100 101 # 反向更新 102 error = target - out_value 103 # 计算输出层误差 104 out_error = out_value * (1 - out_value) * error 105 # 计算隐藏层误差 106 hiden_error = hiden_value * \ 107 (1 - hiden_value) * np.dot(self.w2, out_error) 108 109 # 更新w2,w2是j行k列的矩阵,存储隐藏层到输出层的权值 110 for k in range(self.out_count): 111 # 更新w2第k列的值,连接隐藏层所有节点到输出层的第k个节点的边 112 # 隐藏层学习率×输入层误差×隐藏层的输出值 113 self.w2[:, k] += self.hiden_rate * out_error[k] * hiden_value 114 115 # 更新w1 116 for j in range(self.hiden_count): 117 self.w1[:, j] += self.in_rate * \ 118 hiden_error[j] * train_img_pattern[i] 119 120 # 更新偏置向量 121 self.out_offset += self.hiden_rate * out_error 122 self.hiden_offset += self.in_rate * hiden_error 123 124 def test(self, test_img_pattern): 125 """ 126 测试神经网络的正确率 127 :param test_img_pattern[num][t]表示数字num的第t张图片 128 :return: 129 """ 130 right = np.zeros(10) 131 test_sum = 0 132 for num in range(10): # 10个数字 133 # print("正在识别", num) 134 num_count = len(test_img_pattern[num]) 135 test_sum += num_count 136 for t in range(num_count): # 数字num的第t张图片 137 hiden_value = np.dot( 138 test_img_pattern[num][t], self.w1) + self.hiden_offset 139 hiden_value = sigmoid(hiden_value) 140 out_value = np.dot(hiden_value, self.w2) + self.out_offset 141 out_value = sigmoid(out_value) 142 # print(out_value) 143 if np.argmax(out_value) == num: 144 # 识别正确 145 right[num] += 1 146 print("数字%d的识别正确率%f" % (num, right[num] / num_count)) 147 148 # 平均识别率 149 print("平均识别率为:", sum(right) / test_sum) 150 151 """ 152 def test1: 153 154 155 """ 156 157 158def run(): 159 # 读入训练集 160 train, train_label = get_train_pattern() 161 162 # 读入测试图片 163 test_pattern = get_test_pattern() 164 165 # 神经网络配置参数 166 in_count = 28 * 28 167 hiden_count = 6 168 out_count = 10 169 in_rate = 0.1 170 hiden_rate = 0.1 171 bpnn = BPNetwork(in_count, hiden_count, out_count, in_rate, hiden_rate) 172 bpnn.train(train, train_label) 173 bpnn.test(test_pattern) 174 175 # 单张测试 176 # 识别单独一张图片,返回识别结果 177 """ 178 while True: 179 img_name = input("输入要识别的图片\n") 180 base_url = "/home/lxp/F/developing_folder/intelligence_system/bpneuralnet/" 181 img_url = base_url + img_name 182 img = Image.open(img_url) 183 img = img.convert('1') # 二值化 184 img_array = np.asarray(img, 'i') # 转化为int数组 185 # 得到图片的特征向量 186 img_v = img_array.reshape(img_array.shape[0] * img_array.shape[1]) # 展开成一维数组 187 bpnn.test1(img_v) 188 189 """ 190 191 192if __name__ == "__main__": 193 run() 194 # train, train_label = get_train_pattern() 195 # print(train_label[5][0]) 196# test = get_test_pattern() 197

数据集下载: 链接: https://pan.baidu.com/s/1ldWTSqVUm6l1cc4EDOzHpQ 提取码: mm93

点赞
收藏

评论区

加载中...

相关推荐

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(

MySQL部分从库上面因为大量的临时表tmp_table造成慢查询

背景描述Time:20190124T00:08:14.70572408:00User@Host:@Id:Schema:sentrymetaLast_errno:0Killed:0Query_time:0.315758Lock_

皕杰报表之UUID

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

手写Java HashMap源码

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

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

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