R 中的决策树:分类树示例

⚡ 智能摘要

R 语言中的决策树使用简单的“是/否”规则将数据分割成多个分支,直到每个叶节点都只包含一个主要类别。本教程将演示如何在泰坦尼克号幸存者数据集上构建、绘制、评估和调整 rpart 分类树。

  • 🌳 核心定义: 树递归地划分预测空间,在每个节点选择最大程度减少类别不纯度的分割。
  • 🔄 数据准备: 使用 sample() 对已排序的 Titanic 文件进行随机排序,删除标识符列,转换因子,然后删除 NA 行。
  • 🛠️ 模型语法: 调用 rpart(survived~., data = data_train, method = 'class') 并使用 rpart.plot(fit, extra = 106) 渲染结果。
  • ⚙️ 调整: rpart.control() 公开了 minsplit、minbucket、maxdepth 和 cp,将准确率提高到 79.90%。
  • 📊 评估: 使用 table() 构建混淆矩阵,然后将对角线元素除以总数,得到 78.47% 的测试准确率。

R语言中的决策树

什么是决策树?

决策树 决策树是一种用途广泛的机器学习算法,可以执行分类和回归任务。它们功能强大,能够拟合复杂的数据集。此外,决策树是随机森林的基本组成部分,而随机森林是当今最强大的机器学习算法之一。

在用代码构建树之前,了解树是如何决定分支位置的会很有帮助。

决策树如何工作?

决策树将数据集转换为由三种节点构建的“是/否”问题的流程图: 包含每一次训练观察, 内部节点 询问有关一个预测变量的问题,并将数据分成两部分,以及 停止分裂并返回多数类。

增长遵循一种贪婪的程序,称为 递归二进制分区:

  1. 评估每个候选分割结果。 对于每个预测因子和分割点,衡量由此产生的两个组的纯度如何。
  2. 保留最好的那一个。 在该节点上,需要提出的问题就是:哪种分裂方式能最大程度地减少杂质?
  3. 对每个孩子重复以上步骤 直到控制规则阻止它为止:minsplit、minbucket、maxdepth 或 cp。
  4. 修剪。 然后 cp 会修剪那些没有收益的分支,从而阻止树木记住训练集。

因为每个问题都将一个变量与一个阈值进行比较,所以该算法永远不需要缩放或虚拟编码。

决策树中的基尼指数与熵

可以通过两种方式测量这种杂质,rpart 允许您选择。

标准 基尼指数 熵(信息增益)
公式 1 – 平方类比例之和 -p 乘以 log2(p) 的总和
范围(两类) 0到0.5 0到1
计算 速度更快,无需对数 速度较慢,使用对数
rpart 设置 默认 parms = list(split = “信息”)
fit_entropy <- rpart(survived~., data = data_train, method = 'class',
    parms = list(split = "information"))

实际上,这两个标准大多数时候都会选择相同的分割方式,因此默认的基尼系数是一个安全的选择。

决策树的优点和缺点

权衡利弊可以告诉你什么时候单棵树就足够了,什么时候应该使用集成树。

优势

  • 完全可解读: 拟合模型是一个任何利益相关者都能看懂的图表。
  • 最小预处理: 无需缩放或归一化,因子可直接使用。
  • 能够同时处理这两项任务: method = 'class' 拟合分类器,method = 'anova' 拟合回归树。
  • 训练速度快: 大型数据集可以在几秒钟内处理完毕,因此树状结构是一个有用的初始基线。

缺点

  • 高方差: 训练数据的一个微小变化就可能产生完全不同的决策树。
  • 容易出现过度拟合: 不受约束的树会一直生长,直到每一片叶子都是纯净的,除非 cp 和 maxdepth 对其进行限制。
  • 仅支持轴平行分割: 对角边界需要进行许多阶梯状切割。

前两个弱点的解决方法是取多棵树的平均值,这就是…… 随机森林 一样。

如何在 R 语言中训练和可视化决策树

要在 R 中构建你的第一个决策树,你需要完成以下七个步骤:

  • 步骤 1:导入数据
  • 第 2 步:清理数据集
  • 步骤 3:创建训练/测试集
  • 步骤 4:建立模型
  • 步骤 5:做出预测
  • 第 6 步:衡量绩效
  • 步骤 7:调整超参数

步骤1)导入数据

