/*
 * main.cpp
 *
 *  Created on: 12.06.2012
 *      Author: drayer
 *
 *  Modified on: 14.06.2013
 *      by:     abdulkadir
 *      changes: dependencies from libBlitzHDF5 removed
 *               made compatible with blitz++-0.10
 *
 *  Modified on: 03.07.2014
 *      by:     bensch
 *      changes: replaced hdf5 low-lever interface functions by the simple hdf5 interface library
 *      		BlitzHDF5Helper.hh, developed by abhulkadir
 * 		 added input of parameters from console directly (filenames, lambda, similarity mesure)
 * 		 restructured code and modified naming of functions and variables
 * 		 direct access of the original FAST_PD interface, without using another wrapper
 * 		 added color visualization of deformation fields
 *
 *  Modified on: 08.07.2015
 *      by: ummenhof
 *      changes:
 *       - removed deprecated tinyvec-et.h include
 *       - added examples to command line help.
 *       - similarityMeasure parameter on the command line is now a string
 *       - fixed a bug in ncc patch computations
 *
 * 
 *  g++ -Wall -O3 -g elastic.cc  Fast_PD/graph.cpp Fast_PD/LinkedBlockList.cpp Fast_PD/maxflow.cpp -lgsl -lgslcblas -lblitz -o elastic -lhdf5
 */

// General Includes
#include <string>
#include <vector>
#include <time.h>
#include <ctime>
#include <iostream>
// HDF5
#include "BlitzHDF5Helper.hh"
// Blitz++
#include <blitz/array.h>
// Fast_PD Solver
#include "Fast_PD/Fast_PD.h"
// Visualizations
#include "FlowToImage.hh"

/*
 * compute the dense control point graph
 */
void computeDenseControlPointGraph(  int nRows, int nCols, 
	blitz::Array< blitz::TinyVector<int,2>, 1>& nodes,
	blitz::Array< blitz::TinyVector<int,2>, 1>& edges)
{

  int nNodes = nRows*nCols;
  int nEdges = (nRows-1)*(nCols-1)*2 + (nRows-1) + (nCols-1);
  nodes.resize(nNodes);
  edges.resize(nEdges);

  /*
   * compute nodes and edges
   */
  //take all image pixels as nodes
  //add edges to each pixel's right and lower neighbor
  int idx = 0;
  for(int row=0; row < nRows; ++row){
    for(int col=0; col<nCols; ++col){
      int i = row*nCols+col;
      nodes(i)=blitz::TinyVector<int,2>(row,col);	  
      if (col < nCols-1) {
        edges(idx) = blitz::TinyVector<int,2>(i,i+1);		// add right neighbor
        idx += 1;
      }
      if (row < nRows-1) {
        edges(idx) = blitz::TinyVector<int,2>(i,i+nCols);	// add lower neighbor
        idx += 1;
      }
    }
  } 

}

/*
 * compute dense displacement hypothesis
 */
void computeDenseDisplacementHypotheses( int radius, blitz::Array< blitz::TinyVector<int,2>, 1>& labels)
{
  
  int nLabels=(radius*2+1)*(radius*2+1);
  labels.resize(nLabels);
  int idx=0;
  for(int i = -radius; i <= radius; ++i){
    for(int j = -radius; j <= radius; ++j, ++idx){
      labels(idx) = blitz::TinyVector<int,2>(i,j);
    }
  }

}

/*
 * compute unary costs (SSD)
 */
void computeUnaryCostsSSD( const blitz::Array< blitz::TinyVector<int,2>, 1>& nodes,
			    const blitz::Array< blitz::TinyVector<int,2>, 1>& labels,
			    const blitz::Array<float,2>& srcIm,
			    const blitz::Array<float,2>& trgIm,
			    blitz::Array< float, 2>& unaryCosts )
{   
  int nRows=srcIm.extent(0);
  int nCols=srcIm.extent(1);
  int nNodes=nodes.extent(0);
  int nLabels=labels.extent(0);

  unaryCosts.resize(nLabels,nNodes);

  for(int p=0; p < nNodes; ++p)
    for(int l=0; l < nLabels; ++l){
      blitz::TinyVector<int,2> srcPos = nodes(p);
      blitz::TinyVector<int,2> trgPos= srcPos + labels(l);
      if(trgPos(0)<0) trgPos(0)=0;
      if(trgPos(0)>=nRows) trgPos(0)=nRows-1;
      if(trgPos(1)<0) trgPos(1)=0;
      if(trgPos(1)>=nCols) trgPos(1)=nCols-1;
      float temp = srcIm(srcPos)-trgIm(trgPos);
      float dist = temp*temp;
      unaryCosts(l,p) = dist/255.0;
    }
}

