感知器基础原理及python实现
时间:2019-09-14
本文章向大家介绍感知器基础原理及python实现,主要包括感知器基础原理及python实现使用实例、应用技巧、基本知识点总结和需要注意事项,具有一定的参考价值,需要的朋友可以参考一下。
简单版本,按照李航的《统计学习方法》的思路编写
数据采用了著名的sklearn自带的iries数据,最优化求解采用了SGD算法。
预处理增加了标准化操作。
''' perceptron classifier created on 2019.9.14 author: vince ''' import pandas import numpy import logging import matplotlib.pyplot as plt from sklearn.datasets import load_iris from sklearn.model_selection import train_test_split from sklearn.metrics import accuracy_score ''' perceptron classifier Attributes w: ld-array = weights after training l: list = number of misclassification during each iteration ''' class Perceptron: def __init__(self, eta = 0.01, iter_num = 50, batch_size = 1): ''' eta: float = learning rate (between 0.0 and 1.0). iter_num: int = iteration over the training dataset. batch_size: int = gradient descent batch number, if batch_size == 1, used SGD; if batch_size == 0, use BGD; else MBGD; ''' self.eta = eta; self.iter_num = iter_num; self.batch_size = batch_size; def train(self, X, Y): ''' train training data. X:{array-like}, shape=[n_samples, n_features] = Training vectors, where n_samples is the number of training samples and n_features is the number of features. Y:{array-like}, share=[n_samples] = traget values. ''' self.w = numpy.zeros(1 + X.shape[1]); self.l = numpy.zeros(self.iter_num); for iter_index in range(self.iter_num): for sample_index in range(X.shape[0]): if (self.activation(X[sample_index]) != Y[sample_index]): logging.debug("%s: pred(%s), label(%s), %s, %s" % (sample_index, self.net_input(X[sample_index]) , Y[sample_index], X[sample_index, 0], X[sample_index, 1])); self.l[iter_index] += 1; for sample_index in range(X.shape[0]): if (self.activation(X[sample_index]) != Y[sample_index]): self.w[0] += self.eta * Y[sample_index]; self.w[1:] += self.eta * numpy.dot(X[sample_index], Y[sample_index]); break; logging.info("iter %s: %s, %s, %s, %s" % (iter_index, self.w[0], self.w[1], self.w[2], self.l[iter_index])); def activation(self, x): return numpy.where(self.net_input(x) >= 0.0 , 1 , -1); def net_input(self, x): return numpy.dot(x, self.w[1:]) + self.w[0]; def predict(self, x): return self.activation(x); def main(): logging.basicConfig(level = logging.INFO, format = '%(asctime)s %(filename)s[line:%(lineno)d] %(levelname)s %(message)s', datefmt = '%a, %d %b %Y %H:%M:%S'); iris = load_iris(); features = iris.data[:99, [0, 2]]; # normalization features_std = numpy.copy(features); for i in range(features.shape[1]): features_std[:, i] = (features_std[:, i] - features[:, i].mean()) / features[:, i].std(); labels = numpy.where(iris.target[:99] == 0, -1, 1); # 2/3 data from training, 1/3 data for testing train_features, test_features, train_labels, test_labels = train_test_split( features_std, labels, test_size = 0.33, random_state = 23323); logging.info("train set shape:%s" % (str(train_features.shape))); p = Perceptron(); p.train(train_features, train_labels); test_predict = numpy.array([]); for feature in test_features: predict_label = p.predict(feature); test_predict = numpy.append(test_predict, predict_label); score = accuracy_score(test_labels, test_predict); logging.info("The accruacy score is: %s "% (str(score))); #plot x_min, x_max = train_features[:, 0].min() - 1, train_features[:, 0].max() + 1; y_min, y_max = train_features[:, 1].min() - 1, train_features[:, 1].max() + 1; plt.xlim(x_min, x_max); plt.ylim(y_min, y_max); plt.xlabel("width"); plt.ylabel("heigt"); plt.scatter(train_features[:, 0], train_features[:, 1], c = train_labels, marker = 'o', s = 10); k = - p.w[1] / p.w[2]; d = - p.w[0] / p.w[2]; plt.plot([x_min, x_max], [k * x_min + d, k * x_max + d], "go-"); plt.show(); if __name__ == "__main__": main();
原文地址:https://www.cnblogs.com/thsss/p/11519846.html
- VIM常见用法总结
- Spring Cloud构建微服务架构:服务消费者
- android微信登录,分享
- 注册会计师带你用Python进行探索性风险分析(二)
- Android监听自身卸载,弹出用户反馈调查
- Spring Boot 1.5.x新特性:动态修改日志级别
- XMPP客户端库Smack 4.0.6版开发之二
- Spring Cloud实战小贴士:版本依赖关系
- 如何优雅的用Python做接口自动化测试
- 忘记oracle的sys用户密码怎么修改以及Oracle 11g 默认用户名和密码
- hibernate链接数据库链接池c3p0配置
- Oracle中session和processes的设置
- ssh相关原理学习与常见错误总结
- PyQt5 GUI应用程序工具包入门(1)
- 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 数组属性和方法
- 宿舍(寝室)管理系统设计与实现 | 附 演示、源码地址
- Oracle字符集检查和修改
- Vue3 DOM Diff 核心算法解析
- PHP的LZF压缩扩展工具
- Python函数定义及参数详解
- 代码失而复得心塞往事 - git stash命令
- 如何通过 Shell 监控异常等待事件和活跃会话
- PHP中环境变量的操作
- 一文读懂JAVA并发容器类ConcurrentHashMap
- Creator3D新版本震撼来袭
- SpringBoot源码学习(十)-Spring类级别注解解析原理
- 从安全切面到Security Mesh
- SpringBoot源码学习(十一) - bean的实例化过程
- 每天一杯力扣快乐水
- Typescript的tsconfig.json