如果你对泰坦尼克号的命运感到好奇,你可以观看此视频 Youtube。该数据集的目的是预测哪些人与冰山相撞后更有可能幸存。数据集包含 13 个变量和 1309 个观测值。数据集按变量 X 排序。

set.seed(678)
path <- 'https://raw.githubusercontent.com/guru99-edu/R-Programming/master/titanic_data.csv'
titanic <-read.csv(path)
head(titanic)

输出:

##   X pclass survived                                            name    sex
## 1 1      1        1                   Allen, Miss. Elisabeth Walton female
## 2 2      1        1                  Allison, Master. Hudson Trevor   male
## 3 3      1        0                    Allison, Miss. Helen Loraine female
## 4 4      1        0            Allison, Mr. Hudson Joshua Creighton   male
## 5 5      1        0 Allison, Mrs. Hudson J C (Bessie Waldo Daniels) female
## 6 6      1        1                             Anderson, Mr. Harry   male
##       age sibsp parch ticket     fare   cabin embarked
## 1 29.0000     0     0  24160 211.3375      B5        S
## 2  0.9167     1     2 113781 151.5500 C22 C26        S
## 3  2.0000     1     2 113781 151.5500 C22 C26        S
## 4 30.0000     1     2 113781 151.5500 C22 C26        S
## 5 25.0000     1     2 113781 151.5500 C22 C26        S
## 6 48.0000     0     0  19952  26.5500     E12        S
##                         home.dest
## 1                    St Louis, MO
## 2 Montreal, PQ / Chesterville, ON
## 3 Montreal, PQ / Chesterville, ON
## 4 Montreal, PQ / Chesterville, ON
## 5 Montreal, PQ / Chesterville, ON
## 6                    New York, NY
tail(titanic)

输出:

##         X pclass survived                      name    sex  age sibsp
## 1304 1304      3        0     Yousseff, Mr. Gerious   male   NA     0
## 1305 1305      3        0      Zabour, Miss. Hileni female 14.5     1
## 1306 1306      3        0     Zabour, Miss. Thamine female   NA     1
## 1307 1307      3        0 Zakarian, Mr. Mapriededer   male 26.5     0
## 1308 1308      3        0       Zakarian, Mr. Ortin   male 27.0     0
## 1309 1309      3        0        Zimmerman, Mr. Leo   male 29.0     0
##      parch ticket    fare cabin embarked home.dest
## 1304     0   2627 14.4583              C          
## 1305     0   2665 14.4542              C          
## 1306     0   2665 14.4542              C          
## 1307     0   2656  7.2250              C          
## 1308     0   2670  7.2250              C          
## 1309     0 315082  7.8750              S

从头部和尾部输出中,您可以注意到数据没有被打乱。这是一个大问题!当您将数据拆分为训练集和测试集时,您将选择 仅由 1 类和 2 类乘客(3 类乘客中没有一个进入观测值前 80%),这意味着算法永远不会看到 3 类乘客的特征。这个错误会导致预测不佳。

为了解决这个问题,您可以使用函数sample()。

shuffle_index <- sample(1:nrow(titanic))
head(shuffle_index)

决策树 R 代码解释

  • sample(1:nrow(titanic)):生成从1到1309(即最大行数)的索引随机列表。

输出:

## [1]  288  874 1078  633  887  992

您将使用此索引来重新排列泰坦尼克号数据集。

titanic <- titanic[shuffle_index, ]
head(titanic)

输出:

##         X pclass survived
## 288   288      1        0
## 874   874      3        0
## 1078 1078      3        1
## 633   633      3        0
## 887   887      3        1
## 992   992      3        1
##                                                           name    sex age
## 288                                      Sutton, Mr. Frederick   male  61
## 874                   Humblen, Mr. Adolf Mathias Nicolai Olsen   male  42
## 1078                                 O'Driscoll, Miss. Bridget female  NA
## 633  Andersson, Mrs. Anders Johan (Alfrida Konstantia Brogren) female  39
## 887                                        Jermyn, Miss. Annie female  NA
## 992                                           Mamee, Mr. Hanna   male  NA
##      sibsp parch ticket    fare cabin embarked           home.dest## 288      0     0  36963 32.3208   D50        S     Haddenfield, NJ
## 874      0     0 348121  7.6500 F G63        S                    
## 1078     0     0  14311  7.7500              Q                    
## 633      1     5 347082 31.2750              S Sweden Winnipeg, MN
## 887      0     0  14313  7.7500              Q                    
## 992      0     0   2677  7.2292              C	