/*
 * compute unary costs (NCC)
 */
void computeUnaryCostsNCC( const blitz::Array< blitz::TinyVector<int,2>, 1>& nodes,
			    const blitz::Array< blitz::TinyVector<int,2>, 1>& labels,
			    const blitz::Array<float,2>& srcIm,
			    const blitz::Array<float,2>& trgIm,
			    blitz::Array< float, 2>& unaryCosts,
			    int patch_radius=3 )
{
  int nRows = srcIm.extent(0);
  int nCols = srcIm.extent(1);
  int nNodes = nodes.extent(0);
  int nLabels = labels.extent(0);

  // number of pixels inside a patch
  float inv_patch_size = 1.f/((2*patch_radius+1)*(2*patch_radius+1));

  unaryCosts.resize(nLabels,nNodes);

  //zero padding
  int nRows2 = nRows + patch_radius*2;
  int nCols2 = nCols + patch_radius*2;

  blitz::Array<float,2> srcIm2(nRows2, nCols2);
  blitz::Array<float,2> trgIm2(nRows2, nCols2);
  srcIm2 = 0;
  trgIm2 = 0;
  srcIm2(blitz::Range(patch_radius, nRows2-1-patch_radius), blitz::Range(patch_radius, nCols2-1-patch_radius)) = srcIm;
  trgIm2(blitz::Range(patch_radius, nRows2-1-patch_radius), blitz::Range(patch_radius, nCols2-1-patch_radius)) = trgIm;

  for(int p=0; p<nNodes; ++p)
    for(int l=0; l<nLabels; ++l){
      blitz::TinyVector<int,2> srcPos=nodes(p);
      //offset
      srcPos(0)+=patch_radius;
      srcPos(1)+=patch_radius;

      blitz::TinyVector<int,2> trgPos= srcPos + labels(l);
      if(trgPos(0) < patch_radius) trgPos(0) = patch_radius;
      if(trgPos(0) >= nRows2-patch_radius) trgPos(0) = nRows2-patch_radius-1;
      if(trgPos(1) < patch_radius) trgPos(1) = patch_radius;
      if(trgPos(1) >= nCols2-patch_radius) trgPos(1) = nCols2-patch_radius-1;
      //go for ncc
      float mean1=0.0, mean2=0.0;
      for(int i=-patch_radius; i <= patch_radius; ++i) {
        for(int j=-patch_radius; j <= patch_radius; ++j){
          mean1+=srcIm2(srcPos(0)+i, srcPos(1)+j);
          mean2+=trgIm2(trgPos(0)+i, trgPos(1)+j);
        }
      }
      mean1 *= inv_patch_size;
      mean2 *= inv_patch_size;
      //demean and variance
      float sum=0.0;
      float sigma1=0.0, sigma2=0.0;
      for(int i=-patch_radius; i <= patch_radius; ++i) {
        for(int j=-patch_radius; j <= patch_radius; ++j){
          float val1= srcIm2(srcPos(0)+i, srcPos(1)+j)-mean1;
          float val2= trgIm2(trgPos(0)+i, trgPos(1)+j)-mean2;
          sum+=val1*val2;
          sigma1+=val1*val1;
          sigma2+=val2*val2;
        }
      }
      sum /= sqrt(sigma1*sigma2)+0.00001;
      unaryCosts(l,p) = -1.0*sum;
    }
}

/*======================================================================*/
/*! 
 *   computes the norm of vector 
 *
 *   \param a  vector
 *
 *   \exception none
 *
 *   \return norm of vector
 */
/*======================================================================*/
template< typename T, int I>
inline
T mynorm(const blitz::TinyVector<T,I>& a)
{
  return sqrt(blitz::dot(a,a));
}

float distL2(blitz::TinyVector<float,2> v1, blitz::TinyVector<float,2> v2)
{
  return mynorm( blitz::TinyVector<float,2>(v1-v2));
}

