diff options
Diffstat (limited to 'sourcecodes/parameter_learning')
| -rw-r--r-- | sourcecodes/parameter_learning/drawFigure.m | 3 | ||||
| -rw-r--r-- | sourcecodes/parameter_learning/kfoldCrossValid.m | 322 | ||||
| -rw-r--r-- | sourcecodes/parameter_learning/looCrossValid.m | 250 | ||||
| -rw-r--r-- | sourcecodes/parameter_learning/parameterLearning.m | 5 | ||||
| -rw-r--r-- | sourcecodes/parameter_learning/prepareInput.m | 2 | ||||
| -rw-r--r-- | sourcecodes/parameter_learning/testSetPredictions.m | 361 | ||||
| -rw-r--r-- | sourcecodes/parameter_learning/writeParameters.m | 3 |
7 files changed, 941 insertions, 5 deletions
diff --git a/sourcecodes/parameter_learning/drawFigure.m b/sourcecodes/parameter_learning/drawFigure.m index 7da07a90..c48f5b2d 100644 --- a/sourcecodes/parameter_learning/drawFigure.m +++ b/sourcecodes/parameter_learning/drawFigure.m @@ -123,7 +123,8 @@ for i = 1:nnodes, fprintf(fileID,format,num_child(i),children(1,:)); end - predict = marginal_nodes(engine,i); +% predict = marginal_nodes(engine,i); + predict = marginal_nodes_no_ev(bnet,engine,i); if bnet.node_sizes(i) ~= 1, for j = 1:bnet.node_sizes(i), %%%For discrete nodes, the state and the percent of that state diff --git a/sourcecodes/parameter_learning/kfoldCrossValid.m b/sourcecodes/parameter_learning/kfoldCrossValid.m new file mode 100644 index 00000000..6306c511 --- /dev/null +++ b/sourcecodes/parameter_learning/kfoldCrossValid.m @@ -0,0 +1,322 @@ +function kfoldCrossValid(pre,predict_label,nfolds) +% This function will peform k-fold cross-validation. +% This requires the specification of the name of the variable +% that you want to predict and the number of folds that the +% data should be divided into. + +nfolds = uint8(str2num(nfolds)); + +sfile=strcat(pre,'structure_input.txt'); +dfile=strcat(pre,'continuous_input.txt'); +nnodefile=strcat(pre,'nnode.txt'); +fnnode = fopen(nnodefile,'r'); +nnodes = fscanf(fnnode,'%d'); + +Std_flag=true; +[labels,cases,bnet]=readInput(dfile,sfile,nnodes,Std_flag); + +for i=1:nnodes + if strcmp(labels(i),predict_label) + predict_node = i; + end +end + +predict_cases = bnet.node_sizes(predict_node); + +ncases = size(cases,2); + +kfold_index = kfold_bin(ncases,nfolds); + +if predict_cases == 1 + kfoldCV_continuous(pre,predict_label,nnodes,labels,cases,bnet,predict_node,predict_cases,nfolds,kfold_index); +else + kfoldCV_discrete(pre,predict_label,nnodes,labels,cases,bnet,predict_node,predict_cases,nfolds,kfold_index); +endif + + +end + + +function kfoldCV_continuous(pre,predict_label,nnodes,labels,cases,bnet,predict_node,predict_cases,nfolds,kfold_index) + +ncases = size(cases,2); + +%%Read in original means and standard deviations to report output as +%% untransformed values. +means_orig=cell(1,nnodes); +stdevs_orig=cell(1,nnodes); +labels_orig=cell(1,nnodes); +%Read in original means and standard deviations +mapfile = strcat(pre,'map.txt'); +fmap = fopen(mapfile,'r'); +for i=1:nnodes + buffer = fgetl(mapfile); + temp = cell(1,3); + for j=1:3 + [next,buffer] = strtok(buffer); + temp{j} = next; + end + labels_orig{i} = temp{1}; + means_orig{i} = str2num(temp{3}); + stdevs_orig{i} = str2num(temp{2}); +end +fclose(fmap); +%Need to map the means and stdevs to the correct labels +means = cell(1,nnodes); +stdevs = cell(1,nnodes); +for i = 1:nnodes + for j = 1:nnodes + if strcmp(labels{i},labels_orig{j}) + means{i} = means_orig{j}; + stdevs{i} = stdevs_orig{j}; + break + end + end +end + +kfoldPredictions=zeros(size(cases,2),2); + + +%t=cputime; +%First get predictions +for i =1:nfolds +% i + test_flag = false(ncases,1); + for j=1:ncases + if kfold_index(j) == i + test_flag(j) = true; + end + end + trainData = cases(:,~test_flag); + testData = cases(:,test_flag); + temp_index = zeros(size(testData,2),1); + temp = 1; + for j=1:ncases + if kfold_index(j) == i + temp_index(temp) = j; + temp = temp + 1; + end + end + [bnet] = parameterLearning(bnet,trainData); + for j=1:size(testData,2) + evidence = testData(:,j); + evidence{predict_node} = {}; + engine = jtree_inf_engine(bnet); + [engine,loglik] = enter_evidence(engine,evidence); + predict = marginal_nodes(engine,predict_node); + adj_mu = predict.mu*stdevs{predict_node}+means{predict_node}; + adj_sigma = stdevs{predict_node}*predict.Sigma; + kfoldPredictions(temp_index(j),1) = adj_mu; + kfoldPredictions(temp_index(j),2) = adj_sigma; + end +end +%e=cputime-t; + + +%Open output file. +filename = strcat(pre,'kfoldCV.txt'); +fileID = fopen(filename,'w'); + +fprintf(fileID,'Variable that was predicted: %s\n\n',predict_label); + +%Calculate RMSEP (root mean square error of prediction) and q^2 +%First, calculate TSS (total sum of squares) and +% PRESS (sum of squares of prediction errors) +average = mean(cell2mat(cases'))(predict_node); +%undo standardization +average = average*stdevs{predict_node}+means{predict_node}; +tss = 0; +press = 0; +case_adj=zeros(size(cases,2),1); +for i =1:ncases + case_adj(i) = cases{predict_node,i}*stdevs{predict_node}+means{predict_node}; + tss = (case_adj(i)-average)^2 + tss; + press = (kfoldPredictions(i,1)-case_adj(i))^2 + press; +end +rmsep = sqrt(press/ncases); +q_squared = 1 - press/tss; + +%% Print rmseq and q^2 +fprintf(fileID,'RMS error of predictions: %6.4f\n',rmsep); +fprintf(fileID,'Q^2 of predictions: %6.4f\n\n',q_squared); + + +%%Print the predictions +fprintf(fileID,'Predicted mean and standard deviation for each case:\n'); +fprintf(fileID,'CaseRow\tFoldNumber\tActualValue\tPredictionMean\tPredictionStDev\n'); +for i = 1:ncases + fprintf(fileID,'%i\t%i\t%6.4f\t',i,kfold_index(i),case_adj(i)); + fprintf(fileID,'%6.4f\t%6.4f\n',kfoldPredictions(i,:)); +end + +end + + +function kfoldCV_discrete(pre,predict_label,nnodes,labels,cases,bnet,predict_node,predict_cases,nfolds,kfold_index) + +ncases = size(cases,2); + +%This next section just gets the original names of the levels. +% so they can be written to the output file. +%%Get the maximum_number of states so array will be big enough +%%Add 1 because the input includes the node name +max_states = max(bnet.node_sizes) + 1; +disc_nodes = size(bnet.dnodes,2); + +%%Get mapping of discrete levels. +levelfile = strcat(pre,'nlevels.txt'); +flevels = fopen(levelfile,'r'); +levels = cell(disc_nodes,max_states); +ndisc_nodes = 0; +for i=1:disc_nodes + ndisc_nodes = ndisc_nodes + 1; + buffer = fgetl(flevels); + for j = 1:max_states + [next,buffer] = strtok(buffer); + if j == 1 + levels{i,j} = next; + else + levels{i,j} = next; + end + if length(buffer) < 1 + break + end + end +end + +pred_levels = cell(1,predict_cases); +for i = 1:disc_nodes + if strcmp(levels{i,1},predict_label); + for j = 1:predict_cases + pred_levels{j} = levels{i,j+1}; + end + break + end +end + +kfoldPredictions=zeros(size(cases,2),predict_cases); + +%t=cputime; +%First get kfold CV predictions +for i =1:nfolds + test_flag = false(ncases,1); + for j=1:ncases + if kfold_index(j) == i + test_flag(j) = true; + end + end + trainData = cases(:,~test_flag); + testData = cases(:,test_flag); + temp_index = zeros(size(testData,2),1); + temp = 1; + for j =1:ncases + if kfold_index(j) == i + temp_index(temp) = j; + temp = temp + 1; + end + end + [bnet] = parameterLearning(bnet,trainData); + for j = 1:size(testData,2) + evidence = testData(:,j); + evidence{predict_node} = {}; + engine = jtree_inf_engine(bnet); + [engine,loglik] = enter_evidence(engine,evidence); + predict = marginal_nodes(engine,predict_node); + for k = 1:predict_cases + kfoldPredictions(temp_index(j),k) = predict.T(k); + end + end +end +%e=cputime-t; + +%Now compare with actual outcomes +%actual_states = zeros(1,predict_cases); +%for i=1:ncases +% for j = 1:predict_cases +% if cell2mat(cases(predict_node,i)) == j +% actual_states(j) = actual_states(j) + 1; +% end +% end +%end + +pred_states = zeros(1,ncases); + +for i=1:ncases + max_state = 1; + for j = 2:predict_cases + if kfoldPredictions(i,j) > kfoldPredictions(i,max_state) + max_state = j; + end + end + pred_states(i) = max_state; +end + +correct = 0; +for i=1:ncases + if pred_states(i) == cell2mat(cases(predict_node,i)) + correct = correct + 1; + end +end + +accuracy = correct/ncases; + +%Open output file. +filename = strcat(pre,'kfoldCV.txt'); +fileID = fopen(filename,'w'); + +fprintf(fileID,'Variable that was predicted: %s\n\n',predict_label); + + +%%Print the accuracy +fprintf(fileID,'Fraction of accurate predictions: %6.4f\n\n',accuracy); + +%%Print the predictions +fprintf(fileID,'Predicted likelihood of each state for each case:\n'); +fprintf(fileID,'%s\t%s\t%s\t','CaseRow','FoldNumber','ActualState'); +fprintf(fileID,'%s\t',pred_levels{1:end-1}); +fprintf(fileID,'%s\n',pred_levels{end}); +for i = 1:ncases + fprintf(fileID,'%i\t%i\t',i,kfold_index(i)); + case_level = pred_levels{cases{predict_node,i}}; + fprintf(fileID,'%s\t',case_level); + fprintf(fileID,'%6.4f\t',kfoldPredictions(i,1:end-1)); + fprintf(fileID,'%6.4f\n',kfoldPredictions(i,end)); +end + +end + + +function [kfold_index] = kfold_bin(ncases,nfolds); +%This returns indexes for the different cases to separate data into folds. + +%Get array with random permutation of the number of cases +p = randperm(ncases); + +base_size = idivide(ncases,nfolds); +remainder = rem(ncases,nfolds); + +kfold_index = zeros(ncases,1); +group = 1; +count = 0; +for i=1:ncases + count = count + 1; + kfold_index(p(i)) = group; +%Need to do some checks if groups cannot be exactly equally sized +%If you have already added an extra member to the group, go to next group + if count > base_size + count = 0; + group = group + 1; +%If you have filled the group, check if an extra is needed + elseif count == base_size + if remainder > 0 + remainder = remainder - 1; + else + count = 0; + group = group + 1; + end + end +end + + +end + diff --git a/sourcecodes/parameter_learning/looCrossValid.m b/sourcecodes/parameter_learning/looCrossValid.m new file mode 100644 index 00000000..67b04409 --- /dev/null +++ b/sourcecodes/parameter_learning/looCrossValid.m @@ -0,0 +1,250 @@ +function looCrossValid(pre,predict_label) +% This function will peform leave-one-out cross-validation. +% This requires the specification of the name of the variable +% that you want to predict. + + + +sfile=strcat(pre,'structure_input.txt'); +dfile=strcat(pre,'continuous_input.txt'); +nnodefile=strcat(pre,'nnode.txt'); +fnnode = fopen(nnodefile,'r'); +nnodes = fscanf(fnnode,'%d'); + +Std_flag=true; +[labels,cases,bnet]=readInput(dfile,sfile,nnodes,Std_flag); + +for i=1:nnodes + if strcmp(labels(i),predict_label) + predict_node = i; + end +end + +predict_cases = bnet.node_sizes(predict_node); + +if predict_cases == 1 + looCV_continuous(pre,predict_label,nnodes,labels,cases,bnet,predict_node,predict_cases); +else + looCV_discrete(pre,predict_label,nnodes,labels,cases,bnet,predict_node,predict_cases); +endif + +end + +function looCV_continuous(pre,predict_label,nnodes,labels,cases,bnet,predict_node,predict_cases) + +ncases = size(cases,2); + +%%Read in original means and standard deviations to report output as +%% untransformed values. +means_orig=cell(1,nnodes); +stdevs_orig=cell(1,nnodes); +labels_orig=cell(1,nnodes); +%Read in original means and standard deviations +mapfile = strcat(pre,'map.txt'); +fmap = fopen(mapfile,'r'); +for i=1:nnodes + buffer = fgetl(mapfile); + temp = cell(1,3); + for j=1:3 + [next,buffer] = strtok(buffer); + temp{j} = next; + end + labels_orig{i} = temp{1}; + means_orig{i} = str2num(temp{3}); + stdevs_orig{i} = str2num(temp{2}); +end +fclose(fmap); +%Need to map the means and stdevs to the correct labels +means = cell(1,nnodes); +stdevs = cell(1,nnodes); +for i = 1:nnodes + for j = 1:nnodes + if strcmp(labels{i},labels_orig{j}) + means{i} = means_orig{j}; + stdevs{i} = stdevs_orig{j}; + break + end + end +end + +loopredictions=zeros(size(cases,2),2); + +%t=cputime; +%First get loo predictions +for i =1:ncases +% i + current_data = cases(:,i); + cases_new = cases; + cases_new(:,i) = []; + evidence = current_data; + evidence{predict_node} = {}; + [bnet]=parameterLearning(bnet,cases_new); + engine = jtree_inf_engine(bnet); + [engine,loglik] = enter_evidence(engine,evidence); + predict = marginal_nodes(engine,predict_node); + adj_mu = predict.mu*stdevs{predict_node}+means{predict_node}; + adj_sigma = stdevs{predict_node}*predict.Sigma; + loopredictions(i,1) = adj_mu; + loopredictions(i,2) = adj_sigma; +end +%e=cputime-t; + +%Open output file. +filename = strcat(pre,'looCV.txt'); +fileID = fopen(filename,'w'); + +fprintf(fileID,'Variable that was predicted: %s\n\n',predict_label); + +%Calculate RMSEP (root mean square error of prediction) and q^2 +%First, calculate TSS (total sum of squares) and +% PRESS (sum of squares of prediction errors) +average = mean(cell2mat(cases'))(predict_node); + +%undo standardization +average = average*stdevs{predict_node}+means{predict_node}; +tss = 0; +press = 0; +case_adj=zeros(size(cases,2),1); +for i =1:ncases + case_adj(i) = cases{predict_node,i}*stdevs{predict_node}+means{predict_node}; + tss = (case_adj(i)-average)^2 + tss; + press = (loopredictions(i,1)-case_adj(i))^2 + press; +end +rmsep = sqrt(press/ncases); +q_squared = 1 - press/tss; + +%% Print rmseq and q^2 +fprintf(fileID,'RMS error of predictions: %6.4f\n',rmsep); +fprintf(fileID,'Q^2 of predictions: %6.4f\n\n',q_squared); + + +%%Print the predictions +fprintf(fileID,'Predicted mean and standard deviation for each case:\n'); +fprintf(fileID,'CaseRow\tActualValue\tPredictionMean\tPredictionStDev\n'); +for i = 1:ncases + fprintf(fileID,'%i\t%6.4f\t',i,case_adj(i)); + fprintf(fileID,'%6.4f\t%6.4f\n',loopredictions(i,:)); +end + +end + + +function looCV_discrete(pre,predict_label,nnodes,labels,cases,bnet,predict_node,predict_cases) + +ncases = size(cases,2); + +%This next section just gets the original names of the levels. +% so they can be written to the output file. +%%Get the maximum_number of states so array will be big enough +%%Add 1 because the input includes the node name +max_states = max(bnet.node_sizes) + 1; +disc_nodes = size(bnet.dnodes,2); + +%%Get mapping of discrete levels. +levelfile = strcat(pre,'nlevels.txt'); +flevels = fopen(levelfile,'r'); +levels = cell(disc_nodes,max_states); +ndisc_nodes = 0; +for i=1:disc_nodes + ndisc_nodes = ndisc_nodes + 1; + buffer = fgetl(flevels); + for j = 1:max_states + [next,buffer] = strtok(buffer); + if j == 1 + levels{i,j} = next; + else + levels{i,j} = next; + end + if length(buffer) < 1 + break + end + end +end + +pred_levels = cell(1,predict_cases); +for i = 1:disc_nodes + if strcmp(levels{i,1},predict_label); + for j = 1:predict_cases + pred_levels{j} = levels{i,j+1}; + end + break + end +end + +loopredictions=zeros(size(cases,2),predict_cases); + +%t=cputime; +%First get loo predictions +for i =1:ncases +% i + current_data = cases(:,i); + cases_new = cases; + cases_new(:,i) = []; + evidence = current_data; + evidence{predict_node} = {}; + + [bnet]=parameterLearning(bnet,cases_new); + engine = jtree_inf_engine(bnet); + [engine,loglik] = enter_evidence(engine,evidence); + predict = marginal_nodes(engine,predict_node); + for j = 1:predict_cases + loopredictions(i,j) = predict.T(j); + end +end +%e=cputime-t; + +%Now compare with actual outcomes +%actual_states = zeros(1,predict_cases); +%for i=1:ncases +% for j = 1:predict_cases +% if cell2mat(cases(predict_node,i)) == j +% actual_states(j) = actual_states(j) + 1; +% end +% end +%end + +pred_states = zeros(1,ncases); + +for i=1:ncases + max_state = 1; + for j = 2:predict_cases + if loopredictions(i,j) > loopredictions(i,max_state) + max_state = j; + end + end + pred_states(i) = max_state; +end + +correct = 0; +for i=1:ncases + if pred_states(i) == cell2mat(cases(predict_node,i)) + correct = correct + 1; + end +end + +accuracy = correct/ncases; + +%Open output file. +filename = strcat(pre,'looCV.txt'); +fileID = fopen(filename,'w'); + +fprintf(fileID,'Variable that was predicted: %s\n\n',predict_label); + + +%%Print the accuracy +fprintf(fileID,'Fraction of accurate predictions: %6.4f\n\n',accuracy); + +%%Print the predictions +fprintf(fileID,'Predicted likelihood of each state for each case:\n'); +fprintf(fileID,'%s\t%s\t','CaseRow','ActualState'); +fprintf(fileID,'%s\t',pred_levels{1:end-1}); +fprintf(fileID,'%s\n',pred_levels{end}); +for i = 1:ncases + fprintf(fileID,'%i\t',i); + case_level = pred_levels{cases{predict_node,i}}; + fprintf(fileID,'%s\t',case_level); + fprintf(fileID,'%6.4f\t',loopredictions(i,1:end-1)); + fprintf(fileID,'%6.4f\n',loopredictions(i,end)); +end + +end diff --git a/sourcecodes/parameter_learning/parameterLearning.m b/sourcecodes/parameter_learning/parameterLearning.m index 3ef6c0b3..2c4a1f9f 100644 --- a/sourcecodes/parameter_learning/parameterLearning.m +++ b/sourcecodes/parameter_learning/parameterLearning.m @@ -4,7 +4,7 @@ function [ bnet ] = parameterLearning( bnet,cases,engine_name ) % % This is very basic now. It could be modified to use different engine % types in the future. Now, I always use the 'jtree_inf_engine'. -% +% with dirichlet priors % % parameterLearning is called by runBN_initial.m, % Predictmultiple.m, and Predictmultipleintervention.m @@ -35,7 +35,8 @@ nnodes = size(dnodes,2)+size(cnodes,2); %make dnodes tabular_CPT for i = 1:size(dnodes,2) - bnet.CPD{dnodes(i)} = tabular_CPD(bnet,dnodes(i)); +% bnet.CPD{dnodes(i)} = tabular_CPD(bnet,dnodes(i)); + bnet.CPD{dnodes(i)} = tabular_CPD(bnet,dnodes(i),'prior_type','dirichlet'); end for i = 1:size(cnodes,2) diff --git a/sourcecodes/parameter_learning/prepareInput.m b/sourcecodes/parameter_learning/prepareInput.m index 84d2ae17..5268f872 100644 --- a/sourcecodes/parameter_learning/prepareInput.m +++ b/sourcecodes/parameter_learning/prepareInput.m @@ -274,7 +274,7 @@ for i = 1:nnodes if levels{i} > 1 for j = 1:ncases for k=1:size(states{i},1) - if data{j,i} == states{i}{k} + if strcmp(data{j,i},states{i}{k}) data{j,i} = sprintf('%i',num2cell(k){1});; break end diff --git a/sourcecodes/parameter_learning/testSetPredictions.m b/sourcecodes/parameter_learning/testSetPredictions.m new file mode 100644 index 00000000..0fff197d --- /dev/null +++ b/sourcecodes/parameter_learning/testSetPredictions.m @@ -0,0 +1,361 @@ +function [ ] = testSetPredictions( pre ) + % + % This function will make predictions for the cases included + % in the uploaded data file. + % + % + % Input: ???ts_input.txt + % This is the input file that is uploaded to BNW. + % It is directly written out by the BNW php code with no modification. + % The file format is a header line containing the variable that you want to predict. + % Then, there is a second header line with the variable names + % Finally, the file contains the data, with each case in a row. + % If there is missing data, an "NA" should be entered. + % + % Output: ???ts_output.txt + % + % It is called by the run_test_set script in the 'sourcecodes' directory. + + +% open file for input, include error handling +dfile=strcat(pre,'ts_upload.txt'); + +fin = fopen(dfile,'r'); +if fin < 0 + error(['Could not open ',dfile,' for input']); +end + +% Get the number of cases (the number of rows in the file excluding the header) +ntestcases = fskipl(fin,Inf) - 2; + +frewind(fin); + + + +% Read in first line to get the node label of the variable that should be predicted. +buffer = fgetl(fin); +[predict_label,buffer] = strtok(buffer); + +% Read in second line to get the number of nodes and the node labels. +buffer = fgetl(fin); %get header line as a string +nnodes = numel(strfind(buffer,"\t")) + 1; +labels_test = cell(1,nnodes); +for j=1:nnodes + [next,buffer] = strtok(buffer); + labels_test{j} = next; +end + +% Read in the test_data +data_test_temp = cell(ntestcases,nnodes); +for i = 1:ntestcases + buffer = fgetl(fin); + for j = 1:nnodes + [next,buffer] = strtok(buffer); + data_test_temp{i,j} = next; + end +end + + +% Read in the training (original) data and network structure. +sfile=strcat(pre,'structure_input.txt'); +dfile=strcat(pre,'continuous_input.txt'); +nnodefile=strcat(pre,'nnode.txt'); +fnnode = fopen(nnodefile,'r'); +nnodes = fscanf(fnnode,'%d'); + +Std_flag=true; +[labels,cases,bnet]=readInput(dfile,sfile,nnodes,Std_flag); +[bnet] = parameterLearning(bnet,cases); + +%Get node id in actual network for variable to be predicted +for i=1:nnodes + if strcmp(labels(i),predict_label) + predict_node = i; + end +end + +%Reformat test data so the columns match bnet structure +label_map = cell(2,nnodes); +for i=1:nnodes + label_map{1,i} = labels_test{i}; + for j=1:nnodes + if strcmp(labels_test(i),labels(j)) + label_map{2,i} = j; + break + end + end +end + + +data_test = cell(ntestcases,nnodes); +for i=1:nnodes + data_test(:,label_map{2,i}) = data_test_temp(:,i); +end + +%%Read in training data means and standard deviations +means_orig=cell(1,nnodes); +stdevs_orig=cell(1,nnodes); +labels_orig=cell(1,nnodes); +%Read in original means and standard deviations +mapfile = strcat(pre,'map.txt'); +fmap = fopen(mapfile,'r'); +for i=1:nnodes + buffer = fgetl(mapfile); + temp = cell(1,3); + for j=1:3 + [next,buffer] = strtok(buffer); + temp{j} = next; + end + labels_orig{i} = temp{1}; + means_orig{i} = str2num(temp{3}); + stdevs_orig{i} = str2num(temp{2}); +end +fclose(fmap); +%Need to map the means and stdevs to the correct labels +means = cell(1,nnodes); +stdevs = cell(1,nnodes); +for i = 1:nnodes + for j = 1:nnodes + if strcmp(labels{i},labels_orig{j}) + means{i} = means_orig{j}; + stdevs{i} = stdevs_orig{j}; + break + end + end +end + +%Get mapping of discrete levels. +max_states = max(bnet.node_sizes) + 1; +disc_nodes = size(bnet.dnodes,2); +levelfile = strcat(pre,'nlevels.txt'); +flevels = fopen(levelfile,'r'); +levels = cell(disc_nodes,max_states); +ndisc_nodes = 0; +for i=1:disc_nodes + ndisc_nodes = ndisc_nodes + 1; + buffer = fgetl(flevels); + for j = 1:max_states + [next,buffer] = strtok(buffer); + if j == 1 + levels{i,j} = next; + else + levels{i,j} = next; + end + if length(buffer) < 1 + break + end + end +end + +%Standardize continuous data and map data to levels. +for i=1:nnodes + if bnet.node_sizes(i) == 1 + for j=1:ntestcases + if !strcmp(data_test{j,i},"NA") + data_test{j,i} = (str2num(data_test{j,i}) - means{i})/stdevs{i}; + end + end + else + for j=1:ntestcases + labels{i} + if !strcmp(data_test{j,i},"NA") + for jj = 1:size(levels,1) + if strcmp(labels{i},levels{jj,1}) + break + endif + end + for k = 1:bnet.node_sizes(i) + if strcmp(data_test{j,i},levels{jj,k+1}) + data_test{j,i} = k; + end + end + end + end +end +end + + +%Now make predictions +predict_cases = bnet.node_sizes(predict_node); + +if predict_cases == 1 + ts_continuous(pre,bnet,nnodes,predict_label,predict_node,data_test,means,stdevs) +else + pred_levels = cell(1,predict_cases); + for i = 1:disc_nodes + if strcmp(levels{i,1},predict_label); + for j = 1:predict_cases + pred_levels{j} = levels{i,j+1}; + end + break + end + end + ts_discrete(pre,bnet,nnodes,predict_label,predict_node,data_test,pred_levels,predict_cases) +end + +%delete(dfile) + +end + +function ts_continuous(pre,bnet,nnodes,predict_label,predict_node,data_test,means,stdevs) + +ntestcases = size(data_test,1); + +predictions = zeros(ntestcases,2); + +for i = 1:ntestcases + evidence = data_test(i,:); + evidence{predict_node} = {}; + for j=1:nnodes + if strcmp(evidence{j},"NA") + evidence{j} = {}; + end + end + engine = jtree_inf_engine(bnet); + [engine,loglik] = enter_evidence(engine,evidence); + predict = marginal_nodes(engine,predict_node); + adj_mu = predict.mu*stdevs{predict_node}+means{predict_node}; + adj_sigma = stdevs{predict_node}*predict.Sigma; + predictions(i,1) = adj_mu; + predictions(i,2) = adj_sigma; +end + +%Open output file. +filename = strcat(pre,'ts_output.txt'); +fileID = fopen(filename,'w'); + +fprintf(fileID,'Variable that was predicted: %s\n\n',predict_label); + + +%Calculate RMSEP (root mean square error of prediction) and q^2 +%First, calculate TSS (total sum of squares) and +% PRESS (sum of squares of prediction errors) +%Get rid of 'NA' data for predicted data. +actual_values = []; +predictions_removeNA = []; +for i=1:ntestcases + if !strcmp(data_test(i,predict_node),'NA') + actual_values = [actual_values, cell2num(data_test(i,predict_node))] + predictions_removeNA = [predictions_removeNA,predictions(i,1)] + end +end + +size(actual_values) +size(predictions_removeNA) + +average = mean(actual_values); +average = average*stdevs{predict_node}+means{predict_node}; +tss = 0; +press = 0; +for i=1:length(actual_values) + actual_values(i) = actual_values(i)*stdevs{predict_node}+means{predict_node}; + tss = (actual_values(i)-average)^2 + tss; + press = (predictions_removeNA(i)-actual_values(i))^2 + press; +end +rmsep = sqrt(press/length(actual_values)); +q_squared = 1 - press/tss; + +%% Print rmseq and q^2 +fprintf(fileID,'RMS error of predictions: %6.4f\n',rmsep); +fprintf(fileID,'Q^2 of predictions: %6.4f\n\n',q_squared); + + +%%Print the predictions +fprintf(fileID,'Predicted mean and standard deviation for each case:\n'); +fprintf(fileID,'CaseRow\tActualValue\tPredictionMean\tPredictionStDev\n'); +for i = 1:ntestcases + if strcmp(data_test{i,predict_node},"NA") + fprintf(fileID,'%i\t%s\t',i,data_test{i,predict_node}); + else + temp = data_test{i,predict_node}*stdevs{predict_node}+means{predict_node}; + fprintf(fileID,'%i\t%6.4f\t',i,temp); + end + fprintf(fileID,'%6.4f\t%6.4f\n',predictions(i,:)); +end + + +end + +function ts_discrete(pre,bnet,nnodes,predict_label,predict_node,data_test,pred_levels,predict_cases) + +ntestcases = size(data_test,1); + +predictions=zeros(ntestcases,predict_cases); + +for i = 1:ntestcases + evidence = data_test(i,:); + evidence{predict_node} = {}; + for j=1:nnodes + if strcmp(evidence{j},"NA") + evidence{j} = {}; + end + end + engine = jtree_inf_engine(bnet); + [engine,loglik] = enter_evidence(engine,evidence); + predict = marginal_nodes(engine,predict_node); + for j = 1:predict_cases + predictions(i,j) = predict.T(j); + end +end + + +pred_states = []; +for i=1:ntestcases + if !strcmp(data_test(i,predict_node),'NA') + max_state = 1; + for j = 2:predict_cases + if predictions(i,j) > predictions(i,max_state) + max_state = j; + end + end + pred_states = [pred_states,max_state]; + end +end + + +correct = 0; +j = 0; +for i=1:ntestcases + if !strcmp(data_test(i,predict_node),'NA') + j = j + 1; + if pred_states(j) == cell2mat(data_test(i,predict_node)) + correct = correct + 1; + end + end +end + +accuracy = correct/length(pred_states); + + +%Open output file. +filename = strcat(pre,'ts_output.txt'); +fileID = fopen(filename,'w'); + +fprintf(fileID,'Variable that was predicted: %s\n\n',predict_label); + +%%Print the accuracy +fprintf(fileID,'Fraction of accurate predictions: %6.4f\n\n',accuracy); + +%%Print the predictions +fprintf(fileID,'Predicted likelihood of each state for each case:\n'); +fprintf(fileID,'%s\t%s\t','CaseRow','ActualState'); +fprintf(fileID,'%s\t',pred_levels{1:end-1}); +fprintf(fileID,'%s\n',pred_levels{end}); +for i = 1:ntestcases + fprintf(fileID,'%i\t',i); + if strcmp(data_test{i,predict_node},"NA") + fprintf(fileID,'%s\t',data_test{i,predict_node}); + else + case_level = pred_levels{data_test{i,predict_node}}; + fprintf(fileID,'%s\t',case_level); + end + fprintf(fileID,'%6.4f\t',predictions(i,1:end-1)); + fprintf(fileID,'%6.4f\n',predictions(i,end)); +end + + + + + +end + diff --git a/sourcecodes/parameter_learning/writeParameters.m b/sourcecodes/parameter_learning/writeParameters.m index 42b2a4ef..6efafe8a 100644 --- a/sourcecodes/parameter_learning/writeParameters.m +++ b/sourcecodes/parameter_learning/writeParameters.m @@ -70,7 +70,8 @@ for i = 1:nnodes break end end - predict = marginal_nodes(engine,nodeid); +% predict = marginal_nodes(engine,nodeid); + predict = marginal_nodes_no_ev(bnet,engine,nodeid); %%%Print the name of the node fprintf(fileID,'%s\n',labels{nodeid}); %%%Print the type of node |
