← 计算机

决策树ID3

计算机

文章封面

ID3

信息增益

$$l(x_i)=-log_2p(x_i)$$

信息期望

$$H=-\sum_{i=1}^{n}p(x_i)log_2p(x_i)$$

C4.5

CART

sklearn代码示例

from numpy import *
import operator

x = [[1,1,'yes'],
     [1,1,'yes'],
     [1,0,'no'],
     [0,1,'no'],
     [0,1,'no']
     ]


def chooseBestFeatureToSplit(dataset):
    numFeatures = len(dataset) - 1
    baseEntropy = calcShannonEnt(dataset)
    bestInfoGain = 0.0
    bestFeature = -1
    for i in range(numFeatures):
        featList = [example[i] for example in dataset]
        uniqueVals = set(featList)
        newEntropy = 0.0
        for value in uniqueVals:
            subDataSet = splitDataSet(dataset,i,value)
            prob = len(subDataSet)/float(len(dataset))
            newEntropy += prob * calcShannonEnt(subDataSet)
        infoGain = baseEntropy - newEntropy
        if infoGain > bestInfoGain:
            bestInfoGain = infoGain
            bestFeature = i
    return bestFeature

def majorityCnt(classList):
    classCount = {}
    for vote in classList:
        if vote not in classCount.keys():classCount[vote] = 0
        classCount[vote] += 1
    sortedClassCount = sorted(classCount.items(),
                              key=operator.itemgetter(1),
                              reverse=True
                              )
    return sortedClassCount[0][0]

def splitDataSet(dataset,axis,value):
    retDataSet = []
    for featVec in dataset:
        if featVec[axis] == value:
            reduceFeatVec = featVec[:axis]
            reduceFeatVec.extend(featVec[axis + 1 :])
            retDataSet.append(reduceFeatVec)
    return retDataSet

def calcShannonEnt(dataset):
    numEntries = len(dataset)
    labelCounts = {}
    for featVec in dataset:
        currentLabel = featVec[-1]
        if currentLabel not in labelCounts.keys():
            labelCounts[currentLabel] = 0
    shannon = 0.0
    for key in labelCounts:
        prob = float(labelCounts[key])/numEntries
        shannon -= prob * log(prob,2)
    return shannon

'''
dataset [[1,1,yes],[1,0,no]
lables [no surfacing,filppers]
'''
def creatTree(dataset,labels):
    classList = [example[-1] for example in dataset]
    if classList.count(classList[0]) == len(classList):
        return classList[0]

    if len(dataset[0]) == 1:
        return majorityCnt(classList)

    bestFeat = chooseBestFeatureToSplit(dataset)
    bestFeatLabel = labels[bestFeat]
    myTree = {bestFeatLabel:{}}
    featValues = [example[bestFeat] for example in dataset]
    uniqueVals = set(featValues)
    for value in uniqueVals:
        subLabels = labels[:]
        myTree[bestFeatLabel][value] = creatTree(splitDataSet(dataset,bestFeat,value),subLabels)

    return myTree

'''
{'no surfacing':{0:'no',1:{'flippers':{0:'no',1:'yes'}}}
'''

from sklearn import tree

x = [[1,1],[1,0],[0,1],[0,0]]
y = [1,0,0,0]
tr = tree.DecisionTreeClassifier()
tr.fit(x,y)
读到这里,感谢你的时间。继续阅读 →