You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

复现tidymodels Lasso回归结果不一致问题求助

复现Julia Silge《使用tidymodels进行Lasso回归》教程时结果不一致的问题与解决

问题描述

尝试复现Julia Silge的Lasso回归教程,数据清洗结果与教程一致,但相同回归代码得到完全不同的系数结果,导致预测阶段出错。可复现代码如下:

# Load libraries
library(tidymodels)
library(tidyverse)
library(schrute)
library(janitor)
tidymodels_prefer()

# Load data
ratings_raw <- read_csv("https://raw.githubusercontent.com/rfordatascience/tidytuesday/master/data/2020/2020-03-17/office_ratings.csv")
#> Rows: 188 Columns: 6
#> ── Column specification ────────────────────────────────────────────────────────
#> Delimiter: ","
#> chr  (1): title
#> dbl  (4): season, episode, imdb_rating, total_votes
#> date (1): air_date
#> 
#> ℹ Use `spec()` to retrieve the full column specification for this data.
#> ℹ Specify the column types or set `show_col_types = FALSE` to quiet this message.
schrute::theoffice
#> # A tibble: 55,130 × 12
#>    index season episode episode_name director   writer           character text 
#>    <int>  <int>   <int> <chr>        <chr>      <chr>            <chr>     <chr>
#>  1     1      1       1 Pilot        Ken Kwapis Ricky Gervais;S… Michael   All …
#>  2     2      1       1 Pilot        Ken Kwapis Ricky Gervais;S… Jim       Oh, …
#>  3     3      1       1 Pilot        Ken Kwapis Ricky Gervais;S… Michael   So y…
#>  4     4      1       1 Pilot        Ken Kwapis Ricky Gervais;S… Jim       Actu…
#>  5     5      1       1 Pilot        Ken Kwapis Ricky Gervais;S… Michael   All …
#>  6     6      1       1 Pilot        Ken Kwapis Ricky Gervais;S… Michael   Yes,…
#>  7     7      1       1 Pilot        Ken Kwapis Ricky Gervais;S… Michael   I've…
#>  8     8      1       1 Pilot        Ken Kwapis Ricky Gervais;S… Pam       Well…
#>  9     9      1       1 Pilot        Ken Kwapis Ricky Gervais;S… Michael   If y…
#> 10    10      1       1 Pilot        Ken Kwapis Ricky Gervais;S… Pam       What?
#> # ℹ 55,120 more rows
#> # ℹ 4 more variables: text_w_direction <chr>, imdb_rating <dbl>,
#> #   total_votes <int>, air_date <chr>

# We wish to join up the two data sets. We need to change the way
# season and episodes are counted. So we need to do some work to join them up
remove_regex <- "[:punct:]|[:digit:]|parts |part |the |and"
office_ratings <- ratings_raw |> 
    transmute(episode_name = str_to_lower(title),
              episode_name = str_remove_all(episode_name, remove_regex),
              episode_name = str_trim(episode_name),
              imdb_rating)

office_info <- schrute::theoffice |> 
    mutate(season = as.numeric(season),
           episode = as.numeric(episode),
           episode_name = str_to_lower(episode_name),
           episode_name = str_remove_all(episode_name, remove_regex),
           episode_name = str_trim(episode_name)) |> 
    select(season, episode, episode_name, director, writer, character)

# Let us count how many line every character has per episode
characters <- office_info |> 
    count(episode_name, character) |> 
    add_count(character, wt = n, name = "character_count") |> 
    filter(character_count > 800) |> 
    select(-character_count) |> 
    pivot_wider(names_from = character, 
                values_from = n,
                values_fill = list(n = 0))

# We want to do the same for writer and directors
creators <- office_info |> 
    distinct(episode_name, director, writer) |> 
    pivot_longer(director:writer, names_to = "role", values_to = "person") |> 
    separate_rows(person, sep = ";") |> 
    add_count(person) |> 
    filter(n > 10) |> 
    distinct(episode_name, person) |> 
    mutate(person_value = 1) |> 
    pivot_wider(names_from = person,
                values_from = person_value,
                values_fill = list(person_value = 0))

office <- office_info |> 
    distinct(season, episode, episode_name) |> 
    inner_join(characters) |> 
    inner_join(creators) |> 
    inner_join(office_ratings) |> 
    janitor::clean_names()
#> Joining with `by = join_by(episode_name)`
#> Joining with `by = join_by(episode_name)`
#> Joining with `by = join_by(episode_name)`
#> Warning in inner_join(inner_join(inner_join(distinct(office_info, season, : Detected an unexpected many-to-many relationship between `x` and `y`.
#> ℹ Row 71 of `x` matches multiple rows in `y`.
#> ℹ Row 79 of `y` matches multiple rows in `x`.
#> ℹ If a many-to-many relationship is expected, set `relationship =
#>   "many-to-many"` to silence this warning.


