(1)决策树的定义:
决策树(decision tree)是一个树结构(可以是二叉树或非二叉树)。其每个非叶节点表示一个特征属性上的测试,每个分支代表这个特征属性在某个值域上的输出,而每个叶节点存放一个类别。使用决策树进行决策的过程就是从根节点开始,测试待分类项中相应的特征属性,并按照其值选择输出分支,直到到达叶子节点,将叶子节点存放的类别作为决策结果。
构造决策树的关键步骤是分裂属性。所谓分裂属性就是在某个节点处按照某一特征属性的不同划分构造不同的分支,其目标是让各个分裂子集尽可能地“纯”。尽可能“纯”就是尽量让一个分裂子集中待分类项属于同一类别。分裂属性分为三种不同的情况:
1、属性是离散值且不要求生成二叉决策树。此时用属性的每一个划分作为一个分支。
2、属性是离散值且要求生成二叉决策树。此时使用属性划分的一个子集进行测试,按照“属于此子集”和“不属于此子集”分成两个分支。
3、属性是连续值。此时确定一个值作为分裂点split_point,按照>split_point和<=split_point生成两个分支。
(2)算法:
ID3算法就是在每次需要分裂时,计算每个属性的增益率,然后选择增益率最大的属性进行分裂。
设D为用类别对训练元组进行的划分,则D的熵(entropy)表示为:
其中pi表示第i个类别在整个训练元组中出现的概率,可以用属于此类别元素的数量除以训练元组元素总数量作为估计。熵的实际意义表示是D中元组的类标号所需要的平均信息量。
现在我们假设将训练元组D按属性A进行划分,则A对D划分的期望信息为:
而信息增益即为两者的差值:
(3)说明:
a、在决策树构造过程中可能会出现这种情况:所有属性都作为分裂属性用光了,但有的子集还不是纯净集,即集合内的元素不属于同一类别。在这种 情况下,由于没有更多信息可以使用了,一般对这些子集进行“多数表决”,即使用此子集中出现次数最多的类别作为此节点类别,然后将此节点作 为叶子节点。
b、在实际构造决策树时,通常要进行剪枝,这时为了处理由于数据中的噪声和离群点导致的过分拟合问题。剪枝有两种:
先剪枝——在构造过程中,当某个节点满足剪枝条件,则直接停止此分支的构造。
后剪枝——先构造完成完整的决策树,再通过某些条件遍历树进行剪枝。
(4)代码实现:
#递归构造决策树
import operator
from math import logdef createDataSet():#产生测试数据 dataSet = [[1, 1, 'yes'],[1, 1, 'yes'],[1, 0, 'no'],[0, 1, 'no'],[0, 1, 'no']]labels = ['no surfacing','flippers'] return dataSet, labelsdef calcShannonEnt(dataSet):#计算给定数据集的香农熵numEntries = len(dataSet)labelCounts = {}#统计每个类别出现的次数,保存在字典labelCounts中for featVec in dataSet: currentLabel = featVec[-1]if currentLabel not in labelCounts.keys(): labelCounts[currentLabel] = 0labelCounts[currentLabel] += 1 #如果当前键值不存在,则扩展字典并将当前键值加入字典shannonEnt = 0.0for key in labelCounts:#使用所有类标签的发生频率计算类别出现的概率prob = float(labelCounts[key])/numEntries#用这个概率计算香农熵shannonEnt -= prob * log(prob,2) #取2为底的对数return shannonEntdef splitDataSet(dataSet, axis, value):#按照给定特征划分数据集#dataSet:待划分的数据集#axis: 划分数据集的第axis个特征#value: 特征的返回值(比较值)retDataSet = []#遍历数据集中的每个元素,一旦发现符合要求的值,则将其添加到新创建的列表中for featVec in dataSet:if featVec[axis] == value:reducedFeatVec = featVec[:axis]reducedFeatVec.extend(featVec[axis+1:])retDataSet.append(reducedFeatVec)#extend()和append()方法功能相似,但在处理列表时,处理结果完全不同#a=[1,2,3] b=[4,5,6]#a.append(b) = [1,2,3,[4,5,6]]#a.extend(b) = [1,2,3,4,5,6]return retDataSetdef chooseBestFeatureToSplit(dataSet):#选择最好的数据集划分方式#输入:数据集#输出:最优分类的特征的index#计算特征数量numFeatures = len(dataSet[0]) - 1baseEntropy = calcShannonEnt(dataSet)bestInfoGain = 0.0; bestFeature = -1for i in range(numFeatures):#创建唯一的分类标签列表featList = [example[i] for example in dataSet]uniqueVals = set(featList)#计算每种划分方式的信息熵newEntropy = 0.0for value in uniqueVals:subDataSet = splitDataSet(dataSet, i, value)prob = len(subDataSet)/float(len(dataSet))newEntropy += prob * calcShannonEnt(subDataSet) infoGain = baseEntropy - newEntropy#计算最好的信息增益,即infoGain越大划分效果越好if (infoGain > bestInfoGain):bestInfoGain = infoGainbestFeature = ireturn bestFeaturedef majorityCnt(classList):#投票表决函数#输入classList:标签集合,本例为:['yes', 'yes', 'no', 'no', 'no']#输出:得票数最多的分类名称classCount={}for vote in classList:if vote not in classCount.keys(): classCount[vote] = 0classCount[vote] += 1 #若当前标签不存在,则扩展列表并加入当前标签#把分类结果进行排序,然后返回得票数最多的分类结果sortedClassCount = sorted(classCount.items(), key=operator.itemgetter(1), reverse=True)return sortedClassCount[0][0]def createTree(dataSet,labels):#创建树#输入:数据集和标签列表#输出:树的所有信息# classList为数据集的所有类标签classList = [example[-1] for example in dataSet]# 停止条件1:所有类标签完全相同,直接返回该类标签if classList.count(classList[0]) == len(classList): return classList[0]# 停止条件2:遍历完所有特征时仍不能将数据集划分成仅包含唯一类别的分组,则返回出现次数最多的类标签#if len(dataSet[0]) == 1:return majorityCnt(classList)# 选择最优分类特征bestFeat = chooseBestFeatureToSplit(dataSet)bestFeatLabel = labels[bestFeat]# myTree存储树的所有信息myTree = {bestFeatLabel:{}}# 以下得到列表包含的所有属性值del(labels[bestFeat])#使用列表推导来创建新列表,将所有可能的特征值写入新列表中,使用集合(set)数据类型#从列表创建集合可以得到列表中唯一元素值(不重复)featValues = [example[bestFeat] for example in dataSet]uniqueVals = set(featValues)# 遍历当前选择特征包含的所有属性值for value in uniqueVals:subLabels = labels[:]myTree[bestFeatLabel][value] = createTree(splitDataSet(dataSet, bestFeat, value),subLabels)return myTree## 存储决策树
def storeTree(inputTree,filename):'''使用pickle模块存储决策树'''import picklefw = open(filename,'wb+')pickle.dump(inputTree,fw)fw.close()def grabTree(filename):'''导入决策树模型'''import picklefr = open(filename,'rb')return pickle.load(fr)if __name__== "__main__": dataSet,labels = createDataSet()myTree=createTree(dataSet,labels)print(dataSet)print(myTree)#存取操作storeTree(myTree,'mt.txt')myTree2 = grabTree('mt.txt')print(myTree2)