about summary refs log tree commit diff
path: root/sourcecodes/bnt-master/SLP/scoring/kl_divergence2.m
diff options
context:
space:
mode:
authorziejd22018-03-14 23:23:33 -0500
committerGitHub2018-03-14 23:23:33 -0500
commit1ff6baa44e22b91eefb48aea6f3befa078c0489b (patch)
treee0fd79d2e32fd2aedda2eadaed0f19af3514c520 /sourcecodes/bnt-master/SLP/scoring/kl_divergence2.m
parent6882395afdadf4e982b25b5215071a0932730950 (diff)
parentc80226899f5cdd9f11c163817d59445213f5bef0 (diff)
downloadBNW-1ff6baa44e22b91eefb48aea6f3befa078c0489b.tar.gz
Merge pull request #1 from ziejd2/octave_php_separate
Octave php separate
Diffstat (limited to 'sourcecodes/bnt-master/SLP/scoring/kl_divergence2.m')
-rw-r--r--sourcecodes/bnt-master/SLP/scoring/kl_divergence2.m43
1 files changed, 43 insertions, 0 deletions
diff --git a/sourcecodes/bnt-master/SLP/scoring/kl_divergence2.m b/sourcecodes/bnt-master/SLP/scoring/kl_divergence2.m
new file mode 100644
index 00000000..7b9664ba
--- /dev/null
+++ b/sourcecodes/bnt-master/SLP/scoring/kl_divergence2.m
@@ -0,0 +1,43 @@
+function KLdiv = KL_divergence2(bnetP, bnetQ)
+% KL_DIVERGENCE2 computes the Kullback-Leibler divergence between two BNET distributions
+% KLdiv = KL_divergence2(bnetP, bnetQ)
+%
+% Output :
+%   div = sum_x  P(x).log(P(x)/Q(x))
+%
+% Rem : 
+%   This version is optimized for memory use, but quite slow !!!
+%     ==> if you have no memory problem, use kl_divergence instead
+%
+%   ONLY FOR TABULAR NODES
+%   Make sure that you have done the params learning.
+%
+%   V1.1 : 8 oct 2004 (Ph. Leray - philippe.leray@univ-nantes.fr)
+
+N = size(bnetP.dag,1);
+N2 = size(bnetQ.dag,1);
+ns= bnetP.node_sizes;
+ns2= bnetQ.node_sizes;
+if N~=N2, error('size of dags must be the same'), end
+if ns~=ns2, error('node sizes of dags must be the same'), end
+tiny = exp(-700);
+KLdiv=0;
+
+for i=1:prod(ns),
+  inst = ind2subv(ns, i); % i'th instantiation
+  Px=1; Qx=1;
+  for i=1:N,
+    ps = parents(bnetP.dag, i);
+    e = bnetP.equiv_class(i);
+    [tmp Pxi] = prob_node(bnetP.CPD{e}, inst(i), inst(ps)');
+    Px=Px*Pxi;
+    ps = parents(bnetQ.dag, i);
+    e = bnetQ.equiv_class(i);
+    [tmp Qxi] = prob_node(bnetQ.CPD{e}, inst(i), inst(ps)');
+    Qx=Qx*Qxi;
+  end
+    Px = Px + (Px==0)*tiny; % replace 0s by tiny
+    Qx = Qx + (Qx==0)*tiny; % replace 0s by tiny
+  KLdiv = KLdiv + Px*log(Px/Qx);
+end
+