规则提取决策树

wen IT资讯 24

本文目录导读:

规则提取决策树

  1. 规则提取的基本方法
  2. 规则优化和简化
  3. 高级规则提取技术
  4. 规则评估指标
  5. 实际应用示例
  6. 规则可视化

我来详细介绍如何从决策树中提取规则,这是将决策树模型转化为可解释规则的重要方法。

规则提取的基本方法

从根节点到叶节点的路径提取

每条从根到叶子的路径形成一条规则:

# 示例:从决策树提取规则
from sklearn.tree import DecisionTreeClassifier
from sklearn.datasets import load_iris
import numpy as np
# 训练决策树
iris = load_iris()
X, y = iris.data, iris.target
clf = DecisionTreeClassifier(max_depth=3)
clf.fit(X, y)
# 提取规则
def extract_rules(tree, feature_names=None):
    """
    从决策树中提取规则
    """
    tree_ = tree.tree_
    feature_name = [
        feature_names[i] if feature_names else f"feature_{i}"
        for i in range(tree_.n_features)
    ]
    rules = []
    def get_rule(node, conditions):
        # 叶节点:生成规则
        if tree_.children_left[node] == tree_.children_right[node]:
            class_name = tree_.value[node].argmax()
            rule = {"conditions": conditions.copy(), "class": class_name}
            rules.append(rule)
            return
        # 非叶节点:递归
        if tree_.feature[node] != -2:  # -2 表示叶节点
            feature = feature_name[tree_.feature[node]]
            threshold = tree_.threshold[node]
            # 左分支:feature <= threshold
            cond_left = conditions + [(feature, '<=', threshold)]
            get_rule(tree_.children_left[node], cond_left)
            # 右分支:feature > threshold
            cond_right = conditions + [(feature, '>', threshold)]
            get_rule(tree_.children_right[node], cond_right)
    get_rule(0, [])
    return rules
# 提取规则
rules = extract_rules(clf, iris.feature_names)
for i, rule in enumerate(rules):
    conditions_str = " AND ".join([f"{c[0]} {c[1]} {c[2]:.2f}" for c in rule["conditions"]])
    print(f"Rule {i+1}: IF {conditions_str} THEN class={rule['class']}")

使用sklearn内置方法

from sklearn.tree import export_text
# 使用export_text提取规则
text_representation = export_text(clf, feature_names=iris.feature_names)
print("决策树规则(文本格式):")
print(text_representation)

规则优化和简化

规则覆盖度计算

def calculate_rule_coverage(rule, X):
    """计算规则覆盖的样本数量"""
    mask = np.ones(X.shape[0], dtype=bool)
    for feature_name, operator, threshold in rule["conditions"]:
        feature_idx = iris.feature_names.index(feature_name)
        feature_values = X[:, feature_idx]
        if operator == '<=':
            mask &= feature_values <= threshold
        elif operator == '>':
            mask &= feature_values > threshold
    return np.sum(mask)
# 计算每条规则的覆盖度
for i, rule in enumerate(rules):
    coverage = calculate_rule_coverage(rule, X)
    print(f"Rule {i+1}: 覆盖 {coverage} 个样本")

规则冲突解决

def resolve_conflicts(predictions):
    """解决规则冲突:使用置信度较高的规则"""
    from collections import Counter
    # 统计每个预测的票数
    vote_count = Counter(predictions)
    # 返回得票最多的类别
    return vote_count.most_common(1)[0][0]
# 示例:使用多数投票解决冲突
sample = np.array([[5.1, 3.5, 1.4, 0.2]])
predictions = []
for rule in rules:
    if evaluate_rule(rule, sample):
        predictions.append(rule["class"])
if predictions:
    final_prediction = resolve_conflicts(predictions)
    print(f"最终预测: 类别 {final_prediction}")

规则剪枝

def prune_rules(rules, X, y, min_coverage=5):
    """
    剪枝规则:移除覆盖样本数太少的规则
    合并相似的规则
    """
    pruned_rules = []
    for rule in rules:
        coverage = calculate_rule_coverage(rule, X)
        # 移除覆盖太少的规则
        if coverage >= min_coverage:
            pruned_rules.append(rule)
    # 合并相似规则(简化版本)
    merged_rules = merge_similar_rules(pruned_rules)
    return merged_rules
def merge_similar_rules(rules):
    """合并具有相同结论的相似规则"""
    # 按类别分组
    rules_by_class = {}
    for rule in rules:
        class_label = rule["class"]
        if class_label not in rules_by_class:
            rules_by_class[class_label] = []
        rules_by_class[class_label].append(rule)
    merged_rules = []
    for class_label, class_rules in rules_by_class.items():
        # 对于每个类别,可以进一步优化合并策略
        merged_rules.extend(class_rules)
    return merged_rules

高级规则提取技术

