Skip to content

Commit 3bb3a88

Browse files
Create tree.py
1 parent af751e9 commit 3bb3a88

1 file changed

Lines changed: 131 additions & 0 deletions

File tree

Chapter_5 Random Forest/tree.py

Lines changed: 131 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,131 @@
1+
# coding:UTF-8
2+
'''
3+
Date:20161030
4+
@author: zhaozhiyong
5+
'''
6+
from math import pow
7+
8+
class node:
9+
'''树的节点的类
10+
'''
11+
def __init__(self, fea=-1, value=None, results=None, right=None, left=None):
12+
self.fea = fea # 用于切分数据集的属性的列索引值
13+
self.value = value # 设置划分的值
14+
self.results = results # 存储叶节点所属的类别
15+
self.right = right # 右子树
16+
self.left = left # 左子树
17+
18+
def split_tree(data, fea, value):
19+
'''根据特征fea中的值value将数据集data划分成左右子树
20+
input: data(list):数据集
21+
fea(int):待分割特征的索引
22+
value(float):待分割的特征的具体值
23+
output: (set1,set2)(tuple):分割后的左右子树
24+
'''
25+
set_1 = []
26+
set_2 = []
27+
for x in data:
28+
if x[fea] >= value:
29+
set_1.append(x)
30+
else:
31+
set_2.append(x)
32+
return (set_1, set_2)
33+
34+
def label_uniq_cnt(data):
35+
'''统计数据集中不同的类标签label的个数
36+
input: data(list):原始数据集
37+
output: label_uniq_cnt(int):样本中的标签的个数
38+
'''
39+
label_uniq_cnt = {}
40+
41+
for x in data:
42+
label = x[len(x) - 1] # 取得每一个样本的类标签label
43+
if label not in label_uniq_cnt:
44+
label_uniq_cnt[label] = 0
45+
label_uniq_cnt[label] = label_uniq_cnt[label] + 1
46+
return label_uniq_cnt
47+
48+
def cal_gini_index(data):
49+
'''计算给定数据集的Gini指数
50+
input: data(list):树中
51+
output: gini(float):Gini指数
52+
'''
53+
total_sample = len(data) # 样本的总个数
54+
if len(data) == 0:
55+
return 0
56+
label_counts = label_uniq_cnt(data) # 统计数据集中不同标签的个数
57+
58+
# 计算数据集的Gini指数
59+
gini = 0
60+
for label in label_counts:
61+
gini = gini + pow(label_counts[label], 2)
62+
63+
gini = 1 - float(gini) / pow(total_sample, 2)
64+
return gini
65+
66+
def build_tree(data):
67+
'''构建树
68+
input: data(list):训练样本
69+
output: node:树的根结点
70+
'''
71+
# 构建决策树,函数返回该决策树的根节点
72+
if len(data) == 0:
73+
return node()
74+
75+
# 1、计算当前的Gini指数
76+
currentGini = cal_gini_index(data)
77+
78+
bestGain = 0.0
79+
bestCriteria = None # 存储最佳切分属性以及最佳切分点
80+
bestSets = None # 存储切分后的两个数据集
81+
82+
feature_num = len(data[0]) - 1 # 样本中特征的个数
83+
# 2、找到最好的划分
84+
for fea in range(0, feature_num):
85+
# 2.1、取得fea特征处所有可能的取值
86+
feature_values = {} # 在fea位置处可能的取值
87+
for sample in data: # 对每一个样本
88+
feature_values[sample[fea]] = 1 # 存储特征fea处所有可能的取值
89+
90+
# 2.2、针对每一个可能的取值,尝试将数据集划分,并计算Gini指数
91+
for value in feature_values.keys(): # 遍历该属性的所有切分点
92+
# 2.2.1、 根据fea特征中的值value将数据集划分成左右子树
93+
(set_1, set_2) = split_tree(data, fea, value)
94+
# 2.2.2、计算当前的Gini指数
95+
nowGini = float(len(set_1) * cal_gini_index(set_1) + \
96+
len(set_2) * cal_gini_index(set_2)) / len(data)
97+
# 2.2.3、计算Gini指数的增加量
98+
gain = currentGini - nowGini
99+
# 2.2.4、判断此划分是否比当前的划分更好
100+
if gain > bestGain and len(set_1) > 0 and len(set_2) > 0:
101+
bestGain = gain
102+
bestCriteria = (fea, value)
103+
bestSets = (set_1, set_2)
104+
105+
# 3、判断划分是否结束
106+
if bestGain > 0:
107+
right = build_tree(bestSets[0])
108+
left = build_tree(bestSets[1])
109+
return node(fea=bestCriteria[0], value=bestCriteria[1], \
110+
right=right, left=left)
111+
else:
112+
return node(results=label_uniq_cnt(data)) # 返回当前的类别标签作为最终的类别标签
113+
114+
def predict(sample, tree):
115+
'''对每一个样本sample进行预测
116+
input: sample(list):需要预测的样本
117+
tree(类):构建好的分类树
118+
output: tree.results:所属的类别
119+
'''
120+
# 1、只是树根
121+
if tree.results != None:
122+
return tree.results
123+
else:
124+
# 2、有左右子树
125+
val_sample = sample[tree.fea]
126+
branch = None
127+
if val_sample >= tree.value:
128+
branch = tree.right
129+
else:
130+
branch = tree.left
131+
return predict(sample, branch)

0 commit comments

Comments
 (0)