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])

在这里插入图片描述

Logo

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

更多推荐