基于Python实现的决策树模型; @ V5 W& G8 h0 U5 V' o
l3 m; ]0 T2 Z$ \ t3 J
决策树模型 . X* O# e7 C$ K3 P目录. ^; [, V6 K" Z E' ~
人工智能第五次实验报告 1) t4 K. a; @6 B6 c! s
决策树模型 1 : {# {+ g* Z* Z9 c# I( N一 、问题背景 1 : r; a' c4 Q1 F) ], D" j. R3 |1.1 监督学习简介 1 . v- F' X" l- N, `/ k7 Y# [3 `1.2 决策树简介 1) W. z3 J6 r" F4 ~
二 、程序说明 3+ x* s5 I5 B0 r* T# n0 d
2.1 数据载入 3* A' O' a4 D/ I' u
2.2 功能函数 3" J$ ?& j, z: ?( m( E
2.3 决策树模型 4 * w# I' p% A v6 Q三 、程序测试 57 I& k. m$ @; B. d7 A- I
3.1 数据集说明 5 7 W( ^8 S! H) b3.2 决策树生成和测试 6 7 O# v7 o+ D7 N3 A- Q3.3 学习曲线评估算法精度 7 ' G& K I9 x6 y8 i7 D k四 、实验总结 8 - i$ O; T, `( X附 录 - 程序代码 8 / S* m/ N( i. `& ~7 ~/ n一 、问题背景- Q% X) m( c3 Q( k- |
1.1监督学习简介 ; [0 t. l8 n! S6 p( f/ q机器学习的形式包括无监督学习,强化学习,监督学习和半监督学习;学习任务有分类、聚类和回 归等。# @7 U4 |/ y K( F. @9 M& o
监督学习通过观察“输入—输出”对,学习从输入到输出的映射函数。分类监督学习的训练集为标记 数据,本文转载自http://www.biyezuopin.vip/onews.asp?id=16720每一条数据有对应的”标签“,根据标签可以将数据集分为若干个类别。分类监督学习经训练集生 成一个学习模型,可以用来预测一条新数据的标签。 + a) z4 X% s/ _4 h) z) }7 B常见的监督学习模型有决策树、KNN算法、朴素贝叶斯和随机森林等。2 L9 |* }9 o/ ?8 L9 X
1.2决策树简介5 i; q; i7 H6 X) J
决策树归纳是一类简单的机器学习形式,它表示为一个函数,以属性值向量作为输入,返回一个决策。5 F/ Y8 l! S8 y M9 @
决策树的组成7 g8 [: o" O: c" _# i; Q" }0 h& D( s
决策树由内节点上的属性值测试、分支上的属性值和叶子节点上的输出值组成。0 i3 k& s9 X5 O
7 P$ h* S0 }- x8 ~( b' Y7 z9 y
import numpy as np * S/ x4 d# o( b0 h ^+ ~" e) Efrom matplotlib import pyplot as plt & B( Y8 d! O6 |5 G- ^- vfrom math import log $ j4 F# j/ ]5 Yimport pandas as pd) L& K. Z) o3 S# A7 H- a
import pydotplus as pdp ( T% a; i2 ?" {. ]/ E( R' j* o/ s
""" ) {+ r! w/ }$ @' @' n19335286 郑有为 q3 H7 x+ {( \8 y. e) P: Q: }% X' C
人工智能作业 - 实现ID3决策树# f$ O. Q4 c9 q ^: y
""" 1 A- g$ C2 H$ }: G7 F- V, [1 `6 Z9 I# [6 c1 T4 }" }/ k
nonce = 0 # 用来给节点一个全局ID7 J6 A' c. _4 {" F d( J$ L1 `
color_i = 0 1 K* y" H- B% [# 绘图时节点可选的颜色, 非叶子节点是蓝色的, 叶子节点根据分类被赋予不同的颜色0 X1 m0 x+ u( f) R1 \
color_set = ["#AAFFDD", "#DDAAFF", "#DDFFAA", "#FFAADD", "#FFDDAA"]: [# A- E; e; _4 }# T/ e+ D
1 `/ q* P; t) i* }/ U3 Z
# 载入汽车数据, 判断顾客要不要买 0 n( x s0 q/ {class load_car:; G: m/ O6 F! X2 \% s$ {" b
# 在表格中,最后一列是分类结果! W7 c, h# u5 U# T% P" s+ N! b
# feature_names: 属性名列表, k' y. X4 N* J2 c3 J# a$ \
# target_names: 标签(分类)名7 l" t3 a- |6 D' i+ \0 G5 o& F. X
# data: 属性数据矩阵, 每行是一个数据, 每个数据是每个属性的对应值的列表 - j. Z3 c' o# q4 q ?5 B # target: 目标分类值列表) l7 Q U0 f" T' R1 }! B9 u
def __init__(self):# S |& ?) g D
df = pd.read_csv('../dataset/car/car_train.csv')% o& |4 K `& N+ O% r
labels = df.columns.values! Q& D2 X7 |/ X; I' G
data_array = np.array(df[1:]) - r0 m; g5 ]3 ?5 I self.feature_names = labels[0:-1]5 f! v* r$ Z5 ] ?
self.target_names = labels[-1]0 F3 i% b# M1 g6 q! m; o
self.data = data_array[0:,0:-1] 3 M( o2 t0 M& ~ U Z$ ~ self.target = data_array[0:,-1] / I2 O9 ]6 o' q8 Z) A- k& _# [( h1 Z+ c9 G: ]: v, y
# 载入蘑菇数据, 鉴别蘑菇是否有毒& S0 ~1 D8 t9 }5 v+ C+ k' k
class load_mushroom:* S) Q) p5 F3 I2 O! z- q7 U3 a
# 在表格中, 第一列是分类结果: e 可食用; p 有毒. 5 u [+ p) c) f0 i' q) m # feature_names: 属性名列表 " q/ ~ G: C4 d9 b # target_names: 标签(分类)名 / x# e: d' \9 ~) ^2 c+ Y1 m # data: 属性数据矩阵, 每行是一个数据, 每个数据是每个属性的对应值的列表! I* b+ H* L/ w; G* N+ Z; R& L6 r3 W
# target: 目标分类值列表 & K) F2 d5 X1 @ def __init__(self): 1 k; s, e# |: ^6 T8 @: N1 V1 P df = pd.read_csv('../dataset/mushroom/agaricus-lepiota.data') / f4 K- n/ B0 ^' o data_array = np.array(df)$ [( a1 v8 P b7 \ d; f
labels = ["edible/poisonous", "cap-shape", "cap-surface", "cap-color", "bruises", "odor", "gill-attachment"," u$ M. u; f0 R" g, v4 z& W3 k' y/ I
"gill-spacing", "gill-size", "gill-color", "stalk-shape", "stalk-root", "stalk-surface-above-ring",4 R2 e+ h9 A( Y" \% y4 `- _
"stalk-surface-below-ring", "stalk-color-above-ring", "stalk-color-below-ring", + o8 Q; H9 M7 k0 X% H$ ~ "veil-type", "veil-color", "ring-number", "ring-type", "spore-print-color", "population", "habitat"] : w# t- \9 i7 L9 o8 D& M self.feature_names = labels[1:]/ y2 L- z4 i# p8 z6 A! H
self.target_names = labels[0]% a5 [' K6 b0 n `0 C: g
self.data = data_array[0:,1:] 6 ]/ l' z- A2 o3 y8 s2 J7 f9 H self.target = data_array[0:,0]9 s3 P; r' i) o+ @
: b6 q3 F1 M9 @; P4 L- K7 T# 创建一个临时的子数据集, 在划分测试集和训练集时使用 4 ~5 g0 |) e; k- }# V7 W8 M# Yclass new_dataset: 1 S5 }) G: E( e7 _* c' e # feature_names: 属性名列表2 o- N" O- N9 g2 S: Z1 Z
# target_names: 标签(分类)名 9 n& X) \5 Q, p" j& F9 h; @ # data: 属性数据矩阵, 每行是一个数据, 每个数据是每个属性的对应值的列表 ! `; p" W# r) @' i1 ?$ o # target: 目标分类值列表 D4 t8 q7 L. I
def __init__(self, f_n, t_n, d, t): / }) O* f2 {6 `; p \5 ~% I# M self.feature_names = f_n# g; H# z- w) J8 t/ A
self.target_names = t_n, J* d# {% _; n* V
self.data = d * j# ~( u- i5 @ q self.target = t . S) d0 F0 Z0 Y' `: j2 C1 ^% g& A" N0 y+ Q3 w
# 计算熵, 熵的数学公式为: $H(V) = - \sum_{k} P(v_k) \log_2 P(v_k)$ ( Q+ v8 O! o s8 e/ u: t! @# 其中 P(v_k) 是随机变量 V 具有值 V_k 的概率 , J! G0 j$ O- i; O# target: 分类结果的列表, return: 信息熵 % K0 z+ M2 n$ j& Ydef get_h(target):0 s% |5 ]$ i! ?
target_count = {}# f* n, l- o/ G
for i in range(len(target)):' j& m2 T6 X; m0 ?0 `
label = target+ _- A0 A$ K5 a% A9 T
if label not in target_count.keys():% v+ y8 P! ?% V6 n6 G
target_count[label] = 1.0 : e1 D0 a# }; M9 z. L; A1 o/ x else: + v) D- b& x' A$ M6 v target_count[label] += 1.0/ O# H' k* U: I1 Z1 M
h = 0.0/ h; b% y3 p% K6 v+ O7 A
for k in target_count:1 E3 g! V: {# i' B& t! L; _
p = target_count[k] / len(target) 9 @6 z8 g" ^9 ^2 K: t/ {2 a0 X h -= p * log(p, 2): l5 z! _1 e2 b* q' E/ [0 g3 X
return h l+ b1 O1 b+ t+ {5 R
, c( [% o* A7 D8 ?2 \# 取数据子集, 选择条件是原数据集中的属性 feature_name 值是否等于 feature_value 4 w$ M- q6 D5 e8 g: E# 注: 选择后会从数据子集中删去 feature_name 属性对应的一列) ^1 O' J' S* {+ i* L' y8 g. L
def get_subset(dataset, feature_name, feature_value):; h2 F9 l# H8 O9 t) r& A7 F
sub_data = [] 1 F, f) ^4 d& |3 W sub_target = []9 s% x4 w d3 D9 r
f_index = -1& n# U& J: ?3 r4 Z' ~. t. M: i
for i in range(len(dataset.feature_names)): ! I1 E( J9 k" F2 j2 f* n if dataset.feature_names == feature_name: ' n G! D* B9 \3 O4 S f_index = i 9 p. l2 s5 K* e, |# @. c break/ b- _) ?) k, g% m, N
( Y% @6 x' [4 C4 s for i in range(len(dataset.data)): 3 b7 W* A2 Y9 r8 F; {- a if dataset.data[f_index] == feature_value: X- v% u' d! r9 V0 @
l = list(dataset.data[:f_index]) ! ^5 a& C; A H, P, O& F l.extend(dataset.data[f_index+1:])6 C4 t, H; o" w" B
sub_data.append(l)2 E, T# m" c2 _- R) O
sub_target.append(dataset.target) " h! \6 O: L& K 9 f! W9 I3 f) i: i6 E. Y sub_feature_names = list(dataset.feature_names[:f_index]) ( p- T1 C! C+ A# K sub_feature_names.extend(dataset.feature_names[f_index+1:]) 0 z# e- V8 u) R0 A( Q return new_dataset(sub_feature_names, dataset.target_names, sub_data, sub_target) l/ x" G4 B! H8 n* ^# y/ V1 K
+ F' l% U3 n6 e# d( C0 C
# 寻找并返回信息收益最大的属性划分 . h" f! w, b+ ^& f Z& N# 信息收益值划分该数据集前后的熵减 - w! _' [5 U- L# 计算公式为: Gain(A) = get_h(ori_target) - sum(|sub_target| / |ori_target| * get_h(sub_target))$ ; d+ }2 r! a3 Y U9 d: Q' ~5 Kdef best_spilt(dataset):3 h5 M* a9 [" L) ~
: v3 J; W2 x5 Y: o3 m: P, Y* o base_h = get_h(dataset.target) / u/ V3 |; A7 W6 G( z: O& i# L best_gain = 0.0 + e+ Q1 k% }2 q best_feature = None) X# W6 {( A9 m& P! F" o
for i in range(len(dataset.feature_names)):! l$ }$ c, K8 L. t% m. N
feature_range = []8 w, h, o$ t, y: g. n" m6 ?
for j in range(len(dataset.data)): * U9 Z, R" V% a" Q5 l) B if dataset.data[j] not in feature_range:( I6 m" E0 P/ W9 G( L7 J/ @
feature_range.append(dataset.data[j]) 5 a6 q" w; d+ u0 h; J& a! s / c8 a- }; l1 m' u7 t! K( u9 v# } spilt_h = 0.04 ]& Y0 m; V0 @0 `( F
for feature_value in feature_range:8 c4 J8 {: w1 ]% _& X* N
subset = get_subset(dataset, dataset.feature_names, feature_value) - L Y; r; s. A spilt_h += len(subset.target) / len(dataset.target) * get_h(subset.target) & D" `( W9 H. t$ j . o! r- @/ Z2 [& @% g if best_gain <= base_h - spilt_h: ! q1 G/ W* _ C9 [$ I best_gain = base_h - spilt_h1 t: N7 }+ d P# o+ l" L. f5 M% t
best_feature = dataset.feature_names % c$ R, p$ D! y: r4 W0 ]4 D5 z- L: H$ z9 J4 H
return best_feature6 c1 L7 X6 j( M
% ?* ^7 J% w1 M6 s( R" [$ J7 }# 返回数据集中一个数据最可能的标签 U8 c$ _4 b. y5 mdef vote_most(dataset):2 J( _/ D: q7 R2 X. Y* H2 u2 ?
target_range = {} 5 t. |8 k0 O/ c; \' T* ^" @ best_target = None8 P, b7 T* E0 }
best_vote = 0& `# Y: J& R6 q) d
% l; Z; I! y- o' L! t! D
for t in dataset.target:" U9 t. d5 f& }9 v$ X" A7 b
if t not in target_range.keys():" t& Q, ?0 _# e6 Y F7 X* z& ^2 h
target_range[t] = 1* M. |3 S: h* d; Z+ }3 k- S: G# b- _
else: ( a7 i& C# {( K( C target_range[t] += 1% z/ t D( Z+ q! Q, w% c! |
/ N3 @$ m* K% c& g( E; {
for t in target_range.keys(): 0 X% u- m/ E1 y( q _1 ` if target_range[t] > best_vote:- n0 R6 q; r0 D- i
best_vote = target_range[t]# N$ L* h9 Y9 v: `. }
best_target = t$ i2 x' t# o# v! S; I
' X6 q1 M+ `5 e3 m7 c- P; D: U! r
return best_target' _0 k9 ]: s% ` v' {
. W; \0 ?2 f( }/ q% f# 返回测试的正确率: a7 b. b6 F% R* T
# predict_result: 预测标签列表, target_result: 实际标签列表 - b! a6 A" D! d/ Jdef accuracy_rate(predict_result, target_result): % I0 x v# X1 O3 L/ E) [( U- H # print("Predict Result: ", predict_result). G1 E Q& V y% u/ l) L
# print("Target Result: ", target_result)1 l; u3 y* _6 ~( }9 b9 q
accuracy_score = 0, l: ?% b6 M, w$ G# D) b9 M
for i in range(len(predict_result)): ) g$ s+ r1 I# c2 H; Z if predict_result == target_result:1 _2 R l8 p$ h( M
accuracy_score += 1/ [& m2 Y) c& v/ U/ i
return accuracy_score / len(predict_result), y' W9 W9 }! n1 _' [3 z: Z# p; |
4 j; }& i4 p0 y( A
# 决策树的节点结构 0 E S3 Q( I3 D8 ^class dt_node: N* [7 u! o+ }% c9 N8 u/ i8 F0 j
def __init__(self, content, is_leaf=False, parent=None): 1 `1 K% }- G9 b) A& ` global nonce( ]6 I& x! _) c0 f2 b5 S2 k
self.id = nonce # 为节点赋予一个全局ID, 目的是方便画图( p" A, |0 J- F. v
nonce += 1 % C( j2 i% m1 B self.feature_name = None5 Z% q, ^, r. T( E; p: A; B% u
self.target_value = None% n/ B7 k6 `4 x @
self.vote_most = None # 记录当前节点最可能的标签 # w- K, S( V" j( Q7 G7 s if not is_leaf:( R% Y, U1 I$ r5 h5 ~4 w: Q
self.feature_name = content # 非叶子节点的属性名; g( C/ d- [; y: J* C3 `
else: . O' N9 }5 R2 Y: t1 q self.target_value = content # 叶子节点的标签 3 r8 h6 o+ ^ ]* V- h# W/ r ?2 [" m9 A8 |, h4 H) P
self.parent = parent 0 t4 u' v. G1 a6 r5 W self.child = {} # 以当前节点的属性对应的属性值作为键值 % p; q( y* B6 v7 t1 B! E5 g. I; \% a0 r7 k6 k# @6 @
# 决策树模型 7 B* [+ H" k8 C j9 eclass dt_tree: ! M: U$ r* J! }. R& E! R) ~1 x- w% _/ \. R$ v3 `" a9 B
def __init__(self):, R2 [: w& _ ~/ `2 m
self.tree = None # 决策树的根节点 2 B2 K! c+ j/ j' Y; v; h self.map_str = """ ! Y: ~7 }* ~3 L6 I4 H1 r' Z digraph demo{ * ?8 [9 m. g! G% V& H* ]6 F node [shape=box, style="rounded", color="black", fontname="Microsoft YaHei"];% ?; Z2 G/ H/ j' ~
edge [fontname="Microsoft YaHei"]; * K$ n% b% i5 i% ?! q" @ """ # 用于作图: pydotplus 格式的树图生成代码结构 ( S1 H; P# w/ N* ` ] self.color_dir = {} # 用于作图: 叶子节点可选颜色, 以标签值为键值 % [3 X' j% W9 w6 ]5 q; t- @) h, l/ ^
# 训练模型, train_set: 训练集 4 \& y1 w. n0 R! c1 E8 B def fit(self, train_set): 5 J9 w7 g* Z& p& @ + ]; g3 n4 E+ ~7 `5 x1 C6 |3 a if len(train_set.target) <= 0: # 如果测试集数据为空, 则返回空节点, 结束递归 i4 q( ^' i1 X9 M0 _, G
return None - C) ]& i! d1 P2 M# q4 u% j : e; V# b# j4 ~7 T1 A target_all_same = True+ k- o$ K; c. s% m* X- L1 H; x
for i in train_set.target: 9 s5 H0 X7 D5 Q, ]- y if i != train_set.target[0]:& J O- j* K, o4 G" y, F+ z
target_all_same = False ) h# F( q! a2 k; B break7 G/ P g' m1 k i8 C/ Y2 o
& {2 z5 r/ x: m( K) |) N if target_all_same: # 如果测试集数据中所有数据的标签相同, 则构造叶子节点, 结束递归$ I/ X/ ^$ S, ]3 ~3 D
node = dt_node(train_set.target[0], is_leaf=True)% O) D/ F- O# B4 W
if self.tree == None: # 如果根节点为空,则让该节点成为根节点3 g5 w2 A1 ]3 W$ C1 v
self.tree = node ) q. m( }7 [9 {5 ]4 v- ^3 G5 J' D0 K9 W+ f, l4 P+ s
# 用于作图, 更新 map_str 内容, 为树图增加一个内容为标签值的叶子节点+ D( M. i+ Z. _
node_content = "标签:" + str(node.target_value)6 d& [6 d: q% L- H0 l. X) T7 D7 I
self.map_str += "id" + str(node.id) + "[label=\"" + node_content + "\", fillcolor=\"" + self.color_dir[node.target_value] + "\", style=filled]\n" F8 Z7 z- ?+ u% }" F
# K6 ]/ B* u# F/ c' n+ ~+ y
return node) M7 W# e8 V; G& g& P
elif len(train_set.feature_names) == 0: # 如果测试集待考虑属性为空, 则构造叶子节点, 结束递归/ r: l% x/ G( z8 x6 }, W$ v, G
node = dt_node(vote_most(train_set), is_leaf=True) # 这里让叶子结点的标签为概率上最可能的标签 & a4 L! e1 S: A; M if self.tree == None: # 如果根节点为空,则让该节点成为根节点 ( x7 |& E( k3 U" | S, C x5 I self.color_dir[vote_most(train_set)] = color_set[0]" _+ |6 R# S% ?/ |
self.tree = node ! V6 D$ S3 q* E * U1 W" o9 F+ A8 P5 }8 _# u0 N # 用于作图, 更新 map_str 内容, 为树图增加一个内容为标签值的叶子节点( F( {7 @7 z, U! Z1 G2 S
node_content = "标签:" + str(node.target_value) * Q! K1 W+ Q( V- D2 r% [ s self.map_str += "id" + str(node.id) + "[label=\"" + node_content + "\", fillcolor=\"" + self.color_dir[node.target_value] + "\", style=filled]\n") a6 ?+ f( T, M- U0 _) L
0 [. L. o W9 j! A0 d$ t
return node $ B' \/ r: j9 j: O9 T9 |1 x: M else: # 普通情况, 构建一个内容为属性的非叶子节点 ! J! ]6 `2 }8 f4 U, `2 Q x best_feature = best_spilt(train_set) # 寻找最优划分属性, 作为该结点的值1 j- e# N# G9 j
best_feature_index = -1' p# u' _( ^) _- j: T4 U, ^4 r
for i in range(len(train_set.feature_names)):+ ?3 m u! Z0 t7 h; @8 C
if train_set.feature_names == best_feature:( i: H1 Q: Y% ?! o
best_feature_index = i8 H+ R1 [- o1 \$ d3 b
break 9 u) g( k/ u. P; N: x/ f) }: z2 |% L % t# E8 X0 M) ]( u# X N node = dt_node(best_feature) 0 z- X& ]. q& z \" h: k node.vote_most = vote_most(train_set)/ K! j2 @2 y% f, ~+ N
if self.tree == None: # 如果根节点为空,则让该节点成为根节点+ }8 n6 q0 F: p6 \- C S ]
self.tree = node( p* P3 `7 J! J5 V
# 用于作图, 初始化叶子节点可选颜色 # K& W: X Q @ for i in range(len(train_set.target)): * @2 I9 u" M8 F5 D8 F if train_set.target not in self.color_dir: - `/ F3 L/ E$ r+ f( W global color_i' T) o) D) a2 p8 Y* W+ {
self.color_dir[train_set.target] = color_set[color_i] % ?) A$ `( _8 C2 R1 P* }8 { color_i += 1 6 G$ O% C3 J4 i- X8 |7 Q2 T color_i %= len(color_set)% \8 Z4 |) P2 w# L! S! `, q
3 X3 F5 b" ?5 |$ S- v feature_range = [] # 获取该属性出现在数据集中的可选属性值 " r2 K* y& W' |8 Q3 q' i4 v! h for t in train_set.data: & F! _( x. m4 A! d! x3 X I if t[best_feature_index] not in feature_range:2 [7 }/ r+ J2 N6 Z1 A
feature_range.append(t[best_feature_index]), `% t4 s _& r1 c
' t0 q7 X, n* m- I) `: K
# 用于做图, 创建一个内容为属性的非叶子节点 / y: {+ [! Z+ j4 [4 l+ Y node_content = "属性:" + node.feature_name1 p2 K: Y- z9 O6 W
self.map_str += "id" + str(node.id) + "[label=\"" + node_content + "\", fillcolor=\"#AADDFF\", style=filled]\n"1 {" h. e' s* k
4 q2 }) }3 s C' f
for feature_value in feature_range:1 P1 P. f% J/ B
subset = get_subset(train_set, best_feature, feature_value) # 获取每一个子集& L z) v$ g/ Z# u# j/ O% Q1 N
node.child[feature_value] = self.fit(subset) # 递归调用 fit 函数生成子节点, `$ W+ Q7 ~! x+ D+ r0 W
if node.child[feature_value] == None: 6 ^0 ?& q" b) `6 O( W$ w+ i # 如果创建的子节点为空, 则创建一个叶子节点作为其子节点, 其中标签值为概率上最可能的标签 6 ^5 u% X, l3 c. n- L' l node.child[feature_value] = dt_node(vote_most(train_set), is_leaf=True)8 B/ s& [$ z5 R. D
node.child[feature_value].parent = node1 @1 p5 c: `$ g5 \
9 G8 N% m- b/ i
# 用于做图, 创建当前节点到所有子节点的连线9 t. ^; P7 T; L/ M
self.map_str += "id" + str(node.id) + " -> " + "id" + str(node.child[feature_value].id) + "[label=\"" + str(feature_value) + "\"]\n" % e' d0 u+ f( f+ n & M8 C9 ~6 ?2 ~1 U- m # print("Rest Festure: ", train_set.feature_names)9 a: I. o; N% V
# print("Best Feature: ", best_feature_index, best_feature, "Feature Range: ", feature_range) ( D9 n) y e% y2 U # for feature_value in feature_range:/ n% k3 l# p5 o8 h- F; C. x
# print("Child[", feature_value, "]: ", node.child[feature_value].feature_name, node.child[feature_value].target_value) . A$ F% Q' T$ V" W3 E return node5 ?) W( O( i( `3 T/ I
3 L+ w* k/ s! w; Q; d. E # 测试模型, 对测试集 test_set 进行预测 9 h3 `% P/ K% p2 R! y7 W) p' Q def predict(self, test_set):$ m. P( K9 M' }
test_result = []) \8 Y7 C) O+ f# @% Q! I. a, C h
for test in test_set.data:! |- E1 B$ h' z s8 q( M
node = self.tree # 从根节点一只往下找, 知道到达叶子节点 8 \! w6 A U: L5 U while node.target_value == None: 9 z. ~5 N6 i% c: A7 {% K- A feature_name_index = -14 D3 v1 M3 T" b L) Q- Y, O
for i in range(len(test_set.feature_names)):8 j% Y' q" } C
if test_set.feature_names == node.feature_name:5 y7 A) q \6 Y2 y8 o( u
feature_name_index = i/ \+ Q5 x9 Y% n a- }5 @7 L* U
break4 m0 i, t3 q. N+ ?' [3 `
if test[feature_name_index] not in node.child.keys():( p6 L+ N, p5 T: v7 |4 b
break4 K' a, ?- m: X p, b0 H2 I
else:- K$ c: L6 _ F p& @
node = node.child[test[feature_name_index]] 9 ]4 C- E/ |& J. P7 e- h4 H, q3 t ' p9 `9 c$ a) n1 |6 J2 N. S if node.target_value == None: - {3 n' ^2 h( c L' J& T! h' n& g' W test_result.append(node.vote_most) 0 k$ ^/ _) m: X* z else: # 如果没有到达叶子节点, 则取最后到达节点概率上最可能的标签为目标值 & O: g. Q3 N" Q9 t1 s" `& Z0 d- G; m9 a test_result.append(node.target_value), \4 }3 d v- t3 p% S5 r% X) `
1 Q4 P) y+ R& p return test_result , S7 i0 l" Z! C& W8 n 8 q+ Z0 s5 h i7 M. D1 u4 n6 S # 输出树, 生成图片, path: 图片的位置 N4 T$ E v c3 j' x* ?
def show_tree(self, path="demo.png"):& |% F" k4 [$ O
map = self.map_str + "}"* p" T! G+ l. D" s* r* ?
print(map) % f' }$ T8 @& g* H+ d, e graph = pdp.graph_from_dot_data(map)% W$ W" }2 N |3 f8 O
graph.write_png(path)% H* E# O$ m& l
. ? a' `) A* l7 u/ C9 Y; P
# 学习曲线评估算法精度 dataset: 数据练集, label: 纵轴的标签, interval: 测试规模递增的间隔! l6 J: c$ ]. J7 |7 q( ^- \
def incremental_train_scale_test(dataset, label, interval=1): 5 I" r3 \ [9 @7 \& `, T- Z8 j c = dataset ' a% @8 j) p3 H1 I* h$ n" J r = range(5, len(c.data) - 1, interval) / U3 L$ R3 ~5 u5 W: E9 d. E rates = []/ `7 {- E3 H8 s4 A! [# _0 _
for train_num in r:* V: g. I5 a+ V: [$ R2 _0 {; r
print(train_num)$ S9 g! H" l4 ^2 g8 c0 D
train_set = new_dataset(c.feature_names, c.target_names, c.data[:train_num], c.target[:train_num]) # \# `# y" D8 F5 K0 t test_set = new_dataset(c.feature_names, c.target_names, c.data[train_num:], c.target[train_num:]) * s" o7 l- K1 [0 f0 K dt = dt_tree() 4 E& _" i' b9 N2 _ f( c* \ dt.fit(train_set) & c& i3 F& y! R( Y2 W1 Y( F rates.append(accuracy_rate(dt.predict(test_set), list(test_set.target))) . x, ]1 `$ z' G4 e- V) c. J! z# k1 V9 d1 ~! U- `* V. c
print(rates)& V: H; X/ P4 ^* q6 m# D
plt.plot(r, rates) 3 u- v4 G( f. d3 c. Y plt.ylabel(label) ; T' t. {1 B: S4 m plt.show()0 q8 Y& N. O y
; S% ^2 L0 o: ]) v
if __name__ == '__main__':2 {* U, r7 q* D" O2 L
l1 `" X$ w8 \9 l- i4 H( \( g
c = load_car() # 载入汽车数据集 ! `! B6 l: m( F, k/ V" \ # c = load_mushroom() # 载入蘑菇数据集) |- J/ T6 q6 o8 F
train_num = 1000 # 训练集规模(剩下的数据就放到测试集) 8 a2 S5 p5 N8 y3 P5 f: z0 v train_set = new_dataset(c.feature_names, c.target_names, c.data[:train_num], c.target[:train_num])4 ~" [6 {( E o
test_set = new_dataset(c.feature_names, c.target_names, c.data[train_num:], c.target[train_num:])$ I* p5 Z# Q9 u; ?( u
2 x# o7 O) r- x* Q: [+ H: d* e dt = dt_tree() # 初始化决策树模型+ C: Q1 h- W7 h2 v8 p0 a4 @2 P
dt.fit(train_set) # 训练 + J) B2 Z8 A$ r' T dt.show_tree("../image/demo.png") # 输出决策树图片4 [% p1 ?" K7 L/ I" R7 r% r
print(accuracy_rate(dt.predict(test_set), list(test_set.target))) # 进行测试, 并计算准确率吧 7 Y3 h3 k, Y3 F$ K, g2 O' h 1 t7 g" u: z$ X3 L( q # incremental_train_scale_test(load_car(), "car")1 O# E0 {' V$ |9 U7 _* L9 u' b: ?
# incremental_train_scale_test(load_mushroom(), "mushroom", interval=20) ; r( z, B7 m6 C7 l- r- B# f- y. J$ Y% \" Q3 M) _
* k) |: O$ Z$ C5 N$ ~2 k
4 w. \/ x7 `, w4 |9 ]: O
1 . ]: }2 z+ M0 O( H; |9 T2 ( o1 k1 r; B8 ?+ T- X3 1 D6 @% K$ k% b8 f4 3 \6 k% r6 [- N7 E3 T1 l5 K3 O% U5 G3 a7 R) ?' i# K
6 9 R9 v7 y+ l1 X5 ^70 y" k* h2 F, K! F- b
8 " w" F/ ?$ O1 X' e2 h93 s2 a8 ~% W M
102 }2 N( W5 k3 C5 D
11 ' ~8 D" `; \2 ^# Z, h12 " Q" G) A. D" ]13 7 q) B7 W( N, v( X8 j" C5 p143 x& s ?' [# l( H" s
15$ r9 ? [' F+ S, p: R5 E: F8 C& n
16 4 a8 t) q! l( T6 r7 Z17 $ L! S+ p3 e7 J- B8 l% V184 V( }" a, O6 m9 G: N' {; o6 r- \
190 ?" a) E. ]! w0 {; m5 U
20 }$ n. w3 ?# x( f21- {& ~( B7 g! V
22 , u. i* _9 B& A1 ^23 % O3 R: o4 m# K. W ?24+ l) H3 M X4 g- @( j" b
25 ! ^7 d: P/ f1 _! c j4 V260 I% `9 ~ O X2 B% z% }
271 q% D; E, G1 }2 H
28 4 x6 m% i4 j! C, @, ~" J5 g1 D294 K9 S5 D3 }; F" `+ }
30* [" X/ C; R8 ^4 V/ t! J, B
310 K' z, B S; V2 K, o# Y! E3 n
32 ( a; K) ?7 J# ]- z/ J3 B33& n9 D. h0 e* T/ n9 O
34 6 |. p2 e6 g: l7 \( V |$ D$ X35 9 N5 M: a1 l$ ?1 e: n36) e1 K# `, J4 Y5 P1 ~3 E
37; F% s1 @& k# P! B# }
38: W/ B# `0 g# @( z
39 * v( A7 R& M/ @& {9 A40 : X7 H. c* w* k3 q: ?41, i2 P8 ?4 t5 U6 Q1 ]5 M, ]; c* X
42$ p" K: o* c& [3 o# l. ^0 ~
439 u: L# q$ V* K a+ R# e2 [% ~
44 1 D7 Y @3 p( n" E459 ]3 r% [0 q/ Q" e. b7 o3 M! H- W3 ~
466 v# z/ l. X. h- `' S' x
47 l9 Z( R% m9 G# v$ j; V, U3 O
48( k: s( }8 u$ Q9 K. k$ O T
49! {1 @0 u+ `7 q: H6 l
50 5 S3 R% I/ T) B2 ^% q' {- d, ]51# a4 C( P. }% ~. F m
52 ' ]* \! k2 S% }53 9 _' m; U& W$ p% j( k54, e. O' C2 g# s+ D5 V% j0 m0 \
558 t1 ]; ?& H' Z$ G& f- g
56 3 H; R% L% u M* F2 r57 3 e2 L- n5 C: C5 t58 1 a$ m3 U. N1 B590 Q9 x3 `0 Z( M" A+ c/ i& M
60) t9 _! Z& H+ }7 a& A
61 , Z/ L2 S1 W z% s: M0 F1 |8 F' k62 8 `. u. r6 X( |5 t4 G- z) Z& O' ?63 ( E; U# q& n& `$ G7 N# `645 z: y' m( M% e& o6 z& H; m
65 # z$ W1 Z( Q4 G; F% \66" O& ^7 u1 R! L8 Y r# j, z# [
67# r3 [ \9 X. i/ h
68 9 n% G* G4 z2 O6 T9 L69& h4 Q1 O( N4 {$ T. c( {8 ?& M) d
70 ! l( n" t* V5 P71 8 o7 F2 s j# V9 X6 ?' Q6 a72 - Y/ a& D5 p$ }; r+ ]$ Z733 J2 l k- Z5 l3 D
74! E2 J1 f: h) k! D
75 ! E4 p( H, b% e) H+ e76 0 W3 O" b$ M! |) P, |77% {9 H5 N( Z* j8 Q `
78 - w1 L t' S$ k8 I+ c# D79 ' Q" ~: e3 S/ m3 X( m) n80) O6 D" z/ L: {' z* _( V
81 % t' u O; E. p* ?, q82 6 i- s. S* V& i) S/ ] Y83 6 b. O+ G) c. m2 D841 O; Q, \/ U7 H! O/ b- \! A. a
854 P/ R2 |) V& T
86 - E0 t8 `$ |& F" U87 # t5 ?% n/ Z0 n1 N88 - E) B- ^8 u# R- A" r+ X; Q( G/ D3 @89. ?& |5 r( d8 x" i7 X
90 ' d! W f: w/ E( R" L91 6 H' P. K: U# i923 `6 Z7 F3 [. \! Z! D' H
93 " U1 q0 S- ~/ |6 z! Q, n8 {0 ]942 V3 K6 G7 R. A; j9 m
95 . ^5 N$ C/ ^. ~/ y& p+ M0 K96* ~ o& L& K+ C% D; |. z8 Y- o7 U/ m
97 2 P2 J2 O3 S+ y98) t9 a8 ~0 D" C, l4 H+ Y' J
99 ' u7 o" `: C+ g# B: z3 p6 |" F1009 ~! l) ^0 ^$ c; @ L
101& h9 v+ P+ [ L$ F! T
102 ' Y6 O. q* ?, O& O* \% B9 A/ q1038 B: A* }4 v/ K4 n
1042 \' b9 Q; w7 W' i
105( }7 P- r1 Y' g. `- H0 Y
1063 I$ k5 e( l( K( j3 ?9 y! D1 `
107+ [: j7 @# [ y2 z3 }, F, {
1089 r0 W2 J* D% r" i) u
109 6 U: T- v; R( `110 ! q# F& w: L; h1 o5 {+ J3 S111 ' k$ a* `5 G& h1 g5 f$ V112) u; f1 L+ X% Y+ N5 O
113 $ X0 ~/ B6 N0 B$ J3 x114 ' ^" j$ b* I% k' u115& H! e, \* C3 m) h4 y
116+ E$ y" f+ x4 i- \
117- u; z1 z. u* y7 a+ ]) K% T
118 9 c9 m* Y H% h6 V119 $ j; H5 m& G d1 ^3 u120, _4 H- |* R' V* K; o
121 7 S7 I( a& O8 e9 ?$ l2 [) p122; Q: a* _ n5 Y+ K$ R
123 . |, Y1 U4 J, y124 $ O) c5 M! i/ Z" i9 h& F- e4 r0 K6 D125 P4 t' _( ~/ G" d/ Z7 e0 I7 e$ V126 / \+ {2 }' |1 D2 {127 5 A" e2 \/ G* F2 T( C128 & f. H, X- F8 x1 E129: b# A; n3 Y. C, y/ A1 o8 [
130' C3 t7 F$ O2 E# K
131: H# K" f# [ B! H% X" [9 E8 [
132 7 T% H* G7 ]- s1 c6 e. z133 4 \3 g# n! Z! U- P- H9 f134 " \" S3 W' j4 d; |# f) t; S4 ^135 " }: m, R0 _2 D( n" q. O* Z136 # D+ g% A" y8 S/ d! p! A, T137 7 f: F" ?) l- I! S3 Z; y+ ]: A! ^, {3 t) X138 - J6 S, `' \8 t1 i& Q1 T' k1397 h0 q; g, A* z I( F
140 $ ?) h% R& z2 f% J; m/ u2 g1417 P F: c* D& [+ K! @# f
142 . W" b2 B6 M! S: Q6 K143$ G$ G4 l) _$ S2 Q$ o$ Y
144# U; N) H2 i1 x8 E/ P4 S
1452 Y3 r h0 O4 j2 N# d
1464 c0 K- q% n$ _( E( ]+ A
147 * r1 s8 u7 v+ I9 d) h3 S6 `; ]148 5 D0 \" ?* o% o* N1494 E8 o6 ~: W% \* Q) l8 B, r
150' W. t+ H' k! u6 y+ @% {9 d
151. V4 m5 d7 H+ x( {, W) W+ f
152. v. P. O% s- O. _( Y+ D5 B
153 2 \9 S- t# C1 n/ A; B' p154# L8 l9 q+ t: X8 Q) v/ F3 r
155" D* w9 p6 p( t& a
1560 `: g# @" ^- {- ?' o& c* c
157 , g6 L8 j+ L- g' z" f1 q158 , Q3 O y3 N7 X/ D/ V159/ I* ^8 X3 j2 U6 w3 r& t$ M
1600 n8 w* s: p- K$ U
161 # x5 s( O5 R" ?. Q8 v162) }: ^: t: o7 ~9 e# v3 N5 B
163 ) D9 u7 t5 ]. m5 U1643 d) A7 q: v0 k# C! _
165& M7 N2 O' _5 W# D Z
166 0 L* l! G/ [5 ~) v167 & W3 f, ?/ s( i. M( h% c8 h168* M3 C. h4 l( b; n
169 0 k0 K4 \2 x+ z* I3 Q+ b170 3 b) ?$ n( N0 q* X4 A8 L1712 U6 G, b$ D1 d- Q6 F3 \5 Z, v$ o
172 , y) l$ b/ @; K+ s. c0 V8 p, {6 N5 m1 F1734 f" b- ]9 U. D7 ~& S
174 7 F5 }1 ]) ], \; `175 7 ~$ }+ d; f3 g176 ! D3 Y6 G) b8 M. G* t177# C2 y$ Q! k5 d# C2 U' C- N
178 9 E) U: H4 m& ~( B& n6 H9 }2 l1794 J9 p) {# t [7 \
180 ) b/ A2 v- [2 F; Y" |+ V5 b181 & x, D+ H8 t: s" y182 ' s2 D/ u6 \, \183 ( ?) Y) V0 m6 B2 o) p0 n5 |184 / J. @. G. g7 y3 S( ?1857 J" F* l+ p' r- N
186( H# m& \. b2 J8 F7 u0 h
187 2 h9 C7 [+ f- ?# Q188 8 f8 a# W9 T4 b' ^* j4 A6 C J7 A189) E2 D& m$ N3 \* a4 ~6 u9 O
1904 m1 P& @* E4 ~; c9 d( J( O! g1 X* o
191 # K6 e1 ]+ c. G# Z- ?+ H* m5 ^1927 F6 B* a( T' I) o' ]: E, C
1938 x" p5 h, O+ F% q6 w+ j
1943 R- V. g; L8 {6 {- g" d
195; a) A4 Z9 P8 @( [1 G8 k1 m' p
1969 S# n2 w2 H5 x: U; U
197 3 Z, ?8 {; t, M) K) Z198 ; Z ?# A& h* l199 & E7 I) t# p) N! l% t7 `200 8 i4 x2 c$ n4 @6 j. ^) _1 O" k201, A: P, Z& \( i6 t: M4 e* S6 n
202 - n4 w& j$ P* ?- Z2039 A- A4 X( L: S( c
2044 E! U7 }7 U& B/ w0 O: y8 s3 K
205$ p( j5 [7 g, X1 }) Z5 T* E/ w
206" L5 K6 J) }% z, |' t# s, O
207 1 r- V8 K' ^- d, V, q+ z208- W3 @3 L: v9 y' L; E9 ]/ h, k
209 4 j4 Z. J6 `1 _: b7 @1 \210' K6 [& \, z9 B( \9 l3 v2 \. {
2113 \: H- G8 B0 m% h& \# F' E
212 0 N4 B% w' V. _* U( e* L213- x* d5 [& }% {1 d/ N
214 ! S, \+ F- I" F2 q- b5 @9 M- O215$ Q* {) |* |& E- K
2168 g3 ]9 _. N! u$ D5 ]; n
217; |% K8 C% q9 {% S
218. \4 l0 h/ Y. L8 A! r2 B! t5 Z
219 . [6 H0 ~5 b4 _' I' x220 6 s3 C) L1 N7 f! n& R1 A- |221 - ?. e& l |9 d D5 [& f3 `222 6 h% a- B$ b" \, F. W223 " ]% {7 c/ D ]/ k( K2241 I, K3 J% @! `5 }3 U
225 ( t2 n1 x' j9 z226 * w8 F; t6 j, P5 m' t( X227$ h# W$ L, }) m* |( \
228 : r& z; X0 ?, t' Y v229 8 q9 Q1 ?5 s' w2 D230; C; W# l8 D" s3 g# J5 [
231 6 X B6 c0 M0 ^" X9 M6 N7 Q232 4 s) `1 M1 x, N233 # [% x" ^2 v& d) L) Q$ { P9 P234) C1 S' E" w" z- j ?/ j/ Z$ O
235 ' J5 ~6 s9 t( g( y0 V236& n- J, Z1 G+ a6 }2 m/ c( ^
237 5 l* y9 F6 W2 }9 p* [8 Q2381 |8 W2 F7 I9 b2 Z5 W- E# V
239 $ v v9 [0 S4 s7 S" N240 8 P* i1 d) y* b |% l, W, ?241 ! ` |' j$ K9 R# A; P3 ~$ s& ~' p, h3 S242$ a' u2 w( o6 V0 c1 G1 [% E! t( l
243+ c, A* h* C; h
244% O( }7 |; T1 ~: v5 f3 [8 q
245 / l: o3 ~) |* R# r+ a5 T' f' H246 0 r8 \ h/ l2 F- {' c( f n247, t. M6 d4 q$ z; q
248 8 T: l6 ]* s8 [- g249 + ^6 I! T; F) o$ m. K; x250 2 D% u4 s. B" H+ g251' Y+ L- {: W, h* G4 b$ }; K
2523 x0 s3 ^7 q9 \" t$ A5 ~& h w; a( E
253/ C" p% a- i; K0 b* M/ j6 Q4 Z( l
254- x3 `/ z4 R, W# X5 J" u
255 : b, u# I& C# _3 ?7 @. g256) P0 W, O2 L( `4 X
257 x9 b' f( V7 J258 % L3 d5 `0 |: b X% N: {2597 k$ j& h5 c( n9 B" {. t% G0 k1 y
260 : f6 k2 D6 U* N2614 C: i. T* \0 Z5 h4 s4 X
262% V) ?/ u) w! }
263% K% O. G: g1 _; i: h
264# O1 l; ]. L0 w/ \ q5 U
265 0 u0 S. j: j) E6 I) \. a2 p2667 P8 y4 t6 X9 K4 o, ~7 _
2671 q1 O( e5 g. r8 l( q
268 " T6 J3 i/ [, K) h! H. b7 Q2694 y6 U5 S/ r \& H+ M( a
270. H, ]9 f8 Y/ l a8 K. r* Z
271 : o. A8 O0 D$ m0 w, V* p; c6 o0 Q272 0 G2 l7 ]3 M0 O$ F! P2 W273 6 f: C& O' V' ^274 1 x4 a3 @: t2 U" K5 D275 4 ] v8 a( Y5 H! i276 1 A" q8 c/ q+ @* u0 T1 y D277 8 W1 W4 ? H1 c. B: q* t& _278$ Z& u2 `6 j2 K, e1 k4 P
279 ' k9 v. C( m8 D' f+ A2809 Z/ r: i4 j* x! Y" c8 U, c
281 8 s i2 q, r# |4 o* l2 o" `2828 C1 X' H7 k( S# F# @5 e& a
283 1 Y# o* t9 `' {2 K9 x284 * W$ n) }: q0 P( \; ]: ?& _1 O285* p! |6 C6 @ o) c
286* ]( p. l6 }" [4 S) b
2877 B6 J0 A D. T0 {
288* C# K! w2 P9 T0 K9 s: P$ b; j$ y. @
289 " |. s j2 U) \290 + o) @- X$ m( F0 ^8 D( _, V: B291+ C; o J9 l( v' j7 U3 b3 {/ T; h
292 & Q8 L8 S5 V& r! D293. G& s& n! G$ X1 j* C" H
294 $ B% I9 t% I$ c* m: J8 W6 c295( ]8 N& T! a$ I1 _
296+ \! b% b T, E8 C2 I% X- ~* a& K
297& K: }# Q) e$ V7 q' J
298 ' p1 S0 y: h- p, z299- H7 ]2 X" d x/ h( k
300 ' r" L% v* R- X9 O3014 d! p. i' u7 d( a6 ~' T
302 & q; ? i) A5 U" o x7 K303, F9 B* V4 A0 p3 ]9 V
3045 X& a$ y- Q0 S4 @+ w
305: Q8 r* J4 [+ }8 t8 F! E, K( p Z' [
306 0 U' \/ }, n* A4 u9 v307 2 E5 y/ W( x) u) j0 B% Y% n& D308 5 ~) c1 R# p I9 n: [; r# n( e309 4 x; H- |7 g8 I# w' g# ]310: K2 C5 e* w9 v9 y8 _1 d# F" [. r
3115 K5 q" t% ?5 o$ m
312 ; Z' n q7 T) e3130 v0 g p* N9 H. N# x
314- U- l) w4 J* F
3153 D. a8 \4 {* ~0 s8 Q( ?2 T6 Y' C7 c6 |
316( B% Q% M8 p( c4 x- M
317; b F" b3 c8 S( C# V& J
318 . {. p# j4 c' D, j5 W319 _8 q& P6 R3 u g7 i- w7 s9 v320/ f" H5 c7 C1 O' i$ h' P- H
321 4 Q. V# P& Y- v$ f2 P% k. I322) j/ z5 E9 O7 Y4 S: G& g
323 ; G4 b ?2 L. g# J- [5 l2 ?6 y3242 s% @* y" R( O
325$ n" o/ H0 K: C; F+ Z0 d
326& @: ~3 \6 f9 z( ^! c
327 ( H' k1 S4 e; m7 U [# M6 f328! K1 U! V8 S5 O4 b. J: R* i
329 ) t9 I n& G4 X& \! G9 W330 ' y1 K" {8 [( f331 * P h6 J5 M, N/ ~( G7 d - w M, h7 i- ]& W) a1 L' n / _; I4 J2 _7 _: I2 }7 d: C8 D 5 I& i( L4 T+ T) k. z: m m& E+ F3 u9 x1 l* y! i! X, E" Z7 L& d. ?6 s
# D$ |9 S6 w) ?' W0 [. p8 S+ e0 z: ]0 T4 f
F( B% n, m) F0 Q" ]+ Y - A8 S, ~: T+ K+ y. j4 n6 o" F - v& P) _& |* |* A: ` T+ S! |* }7 z) I& B' O2 X" O$ B1 H