python实现多层感知器
时间:2019-01-18
这篇文章主要为大家详细介绍了python实现多层感知器的相关资料,具有一定的参考价值,感兴趣的小伙伴们可以参考一下
写了个多层感知器,用bp梯度下降更新,拟合正弦曲线,效果凑合。
# -*- coding: utf-8 -*- import numpy as np import matplotlib.pyplot as plt def sigmod(z): return 1.0 / (1.0 + np.exp(-z)) class mlp(object): def __init__(self, lr=0.1, lda=0.0, te=1e-5, epoch=100, size=None): self.learningRate = lr self.lambda_ = lda self.thresholdError = te self.maxEpoch = epoch self.size = size self.W = [] self.b = [] self.init() def init(self): for i in xrange(len(self.size)-1): self.W.append(np.mat(np.random.uniform(-0.5, 0.5, size=(self.size[i+1], self.size[i])))) self.b.append(np.mat(np.random.uniform(-0.5, 0.5, size=(self.size[i+1], 1)))) def forwardPropagation(self, item=None): a = [item] for wIndex in xrange(len(self.W)): a.append(sigmod(self.W[wIndex]*a[-1]+self.b[wIndex])) """ print "-----------------------------------------" for i in a: print i.shape, print for i in self.W: print i.shape, print for i in self.b: print i.shape, print print "-----------------------------------------" """ return a def backPropagation(self, label=None, a=None): # print "backPropagation--------------------begin" delta = [(a[-1]-label)*a[-1]*(1.0-a[-1])] for i in xrange(len(self.W)-1): abc = np.multiply(a[-2-i], 1-a[-2-i]) cba = np.multiply(self.W[-1-i].T*delta[-1], abc) delta.append(cba) """ print "++++++++++++++delta++++++++++++++++++++" print "len(delta):", len(delta) for ii in delta: print ii.shape, print "\n=======================================" """ for j in xrange(len(delta)): ads = delta[j]*a[-2-j].T # print self.W[-1-j].shape, ads.shape, self.b[-1-j].shape, delta[j].shape self.W[-1-j] = self.W[-1-j]-self.learningRate*(ads+self.lambda_*self.W[-1-j]) self.b[-1-j] = self.b[-1-j]-self.learningRate*delta[j] """print "=======================================1234" for ij in self.b: print ij.shape, print """ # print "backPropagation--------------------finish" error = 0.5*(a[-1]-label)**2 return error def train(self, input_=None, target=None, show=10): for ep in xrange(self.maxEpoch): error = [] for itemIndex in xrange(input_.shape[1]): a = self.forwardPropagation(input_[:, itemIndex]) e = self.backPropagation(target[:, itemIndex], a) error.append(e[0, 0]) tt = sum(error)/len(error) if tt < self.thresholdError: print "Finish {0}: ".format(ep), tt return elif ep % show == 0: print "epoch {0}: ".format(ep), tt def sim(self, inp=None): return self.forwardPropagation(item=inp)[-1] if __name__ == "__main__": tt = np.arange(0, 6.28, 0.01) labels = np.zeros_like(tt) print tt.shape """ for po in xrange(tt.shape[0]): if tt[po] < 4: labels[po] = 0.0 elif 8 > tt[po] >= 4: labels[po] = 0.25 elif 12 > tt[po] >= 8: labels[po] = 0.5 elif 16 > tt[po] >= 12: labels[po] = 0.75 else: labels[po] = 1.0 """ tt = np.mat(tt) labels = np.sin(tt)*0.5+0.5 labels = np.mat(labels) model = mlp(lr=0.2, lda=0.0, te=1e-5, epoch=500, size=[1, 6, 6, 6, 1]) print tt.shape, labels.shape print len(model.W), len(model.b) print model.train(input_=tt, target=labels, show=10) sims = [model.sim(tt[:, idx])[0, 0] for idx in xrange(tt.shape[1])] xx = tt.tolist()[0] plt.figure() plt.plot(xx, labels.tolist()[0], xx, sims, 'r') plt.show()
效果图:
以上就是本文的全部内容,希望对大家的学习有所帮助,也希望大家多多支持脚本之家。
- 高级软件工程师(面试题)
- 高级软件工程师 2016-9月更新
- Httpclient 调用 HTTPS 加密通道的Restful服务
- 使用 Jersey 调用 Restful 服务
- 【学术】将吴恩达的第一个深度神经网络应用于泰坦尼克生存数据集
- 使用 HttpClient 调用 Restful 接口
- 元宵佳节:看Oracle技术粉们用SQL画团圆
- java 脚本引擎
- 不怕学不会 使用TensorFlow从零开始构建卷积神经网络
- 微信公众平台增加批量获取用户基本信息接口
- 谈网络适配器
- 【框架】为降低机器学习开发者门槛,苹果发布了Turi Create框架
- 新闻数据库分表案例
- 建立智能的解决方案:将TensorFlow用于声音分类
- JavaScript 教程
- JavaScript 编辑工具
- JavaScript 与HTML
- JavaScript 与Java
- JavaScript 数据结构
- JavaScript 基本数据类型
- JavaScript 特殊数据类型
- JavaScript 运算符
- JavaScript typeof 运算符
- JavaScript 表达式
- JavaScript 类型转换
- JavaScript 基本语法
- JavaScript 注释
- Javascript 基本处理流程
- Javascript 选择结构
- Javascript if 语句
- Javascript if 语句的嵌套
- Javascript switch 语句
- Javascript 循环结构
- Javascript 循环结构实例
- Javascript 跳转语句
- Javascript 控制语句总结
- Javascript 函数介绍
- Javascript 函数的定义
- Javascript 函数调用
- Javascript 几种特殊的函数
- JavaScript 内置函数简介
- Javascript eval() 函数
- Javascript isFinite() 函数
- Javascript isNaN() 函数
- parseInt() 与 parseFloat()
- escape() 与 unescape()
- Javascript 字符串介绍
- Javascript length属性
- javascript 字符串函数
- Javascript 日期对象简介
- Javascript 日期对象用途
- Date 对象属性和方法
- Javascript 数组是什么
- Javascript 创建数组
- Javascript 数组赋值与取值
- Javascript 数组属性和方法
- Go by Example 中文版: 写文件
- PWN:House Of Force
- Windwos10下使用VS2017搭建cocos2d-x 4.0开发环境
- JavaScript 中的函数式编程:函数,组合和柯里化
- 如何设置一个生产级别的高可用etcd集群
- NVIDIA Jetson nano可以处理4K相机吗?来验证编码性能吧(中)
- House Of Lore原理学习
- 使用 rush 进行命令并行处理
- 老生常谈 Spring Aop 日志收集与处理做的工具包,贼好用?
- Kaggle金牌得主的Python数据挖掘框架,机器学习基本流程都讲清楚了
- Go by Example 中文版: 行过滤器
- Elasticsearch重要知识点 | 选举流程详解
- 妹妹问我:Dubbo集群容错负载均衡
- 系统内核溢出提权
- 201312-3 最大的矩形(Python)