/* 
 * File:   CRF.cpp
 * Author: carlos V2
 * 
 */

#include <CRF.h>
#include <fstream>
#include <limits>
#include <vector>

CRF::CRF() {

    m_neigPolicy = -1;
    m_rows = -1;
    m_cols = -1;
    m_priorSatVal = -1;
    m_likeSatVal = -1;
    m_labelCnt = -1;

    m_prior = NULL;
    m_like = NULL;
    m_msgData1 = NULL;
    m_msgData2 = NULL;
    m_msg = NULL;
    m_msgT_1 = NULL;



    //printf("constr1\n");
}

CRF::CRF(int neighbor, int gridRows, int gridCols, int priorSatVal, int likeSatValue, string logFile, int labelCnt, int scale) {

    //printf("constr2\n");
    m_neigPolicy = neighbor;
    m_rows = gridRows;
    m_cols = gridCols;
    m_priorSatVal = priorSatVal;
    m_likeSatVal = likeSatValue;
    m_labelCnt = labelCnt;
    m_scale = scale;


    this->m_meas = cv::Mat(gridRows, gridCols, CV_8UC1);
    this->m_result = cv::Mat(gridRows, gridCols, CV_8UC1);
    //    this->m_meas = cv::Mat(gridRows, gridCols, CV_32FC1);
    //    this->m_result = cv::Mat(gridRows, gridCols, CV_32FC1);

    //PRIOR
    m_prior = (int**) malloc(sizeof (int)*m_labelCnt);
    for (int r = 0; r < m_labelCnt; r++) {
        m_prior[r] = (int*) malloc(sizeof (int)*m_labelCnt);
    }

    //MIN H
    m_minH = (int**) malloc(sizeof (int)*m_rows);
    for (int r = 0; r < m_rows; r++) {
        m_minH[r] = (int*) malloc(sizeof (int)*m_cols);
    }

    //M and H
    m_m = (int***) malloc(sizeof (int)*m_rows);
    m_h = (int***) malloc(sizeof (int)*m_rows);
    for (int r = 0; r < m_rows; r++) {
        m_m[r] = (int**) malloc(sizeof (int)*m_cols);
        m_h[r] = (int**) malloc(sizeof (int)*m_cols);
        for (int c = 0; c < m_cols; c++) {
            m_m[r][c] = (int*) malloc(sizeof (int)*m_labelCnt);
            m_h[r][c] = (int*) malloc(sizeof (int)*m_labelCnt);
        }
    }

    //LIKELIHOOD
    m_like = (int***) malloc(sizeof (int)*m_rows);
    for (int r = 0; r < m_rows; r++) {
        m_like[r] = (int**) malloc(sizeof (int)*m_cols);
        for (int c = 0; c < m_cols; c++) {
            m_like[r][c] = (int*) malloc(sizeof (int)*m_labelCnt);
        }
    }

    //MESSAGES
    m_msgData1 = (int****) malloc(sizeof (int)*m_rows);
    m_msgData2 = (int****) malloc(sizeof (int)*m_rows);
    for (int r = 0; r < m_rows; r++) {
        m_msgData1[r] = (int***) malloc(sizeof (int)*m_cols);
        m_msgData2[r] = (int***) malloc(sizeof (int)*m_cols);
        for (int c = 0; c < m_cols; c++) {
            m_msgData1[r][c] = (int**) malloc(sizeof (int)*m_labelCnt);
            m_msgData2[r][c] = (int**) malloc(sizeof (int)*m_labelCnt);
            for (int l = 0; l < m_labelCnt; l++) {
                m_msgData1[r][c][l] = (int*) malloc(sizeof (int)*m_neigPolicy);
                m_msgData2[r][c][l] = (int*) malloc(sizeof (int)*m_neigPolicy);
            }
        }
    }
    this->m_msg = m_msgData1;
    this->m_msgT_1 = m_msgData2;

    //INIT DATA TO ZERO
    for (int r = 0; r < m_labelCnt; r++) {
        for (int c = 0; c < m_labelCnt; c++) {
            m_prior[r][c] = 0;

        }
    }
    for (int r = 0; r < m_rows; r++) {
        for (int c = 0; c < m_cols; c++) {
            m_minH[r][c] = 0;
        }
    }
    for (int r = 0; r < m_rows; r++) {
        for (int c = 0; c < m_cols; c++) {
            for (int l = 0; l < m_labelCnt; l++) {
                m_like[r][c][l] = 0;
                m_m[r][c][l] = 0;
                m_h[r][c][l] = 0;
            }
        }
    }
    for (int r = 0; r < m_rows; r++) {
        for (int c = 0; c < m_cols; c++) {
            for (int l = 0; l < m_labelCnt; l++) {
                for (int n = 0; n < m_neigPolicy; n++) {
                    m_msg[r][c][l][n] = 0;
                    m_msgT_1[r][c][l][n] = 0;
                }
            }
        }
    }

}

