Plotting a tree graph by using functions like functions like treeplot
Show older comments
Hellow, I have this code of ID3 decision tree, but I have not been able to figure out how to plot the tree structure generated by the code, can some one help please!
ID3.m
function [tree] = ID3(examples, attributes, activeAttributes)
% ID3 Runs the ID3 algorithm on the matrix of examples and attributes
% args:
% examples - matrix of 1s and 0s for trues and falses, the
% last value in each row being the value of the
% classifying attribute
% attributes - cell array of attribute strings (no CLASS)
% activeAttributes - vector of 1s and 0s, 1 if corresponding attr.
% active (no CLASS)
% return:
% tree - the root node of a decision tree
% tree struct:
% value - will be the string for the splitting
% attribute, or 'true' or 'false' for leaf
% left - left pointer to another tree node (left means
% the splitting attribute was false)
% right - right pointer to another tree node (right
% means the splitting attribute was true)
if (isempty(examples));
error('Must provide examples');
end
% Constants
numberAttributes = length(activeAttributes);
numberExamples = length(examples(:,1));
% Create the tree node
tree = struct('value', 'null', 'left', 'null', 'right', 'null');
% If last value of all rows in examples is 1, return tree labeled 'true'
lastColumnSum = sum(examples(:, numberAttributes + 1));
if (lastColumnSum == numberExamples);
tree.value = 'true';
return
end
% If last value of all rows in examples is 0, return tree labeled 'false'
if (lastColumnSum == 0);
tree.value = 'false';
return
end
% If activeAttributes is empty, then return tree with label as most common
% value
if (sum(activeAttributes) == 0);
if (lastColumnSum >= numberExamples / 2);
tree.value = 'true';
else
tree.value = 'false';
end
return
end
% Find the current entropy
p1 = lastColumnSum / numberExamples;
if (p1 == 0);
p1_eq = 0;
else
p1_eq = -1*p1*log2(p1);
end
p0 = (numberExamples - lastColumnSum) / numberExamples;
if (p0 == 0);
p0_eq = 0;
else
p0_eq = -1*p0*log2(p0);
end
currentEntropy = p1_eq + p0_eq;
% Find the attribute that maximizes information gain
gains = -1*ones(1,numberAttributes); %-1 if inactive, gains for all else
% Loop through attributes updating gains, making sure they are still active
for i=1:numberAttributes;
if (activeAttributes(i)) % this one is still active, update its gain
s0 = 0; s0_and_true = 0;
s1 = 0; s1_and_true = 0;
for j=1:numberExamples;
if (examples(j,i)); % this instance has splitting attr. true
s1 = s1 + 1;
if (examples(j, numberAttributes + 1)); %target attr is true
s1_and_true = s1_and_true + 1;
end
else
s0 = s0 + 1;
if (examples(j, numberAttributes + 1)); %target attr is true
s0_and_true = s0_and_true + 1;
end
end
end
% Entropy for S(v=1)
if (~s1);
p1 = 0;
else
p1 = (s1_and_true / s1);
end
if (p1 == 0);
p1_eq = 0;
else
p1_eq = -1*(p1)*log2(p1);
end
if (~s1);
p0 = 0;
else
p0 = ((s1 - s1_and_true) / s1);
end
if (p0 == 0);
p0_eq = 0;
else
p0_eq = -1*(p0)*log2(p0);
end
entropy_s1 = p1_eq + p0_eq;
% Entropy for S(v=0)
if (~s0);
p1 = 0;
else
p1 = (s0_and_true / s0);
end
if (p1 == 0);
p1_eq = 0;
else
p1_eq = -1*(p1)*log2(p1);
end
if (~s0);
p0 = 0;
else
p0 = ((s0 - s0_and_true) / s0);
end
if (p0 == 0);
p0_eq = 0;
else
p0_eq = -1*(p0)*log2(p0);
end
entropy_s0 = p1_eq + p0_eq;
gains(i) = currentEntropy - ((s1/numberExamples)*entropy_s1) - ((s0/numberExamples)*entropy_s0);
end
end
% Pick the attribute that maximizes gains
[~, bestAttribute] = max(gains);
% Set tree.value to bestAttribute's relevant string
tree.value = attributes{bestAttribute};
% Remove splitting attribute from activeAttributes
activeAttributes(bestAttribute) = 0;
% Initialize and create the new example matrices
examples_0 = []; examples_0_index = 1;
examples_1 = []; examples_1_index = 1;
for i=1:numberExamples;
if (examples(i, bestAttribute)); % this instance has it as 1/true
examples_1(examples_1_index, :) = examples(i, :); % copy over
examples_1_index = examples_1_index + 1;
else
examples_0(examples_0_index, :) = examples(i, :);
examples_0_index = examples_0_index + 1;
end
end
% For both values of the splitting attribute
% For value = false or 0, corresponds to left branch
% If examples_0 is empty, add leaf node to the left with relevant label
if (isempty(examples_0));
leaf = struct('value', 'null', 'left', 'null', 'right', 'null');
if (lastColumnSum >= numberExamples / 2); % for matrix examples
leaf.value = 'true';
else
leaf.value = 'false';
end
tree.left = leaf;
else
% Here is were we can recur
tree.left = ID3(examples_0, attributes, activeAttributes);
end
% For value = true or 1, corresponds to right branch
% If examples_1 is empty, add leaf node to the right with relevant label
if (isempty(examples_1));
leaf = struct('value', 'null', 'left', 'null', 'right', 'null');
if (lastColumnSum >= numberExamples / 2); % for matrix examples
leaf.value = 'true';
else
leaf.value = 'false';
end
tree.right = leaf;
else
% Here is were we can recur
tree.right = ID3(examples_1, attributes, activeAttributes);
end
% Now we can return tree
return
end
ClassifyByTree.m
function [classifications] = ClassifyByTree(tree, attributes, instance)
% ClassifyByTree Classifies data instance by given tree
% args:
% tree - tree data structure
% attributes - cell array of attribute strings (no CLASS)
% instance - data including correct classification (end col.)
% return:
% classifications - 2 numbers, first given by tree, 2nd given by
% instance's last column
% tree struct:
% value - will be the string for the splitting
% attribute, or 'true' or 'false' for leaf
% left - left pointer to another tree node (left means
% the splitting attribute was false)
% right - right pointer to another tree node (right
% means the splitting attribute was true)
% Store the actual classification
actual = instance(1, length(instance));
% Recursion with 3 cases
% Case 1: Current node is labeled 'true'
% So trivially return the classification as 1
if (strcmp(tree.value, 'true'));
classifications = [1, actual];
return
end
% Case 2: Current node is labeled 'false'
% So trivially return the classification as 0
if (strcmp(tree.value, 'false'));
classifications = [0, actual];
return
end
% Case 3: Current node is labeled an attribute
% Follow correct branch by looking up index in attributes, and recur
index = find(ismember(attributes,tree.value)==1);
if (instance(1, index)); % attribute is true for this instance
% Recur down the right side
classifications = ClassifyByTree(tree.right, attributes, instance);
else
% Recur down the left side
classifications = ClassifyByTree(tree.left, attributes, instance);
end
return
end
ClassifyByTree.m
function [classifications] = ClassifyByTree(tree, attributes, instance)
% ClassifyByTree Classifies data instance by given tree
% args:
% tree - tree data structure
% attributes - cell array of attribute strings (no CLASS)
% instance - data including correct classification (end col.)
% return:
% classifications - 2 numbers, first given by tree, 2nd given by
% instance's last column
% tree struct:
% value - will be the string for the splitting
% attribute, or 'true' or 'false' for leaf
% left - left pointer to another tree node (left means
% the splitting attribute was false)
% right - right pointer to another tree node (right
% means the splitting attribute was true)
% Store the actual classification
actual = instance(1, length(instance));
% Recursion with 3 cases
% Case 1: Current node is labeled 'true'
% So trivially return the classification as 1
if (strcmp(tree.value, 'true'));
classifications = [1, actual];
return
end
% Case 2: Current node is labeled 'false'
% So trivially return the classification as 0
if (strcmp(tree.value, 'false'));
classifications = [0, actual];
return
end
% Case 3: Current node is labeled an attribute
% Follow correct branch by looking up index in attributes, and recur
index = find(ismember(attributes,tree.value)==1);
if (instance(1, index)); % attribute is true for this instance
% Recur down the right side
classifications = ClassifyByTree(tree.right, attributes, instance);
else
% Recur down the left side
classifications = ClassifyByTree(tree.left, attributes, instance);
end
return
end
decisiontree.m
% George Wheaton
% EECS 349
% Homework 1 Problem 7
% October 7, 2012
% ID3 Decision Tree Algorithm
function[] = decisiontree(inputFileName, trainingSetSize, numberOfTrials,...
verbose)
% DECISIONTREE Create a decision tree by following the ID3 algorithm
% args:
% inputFileName - the fully specified path to input file
% trainingSetSize - integer specifying number of examples from input
% used to train the dataset
% numberOfTrials - integer specifying how many times decision tree
% will be built from a randomly selected subset
% of the training examples
% verbose - string that must be eiher '1' or '0', if '1'
% output includes training and test sets, else
% it will only contain description of tree and
% results for the trials
% Read in the specified text file contain the examples
fid = fopen(inputFileName, 'rt');
dataInput = textscan(fid, '%s');
% Close the file
fclose(fid);
% Reformat the data into attribute array and data matrix of 1s and 0s for
% true or false
i = 1;
% First store the attributes into a cell array
while (~strcmp(dataInput{1}{i}, 'CLASS'));
i = i + 1;
end
attributes = cell(1,i);
for j=1:i;
attributes{j} = dataInput{1}{j};
end
% NOTE: The classification will be the final attribute in the data rows
% below
numAttributes = i;
numInstances = (length(dataInput{1}) - numAttributes) / numAttributes;
% Then store the data into matrix
data = zeros(numInstances, numAttributes);
i = i + 1;
for j=1:numInstances
for k=1:numAttributes
data(j, k) = strcmp(dataInput{1}{i}, 'true');
i = i + 1;
end
end
% Here is where the trials start
for i=1:numberOfTrials;
% Print the trial number
fprintf('TRIAL NUMBER: %d\n\n', i);
% Split data into training and testing sets randomly
% Use randsample to get a vector of row numbers for the training set
rows = sort(randsample(numInstances, trainingSetSize));
% Initialize two new matrices, training set and test set
trainingSet = zeros(trainingSetSize, numAttributes);
testingSetSize = (numInstances - trainingSetSize);
testingSet = zeros(testingSetSize, numAttributes);
% Loop through data matrix, copying relevant rows to each matrix
training_index = 1;
testing_index = 1;
for data_index=1:numInstances;
if (rows(training_index) == data_index);
trainingSet(training_index, :) = data(data_index, :);
if (training_index < trainingSetSize);
training_index = training_index + 1;
end
else
testingSet(testing_index, :) = data(data_index, :);
if (testing_index < testingSetSize);
testing_index = testing_index + 1;
end
end
end
% If verbose, print out training set
if (verbose);
for ii=1:numAttributes;
fprintf('%s\t', attributes{ii});
end
fprintf('\n');
for ii=1:trainingSetSize;
for jj=1:numAttributes;
if (trainingSet(ii, jj));
fprintf('%s\t', 'true');
else
fprintf('%s\t', 'false');
end
end
fprintf('\n');
end
end
% Estimate the expected prior probability of TRUE and FALSE based on
% training set
if (sum(trainingSet(:, numAttributes)) >= trainingSetSize);
expectedPrior = 'true';
else
expectedPrior = 'false';
end
% Construct a decision tree on the training set using the ID3 algorithm
activeAttributes = ones(1, length(attributes) - 1);
new_attributes = attributes(1:length(attributes)-1);
tree = ID3(trainingSet, attributes, activeAttributes);
% Print out the tree
fprintf('DECISION TREE STRUCTURE:\n');
PrintTree(tree, 'root');
% Run tree and expected prior against testing set, recording
% classifications
% The second column is for actual classification, first for calculated
ID3_Classifications = zeros(testingSetSize,2);
ExpectedPrior_Classifications = zeros(testingSetSize,2);
ID3_numCorrect = 0; ExpectedPrior_numCorrect = 0;
for k=1:testingSetSize; %over the testing set
% Call a recursive function to follow the tree nodes and classify
ID3_Classifications(k,:) = ...
ClassifyByTree(tree, new_attributes, testingSet(k,:));
ExpectedPrior_Classifications(k, 2) = testingSet(k,numAttributes);
if (expectedPrior);
ExpectedPrior_Classifications(k, 1) = 1;
else
ExpectedPrior_Classifications(k, 0) = 0;
end
if (ID3_Classifications(k,1) == ID3_Classifications(k, 2)); %correct
ID3_numCorrect = ID3_numCorrect + 1;
end
if (ExpectedPrior_Classifications(k,1) == ExpectedPrior_Classifications(k,2));
ExpectedPrior_numCorrect = ExpectedPrior_numCorrect + 1;
end
end
% If verbose, print the testing data with final two columns ID3 Class
% and Prior Class
if (verbose);
for ii=1:numAttributes;
fprintf('%s\t', attributes{ii});
end
fprintf('%s\t%s\t', 'ID3 Class', 'Prior Class');
fprintf('\n');
for ii=1:testingSetSize;
for jj=1:numAttributes;
if (testingSet(ii, jj));
fprintf('%s\t', 'true');
else
fprintf('%s\t', 'false');
end
end
if (ID3_Classifications(ii,1));
fprintf('%s\t', 'true');
else
fprintf('%s\t', 'false');
end
if (ExpectedPrior_Classifications(ii,1));
fprintf('%s\t', 'true');
else
fprintf('%s\t', 'false');
end
fprintf('\n');
end
end
% Calculate the proportions correct and print out
if (testingSetSize);
ID3_Percentage = round(100 * ID3_numCorrect / testingSetSize);
ExpectedPrior_Percentage = round(100 * ExpectedPrior_numCorrect / testingSetSize);
else
ID3_Percentage = 0;
ExpectedPrior_Percentage = 0;
end
ID3_Percentages(i) = ID3_Percentage;
ExpectedPrior_Percentages(i) = ExpectedPrior_Percentage;
fprintf('\tPercent of test cases correctly classified by an ID3 decision tree = %d\n' ...
, ID3_Percentage);
fprintf('\tPercent of test cases correctly classified by using prior probabilities from the training set = %d\n\n' ...
, ExpectedPrior_Percentage);
end
meanID3 = round(mean(ID3_Percentages));
meanPrior = round(mean(ExpectedPrior_Percentages));
% Print out remaining details
fprintf('example file used = %s\n', inputFileName);
fprintf('number of trials = %d\n', numberOfTrials);
fprintf('training set size for each trial = %d\n', trainingSetSize);
fprintf('testing set size for each trial = %d\n', testingSetSize);
fprintf('mean performance (percentage correct) of decision tree over all trials = %d\n', meanID3);
fprintf('mean performance (percentage correct) of prior probability from training set = %d\n\n', meanPrior);
end
PrintTree.m
function [] = PrintTree(tree, parent)
% Prints the tree structure (preorder traversal)
% Print current node
if (strcmp(tree.value, 'true'));
fprintf('parent: %s\ttrue\n', parent);
return
elseif (strcmp(tree.value, 'false'));
fprintf('parent: %s\tfalse\n', parent);
return
else
% Current node an attribute splitter
fprintf('parent: %s\tattribute: %s\tfalseChild:%s\ttrueChild:%s\n', ...
parent, tree.value, tree.left.value, tree.right.value);
end
% Recur the left subtree
PrintTree(tree.left, tree.value);
% Recur the right subtree
PrintTree(tree.right, tree.value);
end
Any help would be much appreciated
Answers (0)
Categories
Find more on Document and Integrate Toolboxes in Help Center and File Exchange
Community Treasure Hunt
Find the treasures in MATLAB Central and discover how the community can help you!
Start Hunting!