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

如何在mlr3中绘制XGBoost模型的单棵决策树?

在mlr3中绘制XGBoost模型的单棵决策树

我想要在mlr3中绘制训练好的XGBoost模型的单棵决策树,但未找到相关示例。我知道原生xgboost库有实现方法,mlr3底层也调用该库,但无法从mlr3模型中提取树结构。

原生XGBoost绘制决策树示例

# XGBoost tree plotting example
data(iris) # Iris数据集
sppNum <- as.numeric(iris$Species)-1 # 将Species转为数值型(0-2),满足多分类预测要求

# 划分测试/训练集索引
set.seed(1)
testID <- as.vector(sapply(0:2,function(x) sort(sample(1:50,10,replace = FALSE)+(x*50)))) 
trainID <- which(!1:150 %in% testID)

library(xgboost)
library(DiagrammeR)
# 构建测试/训练数据集
datTrain <- xgb.DMatrix(data=as.matrix(iris[trainID,1:4]),label=sppNum[trainID]) 
datTest <- xgb.DMatrix(data=as.matrix(iris[testID,1:4]),label=sppNum[testID])
watchlist <- list(train = datTrain, eval = datTest)

# 参数列表
parList <- list(max_depth = 3, eta = 1, verbose = 0, nthread = 1,
                objective = "multi:softmax", eval_metric = "auc", num_class=3)
xgbMod <- xgb.train(parList, datTrain, nrounds = 100, watchlist) # 拟合模型

xgb.plot.tree(model = xgbMod, trees = 1) # 绘制单棵树

单棵决策树可视化

mlr3尝试绘制的报错示例

# mlr3 示例 ------------------------------------------------------------
library(mlr3)
library(mlr3learners)
library(mlr3viz)

data(iris)

tsk_mlr3 <- as_task_classif(iris,target='Species') # 创建分类任务
# 设置XGBoost学习器
lrn_mlr3 <- lrn('classif.xgboost',nrounds=100,max_depth = 3, eta = 1,
                eval_metric='auc') 
lrn_mlr3$train(tsk_mlr3,row_ids = trainID) # 在训练子集上训练模型
lrn_mlr3$predict(tsk_mlr3,row_ids = testID) # 对测试集进行预测

xgb.plot.tree(lrn_mlr3$model, trees=1) # 执行后报错:
# Error in xgb.plot.tree(lrn_mlr3$model, trees = 1) : 
#  model: Has to be an object of class xgb.Booster

解决方法

mlr3中classif.xgboost学习器训练后的model是一个封装后的列表,其中的booster元素才是原生的xgb.Booster对象,直接提取该元素即可调用xgb.plot.tree:

# 提取原生XGBoost模型并绘制单棵树
xgb.plot.tree(model = lrn_mlr3$model$booster, trees = 1)

原理说明

  • lrn_mlr3$model返回的是mlr3封装的模型容器,包含训练日志、参数等额外信息
  • 真正的原生XGBoost模型存储在$booster字段中,完全符合xgb.plot.tree要求的xgb.Booster类格式

内容的提问来源于stack exchange,提问作者S. Robinson

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 15:43:13