// MBGammaSampler.c	
// Sampling gamma from its posterior distribution.
//
//   The algorithm for MicroBayes is described in the following paper
//
//   D. Zhang, M.T. Wells, C.D. Smart and W.E. Fry (2003). Bayesian
//       Normalization and Identification for Differential Gene
//       Expression Data.
//
//   PLASE CITE THIS PAPER AFTER YOU USE THIS SOFTWARE.
//
//	Dabao Zhang		February 24, 2003
//	Revised by Dabao Zhang
//      March 9, 2003 -- Consider genes' location effect
//      May 4, 2003 -- Error in updating gamma with non-informative priors
//       July 13, 2003 -- Set the prior for the common p.
//
//	Copyright (c) 2003 by Dabao Zhang.

#include <stdlib.h>             // Needed for rand(), srand()
#include <math.h>               // Needed for sqrt() and log()
#include <time.h>               // Needed for time()
#include "mex.h"

//----- Defines -------------------------------------------------------------
#define PI         3.14159265   // The value of pi

//===========================================================================
//  Function to generate one sample of N(0,1) using the Box-Muller method
//    - Input: mean and standard deviation
//    - Output: Returns with normally distributed random variable

double stdnrnd(void)
{
    double   u, r, theta;           // Variables for Box-Muller method
    double   x;                     // Normal(0, 1) rv

    // Generate u
    u = 0.0;
    while (u == 0.0)
        u = ((double)rand()+0.5) / ((double)RAND_MAX+1.0);

    // Compute u
    r = sqrt(-2.0 * log(u));

    // Generate theta
    theta = 0.0;
    while (theta == 0.0)
        theta = 2.0 * PI * ((double)rand()+0.5)/((double)RAND_MAX+1.0);

    // Generate x value
    x = r * cos(theta);

    // Return the normally distributed RV value
    return(x);
}

int bernrnd(double p)
{
    double   u;
    int   x;

    // Generate u
    u = 0.0;
    while (u == 0.0)
        u = ((double)rand()+0.5) / ((double)RAND_MAX+1.0);

    x = (u<p)?1:0;
        
    return(x);
}

void MBGammaSampler(int nobs,int ngenes,double gIdx[],double errNL[],
                    double veps,double vgamma,double gamma[],double p)
{
    int n, idxS, idxCurr;
    double mCurr,vCurr,tmpP;
    
    idxS = 0;
    idxCurr = (int)gIdx[0];
    mCurr = 0;
    vCurr = 0;
    for(n=0; n<nobs; n++)
    {
        if( idxCurr != ((int)gIdx[n]) )
        {
            tmpP = 1-(1-p)/(1-p+p*exp(0.5*mCurr*mCurr/((n-idxS)*veps
                   +veps*veps/vgamma))/sqrt(1+(n-idxS)*vgamma/veps));
        
            mCurr = (mCurr*vgamma)/((n-idxS)*vgamma+veps);
            vCurr = 1/((n-idxS)/veps+1/vgamma);

            // Generate gamma[idxCurr-1]
            gamma[idxCurr-1] = bernrnd(tmpP)?(stdnrnd()*sqrt(vCurr)+mCurr):0;

            //if( idxCurr <= 2 )
            //{
            //    mexPrintf("tmpP(%d): %f\t\n", idxCurr, tmpP);
            //    mexPrintf("gamma(%d): (%f, %f); New value: %f\n", idxCurr, mCurr, vCurr, gamma[idxCurr-1]);
            //}
            
            mCurr = errNL[n];
            idxS = n;
            idxCurr = (int)gIdx[n];
        }
        else
        {
            mCurr = mCurr + errNL[n];
        }

        if( n==(nobs-1) )
        {
            tmpP = 1-(1-p)/(1-p+p*exp(0.5*mCurr*mCurr/((n-idxS+1)*veps
                   +veps*veps/vgamma))/sqrt(1+(n-idxS+1)*vgamma/veps));
        
            mCurr = (mCurr*vgamma)/((n-idxS+1)*vgamma+veps);
            vCurr = 1/((n-idxS+1)/veps+1/vgamma);

            // Generate gamma[idxCurr-1]
            gamma[idxCurr-1] = bernrnd(tmpP)?(stdnrnd()*sqrt(vCurr)+mCurr):0;
        }
    }
}

void mexFunction(int nlhs, mxArray *plhs[], int nrhs, const mxArray *prhs[])
{
    double *gIdx,*errNL,*gamma;
    double veps,vgamma,p;
    int ngenes,nobs;

    // Check for proper number of arguments.
    if( nrhs!=6 )
    {
        mexErrMsgTxt("Six inputs required.");
    }
    else if( nlhs>1 )
    {
        mexErrMsgTxt("Too many output arguments.");
    }
    
    // RHS: config,geneIdx,errNL,veps,vgamma
    ngenes = (int)mxGetScalar(mxGetField(prhs[0],0,"ngenes"));
    nobs = (int)mxGetScalar(mxGetField(prhs[0],0,"nobs"));

    gIdx = mxGetPr(prhs[1]);
    errNL = mxGetPr(prhs[2]);
    veps = mxGetScalar(prhs[3]);
    vgamma = mxGetScalar(prhs[4]);
    p = mxGetScalar(prhs[5]);
    
    plhs[0] = mxCreateDoubleMatrix(ngenes,1,mxREAL);
    gamma = mxGetPr(plhs[0]);

    // Set random number seed
    srand( time(NULL) );
    bernrnd(0.5);
    
    MBGammaSampler(nobs,ngenes,gIdx,errNL,veps,vgamma,gamma,p);
}
