about summary refs log tree commit diff
path: root/sourcecodes/parameter_learning
diff options
context:
space:
mode:
Diffstat (limited to 'sourcecodes/parameter_learning')
-rw-r--r--sourcecodes/parameter_learning/drawFigure.m3
-rw-r--r--sourcecodes/parameter_learning/kfoldCrossValid.m322
-rw-r--r--sourcecodes/parameter_learning/looCrossValid.m250
-rw-r--r--sourcecodes/parameter_learning/parameterLearning.m5
-rw-r--r--sourcecodes/parameter_learning/prepareInput.m2
-rw-r--r--sourcecodes/parameter_learning/testSetPredictions.m361
-rw-r--r--sourcecodes/parameter_learning/writeParameters.m3
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