決策樹歸納(ID3屬性選擇度量)Java實現

來源:互聯網
上載者:User

一般的決策樹歸納架構見之前的博文:http://blog.csdn.net/zhyoulun/article/details/41978381


ID3屬性選擇度量原理

ID3使用資訊增益作為屬性選擇度量。該度量基於香農在研究訊息的值或”資訊內容“的資訊理論方面的先驅工作。該結點N代表或存放分區D的元組。選擇具有最高資訊增益的屬性作為結點N的分裂屬性。該屬性使結果分區中對元祖分類所需要的資訊量最小,並反映這些分區中的最小隨機性或”不純性“。這種方法使得對一個對象分類所需要的期望測試數目最小,並確保找到一顆簡單的(但不必是最簡單的)樹。

對D中的元組分類所需要的期望資訊由下式給出,


其中pi是D忠任意元組屬於類Ci的非零機率。使用以2為底的對數函數是因為資訊用二進位編碼。Info(D)是識別D中元組的類標號所需要的平均資訊量。注意,此時我們所有的資訊只是每個類的元組所佔的百分比。

現在假設我們要按照某屬性A劃分D中的元組,其中屬性A根據訓練資料的觀測具有v個不同的值{a1,a2,...av}。可以用屬性A將D劃分為v個分區或子集{D1,D2,...,Dv},其中Dj包含D中的元組,它們的A值為aj。這些分區對應於從節點N生長出來的分支。理想情況下,我們希望該劃分產生元組的準確分類。即希望每個分區都是純的(實際情況多半是不純的,如分區可能包含來自不同類的元組)。在此劃分之後,為了得到準確的分類,我們還需要多少資訊。這個量由下式度量:


其中|Dj|/|D|充當第j個分區的權重。Info_A(D)是基於按A劃分對D的元組分類所需要的期望值資訊需要的期望資訊越小,分區的純度越高

資訊增益定義為原來的資訊需求(僅基於類比例)與新的資訊需求(對A劃分後)之前的差。即


換言之,Gain(A)告訴我們通過A上的劃分我們得到了多少。它是知道A的值而導致的資訊需求的期望減少。選擇具有最高資訊增益Gain(A)的屬性A作為結點N的分裂屬性。


以下為例子。


資料

data.txt

youth,high,no,fair,noyouth,high,no,excellent,nomiddle_aged,high,no,fair,yessenior,medium,no,fair,yessenior,low,yes,fair,yessenior,low,yes,excellent,nomiddle_aged,low,yes,excellent,yesyouth,medium,no,fair,noyouth,low,yes,fair,yessenior,medium,yes,fair,yesyouth,medium,yes,excellent,yesmiddle_aged,medium,no,excellent,yesmiddle_aged,high,yes,fair,yessenior,medium,no,excellent,no


attr.txt

age,income,student,credit_rating,buys_computer


運算結果

age(1:youth; 2:middle_aged; 3:senior; )credit_rating(1:fair; 2:excellent; )leaf:no()leaf:yes()leaf:yes()student(1:no; 2:yes; )leaf:no()leaf:yes()



最後附上java代碼

DecisionTree.java

