#include <config.h>
#include <sampler/ConjugateMNormal.h>
#include <sampler/ConjugateSampler.h>
#include <graph/AggNode.h>
#include <graph/MixtureNode.h>
#include <graph/NodeError.h>
#include <graph/LogicalNode.h>
#include <graph/StochasticNode.h>
#include <distributions/DMNorm.h>
#include <matrix/lapack.h>
#include <set>
#include <stdexcept>
#include <vector>
#include <cmath>
#include <string>
#include <Rmath.h>
using std::string;
using std::vector;
using std::set;
using std::sqrt;
using std::invalid_argument;
ConjugateMNormal::ConjugateMNormal(StochasticNode *snode, Graph const &graph)
: ConjugateSampler(snode, graph)
{
if (!canSample(snode, graph)) {
string msg("Can't construct ConjugateNormal sampler");
throw invalid_argument(msg);
}
}
ConjugateMNormal::~ConjugateMNormal()
{
}
bool ConjugateMNormal::canSample(StochasticNode *snode, Graph const &graph)
{
if (getDist(snode) != MNORM)
return false;
if (snode->isBounded())
return false;
vector<StochasticNode const*> stoch_nodes;
vector<Node*> dtrm_nodes, extra_nodes;
classifyChildren(vector<StochasticNode*>(1,snode),
graph, stoch_nodes, dtrm_nodes);
/*
Create a set of nodes containing snode and its deterministic
descendants for the checks below.
*/
set<SArray const *> paramset;
paramset.insert(&snode->data);
for (unsigned int j = 0; j < dtrm_nodes.size(); ++j) {
paramset.insert(&dtrm_nodes[j]->data);
}
// Check stochastic children
for (unsigned int i = 0; i < stoch_nodes.size(); ++i) {
if (getDist(stoch_nodes[i]) != MNORM &&
getDist(stoch_nodes[i]) != NORM) {
return false; //Not normal or multivariate normal
}
if (stoch_nodes[i]->isBounded()) {
return false;
}
vector<SArray const *> const ¶m = stoch_nodes[i]->parameters();
if (paramset.count(param[1])) {
return false; //Precision depends on snode
}
}
// Check deterministic descendants
for (unsigned int j = 0; j < dtrm_nodes.size(); ++j) {
if (isLogical(dtrm_nodes[j])) {
LogicalNode const *lnode = asLogical(dtrm_nodes[j]);
unsigned int nfactor = 0;
if (lnode) {
vector<SArray const *> const ¶m = lnode->parameters();
switch(getOp(lnode)) {
case ADD: case SUBTRACT: case NEG:
break;
case DIVIDE:
if (paramset.count(param[1]))
return false; //reciprocal term
break;
case INPROD:
if (paramset.count(param[0]) && paramset.count(param[1]))
return false; //quadratic term
break;
case MULTIPLY:
for (unsigned int k = 0; k < param.size(); ++k) {
nfactor += paramset.count(param[k]);
}
if (nfactor != 1)
return false; //quadratic or higher term
break;
default:
return false;
}
}
else if (isMixture(dtrm_nodes[j])) {
// Check that indices do not depend on snode
vector<Node *> 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;
}
}
}
return true; //We made it!
}
void ConjugateMNormal::update()
{
vector<StochasticNode const*> const &stoch_children = stochasticChildren();
unsigned int nchildren = stoch_children.size();
double const *priormean = node()->parameters()[0]->value();
double const *priorprec = node()->parameters()[1]->value();
int nrow = node()->parameters()[0]->length();
/*
The log of the full conditional density takes the form
-1/2(t(x) %*% A %*% x - 2 * b %*% x)
*/
int N = nrow * nrow;
double *b = new double[nrow];
double *A = new double[N];
for (int i = 0; i < nrow; ++i) {
b[i] = 0;
for (int i2 = 0; i2 < nrow; ++i2) {
b[i] += priorprec[i * nrow + i2] * priormean[i2];
}
}
for (int i = 0; i < N; ++i) {
A[i] = priorprec[i];
}
/* We know that the prior mean of any stochastic child depends on
the value of the current node via a linear function alpha + beta %*% x.
We first create a vector of coefficients.
*/
vector<double *> alpha, beta;
vector<int> nrow_child;
for (vector<StochasticNode const*>::const_iterator p(stoch_children.begin());
p != stoch_children.end(); ++p)
{
int nrow_p = (*p)->data.length();
alpha.push_back(new double[nrow_p]);
beta.push_back(new double[nrow_p * nrow]);
nrow_child.push_back(nrow_p);
}
/*
In order to calculate the value of the coefficients alpha and
beta, we evaluate the prior mean twice: once setting x (value of
snode) to 0 and once setting each element of x to 1
*/
double *xnew = new double[nrow];
for (int i = 0; i < nrow; ++i) {
xnew[i] = 0;
}
setValue(xnew, nrow);
for (unsigned int j = 0; j < nchildren; ++j) {
double const *mu = stoch_children[j]->parameters()[0]->value();
for (int k = 0; k < nrow_child[j]; ++k) {
alpha[j][k] = mu[k];
}
}
for (int i = 0; i < nrow; ++i) {
xnew[i] = 1;
setValue(xnew, nrow);
for (unsigned int j = 0; j < nchildren; ++j) {
double const *mu = stoch_children[j]->parameters()[0]->value();
for (int k = 0; k < nrow_child[j]; ++k) {
beta[j][nrow * k + i] = mu[k] - alpha[j][k];
}
}
xnew[i] = 0;
}
/* Now add the contribution of each term to A, b */
for (unsigned int j = 0; j < nchildren; ++j) {
double const *Y = stoch_children[j]->data.value();
double const *tau = stoch_children[j]->parameters()[1]->value();
for (int i = 0; i < nrow; ++i) {
//for (int i2 = i; i2 < nrow; ++i2) {
for (int i2 = 0; i2 < nrow; ++i2) {
double Aplus = 0;
for (int k = 0; k < nrow_child[j]; ++k) {
for (int k2 = 0; k2 < nrow_child[j]; ++k2) {
Aplus += tau[nrow_child[j] * k + k2]
* beta[j][nrow * k + i]
* beta[j][nrow * k2 + i2];
}
}
A[i * nrow + i2] += Aplus;
/*
if (i != i2) {
A[i2 * nrow + i] += Aplus;
}
*/
}
}
for (int i = 0; i < nrow; ++i) {
double bplus = 0;
for (int k = 0; k < nrow_child[j]; ++k) {
for (int k2 = 0; k2 < nrow_child[j]; ++k2) {
bplus += tau[nrow_child[j] * k + k2]
* beta[j][nrow * k + i]
* (Y[k2] - alpha[j][k2]);
}
}
b[i] += bplus;
}
}
/* We don't need the coefficients any more */
for (unsigned int j = 0; j < nchildren; ++j) {
delete alpha[j];
delete beta[j];
}
/*
Solve the equation A %*% x = b to get the posterior mean.
We have to take a copy of A as it is overwritten during
the call to DSYSV. The result is stored in b
*/
double * Acopy = new double[N];
for (int i = 0; i < N; ++i) {
Acopy[i] = A[i];
}
int *ipiv = new int[nrow];
int one = 1;
double worktest;
int lwork = -1;
int info;
F77_DSYSV ("L", &nrow, &one, Acopy, &nrow, ipiv, b, &nrow, &worktest,
&lwork, &info);
if (info != 0) {
throw NodeError(node(),
"unable to solve linear equations in Conjugate mnorm sampler");
}
lwork = static_cast<int>(worktest) + 1;
double * work = new double[lwork];
F77_DSYSV ("L", &nrow, &one, Acopy, &nrow, ipiv, b, &nrow, work, &lwork,
&info);
if (info != 0) {
throw NodeError(node(),
"unable to solve linear equations in Conjugate MNorm sampler");
}
delete [] work;
delete [] Acopy;
delete [] ipiv;
DMNorm::randomsample(xnew, b, A, nrow);
setValue(xnew, nrow);
delete [] b;
delete [] A;
delete [] xnew;
}
syntax highlighted by Code2HTML, v. 0.9.1