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

如何让tidymodels/workflow模型适配DALEX::explain()与modelStudio?

解决tidymodels Workflow模型适配DALEX的问题

我明白你遇到的困扰:直接用fit()训练的parsnip模型能正常在DALEX的explain()中工作,但从workflow提取的模型却报错无法预测,哪怕两者的类别显示一致。这背后的核心原因是:从workflow提取的model_fit对象绑定了工作流的预处理逻辑(比如配方的特征工程步骤),而DALEX默认的预测函数没有适配这种带预处理绑定的结构。下面给你两个针对性的解决方案:

方案1:提取底层原生xgboost模型

workflow的pull_workflow_fit()得到的是parsnip封装的model_fit对象,它内部包含了原生的xgboost模型(xgb.Booster类型)。你可以用extract_fit_engine()函数提取出这个底层模型,它和你直接训练的xgb对象行为完全一致:

# 从workflow中训练并提取底层xgboost模型
xgb1_engine <- extract_fit_engine(final_xgb %>% fit(data = train))

# 创建DALEX解释器,完全适配你的原有逻辑
explainer <- explain(xgb1_engine, 
                     data = test, 
                     y = test$Y, 
                     predict_function = predict.xgb.Booster, # 可选,DALEX通常能自动识别
                     label = "xgb_workflow")

这个方法的优势是直接拿到了和原生xgboost一致的模型,完全避开了workflow的预处理绑定问题,和你直接训练的xgb的使用逻辑完全对齐。

方案2:为workflow模型指定适配的预测函数

如果你想保留parsnip的model_fit结构(比如需要依赖workflow的自动预处理),可以给explain()指定一个自定义的预测函数。因为DALEX对于分类任务需要模型输出概率值,而parsnip的predict()默认可能输出类别标签,所以需要明确指定输出类型:

# 从workflow中提取model_fit对象
xgb1 <- final_xgb %>% fit(data = train) %>% pull_workflow_fit()

# 自定义预测函数(以二分类为例,提取正类概率)
custom_predict <- function(model, newdata) {
  # 输出概率矩阵,取第二列作为正类概率(根据你的类别顺序调整)
  predict(model, newdata = newdata, type = "prob")[[2]]
}

# 创建适配的DALEX解释器
explainer <- explain(xgb1, 
                     data = test, 
                     y = test$Y, 
                     predict_function = custom_predict, 
                     label = "xgb_workflow")

补充说明为什么直接训练的xgb能正常工作

你直接用boost_tree() %>% fit(Y ~ ., data = train)训练的模型,没有绑定任何额外的预处理配方,它的predict()可以直接接收和训练数据结构一致的test数据,行为和DALEX的预期完全匹配;而从workflow提取的xgb1因为绑定了工作流的配方,调用predict()时会自动对输入数据应用预处理步骤,如果你的test数据已经是预处理好的,就会导致数据格式不匹配,从而触发预测错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.08 21:17:48