如何让tidymodels/workflow模型适配DALEX::explain()与modelStudio?
我明白你遇到的困扰:直接用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

