机器学习基石笔记:Homework #2 decision stump相关习题
原文地址:http://www.jianshu.com/p/4bc01760ac20
问题描述


程序实现
17-18
# coding: utf-8
import numpy as np
import matplotlib.pyplot as plt
def sign(n):
if(n>0):
return 1
else:
return -1
def gen_data():
data_X=np.random.uniform(-1,1,(20,1))# [-1,1)
data_Y=np.zeros((20,1))
idArray=np.random.permutation([i for i in range(20)])
for i in range(20):
if(i<20*0.2):
data_Y[idArray[i]][0]=-sign(data_X[idArray[i]][0])
else:
data_Y[idArray[i]][0] = sign(data_X[idArray[i]][0])
data=np.concatenate((data_X,data_Y),axis=1)
return data
def decision_stump(dataArray):
minErrors=20
min_s_theta_list=[]
num_data=dataArray.shape[0]
data=dataArray.tolist()
data.sort(key=lambda x:x[0])
for s in [-1.0,1.0]:
for i in range(num_data):
if(i==num_data-1):
theta=(data[i][0]+1.0)/2
else:
theta=(data[i][0]+data[i+1][0])/2
errors=0
for i in range(20):
pred=s*sign(data[i][0]-theta)
if(pred!=data[i][1]):
errors+=1
if(minErrors>errors):
minErrors=errors
min_s_theta_list=[]
elif(minErrors<errors):
continue
min_s_theta_list.append((s, theta))
i=np.random.randint(low=0,high=len(min_s_theta_list))
min_s,min_theta=min_s_theta_list[i]
return minErrors,min_s,min_theta
def computeEinEout(minErrors,min_s,min_theta):
Ein=minErrors/20
Eout=0.5+0.3*min_s*(abs(min_theta)-1)
return Ein,Eout
if __name__=="__main__":
Ein_list=[]
Eout_list=[]
for i in range(5000):
dataArray=gen_data()
minErrors,min_s,min_theta=decision_stump(dataArray)
Ein,Eout=computeEinEout(minErrors,min_s,min_theta)
Ein_list.append(Ein)
Eout_list.append(Eout)
# show results
# 17 & 18
print("the average Ein: ",sum(Ein_list)/5000)
print("the average Eout: ",sum(Eout_list)/5000)
plt.figure(figsize=(16,6))
plt.subplot(121)
plt.hist(Ein_list)
plt.xlabel("Ein")
plt.ylabel("frequency")
plt.subplot(122)
plt.hist(Eout_list)
plt.xlabel("Eout")
plt.ylabel("frequency")
plt.savefig("EinEout.png")
19-20
# coding: utf-8
import numpy as np
def read_data(dataFile):
with open(dataFile, 'r') as file:
data_list = []
for line in file.readlines():
line = line.strip().split()
data_list.append([float(l) for l in line])
data_array = np.array(data_list)
return data_array
def predict(s,theta,dataX):
num_data=dataX.shape[0]
res=s*np.sign(dataX-theta)
return res
def decision_stump(dataArray):
min_s_theta_list=[]
num_data=dataArray.shape[0]
minErrors=num_data
data=dataArray.tolist()
data.sort(key=lambda x:x[0])
dataArray=np.array(data)
dataX=dataArray[:,0].reshape(num_data,1)
dataY=dataArray[:,1].reshape(num_data,1)
for s in [-1.0,1.0]:
for i in range(num_data):
if(i==num_data-1):
theta=(dataX[i][0]*2+1)/2
else:
theta=(dataX[i][0]+dataX[i+1][0])/2
pred=predict(s,theta,dataX)
errors=np.sum(pred!=dataY)
if(minErrors>errors):
minErrors=errors
min_s_theta_list=[]
elif(minErrors<errors):
continue
min_s_theta_list.append((s, theta))
i=np.random.randint(low=0,high=len(min_s_theta_list))
min_s,min_theta=min_s_theta_list[i]
return minErrors,min_s,min_theta
def best_of_best(candidate):
candidate.sort(key=lambda x:x[1])
counts=0
for i in range(len(candidate)):
if(candidate[i][1]!=candidate[0][1]):
break
counts+=1
i=np.random.randint(low=0,high=counts)
return candidate[i][0],candidate[i][1],candidate[i][2],candidate[i][3]
if __name__=="__main__":
data_array=read_data("hw2_train.dat")
num_data=data_array.shape[0]
num_dim=data_array.shape[1]-1
candidate=[]
dataY=data_array[:,-1].reshape(num_data,1)
for i in range(num_dim):
dataX=data_array[:,i].reshape(num_data,1)
min_errors,min_s,min_theta=decision_stump(np.concatenate((dataX,dataY),axis=1))
candidate.append([i,min_errors,min_s,min_theta])
min_id,min_errors,min_s,min_theta=best_of_best(candidate)
print("the optimal decision stump:\n","s: ",min_s,"\ntheta: ",min_theta)
print("the Ein of the optimal decision stump:\n",min_errors/num_data)
test_array=read_data("hw2_test.dat")
num_test=test_array.shape[0]
testY=test_array[:,-1].reshape(num_test,1)
num_dim=test_array.shape[1]-1
testX=test_array[:,min_id].reshape(num_test,1)
pred=predict(min_s,min_theta,testX)
print("the Eout of the optimal decision stump by Etest:\n",np.sum(pred!=testY)/num_test)
运行结果
17-18


