一、什么是决策树

      你是否有过这样的经历:纠结要不要买一件衣服时,会依次思考 “价格是否在预算内?”“款式是否喜欢?”“材质是否舒适?”,每一个问题的答案都会引导你走向下一个判断,最终得出 “买” 或 “不买” 的结论。决策树(Decision Tree)正是这样一种模拟人类决策过程的机器学习算法,它将复杂的决策问题拆解为一系列简单的二元或多元判断,最终输出明确的结果。​

从结构上看,决策树就像一棵倒置的 “树”:​

  • 根节点:整个决策的起点,对应最核心的第一个判断问题(比如 “价格≤500 元?”);​
  • 内部节点:中间的判断步骤,每个节点都代表一个特征的筛选(比如 “款式是休闲风?”);​
  • 叶节点:决策的终点,对应最终的分类结果(比如 “买”“不买”);​
  • 分支:每个判断的不同答案(是 / 否、A/B/C),引导着从一个节点走向下一个节点的路径。​

举个更直观的例子:用决策树判断 “某动物是否为猫”。根节点可以是 “是否有毛发?”,若答案为 “否”,直接判定为 “非猫”(叶节点);若答案为 “是”,则进入下一个内部节点 “是否有尾巴?”,继续分支判断,直到最终得出结论。这种 “层层筛选” 的逻辑,让决策树既容易理解,又能高效处理分类问题。

二、决策树的核心:如何 “选对” 判断节点?​

       决策树的关键的是选择最优的特征作为每个节点的判断依据—— 如果选对了特征,能快速缩小范围、提高决策效率;如果选得不好,树会变得冗长,甚至出现 “过拟合”(只适配训练数据,泛化能力差)。​

常用的 “特征选择指标” 有 3 个:​

  1. 信息增益(ID3 算法):衡量 “某个特征能减少多少不确定性”。比如判断 “是否为猫” 时,“是否有毛发” 比 “是否会飞” 的信息增益更高(因为绝大多数猫有毛发,而 “会飞” 和 “是猫” 几乎无关),所以优先选 “是否有毛发” 作为节点。​
  2. 信息增益比(C4.5 算法):解决信息增益 “偏爱多分类特征” 的问题。比如 “身份证号” 是多分类特征(每个值唯一),信息增益极高,但显然没有实际意义,信息增益比会对这种特征进行惩罚。​
  3. 基尼系数(CART 算法):衡量 “数据的纯度”,基尼系数越小,数据越纯(比如全是 “猫” 的节点基尼系数为 0)。CART 算法用基尼系数选择最优特征,且只能生成二叉树(每个节点只有 “是 / 否” 两个分支),应用最广泛

三、决策树的构建步骤

构建决策树的步骤如下:​

  1. 准备数据:收集训练数据,包含 “特征”(比如年龄、收入、是否浏览商品详情)和 “标签”(是否购买),确保数据无缺失、无异常。​
  2. 选择根节点:计算所有特征的信息增益 / 基尼系数,选择最优特征作为根节点。​
  3. 分裂分支:根据根节点的特征值分裂分支(比如 “是”“否” 两个分支),每个分支对应一个子集数据。​
  4. 递归构建子树:对每个分支的子集数据,重复步骤 2-3,选择最优特征作为内部节点,继续分裂,直到满足停止条件。​
  5. 停止条件:​
  • 某个节点的所有数据标签一致(比如全是 “是”),成为叶节点;​
  • 没有更多特征可用于分裂;​
  • 节点的数据量过少(避免过拟合)。​

四、应用

以已下图片作为数据,来用决策树解决这个问题:

将数据集转化为数字标签:

将数据集与代码保存在同一个文件夹中

完整代码:

import math
import pandas as pd
import numpy as np
from collections import Counter

