如何向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
相关产品推荐
相关产品推荐

