Databricks Notebook中LGBMClassifier create_tree_digraph无法显示树形图
解决LightGBM树形图在Databricks Notebook中不显示/保存的问题
1. 先搞定依赖库
LightGBM的create_tree_digraph功能依赖graphviz工具,得同时装系统级工具和Python库:
安装系统级graphviz(适配Databricks集群)
在Notebook里运行Shell命令,根据集群系统选对应的安装方式:
%sh # 针对Ubuntu/Debian系集群 sudo apt-get update && sudo apt-get install -y graphviz # 如果是CentOS/RHEL系集群,替换成下面的命令 # sudo yum install -y graphviz
安装Python版graphviz库
%pip install graphviz
2. 在Databricks里显示树形图
lgb.create_tree_digraph()返回的是graphviz.Digraph对象,直接打印只会显示内存地址,得用Databricks的display()函数渲染展示:
import lightgbm as lgb # 你的模型训练代码(保持原样即可) clf = lgb.LGBMClassifier() clf.fit(x_train, y_train, categorical_feature = x_train.select_dtypes(include = 'category').columns.tolist()) # 生成树形图对象 tree_digraph = lgb.create_tree_digraph(clf, orientation='vertical') # 直接在Notebook里显示图形 display(tree_digraph)
如果上面的方法没生效,可以先存成文件再读取显示:
# 把图形保存为PNG文件,默认存在/databricks/driver/路径下 tree_digraph.render("lightgbm_tree", format="png") # 读取图片并展示 image_df = spark.read.format("image").load("/databricks/driver/lightgbm_tree.png") display(image_df)
3. 把树形图保存为PNG文件
直接调用render()方法指定格式就行,文件会存在Databricks的driver节点默认路径(/databricks/driver/):
# 保存为PNG,生成的文件名是lightgbm_tree.png tree_digraph.render("lightgbm_tree", format="png") # 也可以用更直接的写法,只保存PNG不生成额外的.dot文件 tree_digraph.format = "png" tree_digraph.save("lightgbm_tree.png")
额外小提示
- 如果你用
lightgbm.train()构建模型,直接传训练后的booster对象就行:booster = lgb.train(params, train_set) tree_digraph = lgb.create_tree_digraph(booster, orientation='vertical') display(tree_digraph) - 可以用
num_trees参数指定生成某一棵特定树的图形,比如lgb.create_tree_digraph(clf, num_trees=0, orientation='vertical')(默认生成第一棵树)。
内容的提问来源于stack exchange,提问作者BigDawg007
相关产品推荐
相关产品推荐