class DecisionTree:
    def __init__(self, criterion='info_gain', max_depth=None):
        self.criterion = criterion  # 'info_gain' or 'gain_ratio'
        self.max_depth = max_depth
        self.tree = None
        self.feature_names = ['年龄段', '有工作', '有自己的房子', '信贷情况']
        
    def fit(self, X, y):
        self.tree = self._build_tree(X, y, depth=0)
        
    def predict(self, X):
        return [self._predict_single(x, self.tree) for x in X]
    
    def _entropy(self, y):
        """计算熵"""
        counter = Counter(y)
        entropy = 0.0
        for count in counter.values():
            p = count / len(y)
            entropy -= p * math.log2(p)
        return entropy
    
    def _information_gain(self, X, y, feature_idx):
        """计算信息增益"""
        base_entropy = self._entropy(y)
        
        # 获取该特征的所有取值
        feature_values = [x[feature_idx] for x in X]
        unique_values = set(feature_values)
        
        # 计算按该特征分割后的条件熵
        conditional_entropy = 0.0
        split_info = 0.0
        
        for value in unique_values:
            sub_y = [y[i] for i in range(len(X)) if X[i][feature_idx] == value]
            if len(sub_y) == 0:
                continue
                
            p = len(sub_y) / len(y)
            conditional_entropy += p * self._entropy(sub_y)
            
            # 计算分裂信息(用于增益率)
            if self.criterion == 'gain_ratio':
                split_info -= p * math.log2(p)
        
        info_gain = base_entropy - conditional_entropy
        
        if self.criterion == 'gain_ratio':
            if split_info == 0:  # 避免除零
                return 0
            return info_gain / split_info
        else:
            return info_gain
    
    def _majority_vote(self, y):
        """多数投票"""
        counter = Counter(y)
        return counter.most_common(1)[0][0]
    
    def _build_tree(self, X, y, depth):
        """递归构建决策树"""
        # 终止条件
        if len(set(y)) == 1:  # 所有样本属于同一类别
            return y[0]
        
        if len(X[0]) == 0:  # 没有特征了
            return self._majority_vote(y)
        
        if self.max_depth and depth >= self.max_depth:
            return self._majority_vote(y)
        
        # 选择最佳特征
        best_gain = -1
        best_feature = None
        
        for feature_idx in range(len(X[0])):
            gain = self._information_gain(X, y, feature_idx)
            if gain > best_gain:
                best_gain = gain
                best_feature = feature_idx
        
        if best_gain == 0:  # 没有信息增益
            return self._majority_vote(y)
        
        # 构建子树
        tree = {
            'feature': best_feature,
            'feature_name': self.feature_names[best_feature],
            'children': {}
        }
        
        # 按最佳特征的值分割数据
        feature_values = [x[best_feature] for x in X]
        unique_values = set(feature_values)
        
        for value in unique_values:
            sub_X = [x[:best_feature] + x[best_feature+1:] for i, x in enumerate(X) if X[i][best_feature] == value]
            sub_y = [y[i] for i in range(len(X)) if X[i][best_feature] == value]
            
            if len(sub_X) == 0:
                tree['children'][value] = self._majority_vote(y)
            else:
                tree['children'][value] = self._build_tree(sub_X, sub_y, depth + 1)
        
        return tree
    
    def _predict_single(self, x, tree):
        """预测单个样本"""
        if isinstance(tree, int):  # 叶节点
            return tree
        
        feature_value = x[tree['feature']]
        if feature_value in tree['children']:
            return self._predict_single(x, tree['children'][feature_value])
        else:
            # 如果遇到未见过的特征值,返回训练集中最常见的类别
            return 1  # 默认给贷款(根据数据分布调整)

def load_data(filename):
    """加载数据文件"""
    with open(filename, 'r') as f:
        lines = f.readlines()
    
    data = []
    for line in lines:
        line = line.strip()
        if line:
            values = list(map(int, line.split(',')))
            data.append(values)
    
    return data

def evaluate_model(y_true, y_pred):
    """评估模型性能"""
    accuracy = sum(1 for i in range(len(y_true)) if y_true[i] == y_pred[i]) / len(y_true)
    
    # 计算精确率、召回率、F1分数
    tp = sum(1 for i in range(len(y_true)) if y_true[i] == 1 and y_pred[i] == 1)
    fp = sum(1 for i in range(len(y_true)) if y_true[i] == 0 and y_pred[i] == 1)
    fn = sum(1 for i in range(len(y_true)) if y_true[i] == 1 and y_pred[i] == 0)
    
    precision = tp / (tp + fp) if (tp + fp) > 0 else 0
    recall = tp / (tp + fn) if (tp + fn) > 0 else 0
    f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0
    
    return {
        'accuracy': accuracy,
        'precision': precision,
        'recall': recall,
        'f1_score': f1
    }

