CART回归树后剪枝_李航统计学习方法
·
import numpy as np
import collections
class Node:
def __init__(self,fea = None,val=None,left = None,right = None,res = None,leaf = False,MSE = None,Num = None):
#val:划分值
#fea:划分变量
#res:节点的预测值
self.fea = fea
self.val = val
self.left = left
self.right = right
self.res = res
self.leaf = leaf
self.MSE = MSE
self.Num = Num
class CART:
def __init__(self,epsilon = 0,min_sample=1):
self.epsilon = epsilon
self.min_sample = min_sample#叶节点最少样本数
self.tree = None
self.X_data = None
self.Y_data = None
def MSE(self,y_data):
c = np.mean(y_data)
return np.sum((y_data-c)**2)
def getFeaMSE(self,y1,y2):
return self.MSE(y1)+self.MSE(y2)
def leaf_val(self,y_data):
return y_data.mean()
def split(self,fea,val,X_data):
set1_ind = np.where(X_data[:,fea]<=val)[0]
set2_ind = np.where(X_data[:,fea]>val)[0]
return set1_ind,set2_ind
def getBestSplit(self,X_data,y_data):
n = X_data.shape[1]
best_MSE = self.MSE(y_data)
best_fea = None
best_val = None
best_index = None
for fea in range(n):
for val in X_data[:,fea]:
set1_ind,set2_ind = self.split(fea,val,X_data)
if len(set1_ind)<2 or len(set2_ind)<2:#停止条件1,叶节点数量大小
continue
MSE = self.getFeaMSE(y_data[set1_ind],y_data[set2_ind])
if MSE<best_MSE:
best_MSE = MSE
best_fea = fea
best_val = val
best_index = (set1_ind,set2_ind)
return best_MSE,best_fea,best_val,best_index
def buidTree(self,X_data,Y_data):
self.X_data = X_data
self.Y_data = Y_data
if Y_data.shape[0] < self.min_sample:
return Node(res=self.leaf_val(Y_data),leaf=True,MSE=self.MSE(Y_data))
best_MSE,best_fea,best_val,best_index = self.getBestSplit(X_data,Y_data)
if best_index is None:
return Node(res = self.leaf_val(Y_data),leaf=True,MSE=self.MSE(Y_data))
if best_MSE<self.epsilon:
return Node(res=self.leaf_val(Y_data),leaf=True,MSE=self.MSE(Y_data))
else:
lindex,rindex = best_index
left = self.buidTree(X_data[lindex,:],Y_data[lindex])
right = self.buidTree(X_data[rindex,:],Y_data[rindex])
return Node(fea=best_fea,val=best_val,left=left,right=right,res=self.leaf_val(Y_data),MSE=self.MSE(Y_data))
def RegressionTree(self,X_data,Y_data):
self.tree = self.buidTree(X_data,Y_data)
print("Finished...")
def predict(self,x):
root = self.tree
while (root.leaf is False):
fea = root.fea
val = root.val
if x[fea]<=val:
root = root.left
else:
root = root.right
return root.res
#################################
def countleaf(self,node):#前序遍历
stack = [node]
ans = 0
Cat = 0
while stack:
curnode = stack.pop()
if curnode.leaf is True:
ans += 1
Cat += curnode.MSE
else:
stack.append(curnode.right)
stack.append(curnode.left)
return ans,Cat
def ID(self):
num = 0
Que = collections.deque([self.tree])
while Que:
level = len(Que)
for i in range(level):
curnode = Que.popleft()
if curnode.leaf is False:
curnode.Num = num
num += 1
Que.append(curnode.left)
Que.append(curnode.right)
#写的很不好的剪枝...很多操作都重复了
def prun(self,tree):
root = tree
Que = collections.deque([root])
Tnode = None
alpha = np.Inf
optnum = None
while Que:
level = len(Que)
for i in range(level):
curnode = Que.popleft()
if curnode.leaf is False:
Que.append(curnode.left)
Que.append(curnode.right)
Ct = curnode.MSE
Tt,CTt = self.countleaf(curnode)
gt = (Ct-CTt)/(Tt-1)
if gt<alpha:
alpha = gt
optnum = curnode.Num
return optnum,alpha
def Tk(self,tree,num):
node = tree
copynode = Node()
Que = collections.deque([node])
Que2 = collections.deque([copynode])
while Que:
level = len(Que)
for i in range(level):
curnode = Que.popleft()
copycurnode = Que2.popleft()
copycurnode.res = curnode.res
copycurnode.leaf = curnode.leaf
copycurnode.MSE = curnode.MSE
copycurnode.Num = curnode.Num
copycurnode.fea = curnode.fea
copycurnode.val = curnode.val
if curnode.leaf is False:
Que.append(curnode.left)
Que.append(curnode.right)
copycurnode.left = Node()
copycurnode.right = Node()
Que2.append(copycurnode.left)
Que2.append(copycurnode.right)
Que3 = collections.deque([copynode])
stop = False
while Que3:
if stop is True:
break
level = len(Que3)
for i in range(level):
curnode = Que3.popleft()
if curnode.leaf is False:
if num == curnode.Num:
curnode.left = None
curnode.right = None
curnode.leaf = True
curnode.Num = None
stop = True
break
else:
Que3.append(curnode.left)
Que3.append(curnode.right)
return copynode
def getTreeSequence(self):
root = self.tree
self.ID()
T = [self.tree]
a = [0]
while root.leaf is False:
optnum,alpha = self.prun(root)
a.append(alpha)
root = self.Tk(root,optnum)
T.append(root)
return T,a
写的不是很好,很多操作都有重复计算,因为sklearn库没有后剪枝,只能自己写了一个不太好的。。接下来使用测试集挑选出最优的子树
import matplotlib.pyplot as plt
X_data_raw= np.linspace(-3,3,50)
np.random.shuffle(X_data_raw)
X_data = X_data_raw.reshape(-1,1)
y_data = np.sin(X_data_raw)+0.1*np.random.randn(X_data_raw.shape[0])
model = CART(epsilon=1e-4,min_sample=2)
model.RegressionTree(X_data=X_data,Y_data=y_data)
X_test_raw = np.linspace(-3.2,3.2,50)
Y_test = np.sin(X_test_raw)+0.1*np.random.randn(X_test_raw.shape[0])
X_test = X_test_raw.reshape(-1,1)
Y_pred = []
for test in X_test:
Y_pred.append(model.predict(test))
plt.scatter(X_test,Y_test)
plt.scatter(X_test,Y_pred,marker= '*')
plt.show()

subTrees,alpha = model.getTreeSequence()
mse = []
for subtree in subTrees:
tree = CART(epsilon=1e-4,min_sample=2)
tree.tree = subtree
prediction = []
for test in X_test:
prediction.append(tree.predict(test))
mse.append(np.sum((prediction-Y_test)**2))
#第四颗子树最好
mse

接下来是打印各个子树
def printTree(tree):
node = tree
Que = collections.deque([node])
ans = []
while Que:
level = len(Que)
layer = []
for i in range(level):
curnode = Que.popleft()
layer.append(curnode.Num)
if curnode.leaf is False:
Que.append(curnode.left)
Que.append(curnode.right)
ans.append(layer)
return ans
printTree(subTrees[0])

printTree(subTrees[3])

魔乐社区(Modelers.cn) 是一个中立、公益的人工智能社区,提供人工智能工具、模型、数据的托管、展示与应用协同服务,为人工智能开发及爱好者搭建开放的学习交流平台。社区通过理事会方式运作,由全产业链共同建设、共同运营、共同享有,推动国产AI生态繁荣发展。
更多推荐


所有评论(0)