决策树模型怎么生成?用 scikit-learn 训练并导出规则

用 scikit-learn 从带标签的 Iris 数据训练决策树,导出完整文字规则和逐行预测,通过独立测试、混淆矩阵及真实决策路径检查模型。

决策树模型是从带标签样本中学习“如果某个特征满足条件,就走向哪个分支”的分类规则。用 scikit-learn 可以先划分训练集和测试集,训练一棵受限深度的树,再导出可阅读的文字规则。导出规则只是看清模型如何判断,是否可用仍要核对独立测试结果。

用四个测量值预测鸢尾花类别

本教程使用 scikit-learn 自带的 Iris 数据:150 条样本、4 个数值特征、3 个类别,特征依次是萼片长度、萼片宽度、花瓣长度、花瓣宽度,单位为厘米。load_iris 官方文档给出了数据形状及类别名称。这里是公开教学数据演示,不代表你的业务分类效果。

决策树模型怎么生成?用 scikit-learn 训练并导出规则

本机实际环境为 Windows 11、Python 3.13.14、scikit-learn 1.9.1、NumPy 2.5.3。准备一个空目录并安装:

python -m pip install scikit-learn
python -c "import sklearn; print(sklearn.__version__)"

先留测试集,再训练和导出规则

保存下方完整代码为 tree_rules.py,执行 python tree_rules.py。代码按类别比例划分:105 条训练,45 条测试;测试集不参与 fit。max_depth=3 限制树深,min_samples_leaf=3 要求每个叶节点至少有 3 条训练样本。DecisionTreeClassifier 文档解释了这些参数及随机种子的作用。

from pathlib import Path
import csv, json
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score, confusion_matrix
from sklearn.tree import DecisionTreeClassifier, export_text

iris = load_iris()
X_train, X_test, y_train, y_test = train_test_split(
    iris.data, iris.target, test_size=0.30, stratify=iris.target, random_state=42)
model = DecisionTreeClassifier(max_depth=3, min_samples_leaf=3, random_state=42)
model.fit(X_train, y_train)
rules = export_text(model, feature_names=list(iris.feature_names),
                    class_names=list(iris.target_names), decimals=3,
                    max_depth=model.get_depth())
Path('tree_rules.txt').write_text(rules, encoding='utf-8')
pred = model.predict(X_test)
prob = model.predict_proba(X_test)
assert pred.shape == y_test.shape
assert all(abs(row.sum() - 1) < 1e-10 for row in prob)
leaves = model.apply(X_test)
with open('tree_test_predictions.csv', 'w', newline='', encoding='utf-8-sig') as file:
    writer = csv.writer(file)
    writer.writerow(['row', 'actual', 'predicted', 'leaf_id', *iris.feature_names])
    for i in range(len(y_test)):
        writer.writerow([i, iris.target_names[y_test[i]], iris.target_names[pred[i]],
                         int(leaves[i]), *X_test[i].tolist()])
first = []
for node in model.decision_path(X_test[:1]).indices:
    feature = int(model.tree_.feature[node])
    if feature >= 0:
        value = float(X_test[0, feature]); cut = float(model.tree_.threshold[node])
        first.append({'feature': iris.feature_names[feature], 'value': value,
                      'threshold': cut, 'branch': '<=' if value <= cut else '>'})
result = {'training_rows': len(y_train), 'test_rows': len(y_test),
          'depth': int(model.get_depth()), 'leaf_count': int(model.get_n_leaves()),
          'accuracy': float(accuracy_score(y_test, pred)),
          'confusion_matrix': confusion_matrix(y_test, pred, labels=model.classes_).tolist(),
          'class_order': iris.target_names[model.classes_].tolist(),
          'first_test_row_path': first, 'first_prediction': str(iris.target_names[pred[0]])}
Path('tree_result.json').write_text(json.dumps(result, indent=2), encoding='utf-8')
print(rules)
print(json.dumps(result, indent=2))

运行后得到 tree_rules.txt、tree_test_predictions.csv 和 tree_result.json。第一份便于看分支,第二份逐行对照测试样本真实类别与预测,第三份保存深度、叶节点数量、混淆矩阵和首条测试样本的实际决策路径。

export_text 官方文档说明规则可以带特征名和类别名,且超过导出深度的分支会被截断。本例将导出深度设为实际树深,避免只展示树的前几层就误以为已经拿到完整规则。

如何读规则并核对一条预测

本机规则的第一层为 petal length (cm) <= 2.450 时预测 setosa;否则继续看花瓣宽度等条件。规则中同样写 versicolor 的两个叶节点,并不表示程序重复:它们来自不同分支,保留了不同的样本分区。

首条测试样本为 [7.3, 2.9, 6.3, 1.8]。依次核对:花瓣长度 6.3 大于约 2.45;花瓣宽度 1.8 大于约 1.55;萼片长度 7.3 大于约 6.10,因此进入预测 virginica 的叶节点。本机 first_prediction 也为 virginica,CSV 中该行真实类别同样为 virginica。

文字导出的阈值只显示三位小数,接近分界线的样本不要靠四舍五入后的规则重新实现预测。代码保存的 first_test_row_path 使用模型内部阈值,真正推理应调用原模型的 predict,并保持四个特征的顺序与厘米单位。

结果有错分时,先看具体类别

这次实际测试准确率为 0.888889,即 45 条中判断对 40 条,错 5 条;树深为 3,叶节点为 5。这些是该固定划分和环境下的结果,不是决策树算法的一般准确率。

真实类别 / 预测类别 setosa versicolor virginica
setosa 15 0 0
versicolor 0 12 3
virginica 0 2 13

矩阵按行看真实类别、按列看预测类别,错分集中在 versicolor 与 virginica。打开 CSV 筛选 actual != predicted,观察这些行是否集中在相邻花瓣尺寸;不要只看总准确率,也不要删掉难样本让数字变好。

替换成自己的表格前要补齐什么

  • 每行必须有已确认的标签;没有标签的表不能直接训练分类树。
  • 本例输入是完整数值测量。字符串类别、缺失值、未知编码需先制定处理规则,并让训练和预测使用同一套处理。
  • 同一用户、设备或采集对象的多条样本,需要按对象分组划分,不能照搬随机逐行拆分。
  • 如需调整树深或叶节点大小,在训练数据内部用验证集或交叉验证选参数,保留测试集做最后检查。反复用测试集挑树会使结果失去独立性。

问答:有了规则文本,能直接替代原模型上线吗?

不建议直接复制显示后的条件上线。输出小数精度、特征顺序、缺失值处理、叶节点类别顺序都可能改变判断。规则文本用于说明;若必须转成业务规则,应以模型实际阈值为依据,拿包含边界值的固定样本逐条比对,确认输出一致后再接入业务。

Ai菜鸟网。发布者:AI小管家,转载请注明出处:https://www.alyyhw.com/30303.html

赞 (0)
AI小管家的头像AI小管家
Datasets 批量 map 怎么把一行拆成多行?避免列长不一致并保留来源
上一篇 1天前
PyTorch 训练中断后怎么继续?保存并恢复模型和优化器检查点
下一篇 1天前

相关推荐

联系我们

联系我们

1

在线咨询: QQ交谈

邮件:admin@example.com

工作时间:周一至周五,9:30-18:30,节假日休息

关注微信
关注微信
分享本页
返回顶部