CRF::CRF(const CRF& orig) {

}

CRF::~CRF() {

    if (m_prior != NULL)
        freeData();


}

void CRF::freeData() {
    //        printf("free\n");
    //PRIOR
    for (int r = m_labelCnt - 1; r >= 0; r--) {
        free(m_prior[r]);
    }
    free(m_prior);

    //MIN H
    for (int r = m_rows - 1; r >= 0; r--) {
        free(m_minH[r]);
    }
    free(m_minH);

    //M
    for (int r = m_rows - 1; r >= 0; r--) {
        for (int c = m_cols - 1; c >= 0; c--) {
            free(m_m[r][c]);
            free(m_h[r][c]);
        }
        free(m_m[r]);
        free(m_h[r]);
    }
    free(m_m);
    free(m_h);

    //LIKELIHOOD
    for (int r = m_rows - 1; r >= 0; r--) {
        for (int c = m_cols - 1; c >= 0; c--) {
            free(m_like[r][c]);
        }
        free(m_like[r]);
    }
    free(m_like);

    //MESSAGES
    for (int r = m_rows - 1; r >= 0; r--) {
        for (int c = m_cols - 1; c >= 0; c--) {
            for (int l = m_labelCnt - 1; l >= 0; l--) {
                free(m_msgData1[r][c][l]);
                free(m_msgData2[r][c][l]);
            }
            free(m_msgData1[r][c]);
            free(m_msgData2[r][c]);
        }
        free(m_msgData1[r]);
        free(m_msgData2[r]);
    }
    free(m_msgData1);
    free(m_msgData2);

    this->m_msg = NULL;
    this->m_msgT_1 = NULL;
}
////////////////////////////////////////////////////////////////////////////////

/*pasa de un valor concreto de label a la medida real. 
 en el caso de la imagen, como mis 50 labels están entre 90-140 en nivel de gris, 
 le sumo 90 para que restablezca el valor a niveles de gris correctos */
int CRF::fromLabelToValue(int label) {
    //    return label;
    return label + 90; //img
}

/*pasa de un valor de medida real a un valor de label
 en el caso de la imagen, como mis 50 labels están entre 90-140 en nivel de gris, 
 le resto 90 para que las labels se encuentren entre 0-50*/
int CRF::fromValueToLabel(int value) {
    //    return value;
    return value - 90; //img    
}

//////////////////////////////////////////////////////////////////////////////////
//float CRF::fromLabelToValue(int label) {
//    //printf("label2value\n");
//    return (-0.5) + (float)label * (0.01);//dem
//    //return label + 90; //img
//}
//
//int CRF::fromValueToLabel(float value) {    
//    int ret = (value - (-0.5)) / (0.01);//dem
//    return ret;
//}
//////////////////////////////////////////////////////////////////////////////////

void CRF::updateMeasurements(cv::Mat meas) {
    //printf("update\n");
    this->m_meas = meas;
}

int CRF::calcLineEqu(int u1, int v1, int u2, int v2, float* m, float* n) {
    int ret = 0;
    float mm, nn;
    float num = v2 - v1;
    float denom = u2 - u1;
    if (num == 0) {
        mm = 0;
        nn = v1;
    }
    else {
        mm = num / denom;
        nn = v1 - mm*u1;
    }
    if (denom == 0) {
        ret = -1;
        mm = 0;
        nn = 0;
        printf("Error estimating line equation. Vertical line!\n");
    }
    *m = mm;
    *n = nn;
    return ret;
}