/*
 * compute pairwise costs (L2 norm)
 */
void computePairwiseCostsL2( const blitz::Array< blitz::TinyVector<int,2>, 1>& labels,
			     blitz::Array< float, 2>& pairwiseCosts )
{
  const int nLabels = labels.extent(0);
  pairwiseCosts.resize(nLabels, nLabels);

  for(int l0=0; l0 < nLabels; ++l0){
    for(int l1=0; l1 < nLabels; ++l1){
      blitz::TinyVector<float,2> vec1=labels(l0);
      blitz::TinyVector<float,2> vec2=labels(l1);
      pairwiseCosts(l0,l1) = distL2(vec1,vec2);
    }
  }
    
}

/*****************************************************************************
 *
 * interpolate
 *
 ******************************************************************************/
inline
float nnInterpolation(blitz::TinyVector<int,2> pos, const blitz::Array<float,2>& src)
{
  if(pos(0)<0 || pos(1)<0 || pos(0)>=src.extent(0) || pos(1)>=src.extent(1))
    return 0.0;
  return src(pos);
}

/*****************************************************************************
 *
 * warp image back
 *
 ******************************************************************************/
inline
void warpImage(	blitz::Array<float,2>& warpedBackTrg,
				const blitz::Array<float,2>& trg,
				const blitz::Array<blitz::TinyVector<int,2>,2 >& deformationField)
{  
  for(int row=0; row<trg.extent(0); ++row)
    for(int col=0; col<trg.extent(1);++col){
      blitz::TinyVector<int,2> trgPos = blitz::TinyVector<int,2>(row,col) + deformationField(row,col);
      warpedBackTrg(row,col) = nnInterpolation(trgPos, trg);
    }
}

/*****************************************************************************
 *
 * generate grid
 *
 ******************************************************************************/
inline
blitz::Array<float,2> getGrid(int nRows, int nCols, int d=5)
{
  blitz::Array<float,2> grid(nRows,nCols);
  grid = 0;
  grid(blitz::Range(0,nRows-1,d), blitz::Range::all()) = 255.0;
  grid(blitz::Range::all(), blitz::Range(0,nCols-1,d)) = 255.0;
  return grid;
}