# Training the model
# We are going to use a lasso regression model
set.seed(1234)
office_split <- initial_split(office, strata = season)
office_train <- training(office_split)
office_test <- testing(office_split)

# Let us use a recipe
office_rec <- recipe(imdb_rating ~ ., data = office_train) |> 
    update_role(episode_name, new_role = "ID") |> # It is no longer a predictor
    step_zv(all_numeric(), - all_outcomes()) |> 
    step_normalize(all_numeric(), - all_outcomes())

office_prep <- office_rec |> 
    prep(strings_as_factors = FALSE)

# We can now train the model
lasso_spec <- linear_reg(penalty = 0.1, mixture = 1) |> 
    set_engine("glmnet")

wf <- workflow() |> 
    add_recipe(office_rec)

lasso_fit <- wf |> 
    add_model(lasso_spec) |> 
    fit(data = office_train)

lasso_fit %>%
    pull_workflow_fit() %>%
    tidy()
#> Warning: `pull_workflow_fit()` was deprecated in workflows 0.2.3.
#> ℹ Please use `extract_fit_parsnip()` instead.
#> This warning is displayed once every 8 hours.
#> Call `lifecycle::last_lifecycle_warnings()` to see where this warning was
#> generated.
#> Loading required package: Matrix
#> 
#> Attaching package: 'Matrix'
#> 
#> The following objects are masked from 'package:tidyr':
#> 
#>     expand, pack, unpack
#> 
#> Loaded glmnet 4.1-8
#> # A tibble: 31 × 3
#>    term        estimate penalty
#>    <chr>          <dbl>   <dbl>
#>  1 (Intercept)  8.37        0.1
#>  2 season       0           0.1
#>  3 episode      0           0.1
#>  4 andy         0           0.1
#>  5 angela       0.00234     0.1
#>  6 darryl       0           0.1
#>  7 dwight       0           0.1
#>  8 jim          0.00150     0.1
#>  9 kelly        0           0.1
#> 10 kevin        0           0.1
#> # ℹ 21 more rows

教程中得到的预期结果:
教程中的Lasso回归系数结果

问题原因

  1. 固定penalty值的错误:你直接指定了penalty = 0.1,但教程中是通过交叉验证选择了最优的penalty值,Lasso回归的系数对penalty值极度敏感,不同的penalty会导致大量系数被压缩为0或出现完全不同的非零系数。
  2. 数据连接的多对多问题:代码中inner_join触发了多对多匹配警告,这会导致数据集中出现重复行,改变了模型训练的输入数据基础,进而影响结果。
  3. 包版本差异:tidymodels、glmnet等包的版本更新可能改变了模型拟合的默认行为,比如标准化方式、参数处理逻辑等。

解决方法

1. 通过交叉验证选择最优penalty值

这是核心修正步骤,教程中使用交叉验证确定最佳penalty而非固定值,修改模型训练部分代码如下:

# 创建带调参的Lasso模型规格
lasso_spec <- linear_reg(penalty = tune(), mixture = 1) |> 
    set_engine("glmnet")

# 创建交叉验证折
set.seed(1234)
office_folds <- vfold_cv(office_train, strata = season)

# 定义调参网格
lambda_grid <- grid_regular(penalty(), levels = 50)

# 训练并调参
lasso_res <- wf |> 
    add_model(lasso_spec) |> 
    tune_grid(
        resamples = office_folds,
        grid = lambda_grid,
        metrics = metric_set(rmse, rsq)
    )

# 选择最优penalty
best_lasso <- lasso_res |> select_best("rmse")

# 用最优penalty拟合最终模型
final_lasso <- wf |> 
    finalize_workflow(best_lasso) |> 
    fit(data = office_train)

# 查看系数
final_lasso |> 
    extract_fit_parsnip() |> 
    tidy()

2. 修复数据连接的多对多问题

在数据清洗阶段确保episode_name唯一,避免重复匹配:

# 在创建office_ratings时,确保每个episode_name唯一
office_ratings <- ratings_raw |> 
    transmute(episode_name = str_to_lower(title),
              episode_name = str_remove_all(episode_name, remove_regex),
              episode_name = str_trim(episode_name),
              imdb_rating) |>
    distinct(episode_name, .keep_all = TRUE)

# 同样在office_info去重时确保episode_name唯一
office_info <- schrute::theoffice |> 
    mutate(season = as.numeric(season),
           episode = as.numeric(episode),
           episode_name = str_to_lower(episode_name),
           episode_name = str_remove_all(episode_name, remove_regex),
           episode_name = str_trim(episode_name)) |> 
    select(season, episode, episode_name, director, writer, character) |>
    distinct(season, episode, episode_name, .keep_all = TRUE)

3. 对齐包版本

安装教程发布时的对应包版本(教程发布于2020年左右),避免版本差异导致的行为变化。

内容的提问来源于stack exchange,提问作者Logos

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.05 05:42:04