/*
 * main_plaquette.cpp
 *
 * Lukas Mazur, 10 Apr 2018 (original Author for Plaquette)
 * Vishal Rao , 06 Jan 2024 (generalisation from plaquette to wilson loop)
 *
 * This is just an example how a very very basic program works. Look at src/testing/main_GeneralOperatorTest.cpp to see
 * how to write more advanced GPU code.
 *
 */
#include "../simulateqcd.h"
/* #include "../modules/hyp/hypParameters.h" was already included in ../modules/hypSmearing.h */
/* #include "../modules/hyp/smearParameters.h" */
#include "../modules/hyp/hypSmearing.h"
/* #include "../modules/hyp/hypSmearing.cpp" */
#include <algorithm>
#include <iostream>
#include <string>
#include <vector>
#include <cmath>
#define PREC double
template<class floatT,size_t HaloDepth>
struct CalcWilson{

    //Gauge accessor to access the gauge field
    SU3Accessor<floatT> SU3Accessor;
    int _length=1;
    int _height=1;

    //Constructor to initialize all necessary members.
    CalcWilson(Gaugefield<floatT,true,HaloDepth> &gauge, int &length, int &height) : SU3Accessor(gauge.getAccessor()), _length(length),_height(height){
    }

__device__ __host__ SU3<floatT> HalfLoop(int &length, int &height,gSite  site,int &mu, int &nu)
{
    // we have not taken site variable as reference as if it were take as reference, it will update the values at same memory location
    SU3<floatT> temp;
    typedef GIndexer<All, HaloDepth> GInd;
    temp = SU3Accessor.getLink(GInd::getSiteMu(site, mu));
    for (int len=1;len<length;len++){
    site=GInd::site_up(site, mu);
    temp *= SU3Accessor.getLink(GInd::getSiteMu(site, mu));
    }

    site=GInd::site_up(site,mu);
    for (int hei=1;hei<=height;hei++){
    temp *= SU3Accessor.getLink(GInd::getSiteMu(site, nu));
    site=GInd::site_up(site,nu);
    }
    return temp; 
 }   
    //This is the operator that is called inside the Kernel
    __device__ __host__ floatT operator()(gSite site) {

        /// We need to choose the type of indexer. The first template is the layout of the lattice.
        typedef GIndexer<All, HaloDepth> GInd;
        int l=_length;
        int h=_height;
        /// Define a SU(3) matrix
        SU3<floatT> temp1,temp2;

        floatT result = 0;
        for (int nu = 1; nu < 3; nu++) 
        {
            for (int mu = 0; mu < nu; mu++) 
            {
                temp1=HalfLoop(l,h,site,mu,nu);
                temp2=HalfLoop(h,l,site,nu,mu);
                result+=tr_d(temp1*dagger(temp2));
            }
        }
        return result/3.0;
        }

};

template<class floatT,size_t HaloDepth>
struct CalcWilsonBresenham{
    int _z;
    int _size;
    int *_x_bresenham, *_y_bresenham;
    SU3Accessor<floatT> SU3Accessor;
    CalcWilsonBresenham(Gaugefield<floatT,true,HaloDepth> &gauge, int *x_bresenham_vec,int *y_bresenham_vec,int &size, int &z) : SU3Accessor(gauge.getAccessor()), _x_bresenham(x_bresenham_vec),_y_bresenham(y_bresenham_vec),_size(size),_z(z){}
  //This is the operator that is called inside the Kernel


