
#include <stdio.h>
#include <stdlib.h>
#include <math.h>

/* Computes log(Beta(m, n)) using table look-up.  Table size is
   determined by value of k first time function is called.  Falls back
   on explicit calculation if table lookup fails  */

double logbeta(int m, int n, int tsize)
{
    int i;
    static double *lg; 	/* lg[i] == log(Gamma(i+1)) */
    static int first = 1, done_upto = 0, table_size;
    if (first) {
	lg = (double *) malloc(tsize * sizeof(double));
	table_size = tsize;
	lg[0] = 0; lg[1] = 0;
	first = 0;
	done_upto = 2;
    }
    if (m + n > table_size) {
	/* Fall back on explicit calculation rather than table look-up */
	return lgamma(m) + (lgamma(n) - lgamma(m+n));
    }
    if (m + n > done_upto) {
	for (i = done_upto; i < m + n; i++) {
	    lg[i] = lg[i-1] + log((double) i);
	}
    }
    return lg[m-1] + (lg[n-1] - lg[m+n-1]);
}



double beta_loglik(double sum1, double sum2, double dsize, int m, int n, int k)
{
    return -dsize * logbeta(m, n, k) + (m-1) * sum1 + (n-1) * sum2;
}


void beta_mle(double *x, int dsize, int k, int *mest, int *nest)
{
    int i, m, n, mhat = 1, nhat = 1;
    double ll, maxll, sum1 = 0, sum2 = 0;
    dsize++;
    dsize--;
    for (i = 0; i < dsize; i++) {
	sum1 += log(x[i]);
	sum2 += log(1 - x[i]);
    }
    maxll = beta_loglik(sum1, sum2, dsize, mhat, nhat, k);
    for (m = 1; m < k; m++) 
	for (n = 1; n < k; n++) {
	    ll = beta_loglik(sum1, sum2, dsize, m, n, k);
	    if (ll > maxll) {
		maxll = ll;
		mhat = m;
		nhat = n;
	    }
	}
    /* printf("%d %d %f\n", mhat, nhat, maxll); */
    *mest = mhat;
    *nest = nhat;
    return;
}



