Skip to content

R - 逻辑回归

R 语言:使用 Tidymodels 进行逻辑回归的现代方法

Section titled “R 语言:使用 Tidymodels 进行逻辑回归的现代方法”

逻辑回归是一种基本的分类算法,当响应变量是分类变量时使用(例如,真/假、是/否、0/1)。它不直接建模响应变量的值,而是建模响应变量属于特定类别的概率。对于二元响应,此概率使用逻辑函数(或 Sigmoid 函数)计算。

逻辑回归中概率 (p) 的通用数学方程为:

p = 1 / (1 + e^-(β₀ + β₁x₁ + β₂x₂ + ...))

其中:

  • p 是目标结果的概率(例如,P(y=1))。
  • x₁、x₂、… 是预测变量。
  • β₀、β₁、… 是模型系数,这些系数是从数据中学习得到的。

虽然 R 基础包的 glm() 函数是实现逻辑回归的传统工具,但 tidymodels 框架为建模提供了现代化、整洁且全面的方法。它鼓励采用最佳实践,例如将数据拆分进行验证,并为多种不同模型提供了一致的接口。

我们将使用内置的 mtcars 数据集,根据汽车的引擎规格来建模其是自动挡 (am = 0) 还是手动挡 (am = 1) 变速箱。我们的现代工作流包括以下步骤:

  1. 设置:加载包并准备数据。
  2. 数据拆分:将数据分为训练集和测试集。
  3. 模型规范:定义逻辑回归模型。
  4. 模型拟合:在训练数据上训练模型。
  5. 模型评估:评估模型在测试数据上的性能。

首先,我们安装并加载必要的包。tidyverse 用于数据操作,tidymodels 用于建模。我们还将把结果变量 am 转换为因子(factor),这是 R 语言中表示分类变量的标准方式。

# 如果尚未安装,请安装包
# install.packages(c("tidyverse", "tidymodels"))
library(tidyverse)
library(tidymodels)
# 准备数据
car_data <- mtcars %>%
select(am, cyl, hp, wt) %>%
mutate(am = factor(am, levels = c(0, 1), labels = c("automatic", "manual")))
# 概览准备好的数据
glimpse(car_data)

这会产生以下结果,显示了我们的变量及其类型:

Rows: 32
Columns: 4
$ am <fct> manual, manual, manual, automatic, automatic, automatic, automa…
$ cyl <dbl> 6, 6, 4, 6, 8, 6, 8, 4, 4, 6, 6, 8, 8, 8, 8, 8, 8, 4, 4, 4, 4, …
$ hp <dbl> 110, 110, 93, 110, 175, 105, 245, 62, 95, 123, 123, 180, 180, 1…
$ wt <dbl> 2.620, 2.875, 2.320, 3.215, 3.440, 3.460, 3.570, 3.190, 3.150, …

步骤 2:将数据拆分为训练集和测试集

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

在现代机器学习中,拆分数据是至关重要的一步。我们在训练集上训练模型,并在未见过的测试集上评估其性能。这能为模型在新数据上的表现提供一个真实的估计。我们使用 rsample 包中的 initial_split 函数来完成此操作。

# 设置随机种子以保证结果可复现
set.seed(123)
# 创建拆分对象
data_split <- initial_split(car_data, prop = 0.75, strata = am)
# 创建训练和测试数据框
train_data <- training(data_split)
test_data <- testing(data_split)

使用 parsnip 包,我们可以独立于模型的具体实现来定义模型。我们指定一个 logistic_reg() 模型,并将其计算引擎设置为 "glm"(与原始教程中使用的函数相同)。

# 指定模型
log_reg_spec <- logistic_reg() %>%
set_engine("glm")
# 将模型拟合到训练数据
log_reg_fit <- log_reg_spec %>%
fit(am ~ cyl + hp + wt, data = train_data)
# 以整洁格式查看模型系数
tidy(log_reg_fit)

broom 包的 tidy() 函数为我们的模型提供了简洁的摘要:

# A tibble: 4 × 5
term estimate std.error statistic p.value
<chr> <dbl> <dbl> <dbl> <dbl>
1 (Intercept) 13.9 6.99 1.99 0.0465
2 cyl 0.871 1.13 0.770 0.441
3 hp 0.0218 0.0210 1.04 0.298
4 wt -7.39 3.40 -2.17 0.0298

现在,我们使用拟合好的模型 (log_reg_fit) 对 test_data 进行预测,并查看其性能如何。yardstick 包为此提供了相关函数。

# 在测试集上进行预测
predictions <- predict(log_reg_fit, new_data = test_data, type = "class") %>%
bind_cols(test_data)
# 创建混淆矩阵
conf_mat(predictions, truth = am, estimate = .pred_class)
# 计算准确率
accuracy(predictions, truth = am, estimate = .pred_class)

评估输出将如下所示:

# Confusion Matrix
# Truth
# Prediction automatic manual
# automatic 4 1
# manual 1 2
# Accuracy
# A tibble: 1 × 3
# .metric .estimator .estimate
# <chr> <chr> <dbl>
# 1 accuracy binary 0.75

我们的模型在测试集中正确分类了 8 辆车中的 6 辆,准确率为 75%。混淆矩阵显示它正确识别了 4 辆“自动挡”汽车和 2 辆“手动挡”汽车,同时各自错误分类了一辆。相比仅依赖训练摘要中的 p 值(这可能会产生误导),这种在未见过的数据上进行的评估为模型的预测能力提供了更可靠的衡量标准。tidymodels 工作流强制执行这些现代最佳实践,从而带来更健壮、更可靠的模型。

要改进你的模型,你可以:

  • 使用 recipes 包探索特征工程来预处理你的变量。
  • 使用 vfold_cv 进行交叉验证,以获得更健壮的性能估计。
  • 学习如何使用 ROC 曲线(来自 yardstick 包的 roc_curve 函数)来解释模型结果。

官方的 Tidymodels 网站 是深入学习的极佳资源。