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

如何实现依赖训练数据额外特征的自定义评估指标?

解决依赖额外特征的自定义评估指标适配调参函数问题

要让依赖训练数据额外特征的自定义指标能和tune_grid/workflow_map配合,核心是把额外特征(比如分组变量)保留在建模流程的数据集里,确保调参函数传递给指标的data包含该列。以下是针对组内R²的具体实现方案:

关键改动点

  • 在Recipe中添加分组变量,标记为不参与建模的辅助列,保证它会被带到后续的预测数据中
  • 调整自定义指标函数,直接从传入的data中读取分组列,不再依赖外部参数传递

完整实现代码

# 1. 自定义组内R²指标(适配调参函数版本)
rsq_within_vec <- function(truth, estimate, group, na_rm = TRUE, ...) {
  rsq_within_impl <- function(truth, estimate, group) {
    d <- tibble(truth, estimate, group) %>%
      group_by(group) %>%
      mutate(truth = truth - mean(truth), 
             estimate = estimate - mean(estimate))
    
    if(sd(d$estimate) == 0) return(0)
    yardstick:::yardstick_cor(d$truth, d$estimate)^2
  }

  metric_vec_template(
    metric_impl = rsq_within_impl,
    truth = truth, 
    estimate = estimate,
    na_rm = na_rm,
    cls = "numeric",
    group = group,
    ...
  )
}

rsq_within <- function(data, ...) {
  UseMethod("rsq_within")
}

rsq_within <- new_numeric_metric(rsq_within, direction = "maximize")

# 修改此处:直接从data中引用group列,无需外部传入
rsq_within.data.frame <- function(data, truth, estimate, group, na_rm = TRUE, ...) {
  numeric_metric_summarizer(
    name = "rsq_within",
    fn = rsq_within_vec,
    data = data,
    truth = !! enquo(truth),
    estimate = !! enquo(estimate),
    fn_options = list(group = data[[rlang::as_name(enquo(group))]]),
    na_rm = na_rm,
    ...
  )
}

# 2. 定义针对gear分组的指标包装器
rsq_within_gear <- function(data, truth, estimate, na_rm = TRUE, ...) {
  rsq_within(
    data = data,
    truth = !!rlang::enquo(truth),
    estimate = !!rlang::enquo(estimate),
    group = gear,
    na_rm = na_rm,
    ...
  )
}

rsq_within_gear <- new_numeric_metric(rsq_within_gear, direction = "maximize")

# 3. 修改Recipe:保留gear列作为辅助列
set.seed(6735)
folds <- vfold_cv(mtcars, v = 5)

# 用update_role把gear设为"ID"角色(不参与建模但会被保留)
recipe <- recipes::recipe(mpg ~ cyl, data = mtcars) %>%
  update_role(gear, new_role = "ID")

model <- linear_reg() %>% set_engine("lm")
wf <- workflow() %>% add_recipe(recipe) %>% add_model(model)

# 现在可以正常运行调参
tune_grid(
    object    = wf,
    resamples = folds,
    grid      = 1,
    metrics   = metric_set(rmse, rsq, rsq_within_gear)
) %>%
collect_metrics()

原理说明

  • update_role(gear, new_role = "ID"):把gear标记为ID角色,recipes包会保留该列但不会将其作为预测变量或响应变量,这样在交叉验证的每个折里,预测数据会包含gear列
  • 指标函数修改后,直接从传入的data中提取group列,而tune_grid现在传递的data里已经包含gear,因此可以正常计算组内R²

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 16:55:17