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

如何向tidymodels拟合函数传递额外groups变量?

解决方案

1. 修正模型参数存储逻辑

groups是数据变量而非模型超参数,不能和penalty/fusion一样用enquo()作为模型参数处理,需将其存入eng_args而非args:

joint_lasso <- function(
        mode = "regression", 
        penalty = NULL,
        fusion = NULL,
        groups = NULL) {

    if (mode  != "regression") {
        rlang::abort("`mode` should be 'regression'")
    }
    
    # 超参数存入args
    args <- list(
        penalty = rlang::enquo(penalty),
        fusion = rlang::enquo(fusion)
    )
    
    # 数据变量存入eng_args
    eng_args <- list(groups = rlang::enquo(groups))
    
    new_model_spec(
        "joint_lasso",
        args = args,
        eng_args = eng_args,
        mode = mode,
        method = NULL,
        engine = NULL
    )
}

2. 配置拟合函数的变量提取逻辑

由于使用matrix接口时,parsnip仅自动生成X(预测矩阵)和Y(结果向量),需手动从数据中提取groups并传递给fusedLassoProximal。修改set_fit配置,添加预处理逻辑:

set_fit(
    model = "joint_lasso",
    eng = "fusedLassoProximal",
    mode = "regression",
    value = list(
        interface = "matrix",
        protect = c("X", "Y"),
        func = c(pkg = "fuser", fun = "fusedLassoProximal"),
        defaults = list(),
        # 预处理:从配方模板中提取groups变量
        pre = function(model_spec, recipe, engine = NULL) {
            group_expr <- model_spec$eng_args$groups
            # 从配方模板中获取groups列的值
            model_spec$eng_args$groups <- rlang::eval_tidy(group_expr, recipe$template)
            model_spec
        }
    )
)

3. 调整模型调用与配方

  • 配方中无需给group设置特殊role,保持为普通列即可:
norm_recipe <- recipe(input) %>%
    update_role(matches("Feature"), new_role = "predictor") %>%
    update_role(outcome, new_role = "outcome") %>%
    # 移除update_role(group, new_role = "group")
    prep()
  • 模型调用时直接指定列名,无需用matches():
glmnet_model <- joint_lasso(penalty = tune(), fusion = tune(), groups = group) %>% 
    set_engine("fusedLassoProximal")

4. 完善预测阶段的分组处理

预测时需从新数据中提取groups,匹配系数矩阵的对应列计算预测值:

predict_joint_lasso <- function(object, new_data, group_col) {
    # 提取新数据中的分组信息
    group_vals <- new_data %>% dplyr::pull(!!group_col)
    # 匹配分组到系数矩阵的列(系数矩阵列对应训练集的唯一分组)
    coef_col_indices <- match(group_vals, colnames(object$beta))
    # 计算预测值
    preds <- as.matrix(new_data) %*% object$beta[, coef_col_indices, drop = FALSE]
    # 返回向量格式的预测结果
    tibble::tibble(.pred = as.vector(preds))
}

# 更新预测配置,传递分组列参数
pred_info <- list(
    pre = NULL,
    post = NULL,
    func = c(fun = "predict_joint_lasso"),
    args = list(
        object = quote(object$fit),
        new_data = quote(new_data),
        group_col = quote(object$spec$eng_args$groups)
    )
)

set_pred(
    model = "joint_lasso",
    eng = "fusedLassoProximal",
    mode = "regression",
    type = "numeric",
    value = pred_info
)

5. 重采样的分组保障

如果需要按分组分层采样(避免某组在子集缺失),用group_vfold_cv()替代vfold_cv():

folds <- group_vfold_cv(input, group = group, v = nfolds, repeats = nrepeats)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 05:57:03