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

如何从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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 04:41:03