"基于信息增益准则生成决策树"相关的理论论述,请详见我的博客机器学习(八)决策树-划分选择与剪枝处理https://blog.csdn.net/FluxMelodySun/article/details/158013171?spm=1001.2014.3001.5501(已入选数据结构与算法领域内容榜)

本篇是关于如何使用 Python 编程实现"基于信息增益准则生成决策树"这项任务,具体请参考代码片段中的详细注释,文末附有完整代码,希望对大家有用。

1. 导入数据集

2. 初始化数据集中属性相关的常量设置,以及基础方法

3. 方法:计算根结点(或分支结点)信息熵

4. 方法:计算在当前样本集(dataset)下,某个属性(idx)其分支结点的信息熵

5. 根据"根结点"和分支结点的信息熵,计算其属性的信息增益 (注意:每个分支结点信息熵的权重已计算在信息熵的变量中)

6. 生成决策树的递归方法 (主干逻辑)

7. 开始生成一棵决策树:

8. 控制台输出生成决策树的最优划分属性、属性划分过程与最终划分结果(分类结果):

通过控制台输出的日志可以清晰看出递归操作对每个子树的每个分支结点依次调用的过程。

生成决策树的编程过程很愉快!希望你也能尝试生成自己的决策树 :)

import math
from collections import defaultdict

# 在数据集中各属性对应的坐标 idx:
color = 1
root = 2
knocking = 3
texture = 4
umbilical = 5
touching = 6
lable = 7
features = [color, root, knocking, texture, umbilical, touching]
feature_names = ["编号", "色泽", "根蒂", "敲声", "纹理", "脐部", "触感", "标签"]

def xlog2x(x):
    if x == 0:
        return 0 
    return x*math.log2(x)

# 从样本集(dataset)获取某个属性(idx)的样本数据
def get_data(dataset, idx):
    data = []
    for wm in dataset:
        data.append(wm[idx])
    return data


# 计算根结点(或分支结点)信息熵
def get_root_entropy(label_data):
    total = len(label_data)
    p1 = label_data.count(1)/total
    p0 = label_data.count(0)/total

    root_entropy = -(xlog2x(p1) + xlog2x(p0))
    return root_entropy

# 计算在当前样本集(dataset)下,某个属性(idx)其分支结点的(每个属性值对应的)"信息熵"
def get_feature_entropy(dataset, label_data, idx):
    total = len(dataset)
    feature_data = get_data(dataset, idx)
    value_class = {}
        
    for value, label in zip(feature_data, label_data):
        value_class.setdefault(value, [0,0])[label] = value_class.get(value, [0,0])[label] + 1 
      
    feature_entropy = {}
    for v,c in value_class.items():
        t = c[0]+c[1]
        if t > 0:
            p0 = c[0]/t
            p1 = c[1]/t
            entropy = -(t/total) * (xlog2x(p1) + xlog2x(p0))
            feature_entropy[v] = entropy
    return feature_entropy

# 根据"根结点"和分支结点的信息熵,计算其属性的信息增益
def get_feature_gain(root_entropy, feature_entropy):
    gain = root_entropy
    for e in feature_entropy.values():
        gain -= e
    return gain


def printSubTree(tree, node, idx):
    print(f"creating a sub_tree from D({node}) partitioned by [{feature_names[idx]}] :")
    for value, sub_dataset in tree.items():
        sub_ids = [i[0] for i in sub_dataset]
        print(f"D({value}) Sample IDs : {sub_ids}")
    print()
        
def printTree(tree):
    for value, (dataset, label_data) in tree.items():
        ids = [i[0] for i in dataset]
        print(f"{value} IDs:{ids} labels:{label_data}")


tree = {}
def createTree(dataset, features, node):
    label_data = get_data(dataset, lable) 
    root_entropy = get_root_entropy(label_data) 
    
    if root_entropy == 0.0: 
        print(f"D({node}) is a leaf node.\n")
        tree[node] = (dataset, label_data)
        return None
    
    if root_entropy == 1.0: 
        print(f"D({node})'s entropy is equal to 1, so pre-prune it.\n")
        tree[node] = (dataset, label_data)
        return None
    
    feature_gain = {} 
    for idx in features: 
        feature_entropy = get_feature_entropy(dataset, label_data, idx) 
        gain = get_feature_gain(root_entropy, feature_entropy) 
        feature_gain[idx] = round(gain, 4)
    
    max_idx = max(feature_gain, key=feature_gain.get)
    feature_values = list(set(get_data(dataset, max_idx)))
    
    sub_tree = {} 
    for value in feature_values:
        sub_tree[value] = [] 
    for wm in dataset:
        value = wm[max_idx]
        sub_tree.get(value, []).append(wm)
    printSubTree(sub_tree, node, max_idx)
    features.remove(max_idx)
    
    for value, node_dataset in sub_tree.items():
        createTree(node_dataset, features, value)
   

print("- - - - - - - - - - - - - - - - - - - - - -")
print("The progress of generating a decision tree:")
print("- - - - - - - - - - - - - - - - - - - - - -")
createTree(watermelon_dataset, features, "root")

print("- - - - - - - - - - - - - - - - - - - - - - - - - -")
print("The final partitioning result of the decision tree:")
print("- - - - - - - - - - - - - - - - - - - - - - - - - -")
printTree(tree)

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