#include #include #include #include #include #include #include #include #include #include #include #include using std::vector; using std::set; using std::sqrt; using std::invalid_argument; using std::logic_error; using std::max; using std::min; ConjugateBeta::ConjugateBeta(StochasticNode *snode, Graph const &graph) : ConjugateSampler(snode, graph) { if (!canSample(snode, graph)) { throw invalid_argument("Can't construct ConjugateBeta sampler"); } } ConjugateBeta::~ConjugateBeta() { } bool ConjugateBeta::canSample(StochasticNode *snode, Graph const &graph) { switch(getDist(snode)) { case BETA: break; case UNIF: // dunif(0,1) is equivalent to dbeta(1,1) if(!(*snode->parameters()[0]->value() == 0 && *snode->parameters()[1]->value() == 1 && snode->parameters()[0]->isFixed() && snode->parameters()[1]->isFixed())) return false; break; default: return false; } vector stoch_nodes; vector dtrm_nodes; classifyChildren(vector(1,snode), graph, stoch_nodes, dtrm_nodes); /* Create a set of nodes containing snode and its deterministic descendants for the checks below. */ set paramset; paramset.insert(&snode->data); for (unsigned int j = 0; j < dtrm_nodes.size(); ++j) { paramset.insert(&dtrm_nodes[j]->data); } // Check deterministic descendants // Only Mixture nodes are allowed for (unsigned int j = 0; j < dtrm_nodes.size(); ++j) { if (isMixture(dtrm_nodes[j])) { // Check that indices do not depend on snode vector const &index = asMixture(dtrm_nodes[j])->index(); for (unsigned int i = 0; i < index.size(); ++i) { if (paramset.count(&index[i]->data)) return false; } } else { return false; } } // Check stochastic children for (unsigned int i = 0; i < stoch_nodes.size(); ++i) { if (stoch_nodes[i]->isBounded()) { return false; //Bounded } switch(getDist(stoch_nodes[i])) { case BIN: if (paramset.count(stoch_nodes[i]->parameters()[1])) { return false; //n depends on snode } break; case BERN: break; default: return false; } } return true; //We made it! } void ConjugateBeta::update() { vector const &stoch_children = stochasticChildren(); StochasticNode const *snode = node(); double a, b; switch (_target_dist) { case BETA: a = *snode->parameters()[0]->value(); b = *snode->parameters()[1]->value(); break; case UNIF: a = 1; b = 1; } unsigned int Nchild = stoch_children.size(); /* For mixture models, we count only stochastic children that depend on snode */ double *C = 0; bool is_mix = !deterministicChildren().empty(); if (is_mix) { C = new double[Nchild]; for (unsigned int i = 0; i < Nchild; ++i) { C[i] = *stoch_children[i]->parameters()[0]->value(); } // Perturb current value, keeping in the legal range [0,1] double x = *snode->data.value(); x = x > 0.5 ? x - 0.4 : x + 0.4; setValue(&x, 1); // C[i] == 1 if parameter of child i has changed (so depends on snode) // C[i] == 0 otherwise for (unsigned int i = 0; i < Nchild; ++i) { C[i] = (*stoch_children[i]->parameters()[0]->value() != C[i]); } } for (unsigned int i = 0; i < stoch_children.size(); ++i) { if (!(is_mix && C[i] == 0)) { double y = *stoch_children[i]->data.value(); double n; switch(_child_dist[i]) { case BIN: n = *stoch_children[i]->parameters()[1]->value(); break; case BERN: n = 1; break; default: throw logic_error("Invalid distribution in Conjugate Beta sampler"); } a += y; b += (n - y); } } // Draw the sample double xnew = rbeta(a, b); if (node()->isBounded()) { double lower = 0; SArray const *lb = node()->lowerBound(); if (lb) { lower = max(lower, *lb->value()); } double upper = 1; SArray const *ub = node()->upperBound(); if (ub) { upper = min(upper, *ub->value()); } /* Try 4 more attempts to get random sample within the bounds */ for (int i = 0; i < 4; i++) { if (xnew >= lower && xnew <= upper) { setValue(&xnew, 1); return; } xnew = rbeta(a, b); } /* Failure! Use inversion */ double plower = lb ? pbeta(lower, a, b, 1, 0) : 0; double pupper = ub ? pbeta(upper, a, b, 1, 0) : 1; double p = runif(plower, pupper); xnew = qbeta(p, a, b, 1, 0); } setValue(&xnew, 1); if (is_mix) { delete [] C; } }