19-20

机器学习基石笔记:Homework #2 decision stump相关习题的更多相关文章
- 机器学习基石笔记:Homework #1 PLA&PA相关习题
原文地址:http://www.jianshu.com/p/5b4a64874650 问题描述 程序实现 # coding: utf-8 import numpy as np import matpl ...
- 机器学习基石笔记:Homework #4 Regularization&Validation相关习题
原文地址:https://www.jianshu.com/p/3f7d4aa6a7cf 问题描述 程序实现 # coding: utf-8 import numpy as np import math ...
- 机器学习基石笔记:Homework #3 LinReg&LogReg相关习题
原文地址:http://www.jianshu.com/p/311141f2047d 问题描述 程序实现 13-15 # coding: utf-8 import numpy as np import ...
- 机器学习基石:Homework #0 SVD相关&常用矩阵求导公式
- 林轩田机器学习基石笔记1—The Learning Problem
机器学习分为四步: When Can Machine Learn? Why Can Machine Learn? How Can Machine Learn? How Can Machine Lear ...
- 机器学习基石笔记:01 The Learning Problem
原文地址:https://www.jianshu.com/p/bd7cb6c78e5e 什么时候适合用机器学习算法? 存在某种规则/模式,能够使性能提升,比如准确率: 这种规则难以程序化定义,人难以给 ...
- 机器学习基石笔记:04 Feasibility of Learning
原文地址:https://www.jianshu.com/p/f2f4d509060e 机器学习是设计算法\(A\),在假设集合\(H\)里,根据给定数据集\(D\),选出与实际模式\(f\)最为相近 ...
- 机器学习基石笔记:03 Types of Learning
原文地址:https://www.jianshu.com/p/86b2a9cef742 一.学习的分类 根据输出空间\(Y\):分类(二分类.多分类).回归.结构化(监督学习+输出空间有结构): 根据 ...
- 机器学习技法笔记:09 Decision Tree
Roadmap Decision Tree Hypothesis Decision Tree Algorithm Decision Tree Heuristics in C&RT Decisi ...
随机推荐
- CentOS添加SSH登录提示
前言 使用阿里云服务器的应该都注意到每次ssh登录后都能看见类似下面这样的欢迎语: Last login:xxxxxxxxxxxxx Welcome to Alibaba Cloud Elastic ...
- 记一次tomcat自动退出问题
问题 环境: centos/tomcat8/jdk1.8 最近遇到部署在服务器的tomcat总是过一段时间就自动结束进程 ; 通过监控tomcat 日志文件(tail -f ./logs/catali ...
- MySQL,Oracle,PostgreSQL,mongoDB 通过web方式管理维护, 提高开发及运维效率
在开发及项目运维中,对数据库的操作大家目前都是使用客户端工具进行操作,例如MySQL的客户端工具navicat,Oracle的客户端工具 PL/SQL Developer, MSSQL的客户端工具查询 ...
- cf1060C. Maximum Subrectangle(思维 枚举)
题意 题目链接 Sol 好好读题 => 送分题 不好好读题 => 送命题 开始想了\(30\)min数据结构发现根本不会做,重新读了一遍题发现是个傻逼题... \(C_{i, j} = a ...
- 洛谷P4781 【模板】拉格朗日插值(拉格朗日插值)
题意 题目链接 Sol 记得NJU有个特别强的ACM队叫拉格朗,总感觉少了什么.. 不说了直接扔公式 \[f(x) = \sum_{i = 1}^n y_i \prod_{j \not = i} \f ...
- js 控制页面跳转的5种方法
js 控制页面跳转的5种方法 编程式导航: 点击跳转路由,称编程式导航,用js编写代码跳转. History是bom中的 History.back是回退一页 Histiory.go(1)前进一页 Hi ...
- php中http_build_query函数
http_build_query ( array $formdata [, string $numeric_prefix ] ) 使用给出的关联(或下标)数组生成一个经过 URL-encode 的请求 ...
- 毕向东_Java基础视频教程第19天_IO流(01~05)
第19天-01-IO流(BufferedWriter) 字符流的缓冲区 缓冲区的出现提高了对数据的读写效率. 对应类缓冲区要结合流才可以使用. BufferedWriter BufferedReade ...
- 排查在 Azure 中新建 Windows 虚拟机时遇到的经典部署问题
尝试创建新的 Azure 虚拟机 (VM) 时,遇到的常见错误是预配失败或分配失败. 当由于准备步骤不当,或者在从门户捕获映像期间选择了错误的设置而导致 OS 映像无法加载时,将发生预配失败. 当群集 ...
- Python函数汇总(陆续更新中...)
range的用法 函数原型:range(start, end, scan): 参数含义: start:计数从start开始.默认是从0开始.例如range(5)等价于range(0, 5); end: ...