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