#lywu
#2022.09.12

import argparse
import numpy as np
import pandas as pd

if __name__ == '__main__':
    parser = argparse.ArgumentParser(description='2D free energy of REMD trajectories')
    parser.add_argument('-f', '--filenames', required=True, type=list,
                        help='list of input files with temperature')
    parser.add_argument('-t', '--reference_temperature', default='300.00', type=float,
                        help='reference temperature in reweighting, Default: %(default)s')
    parser.add_argument('-x', '--x_grid_num', default='50', type=int,
                        help='the number of grids on x-axis, Default: %(default)s')
    parser.add_argument('-y', '--y_grid_num', default='50', type=int,
                        help='the number of grids on y-axis, Default: %(default)s')
    parser.add_argument('-o', '--output', default='output.csv', type=str,
                        help='name of output file, Default: %(default)s')                                     
    args = parser.parse_args()

input_list = args.filenames
reference_temp = args.reference_temperature
x_grid_num = args.x_grid_num
y_grid_num = args.y_grid_num
output_ = args.output

def location(m,m_grid):
    if m % m_grid == 0:
        m_loc = m // m_grid - 1
    else:
        m_loc = m // m_grid   
    if m_loc == -1:
        m_loc = 0   
    return int(m_loc)

def count(data,re_factor=1):
    table = [[0 for i in range(x_grid_num)] for i in range(y_grid_num)]
    table_re = [[0 for i in range(x_grid_num)] for i in range(y_grid_num)]    
    total_num =len(data)
    x= data.loc[:,'x']
    y= data.loc[:,'y']
    for i in range(total_num):
        x_loc = location((x[i] - x_min),x_grid)
        y_loc = location((y[i] - y_min),y_grid)
        table[x_loc][y_loc] += 1
    for i in range(x_grid_num):
        for j in range(y_grid_num):
            table_re[i][j] = pow((table[i][j] / total_num),re_factor) * total_num
    return table_re

input_ = pd.read_csv('input_list.csv',sep=',')
xmin = []
ymin = []
xmax = []
ymax = []

for file in input_.loc[:,'file']:
    data_=pd.read_csv(file,sep=',')
    xmin.append(min(data_.loc[:,'x']))
    ymin.append(min(data_.loc[:,'y']))
    xmax.append(max(data_.loc[:,'x']))
    ymax.append(max(data_.loc[:,'y']))

x_min, x_max = min(xmin), max(xmax)
y_min, y_max = min(ymin), max(ymax)
x_grid = (x_max - x_min) / x_grid_num
y_grid = (y_max - y_min) / y_grid_num

count_n = [[0 for i in range(x_grid_num)] for i in range(y_grid_num)]
E =  [[0 for i in range(x_grid_num)] for i in range(y_grid_num)]

for i in range(len(input_)):
    data_=pd.read_csv(input_.loc[:,'file'][i],sep=',')
    temp = input_.loc[:,'temp'][i]
    re_factor = temp / reference_temp
    count_ = count(data_,re_factor)
    for m in range(x_grid_num):
        for n in range(y_grid_num):
            count_n[m][n] += count_[m][n]

    if i == int(len(input_)-1):
        N_sum = 0
        for m in range(x_grid_num):
            for n in range(y_grid_num):
                N_sum += count_n[m][n]

        maxE = 0
        for m in range(x_grid_num):
            for n in range(y_grid_num): 
                if count_n[m][n] != 0:
                    E[m][n] = -1.3717*np.log10(count_n[m][n]/N_sum) # kcal/mol
                    if E[m][n] > maxE:
                        maxE = E[m][n]
        
        count_out = ''
        for m in range(x_grid_num):
            for n in range(y_grid_num):
                x_loc = m * x_grid + x_min
                y_loc = n * y_grid + y_min
                if E[m][n] == 0:    
                    count_out += '%.4f,%.4f,%.8f\n' % (x_loc, y_loc, maxE)
                else:
                    count_out += '%.4f,%.4f,%.8f\n' % (x_loc, y_loc, E[m][n])
        with open(output_,'w') as f:
            f.write('x,y,e\n' + count_out)
        print('Job is done!')
    else:
        print('md' + str(i) + ' is done!')
        continue