import numpy as np
import matplotlib
import matplotlib.pyplot as plt
from sklearn.model_selection import GridSearchCV
from multivariate_gpr_cv import MultivariateGaussianProcessCV
from sklearn.model_selection import train_test_split, KFold
from sklearn.model_selection import PredefinedSplit
import math
import os
import sys
import time
from scipy.spatial.distance import cdist
from sklearn.metrics.pairwise import euclidean_distances
from sklearn.cluster import KMeans
from ase import Atoms
from ase import units
from ase.calculators.interface import Calculator
from ase.calculators.loggingcalc import LoggingCalculator
from ase.md.velocitydistribution import MaxwellBoltzmannDistribution
from ase.optimize import BFGS as BFGS_relaxation
from ase.optimize.sciopt import SciPyFminCG
from ase.md.nvtberendsen import NVTBerendsen
from ase.md.verlet import VelocityVerlet
from ase.md.langevin import Langevin
from scipy.optimize import approx_fprime
from sklearn.metrics.pairwise import euclidean_distances
import dft_utils as du
from joblib import Memory
memory = Memory(cachedir='temp/', verbose=0)

@memory.cache
def compute_dist_matrix(X):
    return euclidean_distances(X, None, squared=True)
charge_to_str = {8: 'O', 1: 'H', 6: 'C'}
str_to_charge = {'O': 8, 'H': 1, 'C': 6}
#load the initial position and velocity
start_pos= du.positions_from_xyz( 'init_pos.xyz', convert_to_bohr=True)[0,:9,:]
start_types = np.array([1,1,1,1,6,6,6,8,8])
velocities = du.positions_from_xyz('vel.xyz', convert_to_bohr=False)[0] /(0.02418884326505*units.fs/units.Bohr)
#load the base position for alignment during dynamics
base_types = np.array([1,1,1,1,6,6,6,8,8])
base_pos =  du.positions_from_xyz('./mp2.xyz', convert_to_bohr=True)[0,:9,:]
base_pos = base_pos + np.ones(base_pos.shape)*10.  # move to center of 20 bohr box.

#load the model
density_alpha_params = [2.21221629e-10]
density_gamma_params = list(1 / np.sqrt(2 * np.logspace(-11, -4, 20)))
energy_alpha_params = [2.21221629e-10]
energy_gamma_params = list(1 / np.sqrt(2 * np.logspace(-15, -5, 20)))
verbose =1
energy_kr = MultivariateGaussianProcessCV(cv_nfolds=5,
                                              krr_param_grid={"alpha": energy_alpha_params,
                                                              "gamma": energy_gamma_params},
                                              id=1,
                                              verbose=verbose,
                                              cluster_params=['-l h_vmem=100G'])
energy_kr.load('./train/pbe0_kr_0')
density_kr = MultivariateGaussianProcessCV(cv_nfolds=5,
                                                   krr_param_grid={"alpha": density_alpha_params,
                                                                   "gamma": density_gamma_params},
                                                   id=1,
                                                   verbose=verbose,
                                                   cluster_params=['-l h_vmem=100G'])
density_kr.load('./train/density_kr_2')


