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

如何在R中为caret构建的RF模型获取SHAP值并导出?

问题:Caret训练的随机森林回归模型的SHAP值计算、可视化与导出

训练数据

data = structure(list(Main_Street = structure(c(2L, 3L, 2L, 1L, 3L, 
2L, 3L, 1L, 2L, 2L), .Label = c("64", "70", "270"), class = "factor"), 
    Blocked_Lanes = c(3L, 4L, 2L, 1L, 1L, 2L, 6L, 3L, 3L, 3L), 
    Total_Vehicle_Count = c(1L, 2L, 2L, 2L, 1L, 4L, 3L, 2L, 2L, 
    1L), Tractor_Trailer_Count = c(0L, 0L, 0L, 0L, 0L, 0L, 0L, 
    0L, 0L, 0L), Weather_Winter_Storm = structure(c(1L, 2L, 1L, 
    1L, 1L, 1L, 1L, 1L, 1L, 1L), .Label = c("No", "Yes"), class = "factor"), 
    Weather_Rain = structure(c(2L, 1L, 1L, 1L, 1L, 1L, 1L, 1L, 
    2L, 1L), .Label = c("No", "Yes"), class = "factor"), Injuries_Count = c(0L, 
    0L, 0L, 0L, 0L, 0L, 1L, 0L, 0L, 1L), Accident_Overturned_Car = structure(c(1L, 
    1L, 1L, 1L, 1L, 1L, 1L, 1L, 1L, 2L), .Label = c("No", "Yes"
    ), class = "factor"), Fatalities_Count = c(0L, 0L, 0L, 0L, 
    0L, 0L, 0L, 0L, 0L, 0L), Speed = c(65L, 46L, 10L, 42L, 40L, 
    21L, 15L, 57L, 59L, 59L), Total_Volume = c(48.7, 22.5, 47.3, 
    102, 138, 75.3, 60.5, 83.3, 18, 26.7), Occupancy = c(3.5, 
    1.7, 40.8, 23.8, 14.1, 31, 27.1, 4.9, 2.6, 2.5), Lanes_Cleared_Duration = c(53L, 
    35L, 32L, 4L, 11L, 35L, 42L, 12L, 36L, 69L)), row.names = c(NA, 
-10L), class = "data.frame")

模型训练代码

fitControl <- trainControl(method = "repeatedcv", 
                           number = 10, 
                           repeats = 10) 

set.seed(2356)
randomforestGrid <- expand.grid(mtry = c(2:sqrt(61))) # 转为数据框格式
set.seed(2356)
rf_model <- train(Lanes_Cleared_Duration~.,
                 data = data, # 原代码中training未定义,替换为提供的data
                 method = "rf", 
                 trControl = fitControl, 
                 metric= "RMSE",
                 verbose = FALSE, 
                 tuneGrid = randomforestGrid,
                 n.trees = c(1:50)*100)

需求

  • 计算模型的SHAP值
  • 生成SHAP可视化图表(如汇总图、依赖图)
  • 导出包含各变量SHAP值的数据框

解决方案

1. 安装并加载所需包

install.packages(c("fastshap", "ggplot2", "randomForest"))
library(fastshap)
library(ggplot2)
library(caret)

2. 提取Caret模型底层的RandomForest对象

Caret的rf模型基于randomForest包实现,提取核心模型:

rf_core <- rf_model$finalModel

3. 计算SHAP值

使用fastshap包的explain函数,传入自定义预测函数适配回归场景:

# 定义预测函数
pred_fun <- function(object, newdata) {
  predict(object, newdata = newdata)
}

# 计算SHAP值(这里用全部数据作为解释对象,可替换为测试集)
shap_values <- explain(rf_core, 
                       X = data[, -which(names(data) == "Lanes_Cleared_Duration")], 
                       pred_wrapper = pred_fun,
                       nsim = 10) # nsim控制蒙特卡洛模拟次数,值越大越准确但速度越慢

4. 导出SHAP值数据框

合并原数据与SHAP值后导出为CSV:

# 合并原数据与SHAP值
shap_df <- cbind(data, shap_values)

# 导出文件
write.csv(shap_df, "shap_values.csv", row.names = FALSE)

5. SHAP可视化

5.1 全局变量重要性汇总图

autoplot(shap_values, type = "summary") +
  labs(title = "SHAP汇总图:变量对车道清理时长的影响",
       x = "SHAP值(对预测结果的影响)")

5.2 单个变量的SHAP依赖图

以Speed为例,查看变量本身与SHAP值的关系:

autoplot(shap_values, type = "dependence", 
         feature = "Speed", 
         X = data[, -which(names(data) == "Lanes_Cleared_Duration")]) +
  labs(title = "SHAP依赖图:Speed对车道清理时长的影响",
       x = "Speed",
       y = "SHAP值")

注意事项

  • 原代码中training对象未定义,已替换为提供的data;实际场景中需替换为对应训练/测试集。
  • 因子变量会被fastshap自动处理,确保模型训练时因子编码正确。
  • nsim参数可根据数据规模调整,小数据集设为20-50,大数据集设为5-10平衡速度与准确性。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 19:05:39