Skip to content

R - 决策树

决策树是一种强大且可解释的机器学习模型,它以树状结构表示选择及其结果。它用于分类(例如,预测电子邮件是否为垃圾邮件)和回归(例如,预测房价)。树中的每个节点表示对特征的测试,每个分支表示测试的结果,每个叶节点表示一个类别标签或一个连续值。

现代 R 机器学习工作流通常使用 tidymodels 框架构建,这是一个包集合,为从数据分割到模型评估的整个建模过程提供了一个统一的、以整洁数据为优先的方法。

首先,我们需要安装必要的包。tidymodels 是核心框架,而 rpart.plot 对于可视化我们创建的树非常出色。

# Install the packages if you don't have them yet
# install.packages(c("tidymodels", "rpart.plot"))
# Load the library
library(tidymodels)

我们将使用 readingSkills 数据集来预测一个人是否为 nativeSpeaker(母语者),基于他们的年龄、鞋码和考试分数。我们将遵循一个标准的、健壮的机器学习工作流。

我们加载数据并确保我们的结果变量 (nativeSpeaker) 是一个因子 (factor),这是分类所必需的。

# Load the data (it's part of the 'party' package, but we can load it directly)
data(readingSkills, package = "party")
# Glimpse the data structure
glimpse(readingSkills)

glimpse 的输出:

Rows: 200
Columns: 4
$ nativeSpeaker <fct> yes, yes, no, yes, yes, yes, no, yes, no, yes, ...
$ age <int> 5, 6, 11, 7, 11, 10, 8, 9, 10, 6, 8, 5, ...
$ shoeSize <dbl> 24.83, 25.95, 30.42, 28.66, 31.88, 30.08, ...
$ score <dbl> 32.29, 36.63, 49.61, 40.28, 55.46, 52.83, ...

步骤 2:将数据分割成训练集和测试集

Section titled “步骤 2:将数据分割成训练集和测试集”

这是关键一步。我们在训练集上训练模型,并在未见的测试集上评估其性能,以获得对其真实世界性能的无偏估计。

# Set a seed for reproducibility
set.seed(123)
# Create the data split object
data_split <- initial_split(readingSkills, prop = 0.80, strata = nativeSpeaker)
# Extract training and testing sets
train_data <- training(data_split)
test_data <- testing(data_split)

使用 tidymodels,我们首先定义一个模型规范(模型的类型),然后将其“拟合”到我们的数据。

# 1. Define the model specification
# We'll use the 'rpart' engine (a classic R decision tree algorithm)
tree_spec <- decision_tree() %>%
set_engine("rpart") %>%
set_mode("classification")
# 2. Fit the model to the training data
tree_fit <- tree_spec %>%
fit(nativeSpeaker ~ ., data = train_data)

可视化决策树是理解它如何做出决策的关键。rpart.plot 包可以创建出色且易读的图表。

# Load the visualization library
library(rpart.plot)
# Create the plot
rpart.plot(tree_fit$fit, box.palette = "RdBu", shadow.col = "gray", nn = TRUE)

输出描述: 这段代码生成了一个显示决策树的图。顶部的根节点可能基于 score 进行分割。例如,score < 40 的分支可能导致另一个基于 age 分割的节点,而 score >= 40 的分支则直接导致一个预测为 yes(母语者)的叶节点。该图有颜色,每个节点中显示百分比和计数,使其高度可解释。

现在,让我们看看我们的模型在未见过的测试数据上的表现如何。

# Make predictions on the test set
predictions <- predict(tree_fit, new_data = test_data, type = "class")
# Combine predictions with the actual truth
results <- bind_cols(test_data, predictions) %>%
select(nativeSpeaker, .pred_class)
# Calculate performance metrics, like a confusion matrix
conf_mat(results, truth = nativeSpeaker, estimate = .pred_class)

混淆矩阵显示了多少预测是正确和不正确的:

Truth
Prediction no yes
no 19 1
yes 1 19

混淆矩阵显示我们的模型在测试集上表现出色。从可视化图表中,我们可以得出简单的规则,例如:“如果一个人的分数大于 48,他们很可能是母语者。”

  • 过拟合: 简单的树可能会过拟合训练数据。tidymodels 提供了强大的工具,用于调优超参数(如树的深度),以找到泛化能力最好的模型。
  • 替代引擎: 除了 rpart,您还可以在相同的 tidymodels 框架内通过更改 set_engine() 使用其他引擎,如 C5.0 或 partykit::ctree。
  • 集成方法: 为了获得更高的性能,单个决策树通常会组合成强大的集成模型,如随机森林 (ranger 引擎) 或梯度提升机 (xgboost 引擎)。