  __device__  __host__  SU3<floatT> HalfLoopBresenhamLower(int *x_bresenham_vec,int *y_bresenham_vec,int &size, int &z, gSite site, int &mu, int &nu)
    {
        /* int n_diagonal_coordinates=0; //could be useful for later improvements in code */
        int delta_x;
        int delta_y;
        SU3<floatT> temp;
        temp=SU3<floatT>(1, 0, 0, 0, 1, 0, 0, 0, 1); //initialize to identity matrix
        int cmax_dir,cmin_dir;
        int *cmax_arr,*cmin_arr;
        if(x_bresenham_vec[size-1]>=y_bresenham_vec[size-1]){cmax_arr=x_bresenham_vec,cmin_arr=y_bresenham_vec,cmax_dir=mu, cmin_dir=nu;}
        else {cmax_arr=y_bresenham_vec,cmin_arr=x_bresenham_vec,cmax_dir=nu, cmin_dir=mu;}
        typedef GIndexer<All, HaloDepth> GInd;

        for (int iy=0;iy<size-1;iy++) //accessing all coordinates, one minus for last link in z direction
        {
            delta_x=cmax_arr[iy+1]-cmax_arr[iy];
            delta_y=cmin_arr[iy+1]-cmin_arr[iy];
            if (delta_x==1 && delta_y==1)
            {
                temp*=SU3Accessor.getLink(GInd::getSiteMu(site,cmax_dir));
                site=GInd::site_up(site,cmax_dir);
                temp*=SU3Accessor.getLink(GInd::getSiteMu(site,cmin_dir));
                site=GInd::site_up(site,cmin_dir);
            }
            else if (delta_x==1 && delta_y==0)
            {
                temp*=SU3Accessor.getLink(GInd::getSiteMu(site,cmax_dir));
                site=GInd::site_up(site,cmax_dir);
            }
            else if (delta_x==0 && delta_y==1) //wrong, for case of (1,2) the bresenham vectors are x={0,1,1} and y={0,1,2}, dx={1,0} and dy={1,1}, for last link, dx=0 and dy=1, we need to take step in cmax and not cmin as intructed below, 
            {
                temp*=SU3Accessor.getLink(GInd::getSiteMu(site,cmin_dir));
                site=GInd::site_up(site,cmin_dir);
            }
        }
        for (int iz=0;iz<z;iz++) //iz is iterator in z direction
        {
        temp*=SU3Accessor.getLink(GInd::getSiteMu(site,3-mu-nu)); //3-cmax_dir-cmin_dir is the fictitious time direction
        site=GInd::site_up(site,3-mu-nu); //half loop will complete at end of loop.
          }
    return temp;
    }


  __device__  __host__  SU3<floatT> HalfLoopBresenhamUpper(int *x_bresenham_vec,int *y_bresenham_vec,int &size, int &z, gSite site, int &mu, int &nu)
    {
        /* int n_diagonal_coordinates=0; //could be useful for later improvements in code */
        int delta_x;
        int delta_y;
        SU3<floatT> temp;
        temp=SU3<floatT>(1, 0, 0, 0, 1, 0, 0, 0, 1); //initialize to identity matrix
        int cmax_dir,cmin_dir;
        int *cmax_arr,*cmin_arr;
        if(x_bresenham_vec[size-1]>=y_bresenham_vec[size-1]){cmax_arr=x_bresenham_vec,cmin_arr=y_bresenham_vec,cmax_dir=mu, cmin_dir=nu;}
        else {cmax_arr=y_bresenham_vec,cmin_arr=x_bresenham_vec,cmax_dir=nu, cmin_dir=mu;}
        typedef GIndexer<All, HaloDepth> GInd;

        for (int iz=0;iz<z;iz++) //iz is iterator in z direction
        {
        temp*=SU3Accessor.getLink(GInd::getSiteMu(site,3-mu-nu)); //3-cmax_dir-cmin_dir is the fictitious time direction
        site=GInd::site_up(site,3-mu-nu); //half loop will complete at end of loop.
          }
        for (int iy=0;iy<size-1;iy++) //accessing all coordinates, one minus for last link in z direction
        {
            delta_x=cmax_arr[iy+1]-cmax_arr[iy];
            delta_y=cmin_arr[iy+1]-cmin_arr[iy];
            if (delta_x==1 && delta_y==1)
            {
                temp*=SU3Accessor.getLink(GInd::getSiteMu(site,cmax_dir));
                site=GInd::site_up(site,cmax_dir);
                temp*=SU3Accessor.getLink(GInd::getSiteMu(site,cmin_dir));
                site=GInd::site_up(site,cmin_dir);
            }
            else if (delta_x==1 && delta_y==0)
            {
                temp*=SU3Accessor.getLink(GInd::getSiteMu(site,cmax_dir));
                site=GInd::site_up(site,cmax_dir);
            }
            else if (delta_x==0 && delta_y==1) //wrong, for case of (1,2) the bresenham vectors are x={0,1,1} and y={0,1,2}, dx={1,0} and dy={1,1}, for last link, dx=0 and dy=1, we need to take step in cmax and not cmin as intructed below, 
            {
                temp*=SU3Accessor.getLink(GInd::getSiteMu(site,cmin_dir));
                site=GInd::site_up(site,cmin_dir);
            }
        }
    return temp;
    }


