如何从fit_resamples与超参数调优中提取训练误差?
在tidymodels中提取交叉验证与超参数调优的训练误差
一、交叉验证(fit_resamples)获取训练误差
默认情况下fit_resamples仅计算验证集指标,要获取训练误差,需在control_resamples中设置extract函数,直接在每个fold的训练集上计算指标:
步骤1:定义包含提取逻辑的控制参数
library(tidymodels) library(dplyr) # 定义控制参数,添加训练集指标提取逻辑 control <- control_resamples( extract = function(x) { # 从workflow结果中获取训练集的特征和目标变量 train_predictors <- x$pre$mold$predictors train_outcomes <- x$pre$mold$outcomes # 生成训练集预测结果 train_preds <- predict(x$fit, train_predictors) # 计算RMSE(可替换为其他自定义指标) rmse(train_outcomes, train_preds) } )
步骤2:运行交叉验证并提取训练误差
# 使用新控制参数执行交叉验证 lr_cv <- lr_wf |> fit_resamples( folds, metrics = metric_set(rmse), control = control ) # 提取每个fold的训练误差,并计算平均值与标准误 train_error <- lr_cv$.extracts %>% unnest(cols = .extracts) %>% select(id, .estimate) %>% rename(train_rmse = .estimate) train_error_summary <- train_error %>% summarise( mean_train_rmse = mean(train_rmse), std_err_train_rmse = sd(train_rmse)/sqrt(n()) ) # 查看训练误差统计结果 train_error_summary
二、超参数调优(tune_grid)获取训练误差
对于超参数调优场景,逻辑与交叉验证一致,通过control_grid设置extract函数,提取每组超参数在每个fold的训练误差:
步骤1:定义控制参数
control <- control_grid( extract = function(x) { train_predictors <- x$pre$mold$predictors train_outcomes <- x$pre$mold$outcomes train_preds <- predict(x$fit, train_predictors) rmse(train_outcomes, train_preds) } )
步骤2:运行调优并提取训练误差
# 使用新控制参数执行超参数调优 tree_res <- tree_wf %>% tune_grid( resamples = folds, grid = tree_grid, metrics = metric_set(rmse), control = control ) # 提取每组超参数的训练误差(按fold计算后取平均) train_grid_error <- tree_res$.extracts %>% unnest(cols = .extracts) %>% select(id, .config, .estimate) %>% rename(train_rmse = .estimate) %>% group_by(.config) %>% summarise( mean_train_rmse = mean(train_rmse), std_err_train_rmse = sd(train_rmse)/sqrt(n()) ) # 合并验证集指标,对比训练/验证误差 valid_grid_error <- collect_metrics(tree_res) %>% select(.config, mean, std_err) %>% rename(mean_valid_rmse = mean, std_err_valid_rmse = std_err) combined_error <- train_grid_error %>% left_join(valid_grid_error, by = ".config") # 查看合并后的误差结果 combined_error
可选:绘制训练/验证误差曲线
library(ggplot2) # 结合超参数数据,绘制学习曲线(以cost_complexity为例) combined_error %>% left_join(tree_grid, by = ".config") %>% ggplot(aes(x = cost_complexity)) + geom_line(aes(y = mean_train_rmse, color = "训练集")) + geom_line(aes(y = mean_valid_rmse, color = "验证集")) + geom_ribbon(aes(ymin = mean_train_rmse - std_err_train_rmse, ymax = mean_train_rmse + std_err_train_rmse), alpha = 0.2, fill = "#1f77b4") + geom_ribbon(aes(ymin = mean_valid_rmse - std_err_valid_rmse, ymax = mean_valid_rmse + std_err_valid_rmse), alpha = 0.2, fill = "#ff7f0e") + labs( x = "复杂度成本", y = "RMSE", color = "数据集", title = "训练集与验证集误差对比" ) + theme_minimal()
内容的提问来源于stack exchange,提问作者Mael Fosso
相关产品推荐
相关产品推荐

