决策树模型是从带标签样本中学习“如果某个特征满足条件,就走向哪个分支”的分类规则。用 scikit-learn 可以先划分训练集和测试集,训练一棵受限深度的树,再导出可阅读的文字规则。导出规则只是看清模型如何判断,是否可用仍要核对独立测试结果。
用四个测量值预测鸢尾花类别
本教程使用 scikit-learn 自带的 Iris 数据:150 条样本、4 个数值特征、3 个类别,特征依次是萼片长度、萼片宽度、花瓣长度、花瓣宽度,单位为厘米。load_iris 官方文档给出了数据形状及类别名称。这里是公开教学数据演示,不代表你的业务分类效果。

本机实际环境为 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