A Simple Decision Tree Implementation

2016-10-11
osdodo

Entropy

I first came across this concept in high-school chemistry, where entropy is a measure of how disordered a system is. Information entropy was proposed by Claude Shannon, the founder of information theory, and it is in fact positively correlated with the entropy above. When we face a highly uncertain event, we need a lot of information to pin it down — that is, the amount of information is directly related to the uncertainty of the event. Let X be a discrete random variable with finitely many states; the relationship between entropy and probability is:

Decision tree

A classmate recently said the cafeteria food is terrible. I go there often and noticed it isn't bad every day — just occasionally bad enough to make you queasy. Since the dorm has only this one cafeteria and everyone is lazy, plenty of people still queue up for it. Let's use that as an example and build a quick predictor for tomorrow:

No. Weather Dish Headcount
1 Rainy Abundant Many
2 Rainy Poor Many
3 Rainy Average Many
4 Sunny Abundant Many
5 Sunny Poor Few
6 Sunny Average Few

Information gain: the degree to which knowing feature A reduces the uncertainty of dataset D.

Generate the decision tree via information gain: First compute Gain(Weather):

Similarly Gain(Dish) = 0.252 bit. Since the largest gain starts the tree, begin with Weather:

Weather?
/    \
Many   Dish?
      / | \
Many  Few Few

Code:

from math import log
from collections import defaultdict
import json

def createDataSet():
    features = ['Weather','Dish']
    dataSet =[['Rainy','Abundant','Many'],
              ['Rainy','Poor','Many'],
              ['Rainy','Average','Many'],
              ['Sunny','Abundant','Many'],
              ['Sunny','Poor','Few'],
              ['Sunny','Average','Few']]
    return dataSet,features

def _entropy(dataSet):
    '''
    Compute the entropy of a dataset
        :param dataSet: the dataset
    '''
    dic = defaultdict(lambda: 0)
    for line in dataSet:
        dic[line[-1]] += 1
    ent = 0.0
    n = float(len(dataSet))
    for v in dic.values():
        p = v / n
        ent = ent - p * log(p,2)
    return ent

def _splitDataSet(dataSet,index,value):
    '''
    Split a dataset
        :param dataSet: the dataset
        :param index: feature index
        :param value: feature value
    '''
    subDataSet = []
    for line in dataSet:
        if line[index] == value:
            subDataSet.append(line[:index] + line[index+1:])
    return subDataSet

def _gain(dataSet,index):
    '''
    Compute information gain
        :param dataSet: the dataset
        :param index: feature index
    '''
    n = float(len(dataSet))
    featureValueSet = set([line[index] for line in dataSet])
    subEnt = 0.0
    for value in featureValueSet:
        subDataSet = _splitDataSet(dataSet,index,value)
        p = len(subDataSet) / n
        subEnt = subEnt + p * _entropy(subDataSet)
    return _entropy(dataSet) - subEnt

def _bestFeatureIndex(dataSet,features):
    '''
    Find the best feature by maximum information gain
        :param dataSet: the dataset
        :param features: all features
    '''
    maxGain , bestFeatureIndex = 0.0 , 0
    for i, _ in enumerate(features):
        g = _gain(dataSet,i)
        if g > maxGain:
            maxGain = g
            bestFeatureIndex = i
    return bestFeatureIndex

def createTree(dataSet,features):
    '''
    Build the tree
        :param dataSet: the dataset
        :param features: all features
    '''
    result = [line[-1] for line in dataSet]
    if len(set(result)) == 1:
        return result[0]
    i = _bestFeatureIndex(dataSet,features)
    bestFeature = features[i]
    tree = {
        bestFeature: {}
    }
    del(features[i])
    bestFeatureValueSet = set([line[i] for line in dataSet])
    for value in bestFeatureValueSet:
        subFeature = features[:]
        tree[bestFeature][value] = \
            createTree(_splitDataSet(dataSet,i,value),subFeature)
    return tree

def testID3(tree,feat,testValue):
    '''
    Test
        :param tree: decision tree
        :param feat: features
        :param testValue: test value (e.g., ['Sunny','Average'])
    '''
    root = ''.join(tree.keys())
    nextDic = tree[root]
    featureIndex = 0
    for i,f in enumerate(feat):
        if f == root:
            featureIndex = i
    for key in nextDic.keys():
        if testValue[featureIndex] == key:
            if isinstance(nextDic[key], dict):
                return  testID3(nextDic[key],feat,testValue)
            else:
                return  nextDic[key]

dataSet,features = createDataSet()
feat = features[:]
tree = createTree(dataSet,features)
data = json.dumps(tree,ensure_ascii=False,indent=1)
print(data)
with open('data.json', 'w') as f:
    json.dump(data, f)
result = testID3(tree,feat,['Sunny','Poor'])
print(f"['Sunny','Poor']---->{result}")
← Back to blog