void CRF::calcMeanStd(int r, int c, int windSize, int u1, int v1, int u2, int v2, float* meanA, float* meanB, float* meanC, float* stdA, float* stdB, float* stdC, vector<float> *ptsA, vector<float> *ptsB, vector<float> *ptsC) {
    float m, n;
    int line = calcLineEqu(u1, v1, u2, v2, &m, &n);
    float x = v1;

    float meana, meanb, stda, stdb, meanc, stdc, ptCnta, ptCntb, ptCntc, suma, sumb, sumc;
    ptsA->clear();
    ptsB->clear();
    ptsC->clear();

    int roiR = r - (windSize - 1) / 2;
    int roiC = c - (windSize - 1) / 2;
    cv::Mat roi = m_meas(cv::Rect(roiR, roiC, windSize, windSize));
    for (int r = 0; r < roi.rows; r++) {
        for (int c = 0; c < roi.cols; c++) {

            float z = roi.at<float>(r, c);
            if (z >= Z_VALID_MIN && z <= Z_VALID_MAX) {
                if (line != -1)
                    x = (r - n) / m;
                if ((int) x >= c) {
                    //current point is in zone A
                    ptCnta++;
                    suma += z;
                    stda += z*z;
                    ptsA->push_back(z);
                }
                else {
                    //current point is in zone B
                    ptCntb++;
                    sumb += z;
                    stdb += z*z;
                    ptsB->push_back(z);
                }
                //zone C contain all valid points
                ptCntc++;
                sumc += z;
                stdc += z*z;
                ptsC->push_back(z);
            }
        }
    }

    meana = suma / ptCnta;
    meanb = sumb / ptCntb;
    meanc = sumc / ptCntc;

    stda = stda / ptCnta - meana*meana;
    stdb = stdb / ptCntb - meanb*meanb;
    stdc = stdc / ptCntc - meanc*meanc;

    *meanA = meana;
    *meanB = meanb;
    *meanC = meanc;
    *stdA = stda;
    *stdB = stdb;
    *stdC = stdc;
}

float CRF::corrExpectedCurb(vector<float> ptsC, float meanA, float meanB) {
    float ret = 0;
    float hits = 0;
    float meanAB = (meanA + meanB)*0.5;
    float numPts = (float) ptsC.size();

    for (int i = 0; i < ptsC.size(); i++) {
        if (ptsC.at(i) >= meanAB)
            hits++;
    }
    ret = (numPts - hits) / numPts + 1 / (fabs(meanA - meanB));
    return ret;
}

float CRF::corrExpectedNonCurb(vector<float> ptsC, float meanC, float stdC, float precision, float meanA, float meanB) {
    float ret = 0;
    float hits = 0;
    float numPts = (float) ptsC.size();
    float variation = precision + stdC;

    for (int i = 0; i < ptsC.size(); i++) {
        if (fabs(ptsC.at(i) - meanC) <= variation)
            hits++;
    }
    ret = (numPts - hits) / numPts + fabs(meanA - meanB);
    return ret;
}

/*calculo de la likelihood para un punto determinado por las fila R y la columna C y una label*/
int CRF::likelihood(int r, int c, int windSize, int u1, int v1, int u2, int v2, float precision) {
    int value = -1;

    float meanA, meanB, meanC, stdA, stdB, stdC;
    vector<float> ptsA, ptsB, ptsC;

    this->calcMeanStd(r, c, windSize, u1, v1, u2, v2, &meanA, &meanB, &meanC, &stdA, &stdB, &stdC, &ptsA, &ptsB, &ptsC);
    float curb = this->corrExpectedCurb(ptsC, meanA, meanB);
    float nonCurb = this->corrExpectedNonCurb(ptsC, meanC, stdC, precision, meanA, meanB);

    printf("curb = %f\nnon-curb = %f\n", curb, nonCurb);
    value = curb >= nonCurb;

    return value;
}