步骤2)清理数据集

多个变量包含缺失值(NA)。清理过程分三个部分进行:

  • 删除变量 home.dest、cabin、name、X 和 ticket。
  • 为 pclass 和 survivors 创建因子变量
  • 放弃 NA
library(dplyr)
# Drop variables
clean_titanic <- titanic %>%
select(-c(home.dest, cabin, name, X, ticket)) %>% 
#Convert to factor level
	mutate(pclass = factor(pclass, levels = c(1, 2, 3), labels = c('Upper', 'Middle', 'Lower')),
	survived = factor(survived, levels = c(0, 1), labels = c('No', 'Yes'))) %>%
na.omit()
glimpse(clean_titanic)

Code 说明

  • select(-c(home.dest, cabin, name, X, ticket)):删除不必要的变量
  • pclass = factor(pclass, levels = c(1,2,3), labels= c('Upper', 'Middle', 'Lower')): 为变量 pclass 添加标签。1 变为 Upper,2 变为 Middle,3 变为 Lower
  • factor(survived, levels = c(0,1), labels = c('No', 'Yes')): 为变量 survived 添加标签。0 变为“No”,1 变为“Yes”。
  • na.omit():删除 NA 观测值

输出:

## Observations: 1,045
## Variables: 8
## $ pclass   <fctr> Upper, Lower, Lower, Upper, Middle, Upper, Middle, U...
## $ survived <fctr> No, No, No, Yes, No, Yes, Yes, No, No, No, No, No, Y...
## $ sex      <fctr> male, male, female, female, male, male, female, male...
## $ age      <dbl> 61.0, 42.0, 39.0, 49.0, 29.0, 37.0, 20.0, 54.0, 2.0, ...
## $ sibsp    <int> 0, 0, 1, 0, 0, 1, 0, 0, 4, 0, 0, 1, 1, 0, 0, 0, 1, 1,...
## $ parch    <int> 0, 0, 5, 0, 0, 1, 0, 1, 1, 0, 0, 1, 1, 0, 2, 0, 4, 0,...
## $ fare     <dbl> 32.3208, 7.6500, 31.2750, 25.9292, 10.5000, 52.5542, ...
## $ embarked <fctr> S, S, S, S, S, S, S, S, S, C, S, S, S, Q, C, S, S, C...		

步骤3)创建训练/测试集

在训练模型之前,您需要执行两个步骤:

  • 创建训练和测试集:在训练集上训练模型,并在测试集(即未见数据)上测试预测
  • 从控制台安装 rpart.plot

常见的做法是将数据分成 80/20,80% 的数据用于训练模型,20% 的数据用于进行预测。您需要创建两个单独的数据框。在完成模型构建之前,您不想接触测试集。您可以创建一个名为 create_train_test() 的函数,它接受三个参数。

create_train_test(df, size = 0.8, train = TRUE)
arguments:
-df: Dataset used to train the model.
-size: Size of the split. By default, 0.8. Numerical value
-train: If set to `TRUE`, the function creates the train set, otherwise the test set. Default value sets to `TRUE`. Boolean value.You need to add a Boolean parameter because R does not allow to return two data frames simultaneously.
create_train_test <- function(data, size = 0.8, train = TRUE) {
    n_row = nrow(data)
    total_row = size * n_row
    train_sample <- 1: total_row
    if (train == TRUE) {
        return (data[train_sample, ])
    } else {
        return (data[-train_sample, ])
    }
}

Code 说明

  • function(data, size=0.8, train = TRUE): 在函数中添加参数
  • n_row = nrow(data): 计算数据集中的行数
  • total_row = size*n_row:返回第 n 行以构建训练集
  • train_sample <- 1:total_row: 选择第一行到第 n 行
  • if (train ==TRUE){ } else { }:如果条件设置为真,则返回训练集,否则返回测试集。

您可以测试您的功能并检查尺寸。

data_train <- create_train_test(clean_titanic, 0.8, train = TRUE)
data_test <- create_train_test(clean_titanic, 0.8, train = FALSE)
dim(data_train)

输出:

## [1] 836   8
dim(data_test)

输出:

## [1] 209   8

训练数据集有 836 行 8 列,而测试数据集有 209 行和 8 列。

您可以使用函数 prop.table() 结合 table() 来验证随机化过程是否正确。

prop.table(table(data_train$survived))

输出:

