复现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
教程中得到的预期结果:
问题原因
- 固定penalty值的错误:你直接指定了
penalty = 0.1,但教程中是通过交叉验证选择了最优的penalty值,Lasso回归的系数对penalty值极度敏感,不同的penalty会导致大量系数被压缩为0或出现完全不同的非零系数。 - 数据连接的多对多问题:代码中
inner_join触发了多对多匹配警告,这会导致数据集中出现重复行,改变了模型训练的输入数据基础,进而影响结果。 - 包版本差异: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
相关产品推荐
相关产品推荐