    __device__ __host__ floatT operator()(gSite site) {

        /// We need to choose the type of indexer. The first template is the layout of the lattice.
        typedef GIndexer<All, HaloDepth> GInd;
        /// Define a SU(3) matrix
        SU3<floatT> temp1,temp2;
        floatT result = 0;
        for (int nu = 1; nu < 3; nu++) 
        {
            for (int mu = 0; mu < nu; mu++) 
            {
                temp1=HalfLoopBresenhamLower(_x_bresenham,_y_bresenham,_size,_z,site,mu,nu); //mu is direction for x coordinates, //nu is for y coordinates 
                temp2=HalfLoopBresenhamUpper(_x_bresenham,_y_bresenham,_size,_z,site,mu,nu); 
                result+=tr_d(temp1*dagger(temp2));

            }
        }
        for (int nu = 1; nu < 3; nu++) 
        {
            for (int mu = 0; mu < nu; mu++) 
            {
                temp1=HalfLoopBresenhamLower(_y_bresenham,_x_bresenham,_size,_z,site,mu,nu); 
                temp2=HalfLoopBresenhamUpper(_y_bresenham,_x_bresenham,_size,_z,site,mu,nu); 
                result+=tr_d(temp1*dagger(temp2));

            }
        }
        return result/6.0;
        }


};

//Function to compute the wilson loop using the above struct CalcWilson.
template<class floatT, size_t HaloDepth>
floatT WilsonLoop(Gaugefield<floatT,true, HaloDepth> &gauge, LatticeContainer<true,floatT> &redBase, int &l, int &h){

    typedef GIndexer<All,HaloDepth> GInd;
    const size_t elems = GInd::getLatData().vol4;
    //Make sure, redBase is large enough
    redBase.adjustSize(elems);
// we will iterate the process of finding wilson loop over whole bulk.
    redBase.template iterateOverBulk<All, HaloDepth>(CalcWilson<floatT, HaloDepth>(gauge, l,h));

    //Do the final reduction
    floatT Wloop;
    redBase.reduce(Wloop, elems);

    //Normalize the result
    const int n_colors=3;
    Wloop /= (GInd::getLatData().globalLattice().mult()*n_colors); 
    return Wloop;
}
//defining function for bresenham
template<class floatT, size_t HaloDepth>
floatT WilsonLoopBresenham(Gaugefield<floatT,true, HaloDepth> &gauge, LatticeContainer<true,floatT> &redBase, int *d_x_bresenham, int *d_y_bresenham,int &size, int &z){

    typedef GIndexer<All,HaloDepth> GInd;
    const size_t elems = GInd::getLatData().vol4;
    //Make sure, redBase is large enough
    redBase.adjustSize(elems);
// we will iterate the process of finding wilson loop over whole bulk.
    redBase.template iterateOverBulk<All, HaloDepth>(CalcWilsonBresenham<floatT, HaloDepth>(gauge,d_x_bresenham,d_y_bresenham,size,z));

    //Do the final reduction
    floatT Wloop;
    redBase.reduce(Wloop, elems);

    //Normalize the result
    const int n_colors=3;
    Wloop /= (GInd::getLatData().globalLattice().mult()*n_colors); //
    return Wloop;
}

//
// I am going to paste the hypSmearing.cpp file content here for the moment.


// also call updateAll()
template<class floatT, bool onDevice, size_t HaloDepth, CompressionType comp>
void HypSmearing<floatT, onDevice, HaloDepth, comp>::Su3Unitarize(Gaugefield<floatT, onDevice, HaloDepth, comp> &gauge_out, Gaugefield<floatT, onDevice, HaloDepth, comp> &gauge_base){

    HypStaple<floatT, HaloDepth, comp, 4> su_3_unitarize(gauge_out.getAccessor(), gauge_base.getAccessor(), gauge_base.getAccessor(), gauge_base.getAccessor());
    gauge_out.iterateOverBulkAllMu(su_3_unitarize);
    if(update_all)gauge_out.updateAll();
}