基于规则学习(OneR)

def oner_rule_extraction(X, y, feature_names):
    """
    OneR (One Rule) 算法:从单个特征学习规则
    """
    n_features = X.shape[1]
    best_rule = None
    best_accuracy = 0
    for feature_idx in range(n_features):
        feature_values = X[:, feature_idx]
        feature_name = feature_names[feature_idx]
        # 对特征值进行离散化
        thresholds = np.percentile(feature_values, [25, 50, 75])
        for threshold in thresholds:
            # 创建规则
            rule_conditions = [(feature_name, '<=', threshold)]
            # 预测
            predictions = np.where(feature_values <= threshold, 
                                   np.argmax(np.bincount(y[feature_values <= threshold])),
                                   np.argmax(np.bincount(y[feature_values > threshold])))
            # 计算准确率
            accuracy = np.mean(predictions == y)
            if accuracy > best_accuracy:
                best_accuracy = accuracy
                best_rule = {
                    "conditions": rule_conditions,
                    "accuracy": accuracy,
                    "predictions": predictions
                }
    return best_rule

从随机森林提取规则

from sklearn.ensemble import RandomForestClassifier
def extract_rules_from_forest(forest, X, y, max_trees=10):
    """
    从随机森林中提取规则
    """
    rules = []
    for i, tree in enumerate(forest.estimators_[:max_trees]):
        tree_rules = extract_rules(tree, iris.feature_names)
        rules.extend(tree_rules)
        if len(rules) > 100:  # 设置规则数量上限
            break
    return rules
# 示例
rf = RandomForestClassifier(n_estimators=10)
rf.fit(X, y)
forest_rules = extract_rules_from_forest(rf, X, y)
print(f"从森林中提取了 {len(forest_rules)} 条规则")

规则评估指标

def evaluate_rules(rules, X, y):
    """
    评估规则集的质量
    """
    from sklearn.metrics import accuracy_score, precision_score, recall_score
    # 生成预测
    predictions = []
    for sample in X:
        # 找到匹配的规则
        matching_rules = []
        for rule in rules:
            if evaluate_rule(rule, sample.reshape(1, -1)):
                matching_rules.append(rule["class"])
        if matching_rules:
            # 投票决定最终预测
            predictions.append(resolve_conflicts(matching_rules))
        else:
            # 无规则匹配,使用默认类
            predictions.append(np.argmax(np.bincount(y)))
    # 计算评估指标
    metrics = {
        "accuracy": accuracy_score(y, predictions),
        "precision": precision_score(y, predictions, average='weighted'),
        "recall": recall_score(y, predictions, average='weighted')
    }
    return metrics
def evaluate_rule(rule, sample):
    """评估规则是否匹配样本"""
    for feature_name, operator, threshold in rule["conditions"]:
        feature_idx = iris.feature_names.index(feature_name)
        feature_value = sample[0, feature_idx]
        if operator == '<=' and feature_value > threshold:
            return False
        elif operator == '>' and feature_value <= threshold:
            return False
    return True

实际应用示例

def build_rule_based_classifier(X_train, y_train, feature_names):
    """
    构建完整的基于规则的分类器
    """
    # 1. 训练决策树
    clf = DecisionTreeClassifier(max_depth=4, min_samples_leaf=10)
    clf.fit(X_train, y_train)
    # 2. 提取规则
    rules = extract_rules(clf, feature_names)
    # 3. 剪枝和优化
    rules = prune_rules(rules, X_train, y_train, min_coverage=5)
    # 4. 评估
    metrics = evaluate_rules(rules, X_train, y_train)
    return {
        "rules": rules,
        "metrics": metrics,
        "num_rules": len(rules)
    }
# 使用示例
classifier = build_rule_based_classifier(
    X_train=X, 
    y_train=y, 
    feature_names=iris.feature_names
)
print(f"提取的规则数量: {classifier['num_rules']}")
print(f"规则准确率: {classifier['metrics']['accuracy']:.3f}")

规则可视化

def visualize_rules(rules, max_rules=10):
    """
    可视化提取的规则
    """
    print("=" * 60)
    print("提取的决策规则:")
    print("=" * 60)
    for i, rule in enumerate(rules[:max_rules]):
        conditions_str = " AND ".join([
            f"({c[0]} {c[1]} {c[2]:.2f})" 
            for c in rule["conditions"]
        ])
        print(f"\n规则 {i+1}:")
        print(f"  IF {conditions_str}")
        print(f"  THEN → 类别 {rule['class']}")
    if len(rules) > max_rules:
        print(f"\n... 还有 {len(rules) - max_rules} 条规则未显示")
# 调用可视化
visualize_rules(rules)

这些方法可以帮助你从决策树中有效地提取、优化和应用规则,使模型更具可解释性和实用性。

抱歉,评论功能暂时关闭!