# Create the calculator object
class MLCalculator(Calculator):
    def __init__(self, clf_E, clf_n, base_pos, base_types, atoms=None, fd_eps=1e-8, fd_method='forward', **kwargs):
        #self.clf_mats = clf_mats
        self.clf_E = energy_kr
        self.clf_n = density_kr
        self.atoms = atoms
        self.fd_eps = fd_eps
        self.base_pos = base_pos
        self.base_types = base_types
        self.fd_method = fd_method
        self.calc = None
        self.mylog = []
        self.verbose = 0
        self.charges = np.loadtxt('a_num.txt')
        
        # Pseudo-pseudo-potentials parameters
        energy_type = 'pbe'
        descriptor_type = 'pot'
        self.gaussian_width = 0.36
        self.grid_spacing = 0.20
        grid_file = None
        density_kernel = 'rbf'
        energy_kernel = 'rbf'
        max_pos = np.asarray([16., 15., 13.])
        min_pos = np.asarray([4., 5., 7.] )
        max_range = max_pos - min_pos

        steps = np.round(max_range / self.grid_spacing).astype(np.int)

        if (grid_file is None):
            grid_range = [np.linspace(ma, mm, s) for ma, mm, s in zip(max_pos, min_pos, steps)]

        Y, X, Z = np.meshgrid(grid_range[0], grid_range[1], grid_range[2])
        if verbose > 1:
            print(X.shape)
        X = X.flatten()[:, np.newaxis]
        Y = Y.flatten()[:, np.newaxis]
        Z = Z.flatten()[:, np.newaxis]

        self.grid = np.concatenate((X, Y, Z), axis=1)

        
    def compute_potentials(self, atom_pos):
        potentials = du.calculate_potential(pos=np.asarray(atom_pos), charges=self.charges, gaussian_width=self.gaussian_width, grid=self.grid, verbose=self.verbose) 
        return potentials
    
    def get_energy_from_potentials(self, potentials):
        
        return self.clf_E.predict(self.clf_n.predict(potentials, verbose=0), verbose=0)

        
    def get_energy_via_ML(self, atom_pos, atom_types):
        """ Computes energy

        Input is atom_pos of a single molecule geometry.

        """
        assert np.all(atom_types == self.base_types)

        #restraints_constant in kcal/(mol*bohr^2)
        kres=10
        
        # 0. rearange positions
        pos = []
        pot_restraint=np.zeros(atom_pos.shape[0])
        for i in range(atom_pos.shape[0]):
            pos.append(atom_pos[i].reshape(-1, 3))
        atom_pos = np.array(pos) * 1.889725989  # From Angstrom to Bohr

        # 1. normalize positions and add restraints
        for i in range(atom_pos.shape[0]):
            heavy = atom_types != 1
            atom_pos[i] = du.transform_molecule(atom_pos[i], self.base_pos, heavy)
            pot_restraint[i] = 1/2*kres*(atom_pos[i][0][2]-10)**2 +  1/2*kres*(atom_pos[i][1][2]-10)**2 \
                                        + 1/2*kres*(atom_pos[i][2][2]-10)**2 +  1/2*kres*(atom_pos[i][3][2]-10)**2 \
                                        + 1/2*kres*(atom_pos[i][4][2]-10)**2 +  1/2*kres*(atom_pos[i][5][2]-10)**2 \
                                        + 1/2*kres*(atom_pos[i][6][2]-10)**2 +  1/2*kres*(atom_pos[i][7][2]-10)**2 \
                                        + 1/2*kres*(atom_pos[i][8][2]-10)**2

        # 2. create potential
        potentials = self.compute_potentials(atom_pos)

        # 3. get energy
        energy = self.get_energy_from_potentials(potentials)+pot_restraint

        return energy / 23.061  # from kcal/mol to eV

    
    def calculation_required(self, atoms, quantities=None):
        return (self.calc is None) or (self.calc['atoms'] != atoms)
   
    
    def _make_calc(self, atoms):
        # Calculate both energy and forces via one call to the ML models
        
        if not np.all(atoms.get_atomic_numbers() == base_types):
            raise RuntimeError('ASE switched atom types around.')
            
        current_pos_flat = atoms.get_positions().flatten()
        
        # We always need the current atoms pos for the energy.
        pos = [current_pos_flat.copy()]
            
        for k in range(len(current_pos_flat)):
            if self.fd_method == 'forward':
                xk2 = current_pos_flat.copy()
                xk2[k] += self.fd_eps
                pos.append(xk2)
            elif self.fd_method == 'center':
                xk2 = current_pos_flat.copy()
                xk2[k] += self.fd_eps / 2.
                pos.append(xk2)
                xk2 = current_pos_flat.copy()
                xk2[k] -= self.fd_eps / 2.
                pos.append(xk2)

        # Calculate energies
        #print('--> Start ML calculation...', end='')
        energies = self.get_energy_via_ML(np.array(pos), atoms.get_atomic_numbers())
        #print(' Done.')

        # Compute Forces. force = negative gradient
        if self.fd_method == 'forward':
            forces = (energies[0] - energies[1:]) / self.fd_eps
        elif self.fd_method == 'center':
            forces = (energies[2::2] - energies[1::2]) / self.fd_eps
        
        self.calc = {'forces': forces.reshape(-1, 3),
                     'energy': energies[0],
                     'atoms': atoms.copy()}
    
    def get_potential_energy(self, atoms=None, force_consistent=False):
        #if force_consistent:
        #    raise NotImplementedError('Asking for force conistent energy?')
        if atoms is None:
            atoms = self.atoms
        if self.calculation_required(atoms):
            self._make_calc(atoms)
            
        return self.calc['energy']
    
    
    def get_forces(self, atoms):
        if self.calculation_required(atoms):
            self._make_calc(atoms)
        return self.calc['forces']
forces = []
calculator = MLCalculator(clf_E=energy_kr, clf_n=density_kr, base_pos=base_pos, base_types=base_types,
                          fd_eps=1e-3, fd_method='center')
atoms = Atoms(positions=start_pos * 0.529177249,  # from Bohr to Angstrom
              numbers=start_types,
              calculator=calculator)
random_state = np.random.RandomState(seed=1235)
atoms.set_velocities(velocities)   

dyn = VelocityVerlet(atoms, 0.25 * units.fs)
ts = 0.25

results = {'energy': [], 'forces': [], 'positions': [], 'time': [], 'kinetic': [], 'temperature': []}
def log(a=atoms, results=results):  # store a reference to atoms in the definition.
    with open('md_logs/potential_energy', 'a') as f:
        f.write('%.18e\n' % a.get_potential_energy())
    with open('md_logs/kinetic_energy', 'a') as f:
        f.write('%.18e\n' % a.get_kinetic_energy())
    with open('md_logs/temperature', 'a') as f:
        f.write('%.18e\n' % a.get_temperature())
    with open('md_logs/forces', 'a') as f:
        forces = a.get_forces()
        for v in forces:
            f.write('%.18e %.18e %.18e\n' % (v[0], v[1], v[2]))
    pos = a.get_positions()
    with open('md_logs/positions.xyz', 'a') as f:
        f.write('%d\n' % len(pos))
        f.write(' generated by ML\n')
        for j in range(len(pos)):
            f.write(' %s %.18e %.18e %.18e\n' % (charge_to_str[a.get_calculator().base_types[j]], pos[j][0],
                                      pos[j][1], pos[j][2]))
    pos_aligned = du.transform_molecule(pos * 1.889725989, a.get_calculator().base_pos, a.get_calculator().base_types != 1) / 1.889725989
    with open('md_logs/positions_aligned.xyz', 'a') as f:
        f.write('%d\n' % len(pos_aligned))
        f.write(' generated by ML\n')
        for j in range(len(pos_aligned)):
            f.write(' %s %.18e %.18e %.18e\n' % (charge_to_str[a.get_calculator().base_types[j]], pos_aligned[j][0],
                                      pos_aligned[j][1], pos_aligned[j][2]))
    print(a.get_potential_energy(), a.get_temperature()) 
    
log()
dyn.attach(log, interval=1)

# we want to run for 60fs
n_steps = int(60. / ts)
n_steps_remaining = n_steps - len(results['energy'])
print('Trajectory with %d steps. (%d remaining)' % (n_steps, n_steps_remaining))
dyn.run(n_steps_remaining)

