Fixed NAN related bugs for Visual C++.
[libdai.git] / src / regiongraph.cpp
index 8449745..8078279 100644 (file)
@@ -71,8 +71,8 @@ RegionGraph::RegionGraph( const FactorGraph &fg, const std::vector<VarSet> &cl )
     
     // Create outer regions, giving them counting number 1.0
     ORs.reserve( cg.size() );
-    for( ClusterGraph::const_iterator alpha = cg.begin(); alpha != cg.end(); alpha++ )
-        ORs.push_back( FRegion(Factor(*alpha, 1.0), 1.0) );
+    foreach( const VarSet &ns, cg.clusters )
+        ORs.push_back( FRegion(Factor(ns, 1.0), 1.0) );
 
     // For each factor, find an outer regions that subsumes that factor.
     // Then, multiply the outer region with that factor.
@@ -91,9 +91,9 @@ RegionGraph::RegionGraph( const FactorGraph &fg, const std::vector<VarSet> &cl )
     
     // Create inner regions - first pass
     set<VarSet> betas;
-    for( ClusterGraph::const_iterator alpha = cg.begin(); alpha != cg.end(); alpha++ )
-        for( ClusterGraph::const_iterator alpha2 = alpha; (++alpha2) != cg.end(); ) {
-            VarSet intersect = (*alpha) & (*alpha2);
+    for( size_t alpha = 0; alpha < cg.clusters.size(); alpha++ )
+        for( size_t alpha2 = alpha; (++alpha2) != cg.clusters.size(); ) {
+            VarSet intersect = cg.clusters[alpha] & cg.clusters[alpha2];
             if( intersect.size() > 0 )
                 betas.insert( intersect );
         }
@@ -114,7 +114,7 @@ RegionGraph::RegionGraph( const FactorGraph &fg, const std::vector<VarSet> &cl )
     // Create inner regions - store them in the bipartite graph
     IRs.reserve( betas.size() );
     for( set<VarSet>::const_iterator beta = betas.begin(); beta != betas.end(); beta++ )
-        IRs.push_back( Region(*beta,NAN) );
+        IRs.push_back( Region(*beta,0.0) );
     
     // Create edges
     vector<pair<size_t,size_t> > edges;
@@ -142,8 +142,9 @@ void RegionGraph::Calc_Counting_Numbers() {
     // Calculates counting numbers of inner regions based upon counting numbers of outer regions
     
     vector<vector<size_t> > ancestors(nrIRs());
+    vector<bool> assigned(nrIRs(), false);
     for( size_t beta = 0; beta < nrIRs(); beta++ ) {
-        IR(beta).c() = NAN;
+        IR(beta).c() = 0.0;
         for( size_t beta2 = 0; beta2 < nrIRs(); beta2++ )
             if( (beta2 != beta) && IR(beta2) >> IR(beta) )
                 ancestors[beta].push_back(beta2);
@@ -153,18 +154,19 @@ void RegionGraph::Calc_Counting_Numbers() {
     do {
         new_counting = false;
         for( size_t beta = 0; beta < nrIRs(); beta++ ) {
-            if( isnan( IR(beta).c() ) ) {
-                bool has_nan_ancestor = false;
-                for( vector<size_t>::const_iterator beta2 = ancestors[beta].begin(); (beta2 != ancestors[beta].end()) && !has_nan_ancestor; beta2++ )
-                    if( isnan( IR(*beta2).c() ) )
-                        has_nan_ancestor = true;
-                if( !has_nan_ancestor ) {
+            if( !assigned[beta] ) {
+                bool has_unassigned_ancestor = false;
+                for( vector<size_t>::const_iterator beta2 = ancestors[beta].begin(); (beta2 != ancestors[beta].end()) && !has_unassigned_ancestor; beta2++ )
+                    if( !assigned[*beta2] )
+                        has_unassigned_ancestor = true;
+                if( !has_unassigned_ancestor ) {
                     double c = 1.0;
                     foreach( const Neighbor &alpha, nbIR(beta) )
                         c -= OR(alpha).c();
                     for( vector<size_t>::const_iterator beta2 = ancestors[beta].begin(); beta2 != ancestors[beta].end(); beta2++ )
                         c -= IR(*beta2).c();
                     IR(beta).c() = c;
+                    assigned[beta] = true;
                     new_counting = true;
                 }
             }
@@ -180,10 +182,10 @@ bool RegionGraph::Check_Counting_Numbers() {
     for( vector<Var>::const_iterator n = vars.begin(); n != vars.end(); n++ ) {
         double c_n = 0.0;
         for( size_t alpha = 0; alpha < nrORs(); alpha++ )
-            if( OR(alpha).vars() && *n )
+            if( OR(alpha).vars().contains( *n ) )
                 c_n += OR(alpha).c();
         for( size_t beta = 0; beta < nrIRs(); beta++ )
-            if( IR(beta) && *n )
+            if( IR(beta).contains( *n ) )
                 c_n += IR(beta).c();
         if( fabs(c_n - 1.0) > 1e-15 ) {
             all_valid = false;
@@ -206,11 +208,11 @@ void RegionGraph::RecomputeORs() {
 
 void RegionGraph::RecomputeORs( const VarSet &ns ) {
     for( size_t alpha = 0; alpha < nrORs(); alpha++ )
-        if( OR(alpha).vars() && ns )
+        if( OR(alpha).vars().intersects( ns ) )
             OR(alpha).fill( 1.0 );
     for( size_t I = 0; I < nrFactors(); I++ )
         if( fac2OR[I] != -1U )
-            if( OR( fac2OR[I] ).vars() && ns )
+            if( OR( fac2OR[I] ).vars().intersects( ns ) )
                 OR( fac2OR[I] ) *= factor( I );
 }