#include <config.h>
#include <distributions/DNormMix.h>
#include <sarray/SArray.h>

#include <cmath>

#include <Rmath.h>

using std::vector;

static double const * MU(vector<SArray const *> const &par)
{
    return par[0]->value();
}

static double const * TAU(vector<SArray const *> const &par)
{
    return par[1]->value();
}

static double const * PROB(vector<SArray const *> const &par)
{
    return par[2]->value();
}

DNormMix::DNormMix()
  : DistReal("dnormmix", 3, DIST_UNBOUNDED, false)
{}

DNormMix::~DNormMix()
{}

#include <iostream>
bool DNormMix::checkParameterDim(vector<SArray const *> const &par) const
{
    if (par[0]->length() < 2) {
	// Must be a mixture
	return false;
    }

    // Parameter dimensions must match (but they need not be vectors)
    Index const &dim = par[0]->dim(true);
    if (par[1]->dim(true) != dim || par[2]->dim(true) != dim)
	return false;

    return true;
}

bool DNormMix::checkParameterValue(vector<SArray const *> const &par) const
{
    double const *tau = TAU(par);
    double const *prob = PROB(par);
    unsigned long Ncat = par[0]->length();
    double sump = 0.0;
    for (unsigned int i = 0; i < Ncat; ++i) {
	if (tau[i] <= 0)
	    return false;
	sump += prob[i];
    }
    return (fabs(sump - 1) < 16 * DBL_EPSILON);
}

double
DNormMix::d(double x, vector<SArray const *> const &par, bool give_log) const
{
    double const *mu = MU(par);
    double const *tau = TAU(par);
    double const *p = PROB(par);
    unsigned long Ncat = par[0]->length();

    double density = 0.0;
    for (unsigned int i = 0; i < Ncat; ++i) {
	density += p[i] * dnorm(x, mu[i], 1/sqrt(tau[i]), 0);
    }
    return give_log ? log(density) : density;
}

double
DNormMix::p(double q, vector<SArray const *> const &par, bool lower, bool give_log)
    const
{
    double const *mu = MU(par);
    double const *tau = TAU(par);
    double const *prob = PROB(par);
    unsigned long Ncat = par[0]->length();
    
    double density = 0.0;
    for (unsigned int i = 0; i < Ncat; ++i) {
	density += prob[i] * pnorm(q, mu[i], 1/sqrt(tau[i]), lower, 0);
    }
    return give_log ? log(density) : density;
}

static double 
pmix(double q, double const *mu, double const *sigma, double const *p,
     unsigned int Ncat) 
{
    /* Probability distribution function of a normal mixture */

    double density = 0.0;
    for (unsigned int i = 0; i < Ncat; ++i) {
	density += p[i] * pnorm(q, mu[i], sigma[i], true, false);
    }
    return density;
}

static double 
qmix (double p, double const *mu, double const *sigma,
      double const *prob, unsigned int Ncat)
{
    /* Calculate quantile of normal mixture by bisection
       Based on an algorithm in the R package ensembleBMA by
       Adrian E. Raftery, J. McLean Sloughter, and Michael Polakowski, 
    */

    const double C = qnorm(p, 0, 1, true, false);
    double lower = mu[0] - C * sigma[0];
    double upper = mu[0] + C * sigma[0];
    double mid = (lower + upper)/2;

    while (upper - lower > 16 * sqrt(DBL_EPSILON)) {
	if (pmix(mid, mu, sigma, prob, Ncat) > p) {
	    upper = mid;
        }
        else {
            lower = mid;
        }
        mid = (lower + upper)/2;
    }
    return mid;
}

double 
DNormMix::q(double p, vector<SArray const *> const &par, bool lower, bool log_p)
    const
{
    if (!lower)
	p = 1 - p;
    if (log_p)
	p = exp(p);

    double const *mu = MU(par);
    double const *tau = TAU(par);
    double const *prob = PROB(par);
    unsigned long Ncat = par[0]->length();
    
    double * sigma = new double[Ncat];
    for (unsigned int i = 0; i < Ncat; ++i) {
	sigma[i] = 1/sqrt(tau[i]);
    }
    double q = qmix(p, mu, sigma, prob, Ncat);
    delete [] sigma;

    return q;
}

double 
DNormMix::r(vector<SArray const *> const &par) const
{
    double const *mu = MU(par);
    double const *tau = TAU(par);
    double const *p = PROB(par);
    unsigned long Ncat = par[0]->length();    
    
    // Select mixture component (r)
    unsigned int r = Ncat - 1;
    double sump = 0;
    double p_rand = runif(0,1);
    for (unsigned int i = 0; i < Ncat - 1; ++i) {
	sump += p[i];
	if (sump > p_rand) {
	    r = i;
	    break;
	}
    }

    // Now sample from conditional distribution of component r
    return rnorm(mu[r], 1/sqrt(tau[r]));
}

double DNormMix::mean(std::vector<SArray const*> const &par) const
{
    double const *mu = MU(par);
    double const *p = PROB(par);
    unsigned long Ncat = par[0]->length();    

    double m = 0.0;
    for (unsigned long i = 0; i < Ncat; ++i) {
	m += p[i] * mu[i];
    }
    return m;
}

double DNormMix::var(std::vector<SArray const*> const &par) const
{
    double const *tau = TAU(par);
    double const *p = PROB(par);
    unsigned long Ncat = par[0]->length();    

    double v = 0.0;
    for (unsigned long i = 0; i < Ncat; ++i) {
	v += p[i] / tau[i];
    }
    return v;
}


syntax highlighted by Code2HTML, v. 0.9.1