tidymodel框架下SVM变量重要性图生成失败及API替代咨询
SVM变量重要性计算问题解决(结合tidymodels与vip)
问题背景
由于SVM本身不包含变量重要性信息,采用变量置换法结合tidymodels框架生成SVM的变量重要性图,原代码如下:
library(vip) library(MASS) library(tidymodels) data(Boston, package = "MASS") df <- Boston #Split the data into train and test set set.seed(7) splits <- initial_split(df) train <- training(splits) test <- testing(splits) #Preprocess with recipe rec <- recipe(medv~.,data=train) %>% step_normalize(all_predictors()) svm_spec <- svm_rbf(margin = 0.0937, cost = 26.7, rbf_sigma = 0.0208) %>% set_engine("kernlab") %>% set_mode("regression") #Putting into workflow svr_fit <- workflow() %>% add_recipe(rec) %>% add_model(svm_spec) %>% fit(data = train) svr_fit %>% pull_workflow_fit() %>% vip(method = "permute", nsim = 5, target = "medv", metric = "rmse", pred_wrapper = kernlab::predict, train = train)
遇到的问题
运行代码后出现两个问题:
- 报错信息:
Error in
metric_fun():
!estimateshould be a numeric vector, not a numeric matrix.
pull_workflow_fit()函数已被弃用
解决方案
1. 替代弃用的pull_workflow_fit()
tidymodels更新后,官方推荐使用extract_fit_parsnip()替代pull_workflow_fit(),用于从工作流中提取拟合完成的模型对象。
2. 解决预测结果格式不匹配问题
kernlab::predict()针对SVM回归模型返回的是矩阵格式,但vip的置换法要求预测结果为数值向量。因此需要自定义一个预测包装函数,将矩阵转换为向量:
kernlab_pred_wrapper <- function(object, newdata) { as.vector(kernlab::predict(object, newdata)) }
修正后的完整代码
library(vip) library(MASS) library(tidymodels) data(Boston, package = "MASS") df <- Boston # 划分训练集和测试集 set.seed(7) splits <- initial_split(df) train <- training(splits) test <- testing(splits) # 数据预处理配方 rec <- recipe(medv~.,data=train) %>% step_normalize(all_predictors()) # SVM回归模型定义 svm_spec <- svm_rbf(margin = 0.0937, cost = 26.7, rbf_sigma = 0.0208) %>% set_engine("kernlab") %>% set_mode("regression") # 构建工作流并拟合 svr_fit <- workflow() %>% add_recipe(rec) %>% add_model(svm_spec) %>% fit(data = train) # 自定义预测包装函数:将矩阵转换为向量 kernlab_pred_wrapper <- function(object, newdata) { as.vector(kernlab::predict(object, newdata)) } # 提取模型并计算变量重要性 svr_fit %>% extract_fit_parsnip() %>% # 替代弃用的pull_workflow_fit() vip(method = "permute", nsim = 5, target = "medv", metric = "rmse", pred_wrapper = kernlab_pred_wrapper, train = train)
内容的提问来源于stack exchange,提问作者UseR10085
相关产品推荐
相关产品推荐

