// MBIotaSampler.c	
// Sampling iota 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
//       March 17, 2003 -- Correct an error.
//       July 12, 2003 -- Assume veps = 4*vxi
//
//	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"

struct arms_parm
{
    double xl,xr,convex;
    int ninit,npoint,dometrop,nsamp,ncent;
};

struct config_parm
{
    int degree,nobs,nknots;
    double *knots,*gIdx;
};

struct post_parm
{
    double M,A,veps;
    double *beta,*gamma,*knots;
    int gIdx,degree,nknots;
};

double MBIotaLogPost(double iota, void *pspost)
{
    struct post_parm *pdata;
    double dlog;
    int k;
    
    pdata = pspost;
    dlog = pdata->M-pdata->gamma[pdata->gIdx-1];
    for(k=0; k<pdata->degree+1; k++)
        dlog = dlog - pow(iota,k)*pdata->beta[k];
        
    for(k=0; k<pdata->nknots; k++)
    {
        if(iota>pdata->knots[k])
            dlog = dlog - pow(iota-pdata->knots[k],
                   pdata->degree)*pdata->beta[k+pdata->degree+1];
    }

    // iota \propto 1
    dlog = -dlog*dlog/(2*pdata->veps)-2*pow(pdata->A-iota,2)/pdata->veps;
    
    return dlog;
}

int arms(double *xinit,int ninit,double *xl,double *xr,
      	 double (*myfunc)(double x,void *mydata),void *mydata,
         double *convex,int npoint,int dometrop,double *xprev, 
         double *xsamp,int nsamp,double *qcent,double *xcent, 
         int ncent,int *neval);

void MBIotaSampler(struct arms_parm sparam,struct config_parm config,
                   double inData[],double beta[],double gamma[],
                   double veps,double iotapre[],double iotanew[])
{
    int err,neval,n;
    double xinit[10];   //double xinit[sparam.ninit];
    double xcent, qcent, tmpIota;
    struct post_parm ddata;     //set up structures for each density function
    struct post_parm tparm;

    // Set up starting values
    if( sparam.ninit > 10 )
        mexWarnMsgTxt("Increase the size of xinit[]...");
    
    for(n=0;n<sparam.ninit;n++)
    {
        xinit[n] = sparam.xl+(n+1)*(sparam.xr-sparam.xl)/(sparam.ninit+1.0);
    }
    
    ddata.veps = veps;
    ddata.beta = beta;
    ddata.gamma = gamma;
    ddata.knots = config.knots;
    ddata.degree = config.degree;
    ddata.nknots = config.nknots;

    for(n=0;n<config.nobs;n++)
    {
        ddata.M = inData[n];
        ddata.A = inData[config.nobs+n];
        ddata.gIdx = (int)config.gIdx[n];
        
        tmpIota = iotapre[n];

        tparm = ddata;
        
        err = arms(xinit,sparam.ninit,&sparam.xl,&sparam.xr,
                   MBIotaLogPost,&tparm,&sparam.convex,sparam.npoint,
                   sparam.dometrop,&tmpIota,&iotanew[n],
                   sparam.nsamp,&qcent,&xcent,sparam.ncent,&neval);

        if(err>0)
        {
            iotanew[n] = tmpIota;
            mexWarnMsgTxt("Errors in ARMS for iota with code:");
            mexPrintf("%i\n", err);
            mexWarnMsgTxt("Please check the error type in MBIotaSampler.c!");
        }

        //if( n==config.nobs-1 )
        //{
        //    mexPrintf("Last observation with ID: %i\t", ddata.gIdx);
        //    mexPrintf("Previous iota: %f\t", iotapre[n]);
        //    mexPrintf("New iota: %f\n", iotanew[n]);
        //}
    }
}

void mexFunction(int nlhs, mxArray *plhs[], int nrhs, const mxArray *prhs[])
{
    double *inData,*beta,*gamma,*iotapre,*iotanew;
    double veps;
    struct config_parm config;
    struct arms_parm sparam;

    // Check for proper number of arguments.
    if( nrhs!=7 )
    {
        mexErrMsgTxt("Seven inputs required.");
    }
    else if( nlhs>1 )
    {
        mexErrMsgTxt("Too many output arguments.");
    }

    // RHS: config,arms,inData,beta,gamma,veps,iotapre
    config.degree = (int)mxGetScalar(mxGetField(prhs[0],0,"degree"));
    config.nobs = (int)mxGetScalar(mxGetField(prhs[0],0,"nobs"));
    config.nknots = (int)mxGetScalar(mxGetField(prhs[0],0,"nknots"));
    config.knots = mxGetPr(mxGetField(prhs[0],0,"knots"));
    config.gIdx = mxGetPr(mxGetField(prhs[0],0,"gIdx"));
    
    sparam.ninit = (int)mxGetScalar(mxGetField(prhs[1],0,"ninit"));
    sparam.xl = mxGetScalar(mxGetField(prhs[1],0,"xl"));
    sparam.xr = mxGetScalar(mxGetField(prhs[1],0,"xr"));
    sparam.convex = (int)mxGetScalar(mxGetField(prhs[1],0,"convex"));
    sparam.npoint = (int)mxGetScalar(mxGetField(prhs[1],0,"npoint"));
    sparam.dometrop = (int)mxGetScalar(mxGetField(prhs[1],0,"dometrop"));
    sparam.nsamp = (int)mxGetScalar(mxGetField(prhs[1],0,"nsamp"));
    sparam.ncent = (int)mxGetScalar(mxGetField(prhs[1],0,"ncent"));
    
    inData = mxGetPr(prhs[2]);
    beta = mxGetPr(prhs[3]);
    gamma = mxGetPr(prhs[4]);

    veps = mxGetScalar(prhs[5]);
    iotapre = mxGetPr(prhs[6]);
    
    plhs[0] = mxCreateDoubleMatrix(config.nobs,1,mxREAL);
    iotanew = mxGetPr(plhs[0]);

    // Set random number seed
    srand( time(NULL) );
    
    MBIotaSampler(sparam,config,inData,beta,gamma,veps,iotapre,iotanew);
}

////////////////////////////////////////////////////////////////
// The following codes are copied from Gilks for ARMS
//
// Adaptive Rejection Metropolis Sampling

#include "arms.c"