/*pre-computo de todas las likelihoods para después acceder directamente al dato*/
void CRF::computeLikelihoods() {
    //printf("allLike\n");

    for (int r = 0; r < m_rows; r++) {
        for (int c = 0; c < m_cols; c++) {
            for (int l = 0; l < m_labelCnt; l++) {
                //                int like = likelihood(r, c, l, m_likeSatVal);
                //                m_like[r][c][l] = like;
                //                                if (r == 15 && c == 15)
                //                                    printf("(%d,%d)L(%d)\n", r, c, m_like[r][c][l]);
            }
        }
    }
}

float CRF::prior2(float meanA, float meanB, float satValue, float scale) {
    float value = fabs(meanA - meanB) * scale;
    if (value > satValue)
        value = satValue;

    return value;
}

int CRF::prior1(float dcp, float sigmaDC, float slope) {
    int ret = 0;
    float expon = exp(-(dcp - sigmaDC) * slope);
    float f1 = 1 / (1 + expon);
    float f0 = expon / (1 + expon);

    if (f1 >= f0)
        ret = 1;
    return ret;
}

float CRF::prior(float dcp, float sigmaDC, float slope, float meanA, float meanB, float satValue, float scale) {
    float pri1 = (float)prior1(dcp, sigmaDC, slope);    
    float pri2 = prior2(meanA, meanB, satValue, scale);
    printf("prior1 = %f\nprior2 = %f\n\nprior = %f\n",pri1,pri2,pri1+pri2);
    return pri1 + pri2;
}

/*pre-computo de todas las priors para después acceder directamente al dato*/
void CRF::computePriors() {
    //printf("allPrior\n");


    //printf("priors (%d , %d)\n",m_prior.rows,m_prior.cols);
    for (int p = 0; p < m_labelCnt; p++) {
        for (int q = 0; q < m_labelCnt; q++) {
            //            m_prior[p][q] = prior(p, q, m_priorSatVal, m_scale);
        }
    }
}

/*función que calcula los indices (canales de la matriz de mensajes) de los vecinos de P que son distintos de Q
 el vecino 0(up), 1(right), 2(down), 3(left), 4(up-left), 5(up-right), 6(down-right), 7(down-left)*/
int CRF::neighborIndex(int pr, int pc, int qr, int qc, int policy, vector<int> *neighbors) {
    //printf("vecinoIdx\n");
    int orderRow[] = {-1, 0, 1, 0, -1, 1, 1, -1};
    int orderCol[] = {0, 1, 0, -1, 1, 1, -1, -1};
    neighbors->clear();

    for (int i = 0; i < policy; i++) {
        int r = pr + orderRow[i];
        int c = pc + orderCol[i];
        if ((r != qr) || (c != qc)) {
            neighbors->push_back(i);
        }
    }
    return neighbors->size();
}

/*funcion que calcula las posiciones de TODOS los vecinos de P. 
 * Devuelve un vector con las filas y otro con las columnas para acceder en la matriz*/
int CRF::neighborPos(int pr, int pc, int policy, vector<int> *neigR, vector<int> *neigC) {
    //printf("vecinoPos\n");

    int orderRow[] = {-1, 0, 1, 0, -1, 1, 1, -1};
    int orderCol[] = {0, 1, 0, -1, 1, 1, -1, -1};
    neigR->clear();
    neigC->clear();

    for (int i = 0; i < policy; i++) {
        int r = pr + orderRow[i];
        int c = pc + orderCol[i];
        neigR->push_back(r);
        neigC->push_back(c);
    }
    return neigR->size();
}

