python实现神经网络感知器算法
作者:海峰-清欢 发布时间:2021-03-06 11:23:39
标签:python,神经网络,感知器
现在我们用python代码实现感知器算法。
# -*- coding: utf-8 -*-
import numpy as np
class Perceptron(object):
"""
eta:学习率
n_iter:权重向量的训练次数
w_:神经分叉权重向量
errors_:用于记录神经元判断出错次数
"""
def __init__(self, eta=0.01, n_iter=2):
self.eta = eta
self.n_iter = n_iter
pass
def fit(self, X, y):
"""
输入训练数据培训神经元
X:神经元输入样本向量
y: 对应样本分类
X:shape[n_samples,n_features]
x:[[1,2,3],[4,5,6]]
n_samples = 2 元素个数
n_features = 3 子向量元素个数
y:[1,-1]
初始化权重向量为0
加一是因为前面算法提到的w0,也就是步调函数阈值
"""
self.w_ = np.zeros(1 + X.shape[1])
self.errors_ = []
for _ in range(self.n_iter):
errors = 0
"""
zip(X,y) = [[1,2,3,1],[4,5,6,-1]]
xi是前面的[1,2,3]
target是后面的1
"""
for xi, target in zip(X, y):
"""
predict(xi)是计算出来的分类
"""
update = self.eta * (target - self.predict(xi))
self.w_[1:] += update * xi
self.w_[0] += update
print update
print xi
print self.w_
errors += int(update != 0.0)
self.errors_.append(errors)
pass
def net_input(self, X):
"""
z = w0*1+w1*x1+....Wn*Xn
"""
return np.dot(X, self.w_[1:]) + self.w_[0]
def predict(self, X):
return np.where(self.net_input(X) >= 0, 1, -1)
if __name__ == '__main__':
datafile = '../data/iris.data.csv'
import pandas as pd
df = pd.read_csv(datafile, header=None)
import matplotlib.pyplot as plt
import numpy as np
y = df.loc[0:100, 4].values
y = np.where(y == "Iris-setosa", 1, -1)
X = df.iloc[0:100, [0, 2]].values
# plt.scatter(X[:50, 0], X[:50, 1], color="red", marker='o', label='setosa')
# plt.scatter(X[50:100, 0], X[50:100, 1], color="blue", marker='x', label='versicolor')
# plt.xlabel("hblength")
# plt.ylabel("hjlength")
# plt.legend(loc='upper left')
# plt.show()
pr = Perceptron()
pr.fit(X, y)
其中数据为
控制台输出为
你们跑代码的时候把n_iter设置大点,我这边是为了看每次执行for循环时方便查看数据变化。
来源:http://blog.csdn.net/u013692888/article/details/76999252


猜你喜欢
- opts, args = getopt.getopt(sys.argv[1:], "t:s:h", ["wal
- 这周心血来潮,翻看了现在比较流行的几个JS脚本框架的底层代码,虽然是走马观花,但也受益良多,感叹先人们的伟大……感叹是为了缓解严肃的气氛并引
- 1、去除一个数组中的重复元素:使用grep函数代码片段: 代码:my @array = ( 'a', 'b'
- 我就废话不多说了,大家还是直接看代码吧~#!/usr/bin/env python# -*- coding: utf-8 -*-import
- ASPJPEG组件是Persits出品的共享软件,试用期为30天,您可以在这里下载:http://www.persits.com/aspjp
- 先来看一个简单的利用python调用sqlplus来输出结果的例子:import osimport sysfrom subprocess i
- 1.从官网下载mysql-5.7.21-windowx64.zip mysql下载页面2.解压到合适的位置(E:\mysql) 这名字是我改
- 找到python3的安装路径python3自带一个把python2代码转换成python3代码的程序,叫"2to3"我们
- 什么是1433端口 1433端口,是SQL Server默认的端口,SQL Server服务使用两个端口:TCP-1433、UDP-1434
- Python在读取文件内容时的路径问题,值得深究一下.我想讨论的重点还是在绝对路径上面.在这之前我们先看一下1:相对路径这张图演示了在相对路
- 我就废话不多说,看代码!import numpy as npimport matplotlib.pyplot as pltimport pa
- 思维导图:效果(语句版):源码:# -*- coding: utf-8 -*-"""Created
- pipenv 是Kenneth Reitz大神的作品,能够有效管理Python多个环境,各种包。过去我们一般用virtualenv搭建虚拟环
- 最近,接手的项目里,提供的数据文件格式简直让人看不下去,使用pandas打不开
- 模仿IE自动完成功能,支持Firefox.支持方向键操作运行代码框<!DOCTYPE HTML PUBLIC "-//W3C
- 前言本博客默认读者对神经网络与Tensorflow有一定了解,对其中的一些术语不再做具体解释。并且本博客主要以图片数据为例进行介绍,如有错误
- eval(“1+2”),-> 3 动态判断源代码中的字符串是一种很强大的语
- 一.CSRF简介CSRF是什么? CSRF(Cross-site request forgery),中文名称:跨站
- 本文完全利用numpy实现一个简单的BP神经网络,由于是做regression而不是classification,因此在这里输出层选取的激励
- 百度有啊2009年情人节logo——大纸袋GG给大纸袋MM送了枝玫瑰花,大纸袋MM奖励了大纸袋GG一个吻,好可爱!淘宝网2009年情人节lo