调用shap.prep()遇维度错误,请求排查LightGBM多分类SHAP分析问题
问题原因与解决方案
错误根源
- 多分类模型SHAP值需指定类别:
SHAPforxgboost的shap.prep函数处理多分类LightGBM模型时,必须通过class_id参数指定要计算SHAP值的类别,否则函数无法匹配输出维度,引发维度不匹配错误。 - 标签类型不符合LightGBM要求:LightGBM多分类任务要求标签为整数型,原代码中将标签转为字符型(
as.character(y)),会导致模型内部输出结构与SHAP工具的预期不兼容。
修正后的代码
library(dplyr) library(ggplot2) library(SHAPforxgboost) library(lightgbm) set.seed(111) x1 <- rnorm(2000) x2 <- rnorm(2000) y <- rnorm(2000) df <- data.frame(x1,x2,y) df <- df |> mutate(y = abs(y), y = round(y, digits = 0), y = ifelse(y >= 2, 2, y), y = as.integer(y)) # 修正:将标签转为整数型 # Define response and features y_col <- "y" x_cols <- c("x1","x2") # random split set.seed(83454) ix <- sample(nrow(df), 0.8 * nrow(df)) dtrain <- lgb.Dataset(data.matrix(df[ix, x_cols]), label = df[ix, y_col]) dvalid <- lgb.Dataset(data.matrix(df[-ix, x_cols]), label = df[-ix, y_col]) params <- list( objective = "multiclass", metric = "multi_error", learning_rate = 0.05, num_leaves = 15, num_class = 3 ) fit_lgb <- lgb.train(params, dtrain, nrounds = 89L, valids = list(valid = dvalid), early_stopping_rounds = 20L ) # 修正:指定class_id(可选0/1/2,对应三个类别),并使用训练集特征矩阵 shap <- shap.prep(fit_lgb, X_train = data.matrix(df[ix, x_cols]), class_id = 0) # 计算类别0的SHAP值,可根据需求修改为1或2
额外说明
- 如果需要查看所有类别的SHAP值,可循环调用
shap.prep分别指定class_id = 0、class_id = 1、class_id = 2,分别获取每个类别的SHAP结果。 X_train建议使用训练集的特征矩阵,而非整个数据集,这样SHAP分析的是模型训练时用到的数据分布,结果更准确。
内容的提问来源于stack exchange,提问作者andre de vera
相关产品推荐
相关产品推荐

