使用tidymodels调用kernlab遇protect()栈溢出的解决方案问询
在tidymodels中使用kernlab后端解决「protect(): protection stack overflow」错误
问题背景
在tidymodels框架下以kernlab为后端建模时,遇到「protect(): protection stack overflow」错误。该错误由变量过多+公式表达式解析导致:单独使用kernlab时,通过x=和y=参数传入矩阵而非data.frame可规避,但tidymodels默认流程会将数据转为data.frame/data.table并使用公式接口,触发栈溢出。已尝试设置--max-ppsize=500000,无效。需保留tidymodels的标准化预处理及调参功能。
复现代码
library(tidymodels) x <- as.data.frame(matrix(rnorm(2000000), nrow = 100, ncol = 20000)) y <- data.frame("y" = rnorm(n = 100)) data <- cbind(y, x) train.cv <- vfold_cv(data,5) svm_model <- svm_rbf( cost = tune(), rbf_sigma = tune(), margin = tune(), engine = "kernlab", mode = "regression" ) svm_recep <- recipe(x = as.matrix(data)) %>% update_role(everything(), new_role = "predictor") %>% update_role(y, new_role = "outcome") svm_wflow <- workflow(preprocessor = svm_recep, spec = svm_model) svm_set <- extract_parameter_set_dials(svm_wflow) hypercube <- grid_latin_hypercube(svm_set, size = 5) svm_initial <- svm_wflow %>% tune_grid(resamples = train.cv, grid = hypercube, metrics = metric_set(rmse))
单独使用kernlab的错误与修复示例
错误代码(触发栈溢出)
library(kernlab) svm.train <- ksvm(y ~ ., type="eps-svr", data = data, kernel ="rbfdot")
修复代码(使用矩阵传入)
library(kernlab) svm.train <- ksvm(x = as.matrix(data[,-1]), y = as.matrix(data[,1]), type="eps-svr", kernel ="rbfdot")
解决方案:自定义parsnip矩阵接口
核心思路是让tidymodels直接将矩阵格式的x和y传给kernlab,绕过公式解析流程。具体步骤如下:
1. 自定义kernlab SVM的矩阵接口
复制默认svm_rbf的实现,修改为使用矩阵输入的接口:
# 定义自定义模型规格 svm_rbf_matrix <- function(mode = "regression", cost = NULL, rbf_sigma = NULL, margin = NULL) { parsnip::new_model_spec( "svm_rbf_matrix", args = list( cost = rlang::enquo(cost), rbf_sigma = rlang::enquo(rbf_sigma), margin = rlang::enquo(margin) ), eng_args = list(), mode = mode, method = NULL, engine = "kernlab" ) } # 设置fit和predict方法,指定矩阵接口 set_fit_engine("svm_rbf_matrix", engine = "kernlab") <- function(object, engine = NULL) { object$method <- parsnip::new_method( fit = parsnip::fit_spec( interface = "matrix", protect = c("x", "y", "type", "kernel"), func = c(pkg = "kernlab", fun = "ksvm"), defaults = list( type = "eps-svr", kernel = "rbfdot" ) ), predict = parsnip::predict_spec( interface = "matrix", func = c(pkg = "kernlab", fun = "predict"), defaults = list(type = "response") ) ) object } # 映射parsnip参数到kernlab原生参数 set_model_arg("svm_rbf_matrix", eng = "kernlab", parsnip = "cost", original = "C") <- function(value) value set_model_arg("svm_rbf_matrix", eng = "kernlab", parsnip = "rbf_sigma", original = "sigma") <- function(value) value set_model_arg("svm_rbf_matrix", eng = "kernlab", parsnip = "margin", original = "epsilon") <- function(value) value
2. 修改workflow的预处理与模型配置
确保预处理后输出矩阵格式,并使用自定义模型:
# 定义包含标准化步骤的recipe(保留你的预处理逻辑) svm_recep <- recipe(y ~ ., data = data) %>% update_role(y, new_role = "outcome") %>% update_role(-y, new_role = "predictor") %>% step_normalize(all_predictors()) # 示例标准化步骤,可替换为你的预处理 # 使用自定义模型 svm_model <- svm_rbf_matrix( cost = tune(), rbf_sigma = tune(), margin = tune(), engine = "kernlab", mode = "regression" ) # 设置蓝图,强制预处理输出矩阵格式 mat_blueprint <- hardhat::default_recipe_blueprint( composition = "matrix", allow_nominal = FALSE ) # 构建workflow svm_wflow <- workflow() %>% add_recipe(svm_recep, blueprint = mat_blueprint) %>% add_model(svm_model) # 执行调参(流程与原代码一致) svm_set <- extract_parameter_set_dials(svm_wflow) hypercube <- grid_latin_hypercube(svm_set, size = 5) svm_initial <- svm_wflow %>% tune_grid(resamples = train.cv, grid = hypercube, metrics = metric_set(rmse))
方案原理
自定义的svm_rbf_matrix模型使用parsnip的matrix接口,直接将预处理后的预测变量矩阵和结果向量传入ksvm的x和y参数,彻底避免了公式解析大量变量时的栈溢出问题,同时完整保留tidymodels的预处理、交叉验证和调参功能。
环境信息
> library(tidymodels) ── Attaching packages ───────────────────────────── tidymodels 1.2.0 ── ✔ broom 1.0.5 ✔ recipes 1.0.10 ✔ dials 1.2.1 ✔ rsample 1.2.1 ✔ dplyr 1.1.4 ✔ tibble 3.2.1 ✔ ggplot2 3.5.0 ✔ tidyr 1.3.1 ✔ infer 1.0.7 ✔ tune 1.2.1 ✔ modeldata 1.3.0 ✔ workflows 1.1.4 ✔ parsnip 1.2.1 ✔ workflowsets 1.1.0 ✔ purrr 1.0.2 ✔ yardstick 1.3.1 > sessionInfo() R version 4.3.3 (2024-02-29 ucrt) Platform: x86_64-w64-mingw32/x64 (64-bit) Running under: Windows 10 x64 (build 19045) Matrix products: default locale: [1] LC_COLLATE=German_Germany.utf8 LC_CTYPE=German_Germany.utf8 [3] LC_MONETARY=German_Germany.utf8 LC_NUMERIC=C [5] LC_TIME=German_Germany.utf8 time zone: Europe/Berlin tzcode source: internal attached base packages: [1] stats graphics grDevices utils datasets methods [7] base other attached packages: [1] kernlab_0.9-32 yardstick_1.3.1 workflowsets_1.1.0 [4] workflows_1.1.4 tune_1.2.1 tidyr_1.3.1 [7] tibble_3.2.1 rsample_1.2.1 recipes_1.0.10 [10] purrr_1.0.2 parsnip_1.2.1 modeldata_1.3.0 [13] infer_1.0.7 ggplot2_3.5.0 dplyr_1.1.4 [16] dials_1.2.1 scales_1.3.0 broom_1.0.5 [19] tidymodels_1.2.0 loaded via a namespace (and not attached): [1] tidyselect_1.2.1 timeDate_4022.108 blob_1.2.3 [4] fastmap_1.1.0 digest_0.6.35 rpart_4.1.23 [7] timechange_0.3.0 lifecycle_1.0.3 ellipsis_0.3.2 [10] survival_3.5-8 RSQLite_2.2.20 magrittr_2.0.3 [13] compiler_4.3.3 rlang_1.1.3 tools_4.3.3 [16] utf8_1.2.2 data.table_1.14.2 bit_4.0.4 [19] DiceDesign_1.9 withr_2.5.0 nnet_7.3-19 [22] grid_4.3.3 fansi_1.0.3 colorspace_2.0-3 [25] future_1.33.2 globals_0.16.2 iterators_1.0.14 [28] MASS_7.3-60.0.1 cli_3.6.2 generics_0.1.2 [31] rstudioapi_0.16.0 future.apply_1.10.0 DBI_1.2.2 [34] cachem_1.0.6 stringr_1.5.0 splines_4.3.3 [37] parallel_4.3.3 vctrs_0.6.5 hardhat_1.3.1 [40] Matrix_1.5-1 bit64_4.0.5 listenv_0.9.0 [43] foreach_1.5.2 gower_1.0.1 glue_1.6.2 [46] parallelly_1.34.0 codetools_0.2-19 lubridate_1.9.3 [49] stringi_1.7.6 gtable_0.3.0 munsell_0.5.0 [52] GPfit_1.0-8 pillar_1.9.0 furrr_0.3.1 [55] ipred_0.9-14 lava_1.7.2.1 R6_2.5.1 [58] lhs_1.1.6 lattice_0.22-5 backports_1.4.1 [61] memoise_2.0.1 class_7.3-22 Rcpp_1.0.12 [64] prodlim_2023.03.31 pkgconfig_2.0.3
内容的提问来源于stack exchange,提问作者R.Pickman
相关产品推荐
相关产品推荐