template<class floatT, bool onDevice, size_t HaloDepth, CompressionType comp>
void HypSmearing<floatT, onDevice, HaloDepth, comp>::SmearAll(Gaugefield<floatT, onDevice, HaloDepth, comp> &gauge_out) {

    // create level 1 fields
    _dummy.iterateOverBulkAllMu(staple3_lvl1_10);
    _gauge_lvl1_10 = (1-params.alpha_3) * _gauge_base + params.alpha_3/2 * _dummy;
    Su3Unitarize(_gauge_lvl1_10, _gauge_base);

    _dummy.iterateOverBulkAllMu(staple3_lvl1_20);
    _gauge_lvl1_20 = (1-params.alpha_3) * _gauge_base + params.alpha_3/2 * _dummy;
    Su3Unitarize(_gauge_lvl1_20, _gauge_base);

    _dummy.iterateOverBulkAllMu(staple3_lvl1_30);
    _gauge_lvl1_30 = (1-params.alpha_3) * _gauge_base + params.alpha_3/2 * _dummy;
    Su3Unitarize(_gauge_lvl1_30, _gauge_base);

    _dummy.iterateOverBulkAllMu(staple3_lvl1_21);
    _gauge_lvl1_21 = (1-params.alpha_3) * _gauge_base + params.alpha_3/2 * _dummy;
    Su3Unitarize(_gauge_lvl1_21, _gauge_base);

    _dummy.iterateOverBulkAllMu(staple3_lvl1_31);
    _gauge_lvl1_31 = (1-params.alpha_3) * _gauge_base + params.alpha_3/2 * _dummy;
    Su3Unitarize(_gauge_lvl1_31, _gauge_base);

    _dummy.iterateOverBulkAllMu(staple3_lvl1_32);
    _gauge_lvl1_32 = (1-params.alpha_3) * _gauge_base + params.alpha_3/2 * _dummy;
    Su3Unitarize(_gauge_lvl1_32, _gauge_base);

    // now that we have level 1 fields, create level 2 staples
    // note:  the order of the gauge fields goes in ascending order (10 < 20 < 30, 10 < 21 < 31, 20 < 21 < 32, 30 < 31 < 32)
    // this is ASSUMED by HypStaple<floatT, HaloDepth, comp, 2>; DO NOT change this order without also modifying HypStaple<floatT, HaloDepth, comp, 2> and threeLinkStaple_second_level
    HypStaple<floatT, HaloDepth, comp, 2> staple3_lvl2_0(_gauge_lvl1_10.getAccessor(), _gauge_lvl1_20.getAccessor(), _gauge_lvl1_30.getAccessor(), _dummy.getAccessor(), 0);
    HypStaple<floatT, HaloDepth, comp, 2> staple3_lvl2_1(_gauge_lvl1_10.getAccessor(), _gauge_lvl1_21.getAccessor(), _gauge_lvl1_31.getAccessor(), _dummy.getAccessor(), 1);
    HypStaple<floatT, HaloDepth, comp, 2> staple3_lvl2_2(_gauge_lvl1_20.getAccessor(), _gauge_lvl1_21.getAccessor(), _gauge_lvl1_32.getAccessor(), _dummy.getAccessor(), 2);
    HypStaple<floatT, HaloDepth, comp, 2> staple3_lvl2_3(_gauge_lvl1_30.getAccessor(), _gauge_lvl1_31.getAccessor(), _gauge_lvl1_32.getAccessor(), _dummy.getAccessor(), 3);

    //second level fields
    _dummy.iterateOverBulkAllMu(staple3_lvl2_0);
    _gauge_lvl2_0 = (1-params.alpha_2) * _gauge_base + params.alpha_2/4 * _dummy;
    Su3Unitarize(_gauge_lvl2_0, _gauge_base);

    _dummy.iterateOverBulkAllMu(staple3_lvl2_1);
    _gauge_lvl2_1 = (1-params.alpha_2) * _gauge_base + params.alpha_2/4 * _dummy;
    Su3Unitarize(_gauge_lvl2_1, _gauge_base);

    _dummy.iterateOverBulkAllMu(staple3_lvl2_2);
    _gauge_lvl2_2 = (1-params.alpha_2) * _gauge_base + params.alpha_2/4 * _dummy;
    Su3Unitarize(_gauge_lvl2_2, _gauge_base);

    _dummy.iterateOverBulkAllMu(staple3_lvl2_3);
    _gauge_lvl2_3 = (1-params.alpha_2) * _gauge_base + params.alpha_2/4 * _dummy;
    Su3Unitarize(_gauge_lvl2_3, _gauge_base);

    // now that we have level 2 fields, create level 3 staple
    HypStaple<floatT, HaloDepth, comp, 1> staple3_lvl3(_gauge_lvl2_0.getAccessor(), _gauge_lvl2_1.getAccessor(), _gauge_lvl2_2.getAccessor(), _gauge_lvl2_3.getAccessor());

    _dummy.iterateOverBulkAllMu(staple3_lvl3);

    // OLD VERSION, (MAYBE) DOES NOT WORK FOR SOME REASON
    //_gauge_lvl2_0 = (1-params.alpha_1) * _gauge_base + params.alpha_1/6 * _dummy; //reused _gauge_lvl2_0
    //Su3Unitarize(_gauge_lvl2_0);

    // NEW VERSION, USES EXTRA FIELD RATHER THAN REUSE _gauge_lvl2_0
    gauge_out = (1-params.alpha_1) * _gauge_base + params.alpha_1/6 * _dummy; //reused _gauge_lvl2_0
    Su3Unitarize(gauge_out, _gauge_base);

}

