使用 tidymodels 进行线性回归

在 R 中使用 tidymodels 建模

David Svancer

Data Scientist

使用 parsnip 进行模型拟合

使用 parsnip 拟合模型

在 R 中使用 tidymodels 建模

线性回归模型

使用 cty 预测 hwy  

$$hwy = \beta_{0} + \beta_{1} cty$$

模型参数

  • $ \beta_{0} $ 为截距
  • $ \beta_{1} $ 为斜率

 

高速 vs 市区油耗

在 R 中使用 tidymodels 建模

线性回归模型

使用 cty 预测 hwy  

$$hwy = \beta_{0} + \beta_{1} cty$$

模型参数

  • $ \beta_{0} $ 为截距
  • $ \beta_{1} $ 为斜率

 

基于训练数据的估计参数

$$\small hwy = 0.77 + 1.35(cty)$$

 

Mpg 数据与线性回归线

在 R 中使用 tidymodels 建模

模型公式

parsnip 中的模型公式

  • 用于指定列角色
    • 结果变量
    • 预测变量

通用形式

outcome ~ predictor_1 + predictor_2 + ...

速记法

outcome ~ .

cty 作为预测变量来预测 hwy

hwy ~ cty
在 R 中使用 tidymodels 建模

parsnip 包

R 中统一的模型规范语法

  1. 指定模型类型

    • 线性回归或其他模型类型
  2. 指定引擎

    • 不同引擎对应不同的 R 包
  3. 指定模式

    • 回归或分类

Parsnip 包

在 R 中使用 tidymodels 建模

拟合线性回归模型

 

使用 parsnip 定义模型规范

  • linear_reg()

 

lm_model 传给 fit() 函数

  • 指定模型公式
  • 拟合所用的 data

 

lm_model <- linear_reg() %>%

set_engine('lm') %>%
set_mode('regression')

 

lm_fit <- lm_model %>% 
  fit(hwy ~ cty, data = mpg_training)
在 R 中使用 tidymodels 建模

获取参数估计值

 

tidy() 函数

  • 接受训练后的 parsnip 模型对象
  • 生成模型摘要 tibble
  • termestimate 列给出参数估计

 

tidy(lm_fit)
# A tibble: 2 x 5
  term        estimate std.error statistic  p.value
  <chr>          <dbl>     <dbl>     <dbl>    <dbl>
1 (Intercept)    0.769    0.528       1.46 1.47e- 1
2 cty            1.35     0.0305     44.2  6.32e-97
在 R 中使用 tidymodels 建模

进行预测

将训练好的 parsnip 模型传给 predict()

  • new_data 指定要预测的数据集

 

predict() 的标准化输出

  1. 返回 tibble
  2. 行顺序与输入 new_data 保持一致
  3. 预测列命名为 .pred
hwy_predictions <- lm_fit %>% 
  predict(new_data = mpg_test)

hwy_predictions
# A tibble: 57 x 1
   .pred
   <dbl>
 1  25.0
 2  27.7
 3  25.0
 4  25.0
 5  22.3
# ... with 47 more rows
在 R 中使用 tidymodels 建模

将预测添加到测试数据

bind_cols() 函数

  • 按列方向合并两个或更多 tibble
  • 用于创建模型结果 tibble

步骤

  • mpg_test 选取 hwycty
  • 传给 bind_cols() 并添加预测列
mpg_test_results <- mpg_test %>%
  select(hwy, cty) %>%

bind_cols(hwy_predictions) mpg_test_results
# A tibble: 57 x 3
     hwy   cty .pred
   <int> <int> <dbl>
 1    29    18  25.0
 2    31    20  27.7
 3    27    18  25.0
 4    26    18  25.0
 5    25    16  22.3
# ... with 47 more rows
在 R 中使用 tidymodels 建模

开始建模!

在 R 中使用 tidymodels 建模

Preparing Video For Download...