##
##        No       Yes 
## 0.5944976 0.4055024
prop.table(table(data_test$survived))

输出:

## 
##        No       Yes 
## 0.5789474 0.4210526

在两个数据集中,幸存者的数量相同,约为 40%。

安装 rpart.plot

conda 库中没有 rpart.plot。您可以从控制台安装它:

install.packages("rpart.plot")

步骤4)建立模型

您已准备好构建模型。rpart() 决策树函数的语法如下:

rpart(formula, data=, method='')
arguments:			
- formula: The function to predict
- data: Specifies the data frame
- method:			
- "class" for a classification tree 			
- "anova" for a regression tree	

您使用类方法是因为您要预测一个类。

library(rpart)
library(rpart.plot)
fit <- rpart(survived~., data = data_train, method = 'class')
rpart.plot(fit, extra = 106)

Code 说明

  • rpart():用于拟合模型的函数。参数为:
    • 幸存下来〜。:决策树的公式
    • 数据=data_train:数据集
    • method = ‘class’:拟合二元模型
  • rpart.plot(fit, extra=106): 绘制决策树。extra 参数设置为 106,表示第二个类别的概率加上每个节点中观测值的百分比。您可以参考…… 小插图 了解有关其他选择的更多信息。

输出:

 在 R 中构建决策树模型

从图的顶部、深度为 0(共 3 层)的根节点开始:

  1. 最上方是总体生存概率。它显示了在事故中幸存的乘客比例。41% 的乘客幸存。
  2. 该节点询问乘客的性别是否为男性。如果是,则向下遍历根节点的左子节点(深度为 1)。其中 63% 为男性,存活概率为 21%。
  3. 在第二个节点中,您询问男性乘客是否年满 3.5 岁。如果是,则生存几率为 19%。
  4. 你继续这样做,就能了解哪些特征会影响生存的可能性。

请注意,决策树的众多特性之一是它们几乎不需要数据准备。特别是,它们不需要特征缩放或居中。

默认情况下,rpart() 函数使用 基尼 使用基尼系数来衡量各节点间的混合程度。基尼系数越高,表示该节点内各类别混合程度越高,因此算法总是选择基尼系数最低的节点进行分割。

步骤 5)做出预测

您可以预测测试数据集。要进行预测,可以使用 predict() 函数。R 决策树预测的基本语法是:

predict(fitted_model, df, type = 'class')
arguments:
- fitted_model: This is the object stored after model estimation. 
- df: Data frame used to make the prediction
- type: Type of prediction			
    - 'class': for classification			
    - 'prob': to compute the probability of each class			
    - 'vector': Predict the mean response at the node level	

现在,请预测测试集中 209 名乘客中的每一位,该模型是否预期他们能在碰撞中幸存下来。

predict_unseen <-predict(fit, data_test, type = 'class')

Code 说明

  • predict(fit, data_test, type = 'class'): 预测测试集的类(0/1)

现在将预测结果与实际结果进行比较。

table_mat <- table(data_test$survived, predict_unseen)
table_mat

Code 说明

  • table(data_test$survived, predict_unseen): 构建预测类别与真实结果的列联表

输出:

##      predict_unseen
##        No Yes
##   No  106  15
##   Yes  30  58

行代表实际值,列代表预测值。该模型正确识别出106名非幸存者和58名幸存者,但将15名非幸存者标记为幸存者,将30名幸存者标记为死亡。

步骤 6)衡量绩效

您可以使用以下方法计算分类任务的准确度度量 混淆矩阵:

混淆矩阵 是评估分类性能的更好选择。一般的想法是计算 True 实例被分类为 False 的次数。

测量 R 中决策树的性能

混淆矩阵中的每一行代表一个实际目标,而每一列代表一个预测目标。该矩阵的第一行考虑的是死亡乘客(负类):106 人被正确分类为死亡(真阴性),而 15 人被错误地归类为幸存者(假阳性第二行考虑的是幸存者:58 人被正确识别(真阳性),而漏检了 30 例(假阴性).

您可以计算 准确度测试 来自混淆矩阵:

测量 R 中决策树的性能

它是真阳性和真阴性占矩阵总和的比例。使用 R,您可以按如下方式编写代码:

accuracy_Test <- sum(diag(table_mat)) / sum(table_mat)

Code 说明

  • sum(diag(table_mat)):对角线的总和
  • sum(table_mat):矩阵的总和。

您可以打印测试集的准确率:

print(paste('Accuracy for test', accuracy_Test))

输出:

## [1] "Accuracy for test 0.784688995215311"

测试集上的准确率为 0.7847,即 78.47%。在训练集上重复上述步骤,以观察模型的过拟合程度。

步骤 7)调整超参数

