about summary refs log tree commit diff
path: root/sourcecodes/bnt-master/SLP/learning/learn_struct_gs.m
blob: 3dd3ef1c144e624c1aa2cc562b030f663ae9de39 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
function [dag,best_score] = learn_struct_gs(data, nodesizes, seeddag, varargin)
%
% LEARN_STRUCT_GS(data,seeddag) learns a structure of Bayesian net by Greedy Search.
% dag = learn_struct_gs(data, nodesizes, seeddag)
%
% dag: the final structurre matrix
% Data : training data, data(i,m) is the m obsevation of node i
% Nodesizes: the size array of different nodes
% seeddag: given seed Dag for hill climbing, optional
%
%
% by Gang Li @ Deakin University (gli73@hotmail.com)

[N ncases] = size(data);
if (nargin < 3 ) 
    seeddag = zeros(N,N); % mk_rnd_dag(N); %call BNT function
elseif ~acyclic(seeddag)
    seeddag = mk_rnd_dag(N); %zeros(N,N);
end;

% set default params
scoring_fn = 'bic';
verbose  = 'yes';

% get params
args = varargin;
nargs = length(args);
if length(args) > 0
    if isstr(args{1})
    	for i = 1:2:nargs
    		switch args{i}
    		case 'scoring_fn', scoring_fn = args{i+1};
    		case 'verbose',  verbose  = strcmp(args{i+1},'yes');
    		end;
    	end;
    end;
end;

done = 0;
best_score = score_dags(data,nodesizes, {seeddag},'scoring_fn',scoring_fn);
while ~done
    [dags,op,nodes] = mk_nbrs_of_dag(seeddag);
    nbrs = length(dags);
    scores = score_dags(data, nodesizes, dags,'scoring_fn',scoring_fn);
    max_score = max(scores);
    new = find(scores == max_score );
    if ~isempty(new) & (max_score > best_score)
        p = sample_discrete(normalise(ones(1, length(new))));
        best_score = max_score;
        seeddag = dags{new(p)};
    else
        done = 1;
    end;
end;

dag = seeddag;

outcount = 0; 
best_score = score_dags(data,nodesizes, {seeddag},'scoring_fn',scoring_fn);
while outcount < 2
    innercount = 0;
    for i=1:N
        for j=1:N
           if i==j, continue;    end;
           if seeddag(i,j) == 0  % No edge i-->j, then try to add it
               tempdag = seeddag;
               tempdag(i,j) = 1;
               if acyclic(tempdag)
                    temp_score = score_dags(data,nodesizes, {tempdag},'scoring_fn',scoring_fn);
                    if temp_score > best_score
                        seeddag = tempdag;
                        best_score= temp_score;
                        innercount = innercount +1;
                    end;
               end
           else  % exists edge i--j, then try reverse it or remove it
               tempdag = seeddag;
               tempdag(i,j) = 0; tempdag(j,i) = 1; 
               if acyclic(tempdag)
                   temp_score = score_dags(data,nodesizes, {tempdag},'scoring_fn',scoring_fn);
                   if temp_score > best_score
                       seeddag = tempdag;
                       best_score = temp_score;
                       innercount = innercount +1;
                   else
                       tempdag = seeddag;
                       tempdag(i,j) = 0;
                       temp_score = score_dags(data,nodesizes, {tempdag},'scoring_fn',scoring_fn);
                       if temp_score > best_score
                           seeddag = tempdag;
                           best_score= temp_score;
                           innercount = innercount +1;
                       end;
                   end;
               else
                   tempdag = seeddag;
                   tempdag(i,j)=0;
                   temp_score = score_dags(data,nodesizes, {tempdag},'scoring_fn',scoring_fn);
                   if temp_score > best_score
                       seeddag = tempdag;
                       best_score= temp_score;
                       innercount = innercount +1;
                   end;
               end;
           end;
        end; % end for j
    end; % end for i
    if innercount == 0
        outcount = outcount +1;
    end;
end;  % end while

dag = seeddag;