def print_tree(tree, indent=0):
    """打印决策树结构"""
    if isinstance(tree, int):
        print('  ' * indent + f'类别: {tree}')
    else:
        print('  ' * indent + f'特征: {tree["feature_name"]}')
        for value, subtree in tree['children'].items():
            print('  ' * (indent + 1) + f'值 {value}:')
            print_tree(subtree, indent + 2)

def main():
    # 加载数据
    train_data = load_data('dataset.txt')
    test_data = load_data('testset.txt')
    
    # 分离特征和标签
    X_train = [sample[:-1] for sample in train_data]
    y_train = [sample[-1] for sample in train_data]
    
    X_test = [sample[:-1] for sample in test_data]
    y_test = [sample[-1] for sample in test_data]
    
    print("=" * 60)
    print("决策树分类器 - 贷款审批预测")
    print("=" * 60)
    
    # 使用信息增益
    print("\n1. 使用信息增益作为特征选择标准:")
    print("-" * 40)
    
    dt_info_gain = DecisionTree(criterion='info_gain')
    dt_info_gain.fit(X_train, y_train)
    
    # 预测
    y_pred_info = dt_info_gain.predict(X_test)
    
    # 评估
    metrics_info = evaluate_model(y_test, y_pred_info)
    
    print(f"准确率: {metrics_info['accuracy']:.3f}")
    print(f"精确率: {metrics_info['precision']:.3f}")
    print(f"召回率: {metrics_info['recall']:.3f}")
    print(f"F1分数: {metrics_info['f1_score']:.3f}")
    
    print("\n决策树结构:")
    print_tree(dt_info_gain.tree)
    
    # 使用增益率
    print("\n\n2. 使用增益率作为特征选择标准:")
    print("-" * 40)
    
    dt_gain_ratio = DecisionTree(criterion='gain_ratio')
    dt_gain_ratio.fit(X_train, y_train)
    
    # 预测
    y_pred_ratio = dt_gain_ratio.predict(X_test)
    
    # 评估
    metrics_ratio = evaluate_model(y_test, y_pred_ratio)
    
    print(f"准确率: {metrics_ratio['accuracy']:.3f}")
    print(f"精确率: {metrics_ratio['precision']:.3f}")
    print(f"召回率: {metrics_ratio['recall']:.3f}")
    print(f"F1分数: {metrics_ratio['f1_score']:.3f}")
    
    print("\n决策树结构:")
    print_tree(dt_gain_ratio.tree)
    
    # 详细预测结果
    print("\n\n3. 测试集详细预测结果:")
    print("-" * 40)
    print("样本\t实际类别\t预测类别(信息增益)\t预测类别(增益率)")
    print("-" * 60)
    
    for i in range(len(X_test)):
        actual = "贷款" if y_test[i] == 1 else "不贷款"
        pred_info = "贷款" if y_pred_info[i] == 1 else "不贷款"
        pred_ratio = "贷款" if y_pred_ratio[i] == 1 else "不贷款"
        
        print(f"{i+1}\t{actual}\t\t{pred_info}\t\t\t{pred_ratio}")
    
    # 性能比较
    print("\n\n4. 性能比较:")
    print("-" * 40)
    print("指标\t\t信息增益\t增益率")
    print("-" * 30)
    print(f"准确率\t\t{metrics_info['accuracy']:.3f}\t\t{metrics_ratio['accuracy']:.3f}")
    print(f"精确率\t\t{metrics_info['precision']:.3f}\t\t{metrics_ratio['precision']:.3f}")
    print(f"召回率\t\t{metrics_info['recall']:.3f}\t\t{metrics_ratio['recall']:.3f}")
    print(f"F1分数\t\t{metrics_info['f1_score']:.3f}\t\t{metrics_ratio['f1_score']:.3f}")

if __name__ == "__main__":
    main()

实验结果:

五、总结

决策树的结构就像流程图,即使不是专业人士也可以看懂决策逻辑,方便解释。但容易过拟合,再过度适配训练集后,面对新数据可能预测准确率会降低。个别异常数据可能导致树的结构发生较大变化。

决策树的实现对于新手来说比较友好,可以快速理解“特征选择"“过拟合”“模型优化”等概念。对于资深专业人员来说,决策树也可以起到分类、回归的作用。

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