diff options
| author | ziejd2 | 2018-09-13 23:59:20 -0500 |
|---|---|---|
| committer | ziejd2 | 2018-09-13 23:59:20 -0500 |
| commit | e3f7237ffcb19f19db3b68777b5a94b89e07f66a (patch) | |
| tree | 554a8013776ebeae3e2976074020c09c2d1af8b0 /sourcecodes/parameter_learning | |
| parent | a7eb61ff7a09f39bee67014bf24b8919eaccfc19 (diff) | |
| download | BNW-e3f7237ffcb19f19db3b68777b5a94b89e07f66a.tar.gz | |
New parameter learning options
The main change here is in the parameter learning methods. The parameters that are learned at first (i.e., if there is no evidence) are the distributions that are found directly in the data. I had to create or significantly modify several BNT files for this. If there is evidence, the parameters are learned using a Dirichlet prior. This only required a couple of small changes to the BNW parameter learning files.
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 |