/*suma todos los mensajes de los vecinos de P que no son Q para una label concreta*/
int CRF::msgSumAllMinusQ_T_1(int pr, int pc, int qr, int qc, int label, int policy) {
    //printf("msgT-1\n");
    int sum = 0;
    vector<int> neighbors;
    //calculo los vecinos de P que no son Q (en indices, es decir, el canal que debo leer la matriz de mensajes)
    if (neighborIndex(pr, pc, qr, qc, policy, &neighbors) == policy - 1) {
        sum = 0;
        //        printf("L%.0f msg(%d,%d)->(%d,%d) ", label, pr, pc, qr, qc);
        for (int i = 0; i < neighbors.size(); i++) {
            int idx = neighbors.at(i);
            //            printf("%d ", idx);
            sum += this->m_msgT_1[pr][pc][label][idx];
        }
        //        printf("\n");
    }
    else {
        //printf("Message Sum in T-1 error\n");
    }
    return sum;
}

/*suma todos los mensajes de los vecinos del punto Q para una determinada label*/
int CRF::msgSumAll_T_1(int qr, int qc, int label, int policy) {
    //printf("summAll_T\n");
    int sum = 0;
    for (int i = 0; i < policy; i++) {
        sum += m_msgT_1[qr][qc][label][i];
    }
    return sum;
}

/*suma todos los mensajes de los vecinos del punto Q para una determinada label*/
int CRF::msgSumAll_T(int qr, int qc, int label, int policy) {
    //printf("summAll_T\n");
    int sum = 0;
    for (int i = 0; i < policy; i++) {
        sum += m_msg[qr][qc][label][i];
        //                if (qr == 30 && qc == 50) {
        //                    printf("(R%d, C%d, L%d) msg[%d]= %d\n", qr, qc, label, i, m_msg[qr][qc][label][i]);
        //                }
    }
    //    if (qr == 30 && qc == 50) {
    //        printf("\n");
    //    }
    return sum;
}

/*cálculo de la belief para una determinada label y un punto concreto*/
int CRF::belief(int qr, int qc, int labelQ) {
    int likeli = getLikelihood(qr, qc, labelQ);
    int msgT = msgSumAll_T(qr, qc, labelQ, m_neigPolicy);
    int bel = likeli + msgT;
    //    printf("Belief(R%d,C%d,L%d) like(%d) msgT(%d) bel(%d)\n", qr, qc, labelQ, likeli, msgT, bel);
    return bel;
}

/*para un punto concreto Q, calcula las belief para todos los posibles 
 * valores de labels y se queda con el minimo*/
int CRF::getMinBelief(int qr, int qc) {
    //printf("minBelief\n");
    int labelMin = -1;
    int minim = numeric_limits<int>::max();
    for (int label = 0; label < m_labelCnt; label++) {
        int bel = belief(qr, qc, label);
        //        if (qr == 30 && qc == 50) {            
        //            
        //            int gris = this->fromValueToLabel((int) m_meas.at<uint8_t>(qr, qc));
        //            int gris0 = this->fromValueToLabel((int) m_meas.at<uint8_t>(qr - 1, qc));
        //            int gris1 = this->fromValueToLabel((int) m_meas.at<uint8_t>(qr, qc + 1));
        //            int gris2 = this->fromValueToLabel((int) m_meas.at<uint8_t>(qr + 1, qc));
        //            int gris3 = this->fromValueToLabel((int) m_meas.at<uint8_t>(qr, qc - 1));
        //
        //            printf("gris=%d vecinos= %d %d %d %d (R%d,C%d,L%d) belief = %d\n",
        //                   gris, gris0, gris1, gris2, gris3,
        //                   qr, qc, label, bel);
        //
        //        }
        if (bel < minim) {
            minim = bel;
            labelMin = label;
        }
    }
    return labelMin;
}

void CRF::twoPassAlgorith(int r, int c) {

    for (int labelQ = 1; labelQ < m_labelCnt - 1; labelQ++) {
        m_m[r][c][labelQ] = min(m_m[r][c][labelQ], m_m[r][c][labelQ - 1] + m_scale);
    }
    for (int labelQ = m_labelCnt - 2; labelQ >= 0; labelQ--) {
        m_m[r][c][labelQ] = min(m_m[r][c][labelQ], m_m[r][c][labelQ + 1] + m_scale);
    }
}

