about summary refs log tree commit diff
path: root/sourcecodes/bnt-master/BNT/examples/dynamic/HHMM/Square/Old/sample_square_hhmm.m
blob: a0f9007e54071c2552c7bd6ea6037e80f1e6c479 (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
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
seed = 0;
rand('state', seed);
randn('state', seed);

discrete_obs = 1;
topright = 0;

Qsizes = [2 4 2];
D = 3;
Qnodes = 1:D;
startprob = cell(1,D);
transprob = cell(1,D);
termprob = cell(1,D);

% LEVEL 1

startprob{1} = 'ergodic';
transprob{1} = 'ergodic';

% LEVEL 2

startprob{2} = zeros(2, 4);
startprob{2}(1, :) = [1 0 0 0];
if topright
  startprob{2}(2, :) = [0 0 1 0];
else
  startprob{2}(2, :) = [0 1 0 0];
end

transprob{2} = zeros(4, 2, 4);

transprob{2}(:,1,:) = [0 1 0 0
		       0 0 1 0
		       0 0 0 1
		       0 0 0 1]; % 4->e
if topright
  transprob{2}(:,2,:) = [0 0 0 1
		    1 0 0 0
		    0 1 0 0
		    0 0 0 1]; % 4->e
else
  transprob{2}(:,2,:) = [0 0 0 1
		    1 0 0 0
		    0 0 1 0 % 3->e
		    0 0 1 0];
end

%termprob{2} = 'rightstop';
termprob{2} = zeros(2,4,2);
pfin = 0.8;
termprob{2}(1,:,2) = [0 0 0 pfin]; % finish in state 4 (DU)
termprob{2}(1,:,1) = 1 - [0 0 0 pfin];
if topright
  termprob{2}(2,:,2) = [0 0 0 pfin];
  termprob{2}(2,:,1) = 1 - [0 0 0 pfin];
else
  termprob{2}(2,:,2) = [0 0 pfin 0];  % finish in state 3 (RL)
  termprob{2}(2,:,1) = 1 - [0 0 pfin 0];
end

% LEVEL 3

startprob{3} = 'leftstart';
transprob{3}  = 'leftright';
termprob{3} = 'rightstop';


% OBS LEVEl

if discrete_obs
  chars = ['L', 'l', 'U', 'u', 'R', 'r', 'D', 'd'];
  L=find(chars=='L'); l=find(chars=='l');
  U=find(chars=='U'); u=find(chars=='u');
  R=find(chars=='R'); r=find(chars=='r');
  D=find(chars=='D'); d=find(chars=='d');
  Osize = length(chars);
  
  obsprob = zeros([4 2 Osize]);
  %       Q2 Q3 O
  obsprob(1, 1, L) =  1.0;
  obsprob(1, 2, l) =  1.0;
  obsprob(2, 1, U) =  1.0;
  obsprob(2, 2, u) =  1.0;
  obsprob(3, 1, R) =  1.0;
  obsprob(3, 2, r) =  1.0;
  obsprob(4, 1, D) =  1.0;
  obsprob(4, 2, d) =  1.0;
  
  Oargs = {'CPT', obsprob};
else
  Osize = 2;
  mu = zeros(2, 4, 2);
  noise = 0;
  scale = 10;
  for q3=1:2
    mu(:, 1, q3) = scale*[1;0] + noise*rand(2,1);
  end
  for q3=1:2
    mu(:, 2, q3) = scale*[0;-1] + noise*rand(2,1);
  end
  for q3=1:2
    mu(:, 3, q3) = scale*[-1;0] + noise*rand(2,1);
  end
  for q3=1:2
    mu(:, 4, q3) = scale*[0;1] + noise*rand(2,1);
  end
  Sigma = repmat(reshape(0.01*eye(2), [2 2 1 1 ]), [1 1 4 2]);
  Oargs = {'mean', mu, 'cov', Sigma};
end

bnet = mk_hhmm('Qsizes', Qsizes, 'Osize', Osize', 'discrete_obs', discrete_obs, ...
	       'Oargs', Oargs, 'Ops', Qnodes(2:3), ...
	       'startprob', startprob, 'transprob', transprob, 'termprob', termprob);

if discrete_obs
  Tmax = 30;
else
  Tmax = 200;
end
usecell = ~discrete_obs;
Q1 = 1; Q2 = 2; Q3 = 3; F3 = 4; F2 = 5; Onode = 6;
Qnodes = [Q1 Q2 Q3]; Fnodes = [F2 F3];

for seqi=1:3
  evidence = sample_dbn(bnet, Tmax, usecell, 'stop_sampling_F2');      
  T = size(evidence, 2)
  if discrete_obs
    pretty_print_hhmm_parse(evidence, Qnodes, Fnodes, Onode, chars);
  else
    pos = zeros(2,T+1);
    delta = cell2num(evidence(Onode,:));
    clf
    hold on
    cols = {'r', 'g', 'k', 'b'};
    boundary = cell2num(evidence(F3,:))-1;
    coli = 1;
    for t=2:T+1
      pos(:,t) = pos(:,t-1) + delta(:,t-1);
      plot(pos(1,t), pos(2,t), sprintf('%c.', cols{coli}));
      if boundary(t-1)
	coli = coli + 1;
	coli = mod(coli-1, length(cols)) + 1;
      end
    end
    %plot(pos(1,:), pos(2,:), '.')
    %pretty_print_hhmm_parse(evidence, Qnodes, Fnodes, Onode, []);
    pause
  end
end

eclass = bnet.equiv_class;
S=struct(bnet.CPD{eclass(Q2,2)});