int main(int argc, char** args) 
{
  /*
   * read console input arguments
   */
  std::string fileName1;
  std::string fileName2;

  // similarityMeasure
  // 0: SSD
  // 1: NCC
  int similarityMeasure = 0;
  float lambda = 0.0001;

  if (argc <= 2 || argc >= 6) {
    std::cout << "usage: ./elastic inputfile1 inputfile2 [lambda] [NCC,SSD]\n"
                 "examples:\n"
                 "  ./elastic square.h5 circle.h5\n"
                 "  ./elastic zebra1.h5 zebra2.h5 0.3 NCC\n";
    return 1;
  } else {
    fileName1 = args[1];
    fileName2 = args[2];
    if (argc > 3) {
      lambda = atof(args[3]);
    }
    if (argc > 4) {
      std::string tmp = args[4];
      if( "SSD" == tmp )
        similarityMeasure = 0;
      else if( "NCC" == tmp )
        similarityMeasure = 1;
      else{
        std::cout << "Option similarityMeasure = " << tmp << " not supported!" << std::endl;
        return 1;
      }
    }
  }
  std::cout << "fileName1: " << fileName1 << std::endl;
  std::cout << "fileName2: " << fileName2 << std::endl;
  std::cout << "lambda: " << lambda << std::endl;
  std::cout << "similarityMeasure: " << (similarityMeasure ? "NCC" : "SSD") << std::endl;


  /*
   * read data
   */
  std::string resultFileName="result.h5";
  std::string datasetName="image";

  blitz::Array<float, 2> srcIm;
  blitz::Array<float, 2> trgIm;
  std::cout << "read images" << std::endl;
  readHDF5toBlitz(fileName1, datasetName, srcIm);
  readHDF5toBlitz(fileName2, datasetName, trgIm);


  /*
   * compute dense control point graph
   */
  const int nRows = trgIm.extent(0);
  const int nCols = trgIm.extent(1);
  blitz::Array< blitz::TinyVector<int,2>, 1> nodes;
  blitz::Array< blitz::TinyVector<int,2>, 1> edges;

  std::cout << "compute dense control point graph" << std::endl;
  computeDenseControlPointGraph( nRows, nCols, nodes, edges);
  int nNodes = nodes.extent(0);
  int nEdges = edges.extent(0);


  /*
   * compute dense displacement hypotheses
   */
  const int radius = 10;
  blitz::Array< blitz::TinyVector<int,2>, 1> labels;

  std::cout << "compute dense displacement hypotheses" << std::endl;
  computeDenseDisplacementHypotheses( radius, labels);
  int nLabels = labels.extent(0);


  /*
   * compute unary costs (SSD)
   */
  blitz::Array< float, 2> unaryCosts;

  std::cout << "compute unary costs" << std::endl;
  if (similarityMeasure == 0) {
    std::cout << "\tSSD" << std::endl;
    computeUnaryCostsSSD( nodes, labels, srcIm, trgIm, unaryCosts );
  } else {
    if (similarityMeasure == 1) {
      std::cout << "\tNCC" << std::endl;
      computeUnaryCostsNCC( nodes, labels, srcIm, trgIm, unaryCosts );
    }
  }


  /*
   * compute pairwise costs (L2 norm)
   */
  blitz::Array< float, 2> pairwiseCosts;

  std::cout << "compute pairwise costs" << std::endl;
  computePairwiseCostsL2( labels, pairwiseCosts );


  /*
   * run FAST_PD solver
   */
  pairwiseCosts *= lambda;

  const int max_iterations = 100;
  blitz::Array<float, 1> edgeWeights(nEdges);
  edgeWeights = 1.0;	

  std::cout << "set up FAST_PD solver" << std::endl;
  float* _unaryCosts = reinterpret_cast<float*>(unaryCosts.dataFirst());
  int* _edges  = reinterpret_cast<int*>(edges.dataFirst());
  float* _pairwiseCosts   = reinterpret_cast<float*>(pairwiseCosts.dataFirst());
  float* _edgeWeights = reinterpret_cast<float*>(edgeWeights.dataFirst());

  CV_Fast_PD pd( nNodes, nLabels, _unaryCosts,
      nEdges, _edges, _pairwiseCosts,
      max_iterations, _edgeWeights );

  std::cout << "run FAST_PD solver" << std::endl;
  pd.run();

  std::cout << "extract optimal labels and deformation field" << std::endl;
  blitz::Array<int, 1> optimalLabels(nNodes);
  for( int i = 0; i < nNodes; ++i ) {
    optimalLabels(i)=pd._pinfo[i].label;
  }

  std::cout << "optimal labels:" << std::endl;
  std::cout << optimalLabels << std::endl;

  blitz::Array<blitz::TinyVector<int,2>,2 > deformationField( srcIm.shape());
  for(int i=0; i < nNodes; ++i){
    deformationField(nodes(i)) = labels(optimalLabels(i));
  }

  // color visualization of the deformation field
  blitz::Array< blitz::TinyVector<float,3>, 2> rgbDeformationField;
  float scaleFactor = 1/sqrt(2*float(radius*radius));	// maximum displacement
  flowToImage( deformationField, rgbDeformationField, scaleFactor); 


  /*
   * back-warp image and grid
   */
  std::cout << "warp image and grid" << std::endl;
  blitz::Array<float,2> warpedBackTrgIm(trgIm.extent(0), trgIm.extent(1));
  warpImage(warpedBackTrgIm, trgIm, deformationField);

  blitz::Array<float,2> grid=getGrid(trgIm.extent(0), trgIm.extent(1));
  blitz::Array<float,2> warpedBackGrid(trgIm.extent(0), trgIm.extent(1));
  warpImage(warpedBackGrid, grid, deformationField);


  /*
   * write result
   */
  std::cout << "write result" << std::endl;
  writeBlitzToHDF5(trgIm, "trgIm", resultFileName);
  writeBlitzToHDF5(srcIm, "srcIm", resultFileName);
  writeBlitzToHDF5(warpedBackTrgIm, "warpedBackTrgIm", resultFileName);
  writeBlitzToHDF5(grid, "grid", resultFileName);
  writeBlitzToHDF5(warpedBackGrid, "warpedBackGrid", resultFileName);
  writeBlitzToHDF5(rgbDeformationField, "rgbDeformationField", resultFileName);

  std::cout << "done." << std::endl;
}
