如何将partykit的ctree对象转为rpart以使用rpart.plot绘图?
问题描述
我在使用partykit的ctree()训练多分类模型时,因为树的深度较大(设置了max_depth=5),默认的绘图方法展示效果很差。我很喜欢rpart.plot包的输出风格,它对深树的展示更清晰。但当我尝试直接将ctree模型传入rpart.plot()时,却收到了错误:
Error: the object passed to prp is not an rpart object
我之前用rpart模型的绘图代码是正常的:
library(partykit) library(rpart) library(rpart.plot) df_test <- cu.summary[complete.cases(cu.summary),] multi.class.model <- rpart(Reliability~., data = df_test) rpart.plot(multi.class.model)
现在想用ctree模型实现同样的输出,请问有没有办法把ctree对象转换成rpart类型?
解决方案
首先得明确:没有官方的直接转换工具,因为ctree(条件推断树)和rpart(递归分区树)的模型底层逻辑、结构存储方式差异很大——ctree基于显著性检验选择分裂节点,而rpart依赖成本复杂度剪枝,两者的节点规则、统计量存储格式完全不同,强行转换很容易出错,甚至得到错误的可视化结果。
不过你有几个替代方案可以达到类似的可视化效果:
1. 优化partykit自带的绘图参数
partykit的默认绘图可以通过调整参数来适配深树的显示:
- 缩小字体大小,避免节点内容拥挤:
plot(multi.class.model, gp = gpar(fontsize = 8)) # 按需调整fontsize数值 - 切换显示模式,比如用
type="simple"简化节点内容,或者type="extended"展示更详细的统计信息:plot(multi.class.model, type = "simple", gp = gpar(fontsize = 7)) - 如果需要查看节点的详细信息,可以搭配
nodes()函数:nodes(multi.class.model, 1) # 查看根节点详情,替换数字查看其他节点
2. 使用ggparty包(基于ggplot2的可视化工具)
ggparty是专门为partykit模型设计的可视化包,基于ggplot2,布局更灵活,对深树的支持更好,还能自定义样式:
library(ggparty) # 基础绘图 ggparty(multi.class.model) + geom_edge() + geom_node_splitvar() + geom_node_label(aes(label = splitvar), size = 3) + geom_node_plot( gglist = list(geom_bar(aes(x = Reliability), stat = "count", fill = "lightblue")), shared_axis = TRUE, width = 0.6 )
你可以根据需求调整字体大小、节点宽度、颜色等参数,完全适配深树的展示。
3. 折中方案:用rpart重新训练模型(谨慎使用)
如果只是单纯喜欢rpart.plot的风格,且可以接受模型结果的差异(因为ctree和rpart的分裂逻辑不同),可以直接用rpart()重新训练一个模型,再用rpart.plot()绘图:
multi.class.rpart <- rpart(Reliability~., data = df_test, control = rpart.control(maxdepth = 5)) rpart.plot(multi.class.rpart)
但要注意,这个模型和你之前的ctree模型不是同一个,预测结果可能会有差异。
内容的提问来源于stack exchange,提问作者Hanjo Odendaal