int CRF::calcH(int pr, int pc, int labelP) {
    int ret = 0;

    int likeli = getLikelihood(pr, pc, labelP);
    int msgP = msgSumAll_T_1(pr, pc, labelP, m_neigPolicy);
    ret = likeli + msgP;

    return ret;
}

void CRF::init_M_H() {
    for (int r = 0; r < m_rows; r++) {
        for (int c = 0; c < m_cols; c++) {
            int minimH = numeric_limits<int>::max();
            int minimLabel = 0;
            for (int l = 0; l < m_labelCnt; l++) {
                int h = calcH(r, c, l);
                //                if ((r == 29 && c == 50) || (r == 30 && c == 51) || (r == 31 && c == 50) || (r == 30 && c == 49)) {
                //                    printf("h(R%d, C%d, L%d) = %d\n", r, c, l, h);
                //                }
                m_m[r][c][l] = h;
                m_h[r][c][l] = h;
                if (h < minimH) {
                    minimH = h;
                    minimLabel = l;
                }
            }
            m_minH[r][c] = minimLabel + m_priorSatVal;
            //m_minH[r][c] = minimH + m_priorSatVal;
            //            if ((r == 29 && c == 50) || (r == 30 && c == 51) || (r == 31 && c == 50) || (r == 30 && c == 49)) {
            //                printf("minLabel(R%d, C%d) = %d\n", r, c, minimLabel);
            //            }
        }
    }
}

/*calcula el mensaje de P a Q para una label concreta*/
int CRF::computeMessage(int pr, int pc, int qr, int qc, int labelQ) {

    int m = this->m_m[pr][pc][labelQ];
    int h = this->m_minH[pr][pc];

    return min(m, h);
}

/*calcula los mensajes de todos los puntos y todas las labels posibles*/
int CRF::computeMessages() {
    //printf("allMsg\n");
    int ret = 0;
    int ****aux = m_msg;
    m_msg = m_msgT_1;
    m_msgT_1 = aux;
    int itr = 0;
    CpuTime cpuTime;

    for (int qr = 2; qr < m_meas.rows - 2; qr++) {
        cpuTime.start();
        for (int qc = 2; qc < m_meas.cols - 2; qc++) {



            vector<int> qrList, qcList;
            neighborPos(qr, qc, m_neigPolicy, &qrList, &qcList);
            for (int i = 0; i < qrList.size(); i++) {
                int pr = qrList.at(i);
                int pc = qcList.at(i);

                this->twoPassAlgorith(pr, pc);
                //                if (qr == 30 && qc == 50) {
                //                    for (int l = 0; l < m_labelCnt; l++) {
                //                        printf("i=%d twoPass(R%d, C%d, L%d) = %d\n", i, pr, pc, l, m_m[pr][pc][l]);
                //                    }
                //                }

                for (int labelQ = 0; labelQ < m_labelCnt; labelQ++) {
                    int msgT = computeMessage(pr, pc, qr, qc, labelQ);
                    m_msg[qr][qc][labelQ][i] = msgT;
                    //                    if (qr == 30 && qc == 50) {
                    //                        printf("n=%d msg(R%d, C%d, L%d) = %d\n", i, qr, qc, labelQ, msgT);
                    //                    }
                }
            }
            //            printf("[%010d / %010d] node(%03d,%03d) %s\n", itr, total, 
            //            pr, pc, cpuTime.getText().data());

        }
        cpuTime.stop();
        //        printf("[%010d / %010d] node(%03d) %s\n", itr, total, pr, cpuTime.getText().data());
        itr++;
    }


    return ret;
}

/*calcula las belief para todos los puntos*/
void CRF::computeBeliefs() {
    //printf("allBeliefs\n");
    for (int r = 2; r < m_meas.rows - 2; r++) {
        for (int c = 2; c < m_meas.cols - 2; c++) {
            int minBel = getMinBelief(r, c);
            if (minBel == -1) {
                printf("Error Belief\n");
            }
            else {
                m_result.at<uint8_t>(r, c) = (uint8_t) fromLabelToValue(minBel);
                //                m_result.at<float>(r, c) = (float) fromLabelToValue(minBel);
            }
        }
    }
}