package com.zhyoulun.decision;import java.io.BufferedReader;import java.io.File;import java.io.FileInputStream;import java.io.FileNotFoundException;import java.io.IOException;import java.io.InputStreamReader;import java.util.ArrayList;import java.util.Map;/** * 負責資料的讀入和寫出,以及產生決策樹 *  * @author zhyoulun * */public class DecisionTree{private ArrayList<ArrayList<String>> allDatas;private ArrayList<String> allAttributes;/** * 從檔案中讀取所有相關資料 * @param dataFilePath * @param attrFilePath */public DecisionTree(String dataFilePath,String attrFilePath){super();try{this.allDatas = new ArrayList<>();this.allAttributes = new ArrayList<>();InputStreamReader inputStreamReader = new InputStreamReader(new FileInputStream(new File(dataFilePath)));BufferedReader bufferedReader = new BufferedReader(inputStreamReader);String line = null;while((line=bufferedReader.readLine())!=null){String[] strings = line.split(",");ArrayList<String> data = new ArrayList<>();for(int i=0;i<strings.length;i++)data.add(strings[i]);this.allDatas.add(data);}inputStreamReader = new InputStreamReader(new FileInputStream(new File(attrFilePath)));bufferedReader = new BufferedReader(inputStreamReader);while((line=bufferedReader.readLine())!=null){String[] strings = line.split(",");for(int i=0;i<strings.length;i++)this.allAttributes.add(strings[i]);}inputStreamReader.close();bufferedReader.close();} catch (FileNotFoundException e){// TODO Auto-generated catch blocke.printStackTrace();} catch (IOException e){// TODO Auto-generated catch blocke.printStackTrace();}//for(int i=0;i<this.allAttributes.size();i++)//{//System.out.print(this.allAttributes.get(i)+" ");//}//System.out.println();////for(int i=0;i<this.allDatas.size();i++)//{//for(int j=0;j<this.allDatas.get(i).size();j++)//{//System.out.print(this.allDatas.get(i).get(j)+" ");//}//System.out.println();//}}/** * @param allDatas * @param allAttributes */public DecisionTree(ArrayList<ArrayList<String>> allDatas,ArrayList<String> allAttributes){super();this.allDatas = allDatas;this.allAttributes = allAttributes;}public ArrayList<ArrayList<String>> getAllDatas(){return allDatas;}public void setAllDatas(ArrayList<ArrayList<String>> allDatas){this.allDatas = allDatas;}public ArrayList<String> getAllAttributes(){return allAttributes;}public void setAllAttributes(ArrayList<String> allAttributes){this.allAttributes = allAttributes;}/** * 遞迴產生決策數 * @return */public static TreeNode generateDecisionTree(ArrayList<ArrayList<String>> datas, ArrayList<String> attrs){TreeNode treeNode = new TreeNode();//如果D中的元素都在同一類C中,thenif(isInTheSameClass(datas)){treeNode.setName(datas.get(0).get(datas.get(0).size()-1));//rootNode.setName();return treeNode;}//如果attrs為空白,then(這種情況一般不會出現,我們應該是要對所有的候選屬性集合構建決策樹)if(attrs.size()==0)return treeNode;CriterionID3 criterionID3 = new CriterionID3(datas, attrs);int splitingCriterionIndex = criterionID3.attributeSelectionMethod();treeNode.setName(attrs.get(splitingCriterionIndex));treeNode.setRules(getValueSet(datas, splitingCriterionIndex));attrs.remove(splitingCriterionIndex);Map<String, ArrayList<ArrayList<String>>> subDatasMap = criterionID3.getSubDatasMap(splitingCriterionIndex);//for(String key:subDatasMap.keySet())//{//System.out.println("===========");//System.out.println(key);//for(int i=0;i<subDatasMap.get(key).size();i++)//{//for(int j=0;j<subDatasMap.get(key).get(i).size();j++)//{//System.out.print(subDatasMap.get(key).get(i).get(j)+" ");//}//System.out.println();//}//}for(String key:subDatasMap.keySet()){ArrayList<TreeNode> treeNodes = treeNode.getChildren();treeNodes.add(generateDecisionTree(subDatasMap.get(key), attrs));treeNode.setChildren(treeNodes);}return treeNode;}/** * 擷取datas中index列的範圍 * @param data * @param index * @return */public static ArrayList<String> getValueSet(ArrayList<ArrayList<String>> datas,int index){ArrayList<String> values = new ArrayList<String>();String r = "";for (int i = 0; i < datas.size(); i++) {r = datas.get(i).get(index);if (!values.contains(r)) {values.add(r);}}return values;}/** * 最後一列是類標號,判斷最後一列是否相同 * @param datas * @return */private static boolean isInTheSameClass(ArrayList<ArrayList<String>> datas){String flag = datas.get(0).get(datas.get(0).size()-1);//第0行,最後一列賦初值for(int i=0;i<datas.size();i++){if(!datas.get(i).get(datas.get(i).size()-1).equals(flag))return false;}return true;}public static void main(String[] args){String dataPath = "files/data.txt";String attrPath = "files/attr.txt";//初始化未經處理資料DecisionTree decisionTree = new DecisionTree(dataPath,attrPath);//產生決策樹TreeNode treeNode = generateDecisionTree(decisionTree.getAllDatas(), decisionTree.getAllAttributes());print(treeNode,0);}private static void print(TreeNode treeNode,int level){for(int i=0;i<level;i++)System.out.print("\t");System.out.print(treeNode.getName());System.out.print("(");for(int i=0;i<treeNode.getRules().size();i++)System.out.print((i+1)+":"+treeNode.getRules().get(i)+"; ");System.out.println(")");ArrayList<TreeNode> treeNodes = treeNode.getChildren();for(int i=0;i<treeNodes.size();i++){print(treeNodes.get(i),level+1);}}}



CriterionID3.java

package com.zhyoulun.decision;import java.util.ArrayList;import java.util.HashMap;import java.util.Map;/** * ID3,選擇分裂準則 *  * @author zhyoulun * */public class CriterionID3{private ArrayList<ArrayList<String>> datas;private ArrayList<String> attributes;private Map<String, ArrayList<ArrayList<String>>> subDatasMap;/** * 計算所有的資訊增益,擷取最大的一項作為分裂屬性 * @return */public int attributeSelectionMethod(){double gain = -1.0;int maxIndex = 0;for(int i=0;i<this.attributes.size()-1;i++){double tempGain = this.calcGain(i);if(tempGain>gain){gain = tempGain;maxIndex = i;}}return maxIndex;}/** * 計算 Gain(age)=Info(D)-Info_age(D) 等 * @param index * @return *//** * @param index * @param isCalcSubDatasMap * @return */private double calcGain(int index){double result = 0;//計算Info(D)int lastIndex = datas.get(0).size()-1;ArrayList<String> valueSet = DecisionTree.getValueSet(this.datas,lastIndex);for(String value:valueSet){int count = 0;for(int i=0;i<datas.size();i++){if(datas.get(i).get(lastIndex).equals(value))count++;}result += -(1.0*count/datas.size())*Math.log(1.0*count/datas.size())/Math.log(2);//System.out.println(result);}//System.out.println("==========");//計算Info_a(D)valueSet = DecisionTree.getValueSet(this.datas,index);//for(String temp:valueSet)//System.out.println(temp);//System.out.println("==========");for(String value:valueSet){ArrayList<ArrayList<String>> subDatas = new ArrayList<>();for(int i=0;i<datas.size();i++){if(datas.get(i).get(index).equals(value))subDatas.add(datas.get(i));}//for(ArrayList<String> temp:subDatas)//{//for(String temp2:temp)//System.out.print(temp2+" ");//System.out.println();//}ArrayList<String> subValueSet = DecisionTree.getValueSet(subDatas, lastIndex);//System.out.print("subValueSet:");//for(String temp:subValueSet)//System.out.print(temp+" ");//System.out.println();for(String subValue:subValueSet){//System.out.println("+++++++++++++++");//System.out.println(subValue);int count = 0;for(int i=0;i<subDatas.size();i++){if(subDatas.get(i).get(lastIndex).equals(subValue))count++;}//System.out.println(count);result += -1.0*subDatas.size()/datas.size()*(-(1.0*count/subDatas.size())*Math.log(1.0*count/subDatas.size())/Math.log(2));//System.out.println(result);}}return result;}public CriterionID3(ArrayList<ArrayList<String>> datas,ArrayList<String> attributes){super();this.datas = datas;this.attributes = attributes;}public ArrayList<ArrayList<String>> getDatas(){return datas;}public void setDatas(ArrayList<ArrayList<String>> datas){this.datas = datas;}public ArrayList<String> getAttributes(){return attributes;}public void setAttributes(ArrayList<String> attributes){this.attributes = attributes;}public Map<String, ArrayList<ArrayList<String>>> getSubDatasMap(int index){ArrayList<String> valueSet = DecisionTree.getValueSet(this.datas, index);this.subDatasMap = new HashMap<String, ArrayList<ArrayList<String>>>();for(String value:valueSet){ArrayList<ArrayList<String>> subDatas = new ArrayList<>();for(int i=0;i<this.datas.size();i++){if(this.datas.get(i).get(index).equals(value))subDatas.add(this.datas.get(i));}for(int i=0;i<subDatas.size();i++){subDatas.get(i).remove(index);}this.subDatasMap.put(value, subDatas);}return subDatasMap;}public void setSubDatasMap(Map<String, ArrayList<ArrayList<String>>> subDatasMap){this.subDatasMap = subDatasMap;}}



TreeNode.java

package com.zhyoulun.decision;import java.util.ArrayList;public class TreeNode{private String name; // 該結點的名稱(分裂屬性)private ArrayList<String> rules; // 結點的分裂規則(假設均為離散值)//private ArrayList<ArrayList<String>> datas; // 劃分到該結點的訓練元組(datas.get(i)表示一個訓練元組)//private ArrayList<String> candidateAttributes; // 劃分到該結點的候選屬性(與訓練元組的個數一致)private ArrayList<TreeNode> children; // 子結點public TreeNode(){this.name = "";this.rules = new ArrayList<String>();this.children = new ArrayList<TreeNode>();//this.datas = null;//this.candidateAttributes = null;}public String getName(){return name;}public void setName(String name){this.name = name;}public ArrayList<String> getRules(){return rules;}public void setRules(ArrayList<String> rules){this.rules = rules;}public ArrayList<TreeNode> getChildren(){return children;}public void setChildren(ArrayList<TreeNode> children){this.children = children;}//public ArrayList<ArrayList<String>> getDatas()//{//return datas;//}////public void setDatas(ArrayList<ArrayList<String>> datas)//{//this.datas = datas;//}////public ArrayList<String> getCandidateAttributes()//{//return candidateAttributes;//}////public void setCandidateAttributes(ArrayList<String> candidateAttributes)//{//this.candidateAttributes = candidateAttributes;//}}



參考:《資料採礦概念與技術(第3版)》

轉載請註明出處:

聯繫我們

該頁面正文內容均來源於網絡整理,並不代表阿里雲官方的觀點,該頁面所提到的產品和服務也與阿里云無關,如果該頁面內容對您造成了困擾,歡迎寫郵件給我們,收到郵件我們將在5個工作日內處理。

如果您發現本社區中有涉嫌抄襲的內容,歡迎發送郵件至: info-contact@alibabacloud.com 進行舉報並提供相關證據,工作人員會在 5 個工作天內聯絡您,一經查實,本站將立刻刪除涉嫌侵權內容。

A Free Trial That Lets You Build Big!

Start building with 50+ products and up to 12 months usage for Elastic Compute Service

  • Sales Support

    1 on 1 presale consultation

  • After-Sales Support

    24/7 Technical Support 6 Free Tickets per Quarter Faster Response

  • Alibaba Cloud offers highly flexible support services tailored to meet your exact needs.