How to visualize decision tree model/object in pyspark?

Viewed 3080

Is there any way to visualize/plot decision tree created using either mllib or ml library in pyspark. Also how to get information like number of records in leaf nodes. Thanks

2 Answers

You can get the number of statistics of all the leaf nodes, like impurity, gain, gini, Array of element classified into each label by the model data file.

The data file is located where you save the model/ data/

model.save(location)
modeldf = spark.read.parquet(location+"data/*")

This file contains much of the needed meta data for the decision tree or even randomForest. You can extract all the needed information like.

noderows = modeldf.select("id","prediction","leftChild","rightChild","split").collect()
df = pd.Dataframe([[rw['id'],rw['gain],rw['impurity'],rw['gini']] for rw in noderows if rw['leftChild'] < 0 and rw['rightChild'] < 0])
df.show()
Related