R 中的决策树有各种参数来控制拟合的各个方面。在 rpart 决策树库中,您可以使用 rpart.control() 函数控制参数。在下面的代码中,您将介绍要调整的参数。您可以参考 小插图 其他参数。

rpart.control(minsplit = 20, minbucket = round(minsplit/3), maxdepth = 30)
Arguments:
-minsplit: Set the minimum number of observations in the node before the algorithm perform a split
-minbucket: Set the minimum number of observations in a terminal node, i.e. the leaf
-maxdepth: Set the maximum depth of any node of the final tree. The root node is treated as depth 0

我们将按照以下步骤进行:

  • 构造函数来返回准确率
  • 调整最大深度
  • 调整节点分裂前必须拥有的最小样本数量
  • 调整叶节点必须具有的最小样本数量

你可以编写一个函数来显示准确率。你只需包装一下之前使用的代码:

  1. 预测:predict_unseen <- 预测(fit,data_test,type ='class')
  2. 生成表:table_mat <- table(data_test$survived, predict_unseen)
  3. 计算准确度:accuracy_Test <- sum(diag(table_mat))/sum(table_mat)
accuracy_tune <- function(fit) {
    predict_unseen <- predict(fit, data_test, type = 'class')
    table_mat <- table(data_test$survived, predict_unseen)
    accuracy_Test <- sum(diag(table_mat)) / sum(table_mat)
    accuracy_Test
}

现在调整参数,看看能否改进默认模型。提醒一下,你需要达到比 0.7847 更低的准确率。

control <- rpart.control(minsplit = 4,
    minbucket = round(5 / 3),
    maxdepth = 3,
    cp = 0)
tune_fit <- rpart(survived~., data = data_train, method = 'class', control = control)
accuracy_tune(tune_fit)

输出:

## [1] 0.7990431

使用以下参数:

minsplit = 4
minbucket = round(5/3)
maxdepth = 3
cp = 0

准确率从 0.7847 提高到 0.7990,因此调整后的树比默认配置高出约 1.4 个百分点。

R 语言中的决策树:快速函数参考

下表列出了上述七个步骤中使用的每个函数,以及提供该函数的包和它所需的参数。 R.

图书馆 目的 功能 增益级 参数 信息
部分 在 R 中训练分类树 rpart() 公式、df、方法
部分 训练回归树 rpart() 方差分析 公式、df、方法
部分 绘制树木 rpart.plot() 拟合模型
基地 预测 预测() 拟合模型,类型
基地 预测 预测() 概率 拟合模型,类型
基地 预测 预测() 向量 拟合模型,类型
部分 控制参数 rpart.控制() 最小分割 在算法进行拆分之前,设置节点中的最小观察数
最小桶 设置终端节点(即叶节点)的最小观测值数量。
最大深度 设置最终树中所有节点的最大深度。根节点的深度为 0。
部分 使用控制参数训练模型 rpart() 公式、df、方法、控制

注意:在训练数据上训练模型,并在看不见的数据集(即测试集)上测试性能。

常见问题

两者都适用于拟合分类树和回归树。rpart() 函数实现了 CART 模型,并通过 cp 函数内置了交叉验证剪枝功能,并且与 rpart.plot 函数配合使用可以生成清晰的图表,因此它是更常用的选择。

cp 参数设置了分支必须带来的最小改进值才能被保留。较大的值会进行更激进的剪枝,生成更小的决策树。使用 printcp() 和 plotcp() 函数可以找到交叉验证误差最小的 cp 值。

是的。rpart() 使用代理分割,将缺少预测变量的观测值路由到最相似的分支。本教程使用 na.omit() 代替,纯粹是为了简化示例数据集。

决策树为信贷、保险和医疗保健等行业的可解释人工智能提供支持,监管机构可能会要求提供决策背后的确切理由。它们也是梯度提升和随机森林模型中的基础学习器。

是的。AI 助手可以将分割规则翻译成通俗易懂的语言,建议待测试的 cp 值,并在 printcp() 输出中标记过拟合情况。在采纳任何建议之前,请务必将其与您自己的交叉验证结果进行比对。

总结一下这篇文章: