How do you find feature names for Decision Tree Classification?

Viewed 389

I am trying to find the feature information for my decision trees. More specifically, I want to be able to tell what feature 183 is if it appears in my tree visualization. I have tried dtModel.getInputCol() but receive the following error.

AttributeError: 'DecisionTreeClassificationModel' object has no attribute 'getInputCol'

This is my current code:

from pyspark.ml.classification import DecisionTreeClassifier

# Create initial Decision Tree Model
dt = DecisionTreeClassifier(labelCol="label", featuresCol="features", maxDepth=3)

# Train model with Training Data
dtModel = dt.fit(trainingData)
display(dtModel)

If you can help or need more information, please let me know. Thank you.

1 Answers

See this example taken from Spark doc (I try to have the name consistent with your code, especially featuresCol="features").

I assume you have some code like this (before the code you posted in the question):

featureIndexer = VectorIndexer(inputCol="inputFeatures", outputCol="features", maxCategories=4).fit(data)

After this step, you have the "features" as indexed features, and then you feed to the DecisionTreeClassifier (like your posted code):

# Train a DecisionTree model.
dt = DecisionTreeClassifier(labelCol="indexedLabel", featuresCol="features")

What you're looking for is inputFeatures above, which is the original features before being indexed. If you want to print it, simply do something like:

sc.parallelize(inputFeatures, 1).saveAsTextFile("absolute_path") 
Related