#define CLASS_INIT(floatT,HALO) \
  template class HypSmearing<floatT,true,HALO,R18>;


INIT_PH(CLASS_INIT)

std::vector<std::vector<int>> bresenham_paper(int x,int y)
{
    std::vector<std::vector<int>> result;
    std::vector<int> result_x;
    std::vector<int> result_y;
    int cmax, cmin;
    if(x>=y){cmax=x;cmin=y;}
    else{cmax=y;cmin=x;}
    int cmax2=2*cmax;
    int cmin2=2*cmin;
    int chi=cmin2-cmax;
    int x_update=0;
    int y_update=0;
    result_x.push_back(0);
    result_y.push_back(0);
    for(int i=0;i<cmax;i++)
    {
        x_update+=1;     
        result_x.push_back(x_update);
        if(chi>=0)
        {
            chi-=cmax2;
            y_update+=1;
            result_y.push_back(y_update);
        }
        else {
            result_y.push_back(y_update);
        }
        chi+=cmin2;

    }
    if(x>=y){
    result.push_back(result_x);
    result.push_back(result_y);
    }
    else{
    result.push_back(result_y);
    result.push_back(result_x);
    }
    return result;
}

int main(int argc, char *argv[]) {

    stdLogger.setVerbosity(DEBUG);

    /// Initialize parameter class. This class can also read parameter from textfiles!
    LatticeParameters param;
    /// Initialize the Lattice dimension
    const int LatDim[] = {32, 32, 32, 8};

    const int NodeDim[] = {1, 1, 1, 1};

    /// Just pass these dimensions to the parameter class
    param.latDim.set(LatDim);
    param.nodeDim.set(NodeDim);

    /// Initialize a timer
    StopWatch<true> timer;

    /// Initialize the CommunicationBase. This class handles the communitation between different Cores/GPU's.
    CommunicationBase commBase(&argc, &argv, true);
    commBase.init(param.nodeDim());

    const size_t HaloDepth = 1; //since it is not multi-gpu code


    /// highlight the output differently.
    rootLogger.info("Initialize Lattice");
    /// Initialize the Indexer on GPU and CPU.
    initIndexer(HaloDepth,param,commBase);
    typedef GIndexer<All,HaloDepth> GInd;


    rootLogger.info("Initialize Gaugefield");
    Gaugefield<PREC, true,HaloDepth> gauge(commBase);

    typedef double floatT;
    GaugeAction<floatT, true, HaloDepth, R18> gaugeaction(gauge);
    /// Initialize gaugefield with unity-matrices.
    gauge.one();

    /// Initialize LatticeContainer. This is in principle the "array", where the values of the plaquette are
    /// stored which are summed up in the end
    LatticeContainer<true,PREC> redBase(commBase);
    /// We need to tell the Reductionbase how large our Array will be
    redBase.adjustSize(GInd::getLatData().vol4);
    int eqm_point=155;
    int last_point=1150;
    int node;
    sscanf(argv[1], "%d", &node);
    //sscanf(argv[2], "%d", &height);
    for (int i=eqm_point; i<=last_point;i+=5)
        {
            rootLogger.info("Read configuration");
            gauge.readconf_nersc("/root/project1/build_SIMULATeQCD/applications/try_20_output_dir/after_eqm/node"+std::to_string(node)+"/l328f21b6285m0039185m0783706a_"+std::to_string(node)+"."+std::to_string(i));

            gauge.updateAll();
            Gaugefield<PREC, true,HaloDepth> gauge_out(commBase);
            gauge_out.one(); //initialize the gauge field variable
        for(int no_smearing_steps=0;no_smearing_steps<=30;no_smearing_steps+=30)
            {
                //we shall make a variable which stores gauge configuration
                                 //doing smearing 5,10,15.... times, (loop will run just 5 times where we left of using gauge=gauge_out)
                if (no_smearing_steps!=0){
                    for(int j=1;j<=30;j++)
                    {
                        HypSmearing<floatT, true, HaloDepth, R18> hypsmearing(gauge); //smearing of gauge
                        hypsmearing.SmearAll(gauge_out); //and storing at gauge_out
                        gauge=gauge_out;
                    }
                    }
                const int Ns= LatDim[0];
                for (int length=1; length<Ns; length++)
                {      
                    for (int height=1; height<Ns; height++)
                    {
                        PREC Wloop = 0;
                        /* timer.start(); */
                        /// compute wilson loop with smeared gauge_out
                        Wloop = WilsonLoop<PREC,HaloDepth>(gauge, redBase, length,height );
                        printf("Node:%d,Smear_count:%d,Length=%d,Height=%d,Wilson Loop:%1.15e\n",node,no_smearing_steps,length, height,Wloop);

                        /* rootLogger.info("Reduced Rectangle from rhmc : " ,  gaugeaction.rectangle()); */
                        /* rootLogger.info("Reduced plaquette from rhmc : " ,  gaugeaction.plaquette()); */
                    }
                }
                // we shall now implement bresenham algorithm to calculate the new wilson loops.int x_coord=length;
                for (int x=1;x<Ns;x++)
                {
                    for(int y=x;y<Ns;y++)
                    {
                        std::vector<std::vector<int>> xy_bresenham_vec;
                        xy_bresenham_vec=bresenham_paper(x,y);
                        int size_bresenham_vectors=xy_bresenham_vec[0].size();
                        std::vector<int> x_bresenham_vec=xy_bresenham_vec[0];
                        std::vector<int> y_bresenham_vec=xy_bresenham_vec[1];

                        //allocating memory for x_bresenham_vec  on gpu
                        int *d_x_bresenham,*d_y_bresenham;
                        cudaMalloc((void**)&d_x_bresenham,size_bresenham_vectors*sizeof(int));
                        cudaMalloc((void**)&d_y_bresenham,size_bresenham_vectors*sizeof(int));

                        //copying the bresenham vector to device
                        cudaMemcpy(d_x_bresenham,&x_bresenham_vec[0],size_bresenham_vectors*sizeof(int),cudaMemcpyHostToDevice);
                        cudaMemcpy(d_y_bresenham,&y_bresenham_vec[0],size_bresenham_vectors*sizeof(int),cudaMemcpyHostToDevice);
                        
                        for(int z=1;z<Ns;z++)
                        {
                            
                        PREC Wloop=WilsonLoopBresenham<PREC,HaloDepth>(gauge, redBase,d_x_bresenham,d_y_bresenham,size_bresenham_vectors,z );
                        printf("Node:%d,Smear_count:%d,Length=%f,Height=%d,Wilson Loop:%1.15e\n",node,no_smearing_steps,pow((x*x+y*y),0.5), z,Wloop);
                        }
                        cudaFree(d_x_bresenham);
                        cudaFree(d_y_bresenham);
                    }
                }



            }
        }
    return 0;
}
//