diff options
| author | ziejd2 | 2018-03-14 23:19:16 -0500 |
|---|---|---|
| committer | ziejd2 | 2018-03-14 23:19:16 -0500 |
| commit | c80226899f5cdd9f11c163817d59445213f5bef0 (patch) | |
| tree | e0fd79d2e32fd2aedda2eadaed0f19af3514c520 /BNW_parameter_learning | |
| parent | 324ebc8ddacab8e154047b8518afd5cbb5bb2fa4 (diff) | |
| download | BNW-c80226899f5cdd9f11c163817d59445213f5bef0.tar.gz | |
Separating Octave and php calculations
Diffstat (limited to 'BNW_parameter_learning')
| -rw-r--r-- | BNW_parameter_learning/Predictmultiple.m | 87 | ||||
| -rw-r--r-- | BNW_parameter_learning/Predictmultipleintrvention.m | 76 | ||||
| -rw-r--r-- | BNW_parameter_learning/drawFigure.m | 17 | ||||
| -rw-r--r-- | BNW_parameter_learning/drawFigureM.m | 60 | ||||
| -rw-r--r-- | BNW_parameter_learning/prepareInput.m | 281 | ||||
| -rw-r--r-- | BNW_parameter_learning/readInput.m | 9 | ||||
| -rw-r--r-- | BNW_parameter_learning/runBN_initial.m | 87 | ||||
| -rw-r--r-- | BNW_parameter_learning/standardizeData.m | 2 | ||||
| -rw-r--r-- | BNW_parameter_learning/writeParameters.m | 106 | ||||
| -rw-r--r-- | BNW_parameter_learning/writeParameters_ev.m | 151 | ||||
| -rw-r--r-- | BNW_parameter_learning/writeParameters_int.m | 186 |
11 files changed, 876 insertions, 186 deletions
diff --git a/BNW_parameter_learning/Predictmultiple.m b/BNW_parameter_learning/Predictmultiple.m index c14b516c..781c64a4 100644 --- a/BNW_parameter_learning/Predictmultiple.m +++ b/BNW_parameter_learning/Predictmultiple.m @@ -1,57 +1,72 @@ function Predictmultiple(pre) - dfile=strcat(pre,'structure_input.txt'); sfile=dfile; dfile=strcat(pre,'continuous_input.txt'); nnodefile=strcat(pre,'nnode.txt'); + fnnode = fopen(nnodefile,'r'); nnodes = fscanf(fnnode,'%d'); - - -%nnodes=5; Std_flag=true; [labels,cases,bnet]=readInput(dfile,sfile,nnodes,Std_flag); - -%name -%labels -%map - [bnet]=parameterLearning(bnet,cases); -%[predict_mean,predict_sd,q_sq]=looCrossValid(bnet,cases); -fvarfile=strcat(pre,'var.txt'); -fvar = fopen(fvarfile,'r'); - +fvarfile=strcat(pre,'var.txt'); +fvar = fopen(fvarfile,'r'); select_var_new = fscanf(fvar,'%d'); fvardfile=strcat(pre,'vardata.txt'); - fvard = fopen(fvardfile,'r'); - select_var_data_new = fscanf(fvard,'%f'); +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,4); + for j=1:4 + [next,buffer] = strtok(buffer); + temp{j} = next; + end + labels_orig{i} = temp{1}; + means_orig{i} = str2num(temp{4}); + stdevs_orig{i} = str2num(temp{3}); +end +fclose(fmap); + +%Need to map the means and stdevs to the correct labels +means = cell(1,nnodes); +stdevs = cell(1,nnodes); +%Read in labels in new order. +labelsnew = cell(1,nnodes); +mapdatafile = strcat(pre,'mapdata.txt'); +fmapdata = fopen(mapdatafile,'r'); +buffer = fgetl(fmapdata); +for i = 1:nnodes + [next,buffer ] = strtok(buffer); + labelsnew{i} = next; +end +fclose(fmapdata); +for i = 1:nnodes + for j = 1:nnodes + if strcmp(labelsnew{i},labels_orig{j}) + means{i} = means_orig{j}; + stdevs{i} = stdevs_orig{j}; + break + end + end +end + + filename=strcat(pre,'net_figure_new.txt'); -drawFigureM(nnodes,bnet,labels,filename,cases,select_var_new,select_var_data_new); - -%quit force; -%marginal_nodes(engine,2) -%marginal_nodes(engine,3) -%marginal_nodes(engine,4) -%marginal_nodes(engine,5) -%evidence{1}=2; -%[engine,loglik]=enter_evidence(engine,evidence) -%marginal_nodes(engine,1) -%marginal_nodes(engine,2) -%marginal_nodes(engine,3) -%marginal_nodes(engine,4) -%marginal_nodes(engine,5) -%evidence{2}=0.6; -%evidence{1}=[]; -%[engine,loglik]=enter_evidence(engine,evidence); -%marginal_nodes(engine,3); -%marginal_nodes(engine,4); -%marginal_nodes(engine,5); -end \ No newline at end of file +drawFigureM(nnodes,bnet,labels,filename,cases,stdevs,means,select_var_new,select_var_data_new); + +writeParameters_ev(pre,bnet,nnodes,labels,cases,stdevs,means,select_var_new,select_var_data_new); + +end diff --git a/BNW_parameter_learning/Predictmultipleintrvention.m b/BNW_parameter_learning/Predictmultipleintrvention.m index 4675b846..eaec60dc 100644 --- a/BNW_parameter_learning/Predictmultipleintrvention.m +++ b/BNW_parameter_learning/Predictmultipleintrvention.m @@ -11,17 +11,13 @@ fvarnamefile=strcat(pre,'varname.txt'); varfile = fopen(fvarnamefile,'r'); -%nnodes=5; Std_flag=true; [labels,cases,bnet]=readInput(dfile,sfile,nnodes,Std_flag); -[bnet]=parameterLearning(bnet,cases); - -%[predict_mean,predict_sd,q_sq]=looCrossValid(bnet,cases); +[bnet]=parameterLearning(bnet,cases); fvarfile=strcat(pre,'var.txt'); -fvar = fopen(fvarfile,'r'); - +fvar = fopen(fvarfile,'r'); select_var_new = fscanf(fvar,'%d'); nm = numel(select_var_new); @@ -48,26 +44,52 @@ fvard = fopen(fvardfile,'r'); select_var_data_new = fscanf(fvard,'%f'); +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,4); + for j=1:4 + [next,buffer] = strtok(buffer); + temp{j} = next; + end + labels_orig{i} = temp{1}; + means_orig{i} = str2num(temp{4}); + stdevs_orig{i} = str2num(temp{3}); +end +fclose(fmap); + +%Need to map the means and stdevs to the correct labels +means = cell(1,nnodes); +stdevs = cell(1,nnodes); +%Read in labels in new order. +labelsnew = cell(1,nnodes); +mapdatafile = strcat(pre,'mapdata.txt'); +fmapdata = fopen(mapdatafile,'r'); +buffer = fgetl(fmapdata); +for i = 1:nnodes + [next,buffer ] = strtok(buffer); + labelsnew{i} = next; +end +fclose(fmapdata); +for i = 1:nnodes + for j = 1:nnodes + if strcmp(labelsnew{i},labels_orig{j}) + means{i} = means_orig{j}; + stdevs{i} = stdevs_orig{j}; + break + end + end +end + filename=strcat(pre,'net_figure_new.txt'); -drawFigureM(nnodes,bnet,labels,filename,cases,select_var_new,select_var_data_new); - -%quit force; -%marginal_nodes(engine,2) -%marginal_nodes(engine,3) -%marginal_nodes(engine,4) -%marginal_nodes(engine,5) -%evidence{1}=2; -%[engine,loglik]=enter_evidence(engine,evidence) -%marginal_nodes(engine,1) -%marginal_nodes(engine,2) -%marginal_nodes(engine,3) -%marginal_nodes(engine,4) -%marginal_nodes(engine,5) -%evidence{2}=0.6; -%evidence{1}=[]; -%[engine,loglik]=enter_evidence(engine,evidence); -%marginal_nodes(engine,3); -%marginal_nodes(engine,4); -%marginal_nodes(engine,5); -end \ No newline at end of file +drawFigureM(nnodes,bnet,labels,filename,cases,stdevs,means,select_var_new,select_var_data_new); + +writeParameters_int(pre,bnet,nnodes,labels,cases,stdevs,means,select_var_new,select_var_data_new); + +end diff --git a/BNW_parameter_learning/drawFigure.m b/BNW_parameter_learning/drawFigure.m index a5de0f2d..fa963a4b 100644 --- a/BNW_parameter_learning/drawFigure.m +++ b/BNW_parameter_learning/drawFigure.m @@ -1,18 +1,19 @@ -function [] = drawFigure(nnodes,bnet,labels,filename,cases,selectvar,selectdata) +function [] = drawFigure(nnodes,bnet,labels,filename,cases,stdevs,means,selectvar,selectdata) %drawFigure writes the parameters and data that are needed to draw the %structure of a Bayesian network. -if nargin < 6, - drawFigureNoEv(nnodes,bnet,labels,filename,cases); + +if nargin < 8, + drawFigureNoEv(nnodes,bnet,labels,filename,cases,stdevs,means); else - drawFigureEv(nnodes,bnet,labels,filename,cases,selectvar,selectdata); + drawFigureEv(nnodes,bnet,labels,filename,cases,stdevs,means,selectvar,selectdata); end; end -function [] = drawFigureEv(nnodes,bnet,labels,filename,cases,selectvar,selectdata) +function [] = drawFigureEv(nnodes,bnet,labels,filename,cases,stdevs,means,selectvar,selectdata) %Function to use if there is no entered evidence. % % @@ -152,6 +153,8 @@ for i = 1:nnodes, [x_vals,y_vals] = calcGaussian(predict.mu,predict.Sigma,Amax(i),Amin(i)); %%%For continuous nodes, print x and the pdf of a normal curve. for j = 1:101, + %%Undo standardization + xvals(j,1) = xvals(j,1)*stdevs{i}+means{i} fprintf(fileID,'%6.4f\t%6.4f\n',x_vals(j,1),y_vals(j,1)); end; end; @@ -172,7 +175,7 @@ end -function [] = drawFigureNoEv(nnodes,bnet,labels,filename,cases) +function [] = drawFigureNoEv(nnodes,bnet,labels,filename,cases,stdevs,means) %Function to use if there is no entered evidence. % % @@ -302,6 +305,8 @@ for i = 1:nnodes, [x_vals,y_vals] = calcGaussian(predict.mu,predict.Sigma,Amax(i),Amin(i)); %%%For continuous nodes, print x and the pdf of a normal curve. for j = 1:101, + %%Undo standardization + x_vals(j,1) = x_vals(j,1)*stdevs{i}+means{i}; fprintf(fileID,'%6.4f\t%6.4f\n',x_vals(j,1),y_vals(j,1)); end; end; diff --git a/BNW_parameter_learning/drawFigureM.m b/BNW_parameter_learning/drawFigureM.m index bca223ce..9aa77d35 100644 --- a/BNW_parameter_learning/drawFigureM.m +++ b/BNW_parameter_learning/drawFigureM.m @@ -1,21 +1,6 @@ -function [] = drawFigureM(nnodes,bnet,labels,filename,cases,selectvar,selectdata) -%drawFigure writes the parameters and data that are needed to draw the -%structure of a Bayesian network. - - -%Function to use if there is no entered evidence. -% -% -%Before each printed line, I will have a line that starts with %%% -% that describes what will be on that line - -%Create an empty evidence cell array. - -%val=cases; -%for i = 1:nnodes -% val(i,1)=val(i,2); - -%end +function [] = drawFigureM(nnodes,bnet,labels,filename,cases,stdevs,means,selectvar,selectdata) +%drawFigureM writes the parameters and data that are needed to draw the +%structure of a Bayesian network after added evidence or intervention fileID = fopen(filename,'w'); @@ -30,49 +15,35 @@ engine = jtree_inf_engine(bnet); m = size(selectvar,1); - ev_dat = zeros(1,nnodes); -%For parents, sum down columns for i = 1:m, di=selectvar(i,1); ev_dat(di)=selectdata(i,1); - +%Need to standardized evidence for continuous nodes. + if bnet.node_sizes(di) == 1 + ev_dat(di) = (ev_dat(di) - means{di}) / stdevs{di}; + end evidence{di}=ev_dat(di); fprintf(fileID,'%i\t',di); end -fprintf(fileID,'\n'); - -%ev_dat -% select_var = selectvar(1,1) -% -% select_var_data = selectdata(1,1) -% -% -% evidence{select_var}=select_var_data; +fprintf(fileID,'\n'); [engine,loglik]=enter_evidence(engine,evidence); %Open the file, and write the nodes to a file. - - -%%%%Evidence node - %%% The number of nodes fprintf(fileID,'%i\n',nnodes); %Get canvas size labels_temp = cellstr(labels); [x,y] = make_layout(bnet.dag); - x = x - min(x); y = 1 - y; y = y - min(y); - [x_dim,y_dim] = canvasSize(nnodes,x,y); %%% The dimensions of the canvas for the javascript code -fprintf(fileID,'%i\t%i\t\n',x_dim,y_dim) - +fprintf(fileID,'%i\t%i\t\n',x_dim,y_dim); x = x*x_dim; y = y*y_dim; for i = 1:nnodes, @@ -99,7 +70,6 @@ for i = 1:nnodes, end end - for i = 1:nnodes, %%% The name and type of each node (1=continuous, the number of states %%% if it is discrete @@ -162,23 +132,25 @@ for i = 1:nnodes, fprintf(fileID,'%i\t%6.4f\n',j,predict.T(j)); end; else - [x_vals,y_vals] = calcGaussian(predict.mu,predict.Sigma,Amax(i),Amin(i)); %%%For continuous nodes, print x and the pdf of a normal curve. for j = 1:101, + %%Undo standardization + x_vals(j,1) = x_vals(j,1)*stdevs{i}+means{i}; fprintf(fileID,'%6.4f\t%6.4f\n',x_vals(j,1),y_vals(j,1)); end; end; else - fprintf(fileID,'%6.4f\t%6.4f\n',ev_dat(i),1); + if bnet.node_sizes(i) == 1, + fprintf(fileID,'%6.4f\t%6.4f\n',ev_dat(i)*stdevs{i}+means{i},1); + else + fprintf(fileID,'%6.4f\t%6.4f\n',ev_dat(i),1); + endif end end -%fprintf(fileID,'%s\t %\n',labels_temp{:}); - fclose(fileID); - end diff --git a/BNW_parameter_learning/prepareInput.m b/BNW_parameter_learning/prepareInput.m new file mode 100644 index 00000000..28a06c15 --- /dev/null +++ b/BNW_parameter_learning/prepareInput.m @@ -0,0 +1,281 @@ +function [ ] = prepareInput( pre ) + % + % This function takes files that are uploaded to BNW and creates output + % files that can be used for structure and parameter learning. + % It replaces php code that was previously in bn_file_load_gom.php. + % There are several improvements in performance and ease of use: + % 1) Loading files is significantly (~5x) faster for large input files. + % 2) The allowed values for discrete variables are more flexible. + % (e.g., A genotype variable be 'B' and 'D' instead of having + % to replace to make them '1' and '2'.) + % 3) Continuous variables may be identified as continuous in some cases + % even if there is not a period. + % 4) The states of discrete variables should be correctly ordered in + % almost all cases. + % 5) An additional output file is written that will let users check if + % the input file has been uploaded and parsed correctly. + % 6) Future updates to this code should be easier than updating the php. + % + % + % Input: ???continuous_input_orig.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 names + % followed by the data, with each case in a row. + % + % Output: There are many output files. + % 1) The main output file is ???continuous_input.txt that can be + % used by the structure learning code and parameter learning codes. + % The first line is variable names, the second line is the node type + % (continuous nodes should have 1, discrete nodes have the number + % of states), and the rest is the data. + % 2) A new output file is ???input_desc.txt, a file that describes the + % data so users can check that it has been parsed correctly. + % 3) ???nlevels.txt: The states of discrete variables. + % 4) ???name.txt: The names of the variables as uploaded. + % 5) ???type.txt: The number of states for each variables + % (1 indicates a continuous variable.) + % 6/7) ???nnode.txt and ???nrows.txt: number of nodes and cases + % 8-12) ???ban.txt, ???white.txt, ???k.txt, ???thr.txt, and + % ???parent.txt: Files with default values for structure learning. + % + +% open file for input, include error handling +dfile=strcat(pre,'continuous_input_orig.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) +ncases = fskipl(fin,Inf) - 1; + +frewind(fin); + +% Read in first 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 = cell(1,nnodes); +for j=1:nnodes + [next,buffer] = strtok(buffer); + labels{j} = next; +end + +% Read in the data +data = cell(ncases,nnodes); +for i = 1:ncases + buffer = fgetl(fin); + for j = 1:nnodes + [next,buffer] = strtok(buffer); + data{i,j} = next; + end +end + +% Determine whether or not the nodes are continuous or discrete. +% First, treat them as all discrete and get the states and number of stats(levels). +levels = cell(1,nnodes); +states = []; +for j = 1:nnodes + states{end+1} = unique(data(:,j)); + levels{j} = size(states{j},1); +end + +%Now do some checks to see if nodes are discrete or continuous +for j = 1:nnodes + % If there are 3 or less unique values, I will assume that the node is discrete. + if levels{j} < 4; + continue + % If there are as many unique values as a third of the number of cases, + % I will assume that the node is continuous. + elseif levels{j} > ncases/3; + levels{j} = 1; + % If there are more than twenty unique values, + % I will assume that the node is continuous. + elseif levels{j} > 20; + levels{j} = 1; + % Otherwise, I will scan through the individual values. + % If any of the values contain a '.', I will assume it is continuous. + else + period_test = 0; + column = data(:,j); + k = 1; + while period_test == 0 + period_test = sum(cell2mat(strfind(column(k),"."))); + if period_test != 0; + levels{j} = 1; + end + k++; + if k > ncases + break + end + end + end +end + +%I need to check if any discrete nodes are listed after continuous nodes. +%If so, I need to rearrange the columns. +max_disc = 0; +min_cont = nnodes + 1; +for i = 1:nnodes + if levels{i} > 1 + max_disc = i; + elseif min_cont == nnodes+1 + min_cont = i; + end +end +%If max_disc > min_cont, you need to rearrange the nodes +% to put the discrete nodes first. +if max_disc > min_cont + levels_old = levels; + labels_old = labels; + data_old = data; + states_old = states; + new_order = {}; + for i=1:nnodes + if levels_old{i} > 1 + new_order{end+1} = i; + end + end + for i=1:nnodes + if levels_old{i} == 1 + new_order{end+1} = i; + end + end + labels = {}; + levels = {}; + states = {}; + for i =1:nnodes + labels{i} = labels_old{new_order{i}}; + levels{i} = levels_old{new_order{i}}; + states{i} = states_old{new_order{i}}; + for j=1:ncases + data{j,i} = data_old{j,new_order{i}}; + end + end + +endif + + +%Write other files that are used by BNW for this key. +%The first group of files establish default settings for structure learning. +outfile = strcat(pre,'white.txt'); +fout = fopen(outfile,'w'); +fprintf(fout,'From\tTo\n'); +fclose(fout); + +outfile = strcat(pre,'ban.txt'); +fout = fopen(outfile,'w'); +fprintf(fout,'From\tTo\n'); +fclose(fout); + +outfile = strcat(pre,'k.txt'); +fout = fopen(outfile,'w'); +fprintf(fout,'1\n'); +fclose(fout); + +outfile = strcat(pre,'parent.txt'); +fout = fopen(outfile,'w'); +fprintf(fout,'4\n'); +fclose(fout); + +outfile = strcat(pre,'thr.txt'); +fout = fopen(outfile,'w'); +fprintf(fout,'0.5\n'); +fclose(fout); + + +%The next group of files have information about the uploaded file. +outfile = strcat(pre,'name.txt'); +fout = fopen(outfile,'w'); +fprintf(fout,'%s\t',labels{1:end-1}); +fprintf(fout,'%s\n',labels{end}); +fclose(fout); + +outfile = strcat(pre,'nnode.txt'); +fout = fopen(outfile,'w'); +fprintf(fout,'%i\n',nnodes); +fclose(fout); + +outfile = strcat(pre,'nrows.txt'); +fout = fopen(outfile,'w'); +fprintf(fout,'%i\n',ncases); +fclose(fout); + +outfile = strcat(pre,'type.txt'); +fout = fopen(outfile,'w'); +fprintf(fout,'%s\t',labels{1:end-1}); +fprintf(fout,'%s\n',labels{end}); +fprintf(fout,'%i\t',levels{1:end-1}); +fprintf(fout,'%i\n',levels{end}); +fclose(fout); + +%This output file contains the states for discrete nodes. +% The unique matlab function already sorts the states. +outfile = strcat(pre,'nlevels.txt'); +fout = fopen(outfile,'w'); +for i = 1:nnodes + if levels{i} > 1 + fprintf(fout,'%s\t',labels{i},states{i}{1:end-1}); + fprintf(fout,'%s\n',states{i}{end}); + end +end +fclose(fout); + + +%Print a file with a short description of the input. +descfile = strcat(pre,'input_desc.txt'); +dout = fopen(descfile,'w'); +fprintf(dout,['As loaded, the input file had the following properties:\n\n']); +dout = fopen(descfile,'a'); +fprintf(dout,'There are %i variables and %i cases(rows)\n',size(labels,2),ncases); +fprintf(dout,'The variable names are:\n'); +fprintf(dout,'%s\t',labels{1:end-1}); +fprintf(dout,'%s\n\n',labels{end}); +for i=1:nnodes + if levels{i} == 1 + fprintf(dout,'%s is a continuous variable\n',labels{i}); + column = str2double(data(:,i)); + colmean = mean(column); + colstd = std(column); + fprintf(dout,'It has a mean of %6.3f and a standard deviation of %6.3f\n\n',mean(column),std(column)) + else + fprintf(dout,'%s is a discrete variable with %i states\n',labels{i},levels{i}); + fprintf(dout,'The states are: '); + fprintf(dout,'%s ',states{i}{1:end-1}); + fprintf(dout,'%s\n\n',states{i}{end}); + end +end +fclose(fout); + +outfile = strcat(pre,'continuous_input.txt'); +fout = fopen(outfile,'w'); +fprintf(fout,'%s\t',labels{1:end-1}); +fprintf(fout,'%s\n',labels{end}); +fprintf(fout,'%i\t',levels{1:end-1}); +fprintf(fout,'%i\n',levels{end}); +%Need to replace states in discrete variables with integers for BNT +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} + data{j,i} = sprintf('%i',num2cell(k){1});; + break + end + end + end + end +end +for i = 1:ncases + fprintf(fout,'%s\t',data{i,1:end-1}); + fprintf(fout,'%s\n',data{i,end}); +end +fclose(fout); + + + + + +end +% end of prepareInput.m \ No newline at end of file diff --git a/BNW_parameter_learning/readInput.m b/BNW_parameter_learning/readInput.m index c2c16331..2be0af29 100644 --- a/BNW_parameter_learning/readInput.m +++ b/BNW_parameter_learning/readInput.m @@ -25,19 +25,12 @@ end [labelsold,node_sizes,cases, data] = readInputData(dfile,nnodes); - - % read in the file with the structure [dag] = readInputStructure(sfile,labelsold); % check the ordering of the nodes and reorder if necessary [labels,cases,dag,node_sizes,ord_flag] = checkStructure(labelsold,cases,dag,node_sizes); -%draw_graph(dag,labels); -%if ord_flag == 1 -% fprintf(['Order of nodes was changed to agree with topological order\n']) -%end -%fprintf(['The structure of the network should be correctly displayed in a figure\n']) dcount = 0; for i = 1:nnodes @@ -67,4 +60,4 @@ end end -% end of readInput.m \ No newline at end of file +% end of readInput.m diff --git a/BNW_parameter_learning/runBN_initial.m b/BNW_parameter_learning/runBN_initial.m index c3f2a34b..43caf402 100644 --- a/BNW_parameter_learning/runBN_initial.m +++ b/BNW_parameter_learning/runBN_initial.m @@ -15,84 +15,43 @@ mapfile = fopen(mapfilename,'w'); mapval = fopen(mapvalfilename,'w'); -%nnodes=5; Std_flag=true; [labels,cases,bnet,node_sizes,data,labelsold]=readInput(dfile,sfile,nnodes,Std_flag); -s = std(data,0,1); +s=std(data,0,1); m=mean(data); - for i=1:nnodes fprintf(mapval,'%s\t%d\t%f\t%f\n',labelsold{i},node_sizes(i),s(i),m(i)); end - -% for j=1:nnodes -% [next,buffer] = strtok(buffer); -% name{j}=next; -% for i=1:nnodes -% if strcmp(name{j},labels{i}) -% map{j}=i; -% fprintf(mapfile,'%d\t',i); -% end -% end -% end -%name -%labels -%map fprintf(mapfile,'%s',labels{1}); for i=2:nnodes fprintf(mapfile,'\t%s',labels{i}); end fprintf(mapfile,'\n'); +fclose(mapval); +fclose(mapfile); + +%Need to rearrange the means and stdevs to match the new labeling. +means = cell(1,nnodes); +stdevs = cell(1,nnodes); +for i = 1:nnodes + for j = 1:nnodes + if strcmp(labels{i},labelsold{j}) + means{i} = m(j); + stdevs{i} = s(j); + break + end + end +end + [bnet]=parameterLearning(bnet,cases); -%[predict_mean,predict_sd,q_sq]=looCrossValid(bnet,cases); -%engine=jtree_inf_engine(bnet); -%evidence=cell(1,nnodes); - -%varfile='var.txt'; -%fvar = fopen(varfile,'r'); -%select_var = fscanf(fvar,'%d'); -%select_var=map{select_var}; -%varfiled='vardata.txt'; -%fvard = fopen(varfiled,'r'); -%select_var_data = fscanf(fvard,'%f'); - -%evidence{select_var}=select_var_data; -%[engine,loglik]=enter_evidence(engine,evidence); - -%outdata='prediction.txt'; -%fout = fopen(outdata,'w'); - -%for ii = 1:nnodes - % i=map{ii}; - % data=marginal_nodes(engine,i); - % fprintf(fout,'%d\t%d\t%f\t%f\t%f\n',ii,data.domain,data.T,data.mu,data.Sigma); - %fprintf(1,'%d\n',i); -% end filename=strcat(pre,'net_figure.txt'); -drawFigure(nnodes,bnet,labels,filename,cases); - -%quit force; -%marginal_nodes(engine,2) -%marginal_nodes(engine,3) -%marginal_nodes(engine,4) -%marginal_nodes(engine,5) -%evidence{1}=2; -%[engine,loglik]=enter_evidence(engine,evidence) -%marginal_nodes(engine,1) -%marginal_nodes(engine,2) -%marginal_nodes(engine,3) -%marginal_nodes(engine,4) -%marginal_nodes(engine,5) -%evidence{2}=0.6; -%evidence{1}=[]; -%[engine,loglik]=enter_evidence(engine,evidence); -%marginal_nodes(engine,3); -%marginal_nodes(engine,4); -%marginal_nodes(engine,5); -fclose(mapval); -fclose(mapfile); -end \ No newline at end of file + +drawFigure(nnodes,bnet,labels,filename,cases,stdevs,means); + +writeParameters(pre,nnodes,bnet,labels,cases,labelsold,s,m); + +end diff --git a/BNW_parameter_learning/standardizeData.m b/BNW_parameter_learning/standardizeData.m index ba4ef5dd..449aa440 100644 --- a/BNW_parameter_learning/standardizeData.m +++ b/BNW_parameter_learning/standardizeData.m @@ -1,7 +1,7 @@ function [ cases ] = standardizeData( labels, node_sizes, cases ) %standardizeData standardizes continuous nodes so they have a mean = 0 % and standard deviation = 1 -% Detailed explanation goes here + nnodes = size(labels,2); diff --git a/BNW_parameter_learning/writeParameters.m b/BNW_parameter_learning/writeParameters.m new file mode 100644 index 00000000..0790a8e2 --- /dev/null +++ b/BNW_parameter_learning/writeParameters.m @@ -0,0 +1,106 @@ +function [] = writeParameters(pre,nnodes,bnet,labels,cases,labelsold,s,m) +%Writes a file that contains the parameters of the network with no evidence. + + +%%Get the types of the nodes. +typefile = strcat(pre,'type.txt'); +ftype = fopen(typefile,'r'); +types = cell(1,nnodes); +buffer = fgetl(ftype); +buffer = fgetl(ftype); +for j = 1:nnodes + [next,buffer] = strtok(buffer); + types{j} = uint16(str2num(next)); +end + +max_states = 0; +disc_nodes = 0; +for j = 1:nnodes + if types{j} > max_states + max_states = types{j}; + end + if types{j} > 1 + disc_nodes = disc_nodes + 1; + end +end + +%Add 1 to max_states to account for node name +max_states = max_states + 1; + +%%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} = uint16(str2num(next)); + levels{i,j} = next; + end + if length(buffer) < 1 + break + end + end +end + + +evidence = cell(1,nnodes); +engine = jtree_inf_engine(bnet); +[engine,loglik] = enter_evidence(engine,evidence); + +%Open output file. +filename = strcat(pre,'parameters.txt'); +fileID = fopen(filename,'w'); + +for i = 1:nnodes + for j = 1:nnodes + if strcmp(labelsold{i},labels{j}); + nodeid = j; + break + end + end + predict = marginal_nodes(engine,nodeid); + %%%Print the name of the node + fprintf(fileID,'%s\n',labels{nodeid}); + %%%Print the type of node + if bnet.node_sizes(nodeid) == 1; + line = 'Continuous node\n'; + fprintf(fileID,line); + %%% 'i' in the line below is correct: m and s are had original node labeling + adj_mu = predict.mu*s(i)+m(i); + adj_sigma = s(i)*predict.Sigma; + fprintf(fileID,'%6.4f\t%6.4f\n\n',adj_mu,adj_sigma); + else + line = 'Discrete node with %i states\n'; + fprintf(fileID,line,bnet.node_sizes(nodeid)); + %line = 'Probability of each state\n'; + %fprintf(fileID,line); + nodeid2 = 0; + for k = 1:ndisc_nodes, + if strcmp(levels{k,1},labels{nodeid}), + nodeid2 = k; + break + end + end + for j = 1:bnet.node_sizes(nodeid), + %%%For discrete nodes, the state and the percent of that state +% fprintf(fileID,'%i\t%6.4f\n',levels{nodeid2,j+1},predict.T(j)); + fprintf(fileID,'%s\t%6.4f\n',levels{nodeid2,j+1},predict.T(j)); + end; + fprintf(fileID,'\n') + + end +end + + + +fclose(fileID); + +end + diff --git a/BNW_parameter_learning/writeParameters_ev.m b/BNW_parameter_learning/writeParameters_ev.m new file mode 100644 index 00000000..fc24e2e5 --- /dev/null +++ b/BNW_parameter_learning/writeParameters_ev.m @@ -0,0 +1,151 @@ +function [] = writeParameters_ev(pre,bnet,nnodes,labels,cases,stdevs,means,selectvar,selectdata) +%Writes a file that contains the parameters of the network after entering evidence. + +%Read in original node labels to get node IDs. +infile = strcat(pre,'continuous_input.txt'); +fin = fopen(infile,'r'); +labelsold = cell(1,nnodes); +buffer = fgetl(fin); +for j = 1:nnodes + [next,buffer] = strtok(buffer); + labelsold{j} = next; +end +fclose(fin); + + +evidence = cell(1,nnodes); +engine = jtree_inf_engine(bnet); + +m = size(selectvar,1); + +%%Get the types of the nodes. +typefile = strcat(pre,'type.txt'); +ftype = fopen(typefile,'r'); +types = cell(1,nnodes); +buffer = fgetl(ftype); +buffer = fgetl(ftype); +for j = 1:nnodes + [next,buffer] = strtok(buffer); + types{j} = uint16(str2num(next)); +end + +max_states = 0; +disc_nodes = 0; +for j = 1:nnodes + if types{j} > max_states + max_states = types{j}; + end + if types{j} > 1 + disc_nodes = disc_nodes + 1; + end +end + +%Add 1 to max_states to account for node name +max_states = max_states + 1; + +%%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} = uint16(str2num(next)); + levels{i,j} = next; + end + if length(buffer) < 1 + break + end + end +end + + +ev_dat = zeros(1,nnodes); +for i = 1:m, + di=selectvar(i,1); + ev_dat(di)=selectdata(i,1); +%Need to standardize evidence for continuous nodes. + if bnet.node_sizes(di) == 1, + ev_dat(di) = (ev_dat(di) - means{di})/stdevs{di}; + end + evidence{di} = ev_dat(di); +end + +[engine,loglik]=enter_evidence(engine,evidence); + +%Open output file. +filename = strcat(pre,'parameters_ev.txt'); +fileID = fopen(filename,'w'); + +for i = 1:nnodes + for j = 1:nnodes + if strcmp(labelsold{i},labels{j}); + nodeid = j; + break + end + end + %%%Print the name of the node + fprintf(fileID,'%s\n',labels{nodeid}); + predict = marginal_nodes(engine,nodeid); + if isempty(evidence{nodeid}) + %%%Print the type of node + if bnet.node_sizes(nodeid) == 1; + line = 'Continuous parameters considering evidence:\n'; + fprintf(fileID,line); + %line = 'Mean and standard deviation of Gaussian distribution\n'; + %fprintf(fileID,line); + adj_mu = predict.mu*stdevs{nodeid}+means{nodeid}; + adj_sigma = stdevs{nodeid}*predict.Sigma; + fprintf(fileID,'%6.4f\t%6.4f\n\n',adj_mu,adj_sigma); + else + line = 'Probability of states considering evidence:\n'; + fprintf(fileID,line); + nodeid2 = 0; + for k = 1:ndisc_nodes, + if strcmp(levels{k,1},labels{nodeid}), + nodeid2 = k; + break + end + end + for j = 1:bnet.node_sizes(nodeid), + %%%For discrete nodes, the state and the percent of that state +% fprintf(fileID,'%i\t%6.4f\n',levels{nodeid2,j+1},predict.T(j)); + fprintf(fileID,'%s\t%6.4f\n',levels{nodeid2,j+1},predict.T(j)); + end; + fprintf(fileID,'\n') + end + else + if bnet.node_sizes(nodeid) == 1; + line = 'Evidence was observed for this node. The observed value was:\n'; + fprintf(fileID,line); + adj_mu = ev_dat(nodeid)*stdevs{nodeid}+means{nodeid}; + fprintf(fileID,'%6.4f\n\n',adj_mu); + else + nodeid2 = 0; + for k = 1:ndisc_nodes, + if strcmp(levels{k,1},labels{nodeid}), + nodeid2 = k; + break + end + end + line = 'Evidence was observed for this node. The observed state was:\n'; + fprintf(fileID,line); + state_ev = uint16(ev_dat(nodeid)); +% fprintf(fileID,'%i\n\n',levels{nodeid2,state_ev+1}); + fprintf(fileID,'%s\n\n',levels{nodeid2,state_ev+1}); + end + end +end + + + +fclose(fileID); + +end + diff --git a/BNW_parameter_learning/writeParameters_int.m b/BNW_parameter_learning/writeParameters_int.m new file mode 100644 index 00000000..ed92d593 --- /dev/null +++ b/BNW_parameter_learning/writeParameters_int.m @@ -0,0 +1,186 @@ +function [] = writeParameters_int(pre,bnet,nnodes,labels,cases,stdevs,means,selectvar,selectdata) +%Writes a file that contains the parameters of the network after intervention. + + +%First read input file to get node labels to get node IDs. +infile = strcat(pre,'continuous_input.txt'); +fin = fopen(infile,'r'); +labelsold = cell(1,nnodes); +buffer = fgetl(fin); +for j = 1:nnodes + [next,buffer] = strtok(buffer); + labelsold{j} = next; +end + +evidence = cell(1,nnodes); +engine = jtree_inf_engine(bnet); + +m = size(selectvar,1); + +%%Get the types of the nodes. +typefile = strcat(pre,'type.txt'); +ftype = fopen(typefile,'r'); +types = cell(1,nnodes); +buffer = fgetl(ftype); +buffer = fgetl(ftype); +for j = 1:nnodes + [next,buffer] = strtok(buffer); + types{j} = uint16(str2num(next)); +end + +max_states = 0; +disc_nodes = 0; +for j = 1:nnodes + if types{j} > max_states + max_states = types{j}; + end + if types{j} > 1 + disc_nodes = disc_nodes + 1; + end +end + +%Add 1 to max_states to account for node name +max_states = max_states + 1; + +%%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} = uint16(str2num(next)); + levels{i,j} = next; + end + if length(buffer) < 1 + break + end + end +end + + +ev_dat = zeros(1,nnodes); +for i = 1:m, + di=selectvar(i,1); + ev_dat(di)=selectdata(i,1); +%Need to standardize evidence for continuous nodes. + if bnet.node_sizes(di) == 1, + ev_dat(di) = (ev_dat(di) - means{di})/stdevs{di}; + end + evidence{di} = ev_dat(di); +end + +[engine,loglik]=enter_evidence(engine,evidence); + +%Get list of nodes that are children, grandchildren, etc. of intervened nodes +%int_nodes contains the list of these children nodes +int_nodes = zeros(1,nnodes); +%new_nodes is just a temporary array to know when to keep looking +new_nodes = zeros(1,nnodes); +for i = 1:nnodes + if !isempty(evidence{i}); + new_nodes(i) = 1; + int_nodes(i) = 1; + end +end +while sum(new_nodes) != 0 + new_nodes_old = new_nodes; + new_nodes = zeros(1,nnodes); + for i = 1:nnodes + if new_nodes_old(i) == 1 + for j = 1:nnodes + if int_nodes(j) == 0 + if bnet.dag(i,j) == 1, + new_nodes(j) = 1; + end + end + end + end + end + for i = 1:nnodes + if new_nodes(i) == 1; + int_nodes(i) = 1; + end + end +end + + +%Open output file. +filename = strcat(pre,'parameters_ev.txt'); +fileID = fopen(filename,'w'); + +for i = 1:nnodes + for j = 1:nnodes + if strcmp(labelsold{i},labels{j}); + nodeid = j; + break + end + end + %check to see if this is a node impacted by intervention + if int_nodes(nodeid) == 1 + %%%Print the name of the node + fprintf(fileID,'%s\n',labels{nodeid}); + predict = marginal_nodes(engine,nodeid); + if isempty(evidence{nodeid}) + %%%Print the type of node + if bnet.node_sizes(nodeid) == 1; + line = 'Continuous parameters considering intervention:\n'; + fprintf(fileID,line); + %line = 'Mean and standard deviation of Gaussian distribution\n'; + %fprintf(fileID,line); + adj_mu = predict.mu*stdevs{nodeid}+means{nodeid}; + adj_sigma = stdevs{nodeid}*predict.Sigma; + fprintf(fileID,'%6.4f\t%6.4f\n\n',adj_mu,adj_sigma); + else + line = 'Probability of states considering intervention:\n'; + fprintf(fileID,line); + nodeid2 = 0; + for k = 1:ndisc_nodes, + if strcmp(levels{k,1},labels{nodeid}), + nodeid2 = k; + break + end + end + for j = 1:bnet.node_sizes(nodeid), + %%%For discrete nodes, the state and the percent of that state +% fprintf(fileID,'%i\t%6.4f\n',levels{nodeid2,j+1},predict.T(j)); + fprintf(fileID,'%s\t%6.4f\n',levels{nodeid2,j+1},predict.T(j)); + end; + fprintf(fileID,'\n') + end + else + if bnet.node_sizes(nodeid) == 1; + line = 'Intervention on this node assigned the following value:\n'; + fprintf(fileID,line); + adj_mu = ev_dat(nodeid)*stdevs{nodeid}+means{nodeid}; + fprintf(fileID,'%6.4f\n\n',adj_mu); + else + nodeid2 = 0; + for k = 1:ndisc_nodes, + if strcmp(levels{k,1},labels{nodeid}), + nodeid2 = k; + break + end + end + line = 'Intervention on this node assigned the following state:\n'; + fprintf(fileID,line); + state_ev = uint16(ev_dat(nodeid)); +% fprintf(fileID,'%i\n\n',levels{nodeid2,state_ev+1}); + fprintf(fileID,'%s\n\n',levels{nodeid2,state_ev+1}); + end + end + end +end + + + +fclose(fileID); + +end + |
