#!/usr/bin/env python

import gc
import glob
import os

import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
import seaborn as sns
from gridData import Grid
from MDAnalysis import Universe
from MDAnalysis.analysis import rms
from scipy.spatial import cKDTree
from scipy.spatial import distance
from scipy.cluster.hierarchy import linkage, fcluster
from scipy.stats import pearsonr
from scipy.stats import spearmanr
from scipy.stats import kendalltau
from scipy.stats import linregress
from pymol import cmd
from sklearn.neighbors import KDTree

from waterkit.analysis import HydrationSites
from waterkit.analysis import blur_map

# Turn interactive mode off
plt.ioff()


def select_water_molecules(df, criteria):
    df_copy = df.copy()

    for name, parameter in criteria.items():
        df_copy = df_copy[df_copy[name] == parameter]
    df_copy.reset_index(drop=True, inplace=True)
    
    return df_copy


def cluster(coordinates, distance, method='median'):
    Z = linkage(coordinates, method=method, metric='euclidean')
    clusters = fcluster(Z, distance, criterion='distance')
    return clusters


def average_position_cluster(coordinates, clusters):
    average_coordinates = []

    for i in range(1, np.max(clusters) + 1):
        tmp = coordinates[clusters == i]
        average_coordinates.append(np.mean(tmp, axis=0))

    average_coordinates = np.array(average_coordinates)

    return average_coordinates


def hydration_sites(coordinates, values, density=2, min_cutoff=1.4, max_cutoff=2.6):
    # Keep coordinates with a certain density
    tmp_coordinates = coordinates[values >= density]
    tmp_values = values[values >= density]
    mask = np.ones(tmp_values.shape, dtype=np.bool)

    centers = []
    isocontour = []
    labels = []
    i = 1

    while mask.any():
        center_idx = np.argmax(tmp_values)

        d = distance.cdist([tmp_coordinates[center_idx]], tmp_coordinates, 'euclidean')[0]
        close_points = np.where((d <= max_cutoff) & (mask == True))[0]
        extra_close_points = np.where((d <= min_cutoff) & (mask == True))[0]
        
        # Add center
        centers.extend([tmp_coordinates[center_idx]])
        isocontour.extend([tmp_coordinates[center_idx]])
        labels.extend([i])
        
        # Remove center from the pool
        mask[center_idx] = False
        tmp_values[center_idx] = -1

        # Add the closest points and remove them from the pool
        if close_points.size > 0:
            isocontour.extend([tmp_coordinates[extra_close_points]])
            labels.extend([i] * extra_close_points.shape[0])

            mask[close_points] = False
            tmp_values[close_points] = -1.

        i += 1

    isocontour = np.vstack(isocontour)
    centers = np.vstack(centers)
    labels = np.array(labels)
    
    return centers, isocontour, labels


def r_squared(y_pred, y_true):
    ss_res = np.sum((y_true - y_pred)**2)
    ss_total = np.sum((y_true - np.mean(y_true))**2)
    return 1. - (ss_res / ss_total)


def bootstrap_r_squared(x, y, alpha=0.05, num_iterations=1000, y_uncertainty=0):
    low = 0
    high = 0

    if num_iterations > 1:
        ids = np.random.choice(len(x), (len(x), num_iterations), replace=True)
        bootstrap_dist = [pearsonr(x[i], y[i] + np.random.normal(0, y_uncertainty, len(i)))[0]**2 for i in ids]
        #bootstrap_dist = [r_squared(x[i], y[i] + np.random.normal(0, y_uncertainty, len(i))) for i in ids]
        #bootstrap_dist = [linregress(y[i] + np.random.normal(0, y_uncertainty, len(i)), x[i])[2]**2 for i in ids]
        val = np.mean(bootstrap_dist)
        std = np.std(bootstrap_dist)

        low = 2 * val - np.percentile(bootstrap_dist, 100 * (1 - alpha / 2.))
        high = 2 * val - np.percentile(bootstrap_dist, 100 * (alpha / 2.))
    else:
        val = pearsonr(x, y)[0]**2

    return val, low, high, std


def bootstrap_rmsd(x, y, alpha=0.05, num_iterations=1000, y_uncertainty=0):
    low = 0
    high = 0

    if num_iterations > 1:
        ids = np.random.choice(len(x), (len(x), num_iterations), replace=True)
        bootstrap_dist = [rms.rmsd(x[i], y[i] + np.random.normal(0, y_uncertainty, len(i))) for i in ids]
        val = np.mean(bootstrap_dist)
        std = np.std(bootstrap_dist)

        low = 2 * val - np.percentile(bootstrap_dist, 100 * (1 - alpha / 2.))
        high = 2 * val - np.percentile(bootstrap_dist, 100 * (alpha / 2.))
    else:
        val = rms.rmsd(x, y)[0]

    return val, low, high, std


def bootstrap_spearmanr(x, y, alpha=0.05, num_iterations=1000, y_uncertainty=0):
    low = 0
    high = 0

    if num_iterations > 1:
        ids = np.random.choice(len(x), (len(x), num_iterations), replace=True)
        bootstrap_dist = [spearmanr(x[i], y[i] + np.random.normal(0, y_uncertainty, len(i))) for i in ids]
        val = np.mean(bootstrap_dist)
        std = np.std(bootstrap_dist)

        low = 2 * val - np.percentile(bootstrap_dist, 100 * (1 - alpha / 2.))
        high = 2 * val - np.percentile(bootstrap_dist, 100 * (alpha / 2.))
    else:
        val = spearmanr(x, y)[0]

    return val, low, high, std


def bootstrap_kendalltau(x, y, alpha=0.05, num_iterations=1000, y_uncertainty=0):
    low = 0
    high = 0

    if num_iterations > 1:
        ids = np.random.choice(len(x), (len(x), num_iterations), replace=True)
        bootstrap_dist = [kendalltau(x[i], y[i] + np.random.normal(0, y_uncertainty, len(i))) for i in ids]
        val = np.mean(bootstrap_dist)
        std = np.std(bootstrap_dist)

        low = 2 * val - np.percentile(bootstrap_dist, 100 * (1 - alpha / 2.))
        high = 2 * val - np.percentile(bootstrap_dist, 100 * (alpha / 2.))
    else:
        val = kendalltau(x, y)[0]

    return val, low, high, std


def write_clusters(fname, coordinates, clusters):
    template = 'ATOM  %5d  %-4s%4s%5d    %8.3f%8.3f%8.3f  1.00  0.00      %1s   \n'
    clusters = np.array(clusters)

    with open(fname, 'w') as w:
        n_atom = 1

        for i in range(1, np.max(clusters) + 1):
            tmp = coordinates[clusters == i]

            for j, t in enumerate(tmp):
                w.write(template % (n_atom, 'OH2', 'HOH', i, t[0], t[1], t[2], 'O'))
                n_atom += 1


def rmsd(a, b):
    """
    Return euclidean distance a (can be multiple coordinates) and b
    """
    return np.sqrt(np.mean(np.sum(np.power(a - b, 2), axis=1)))


def waters_in_pocket(water_xyzs, xray_ligands, cutoff=3.4):
    water_in_pocket_idxs = []
    in_pocket = np.zeros(water_xyzs.shape[0], dtype=np.bool)
    KDtree_water = cKDTree(water_xyzs)

    for xray_ligand in xray_ligands:
        u = Universe(xray_ligand)
        water_idxs = KDtree_water.query_ball_point(u.select_atoms("not name H*").positions, cutoff, p=2)
        water_in_pocket_idxs.extend(water_idxs)

    water_in_pocket_idxs = np.concatenate(water_in_pocket_idxs).ravel()
    water_in_pocket_idxs = np.unique(water_in_pocket_idxs).astype(np.int)
    in_pocket[water_in_pocket_idxs] = True

    return in_pocket


def is_close_to_edge(grid, xyz, distance):
    xyz = np.atleast_2d(xyz)
    x, y, z = xyz[:, 0], xyz[:, 1], xyz[:, 2]

    xmin, xmax = grid.edges[0][0], grid.edges[0][-1]
    ymin, ymax = grid.edges[1][0], grid.edges[1][-1]
    zmin, zmax = grid.edges[2][0], grid.edges[2][-1]

    x_close = np.logical_or(np.abs(xmin - x) <= distance, np.abs(xmax - x) <= distance)
    y_close = np.logical_or(np.abs(ymin - y) <= distance, np.abs(ymax - y) <= distance)
    z_close = np.logical_or(np.abs(zmin - z) <= distance, np.abs(zmax - z) <= distance)
    close_to = np.any((x_close, y_close, z_close), axis=0)

    return close_to


def normalize(a):
    """
    Return a normalized vector
    """
    return a / np.sqrt(np.sum(np.power(a, 2)))


def is_bulk_accessible(all_coordinates, msms_vertices_file, radius=25):
    xyz = []
    nxyz = []
    all_coordinates = np.array(all_coordinates)
    accessible = [False] * len(all_coordinates)
    
    with open(msms_vertices_file) as f:
        lines = f.readlines()[3:]
        for line in lines:
            sline = line.split()

            # We do not want edgy norm vectors
            if np.int(sline[6]) >= 0:
                xyz.append(sline[0:3])
                nxyz.append(sline[3:6])
    
    xyz = np.array(xyz).astype(np.float)
    nxyz = np.array(nxyz).astype(np.float)
    
    kdtree = cKDTree(xyz)
    
    for i, coordinate in enumerate(all_coordinates):
        index = kdtree.query_ball_point(coordinate, radius, p=2)
        
        if index:
            d = distance.cdist(xyz[index], [coordinate], 'euclidean').flatten()
            i_xyz = xyz[index[np.argmin(d)]]
            i_nxyz = nxyz[index[np.argmin(d)]]
            
            sp = np.dot(i_nxyz, normalize(coordinate - i_xyz))
            
            if sp >= 0:
                accessible[i] = True
            
        else:
            accessible[i] = np.nan
            print("Warning: neighborhood vertices not found.")
            
    return accessible


def extract_positions_energies_from_md_simulations(references, protein, xray_ligands, msms_vertices_file, 
                                                   density=2, distance_from_protein=5, distance_from_ligand=4, distance_from_edges=2.,
                                                   cutoffs=None):
    all_hydration_sites = []
    data = []
    columns = ["reference", "water_id", "x", "y", "z", "d", "gO", "esw", "eww", "tst", "tso", "dG"]

    protein_positions = protein.select_atoms("not name H*").positions
    protein_kdtree = KDTree(protein_positions)

    if cutoffs is not None:
        cutoffs = np.atleast_1d(cutoffs)
    else:
        cutoffs = [0.10, 0.25, 0.50, 1.0, 1.20, 1.40, 1.60]

    for reference in references:
        reference_name = reference.split("/")[-1]

        try:
            grid_gO = Grid('%s/gist_gO.dx' % (reference))
            grid_esw = Grid('%s/gist_Esw-dens.dx' % (reference))
            grid_eww = Grid('%s/gist_Eww-dens.dx' % (reference))
            grid_tst = Grid('%s/gist_dTStrans-dens.dx' % (reference))
            grid_tso = Grid('%s/gist_dTSorient-dens.dx' % (reference))
        except:
            grid_gO = Grid('%s/gist-gO.dx' % (reference))
            grid_esw = Grid('%s/gist-Esw-dens.dx' % (reference))
            grid_eww = Grid('%s/gist-Eww-dens.dx' % (reference))
            grid_tst = Grid('%s/gist-dTStrans-dens.dx' % (reference))
            grid_tso = Grid('%s/gist-dTSorient-dens.dx' % (reference))

        # Normalized Eww
        #grid_eww = 2 * grid_eww
        # Compute DeltaG
        grid_dg = (grid_esw + grid_eww) - (grid_tst + grid_tso)

        hs = HydrationSites(gridsize=0.5, water_radius=1.4, min_water_distance=2.5, min_density=density)
        hydration_sites = hs.find(grid_gO) # can pass "gist-gO.dx" directly also

        values_gO = hs.hydration_sites_energy(grid_gO, gridsize=0, water_radius=0)
        values_esw = hs.hydration_sites_energy(grid_esw)
        values_eww = hs.hydration_sites_energy(grid_eww)
        values_tst = hs.hydration_sites_energy(grid_tst)
        values_tso = hs.hydration_sites_energy(grid_tso)
        values_dg = hs.hydration_sites_energy(grid_dg)

        # We want to the distances from the closest atoms, to know in which layer the water is in.
        distances, _ = protein_kdtree.query(hydration_sites, k=1, return_distance=True)
        distances = distances.flatten()
        selected_hydration_site_ids = np.where(distances <= distance_from_protein)[0]

        print(reference, selected_hydration_site_ids.shape[0])

        for i, hs_id in enumerate(selected_hydration_site_ids):
            data.append((reference_name, i + 1, 
                         hydration_sites[hs_id][0], hydration_sites[hs_id][1], hydration_sites[hs_id][2], distances[hs_id],
                         values_gO[hs_id], values_esw[hs_id], values_eww[hs_id], values_tst[hs_id], values_tso[hs_id], values_dg[hs_id]))

        all_hydration_sites.extend(hydration_sites[selected_hydration_site_ids])
        hs.export_to_pdb('cluster_average_reference_%s.pdb' % reference_name, hydration_sites[selected_hydration_site_ids])

    df = pd.DataFrame(data=data, columns=columns)
    all_hydration_sites = np.vstack(all_hydration_sites)

    for cutoff in cutoffs:
        clusters = cluster(all_hydration_sites, cutoff)
        df["cluster_%3.2f" % cutoff] = clusters

    df["pocket"] = waters_in_pocket(all_hydration_sites, xray_ligands, distance_from_ligand)
    df["edges"] = is_close_to_edge(grid_gO, all_hydration_sites, distance_from_edges)
    df["accessible"] = is_bulk_accessible(all_hydration_sites, msms_vertices_file)

    return df


def extract_positions_energies_from_waterkit(directories, protein, xray_ligands, msms_vertices_file, 
                                             density=8, distance_from_protein=5, distance_from_ligand=4, distance_from_edges=2.):
    all_hydration_sites = []
    data = []
    columns = ["reference", "water_id", "x", "y", "z", "d", "gO", "esw", "eww", "tst", "tso", "dG"]

    protein_positions = protein.select_atoms("not name H*").positions
    protein_kdtree = KDTree(protein_positions)
        
    for directory in directories:
        directory_name = directory.split("/")[-1]

        try:
            grid_gO = Grid('%s/gist_gO.dx' % (directory))
            grid_esw = Grid('%s/gist_Esw-dens.dx' % (directory))
            grid_eww = Grid('%s/gist_Eww-dens.dx' % (directory))
            grid_tst = Grid('%s/gist_dTStrans-dens.dx' % (directory))
            grid_tso = Grid('%s/gist_dTSorient-dens.dx' % (directory))
        except:
            grid_gO = Grid('%s/gist-gO.dx' % (directory))
            grid_esw = Grid('%s/gist-Esw-dens.dx' % (directory))
            grid_eww = Grid('%s/gist-Eww-dens.dx' % (directory))
            grid_tst = Grid('%s/gist-dTStrans-dens.dx' % (directory))
            grid_tso = Grid('%s/gist-dTSorient-dens.dx' % (directory))

        # Normalized Eww
        #grid_eww = 2 * grid_eww
        # Compute DeltaG
        grid_dg = (grid_esw + grid_eww) - (grid_tst + grid_tso)

        hs = HydrationSites(gridsize=0.5, water_radius=1.4, min_water_distance=2.5, min_density=density)
        hydration_sites = hs.find(grid_gO) # can pass "gist-gO.dx" directly also

        values_gO = hs.hydration_sites_energy(grid_gO, gridsize=0, water_radius=0)
        values_esw = hs.hydration_sites_energy(grid_esw)
        values_eww = hs.hydration_sites_energy(grid_eww)
        values_tst = hs.hydration_sites_energy(grid_tst)
        values_tso = hs.hydration_sites_energy(grid_tso)
        values_dg = hs.hydration_sites_energy(grid_dg)

        # We want to the distances from the closest atoms, to know in which layer the water is in.
        distances, _ = protein_kdtree.query(hydration_sites, k=1, return_distance=True)
        distances = distances.flatten()
        selected_hydration_site_ids = np.where(distances <= distance_from_protein)[0]

        print(directory_name, selected_hydration_site_ids.shape[0])

        for i, hs_id in enumerate(selected_hydration_site_ids):
            data.append((directory_name, i + 1, 
                         hydration_sites[hs_id][0], hydration_sites[hs_id][1], hydration_sites[hs_id][2], distances[hs_id],
                         values_gO[hs_id], values_esw[hs_id], values_eww[hs_id], values_tst[hs_id], values_tso[hs_id], values_dg[hs_id]))

        all_hydration_sites.extend(hydration_sites[selected_hydration_site_ids])
        hs.export_to_pdb('%s/cluster.pdb' % directory, hydration_sites[selected_hydration_site_ids])

    df = pd.DataFrame(data=data, columns=columns)
    all_hydration_sites = np.vstack(all_hydration_sites)

    df["pocket"] = waters_in_pocket(all_hydration_sites, xray_ligands, distance_from_ligand)
    df["edges"] = is_close_to_edge(grid_gO, all_hydration_sites, distance_from_edges)
    df["accessible"] = is_bulk_accessible(all_hydration_sites, msms_vertices_file)
    
    return df


def assign_cluster(df_exp, df_pred, ignore=None, cutoffs=None):
    if ignore is None:
        ignore = []

    if cutoffs is not None:
        cutoffs = np.atleast_1d(cutoffs)
    else:
        cutoffs = [0.10, 0.25, 0.50, 1.0, 1.20, 1.40, 1.60]

    df_exp = df_exp[~df_exp["reference"].isin(ignore)].copy()

    water_exp_xyzs = df_exp[["x", "y", "z"]].values
    water_pred_xyzs = df_pred[["x", "y", "z"]].values
    keys = df_exp.index.values

    for cutoff in cutoffs:
        clusters = []
        
        for water_pred_xyz in water_pred_xyzs:
            d = distance.cdist([water_pred_xyz], water_exp_xyzs, 'euclidean')[0]

            if np.min(d) <= cutoff:
                i = keys[np.argmin(d)]
                clusters.append(df_exp.loc[i]["cluster"])
            else:
                clusters.append("")

        df_pred["cluster_%3.2f" % cutoff] = clusters

    return df_pred


def exchange_hydration_site_properties(df_exp, df_pred, ignore=None, cutoffs=None):
    df_pred = assign_cluster(df_exp, df_pred, ignore, cutoffs)
    
    cluster_ids = df_exp['cluster'].unique()
    max_cutoff = np.max([float(col.split('_')[1]) for col in df_pred.columns if 'cluster_' in col])
    
    for cluster_id in cluster_ids:
        df_exp_select = df_exp[df_exp['cluster'] == cluster_id]
        df_pred_select = df_pred[df_pred['cluster_%.2f' % max_cutoff] == cluster_id]
        
        if not df_pred_select.empty:
            if any(df_exp_select['accessible']) or any(df_pred_select['accessible']):
                df_exp.loc[df_exp_select.index, 'accessible'] = True
                df_pred.loc[df_pred_select.index, 'accessible'] = True
            
            if any(df_exp_select['pocket']) or any(df_pred_select['pocket']):
                df_exp.loc[df_exp_select.index, 'pocket'] = True
                df_pred.loc[df_pred_select.index, 'pocket'] = True
            
            if any(df_exp_select['edges']) or any(df_pred_select['edges']):
                df_exp.loc[df_exp_select.index, 'edges'] = True
                df_pred.loc[df_pred_select.index, 'edges'] = True
    
    return df_exp, df_pred


def compare_md_waterkit_placement(directories, df_exp, df_pred, cutoffs=None, ignore=None, 
                                  accessible_only=False, in_pocket_only=False, close_to_edges=False):
    data = []
    if ignore is None:
        ignore = []

    if cutoffs is not None:
        cutoffs = np.atleast_1d(cutoffs)
    else:
        cutoffs = [0.10, 0.25, 0.50, 1.0, 1.20, 1.40]

    criteria = {'accessible': accessible_only,
                'edges': close_to_edges}

    if in_pocket_only:
        criteria['pocket'] = in_pocket_only

    df_exp = select_water_molecules(df_exp, criteria)
    df_pred = select_water_molecules(df_pred, criteria)
    
    if ignore:
        df_exp = df_exp[~df_exp["reference"].isin(ignore)]
    
    n_references = len(df_exp["reference"].unique())
    n_water = np.unique(np.array([1, n_references]))

    df_exp.reset_index(drop=True, inplace=True)
    df_pred.reset_index(drop=True, inplace=True)
    
    columns = ["model"]
    for cutoff in cutoffs:
        for n in n_water:
            columns += ["TPR_%d_%3.2f" % (n, cutoff), "PPV_%d_%3.2f" % (n, cutoff), "F1_%d_%3.2f" % (n, cutoff), 
                        "TP_%d_%3.2f" % (n, cutoff), "FP_%d_%3.2f" % (n, cutoff), "FN_%d_%3.2f" % (n, cutoff)]

    for directory in directories:
        data_tmp = []
        directory_name = directory.split("/")[-1]
        df_tmp = df_pred[df_pred['reference'] == directory_name].copy()

        print(directory_name)

        for cutoff in cutoffs:
            # Count number of water molecules per cluster in exp
            # If there is 'cluster' column, it means that we are comparing MD with WK
            # ... otherwise, it means that we are comparing MD with MD
            if 'cluster' in df_exp.columns:
                occ = df_exp["cluster"].value_counts(sort=True)
            else:
                occ = df_exp["cluster_%3.2f" % cutoff].value_counts(sort=True)

            for n in n_water:
                df_pred_selected = df_tmp[df_tmp["cluster_%3.2f" % cutoff].isin(occ[occ >= n].index)]
                
                # False positive is: number of water with no cluster ID + number of clusters that are not
                # in the list reference clusters
                false_positive = df_tmp[df_tmp["cluster_%3.2f" % cutoff].isnull()].shape[0]
                false_positive += len(set(df_tmp["cluster_%3.2f" % cutoff].values.astype(int)).difference(occ[occ >= n].index.astype(int)))
                # True positive is: number of identified clusters
                true_positive = df_pred_selected.shape[0]
                # False negative is: number of reference clusters - number of identified clusters
                false_negative = occ[occ >= n].shape[0] - true_positive
                
                if true_positive > 0:
                    sensitivity = np.float(true_positive) / np.float(true_positive + false_negative)
                    precision = np.float(true_positive) / np.float(true_positive + false_positive)
                else:
                    sensitivity = 0.
                    precision = 0.

                try:
                    f1 = 2. * ((precision * sensitivity) / (precision + sensitivity))
                except ZeroDivisionError:
                    f1 = 0.

                data_tmp.extend([sensitivity, precision, f1, true_positive, false_positive, false_negative])

        data.append([directory_name] + data_tmp)

    df = pd.DataFrame(data=data, columns=columns)
    
    return df


def compare_md_placement(directories, df_exp, cutoffs=None, ignore=None,
                         accessible_only=False, in_pocket_only=False, close_to_edges=False):
    dfs = []
    
    if ignore is None:
        ignore = []
    
    # We remove cluster info we use when comparing with WK
    try:
        df_exp = df_exp.copy().drop(columns=['cluster'])
    except:
        pass

    for directory in directories:
        directory_name = directory.split("/")[-1]
        
        if ignore:
            df_exp_tmp = df_exp[(df_exp["reference"] != directory_name) & ~df_exp["reference"].isin(ignore)].copy()
        else:
            df_exp_tmp = df_exp[(df_exp["reference"] != directory_name)].copy()
        df_pred_tmp = df_exp[df_exp["reference"] == directory_name].copy()
        
        dfs.append(compare_md_waterkit_placement([directory], df_exp_tmp, df_pred_tmp, cutoffs, None,
                                                 accessible_only, in_pocket_only, close_to_edges))
    
    df = pd.concat(dfs, sort=False)
    
    return df


def compare_md_waterkit_energy(directories, df_exp, df_pred, colors, cutoff=1.0, ignore=None, 
                               accessible_only=False, in_pocket_only=False, close_to_edges=False,
                               plot_energies=True, bootstrap_iterations=1000):
    data = []
    terms = ["gO", "esw", "eww", "tst", "tso", "dG"]
    if ignore is None:
        ignore = []

    criteria = {'accessible': accessible_only,
                'edges': close_to_edges}

    if in_pocket_only:
        criteria['pocket'] = in_pocket_only

    df_exp = select_water_molecules(df_exp, criteria)
    df_pred = select_water_molecules(df_pred, criteria)
    
    df_exp = df_exp[~df_exp["reference"].isin(ignore)]
    
    # Easier to rename the cluster_<cutoff> column to cluster
    if not 'cluster' in df_exp.columns:
        df_exp.rename(columns={"cluster_%3.2f" % cutoff: "cluster"}, inplace=True)
    df_pred.rename(columns={"cluster_%3.2f" % cutoff: "cluster"}, inplace=True)
    
    # We want to calculate results on ALL and only COMMON water molecules
    n_references = len(df_exp["reference"].unique())
    n_water = np.unique(np.array([1, n_references - 1, n_references]))
    n_water = n_water[n_water >= 1]
    
    occ = df_exp["cluster"].value_counts(sort=True)

    df_exp.reset_index(drop=True, inplace=True)
    df_pred.reset_index(drop=True, inplace=True)

    columns = ["model"]
    for n in n_water:
        for term in terms:
            columns += [s % (term, n) for s in ["rmsd_%s_%d", "rmsd_std_%s_%d"]]
            columns += [s % (term, n) for s in ["rs_%s_%d", "rs_std_%s_%d"]]
            columns += [s % (term, n) for s in ["sr_%s_%d", "sr_std_%s_%d"]]
            columns += [s % (term, n) for s in ["kt_%s_%d", "kt_std_%s_%d"]]

    for directory in directories:
        data_tmp = []
        directory_name = directory.split("/")[-1]
        print(directory_name)

        df_tmp = df_pred[df_pred['reference'] == directory_name].copy()
        
        for n in n_water:
            df_pred_selected = df_tmp[df_tmp["cluster"].isin(occ[occ >= n].index)]
            df_pred_selected = df_pred_selected[~df_pred_selected["cluster"].isnull()]

            df_merge = df_pred_selected.merge(df_exp, on="cluster")

            try:
                for term in terms:
                    rmsd, _, _, rmsd_std = bootstrap_rmsd(df_merge["%s_x" % term], df_merge["%s_y" % term], 0.05, bootstrap_iterations)
                    rs, _, _, rs_std = bootstrap_r_squared(df_merge["%s_x" % term], df_merge["%s_y" % term], 0.05, bootstrap_iterations)
                    sr, _, _, sr_std = bootstrap_spearmanr(df_merge["%s_x" % term], df_merge["%s_y" % term], 0.05, bootstrap_iterations)
                    kt, _, _, kt_std = bootstrap_kendalltau(df_merge["%s_x" % term], df_merge["%s_y" % term], 0.05, bootstrap_iterations)
                    data_tmp.extend([rmsd, rmsd_std, rs, rs_std, sr, sr_std, kt, kt_std])
            except:
                for term in terms:
                    data_tmp.extend([0] * 8)

        if plot_energies:
            figure_filename = "%s/results" % directory

            if in_pocket_only:
                figure_filename += "_pocket"
            if accessible_only:
                figure_filename += "_accessible"

            figure_filename += ".png"
                                
            plot_water_energies(figure_filename, df_exp, colors, df_tmp[df_tmp["cluster"] != ""], False, accessible_only, in_pocket_only)

        data.append([directory_name] + data_tmp)

    df = pd.DataFrame(data=data, columns=columns)
    
    return df


def compare_md_energy(directories, df_exp, colors, cutoff=1.0, ignore=None, 
                      accessible_only=False, in_pocket_only=False, close_to_edges=False, 
                      plot_energies=True, bootstrap_iterations=1000):
    dfs = []
    if ignore is None:
        ignore = []
    
    # We remove cluster info we use when comparing with WK
    try:
        df_exp = df_exp.copy().drop(columns=['cluster'])
    except:
        pass
    
    for directory in directories:
        directory_name = directory.split("/")[-1]
        
        df_pred_tmp = df_exp[df_exp["reference"] == directory_name].copy()
        df_exp_tmp = df_exp[(df_exp["reference"] != directory_name) & ~df_exp["reference"].isin(ignore)].copy()
        
        dfs.append(compare_md_waterkit_energy([directory], df_exp_tmp, df_pred_tmp, colors, cutoff, ignore, 
                                              accessible_only, in_pocket_only, close_to_edges,
                                              plot_energies, bootstrap_iterations))
    
    df = pd.concat(dfs, sort=False)
    
    return df


def convergence(filename, directory, frames, voxel_size=0.5):
    tso = []
    tst = []
    go = []
    esw = []
    eww = []
    
    fig, axarr = plt.subplots(5, figsize=(15, 15), sharex=True)
    
    for frame in frames:
        esw.append(np.sum(Grid('%s/gist_%d-Esw-dens.dx' % (directory, frame)).grid) * voxel_size**3)
        eww.append(np.sum(Grid('%s/gist_%d-Eww-dens.dx' % (directory, frame)).grid) * voxel_size**3)
        tso.append(np.sum(Grid('%s/gist_%d-dTSorient-dens.dx' % (directory, frame)).grid) * voxel_size**3)
        tst.append(np.sum(Grid('%s/gist_%d-dTStrans-dens.dx' % (directory, frame)).grid) * voxel_size**3)
        go.append(np.sum(Grid('%s/gist_%d-gO.dx' % (directory, frame)).grid))

    axarr[0].plot(frames, go)
    axarr[0].scatter(frames, go)
    axarr[1].plot(frames, esw)
    axarr[1].scatter(frames, esw)
    axarr[2].plot(frames, eww)
    axarr[2].scatter(frames, eww)
    axarr[3].plot(frames, tso)
    axarr[3].scatter(frames, tso)
    axarr[4].plot(frames, tst)
    axarr[4].scatter(frames, tst)
    
    axarr[0].set_ylabel(r"gO", fontsize=15)
    axarr[1].set_ylabel(r"Esw (kcal/mol)", fontsize=15)
    axarr[2].set_ylabel(r"Eww (kcal/mol)", fontsize=15)
    axarr[3].set_ylabel(r"TStrans (kcal/mol)", fontsize=15)
    axarr[4].set_ylabel(r"TSorient (kcal/mol)", fontsize=15)        
    axarr[4].set_xlabel("Number of frames", fontsize=15)

    plt.savefig(filename, dpi=300, bbox_inches="tight")

    plt.show()


def convergence_md(filename, patterns, frames, voxel_size=0.5):    
    fig, axarr = plt.subplots(5, figsize=(15, 15), sharex=True)

    for pattern in patterns:
        tso = []
        tst = []
        go = []
        esw = []
        eww = []
    
        for frame in frames:
            esw.append(np.sum(Grid('%s/gist-Esw-dens.dx' % (pattern % frame)).grid) * voxel_size**3)
            eww.append(np.sum(Grid('%s/gist-Eww-dens.dx' % (pattern % frame)).grid) * voxel_size**3)
            tso.append(np.sum(Grid('%s/gist-dTSorient-dens.dx' % (pattern % frame)).grid) * voxel_size**3)
            tst.append(np.sum(Grid('%s/gist-dTStrans-dens.dx' % (pattern % frame)).grid) * voxel_size**3)
            go.append(np.sum(Grid('%s/gist-gO.dx' % (pattern % frame)).grid))

        axarr[0].plot(frames, go)
        axarr[0].scatter(frames, go)
        axarr[1].plot(frames, esw)
        axarr[1].scatter(frames, esw)
        axarr[2].plot(frames, eww)
        axarr[2].scatter(frames, eww)
        axarr[3].plot(frames, tso)
        axarr[3].scatter(frames, tso)
        axarr[4].plot(frames, tst)
        axarr[4].scatter(frames, tst)
    
    axarr[0].set_ylabel(r"gO", fontsize=15)
    axarr[1].set_ylabel(r"Esw (kcal/mol)", fontsize=15)
    axarr[2].set_ylabel(r"Eww (kcal/mol)", fontsize=15)
    axarr[3].set_ylabel(r"TStrans (kcal/mol)", fontsize=15)
    axarr[4].set_ylabel(r"TSorient (kcal/mol)", fontsize=15)        
    axarr[4].set_xlabel("Number of frames", fontsize=15)

    plt.savefig(filename, dpi=300, bbox_inches="tight")

    plt.show()


def convergence_md_wk(filename, md_patterns, wk_patterns, md_frames, wk_frames, voxel_size=0.5):    
    fig, axarr = plt.subplots(5, figsize=(15, 15), sharex=True)

    time_ns = np.max(md_frames) / 1000
    
    with sns.axes_style("ticks"):
        for pattern in md_patterns:
            tso = []
            tst = []
            go = []
            esw = []
            eww = []

            for frame in md_frames:
                esw.append(np.sum(Grid('%s/gist-Esw-dens.dx' % (pattern % frame)).grid) * voxel_size**3)
                eww.append(np.sum(Grid('%s/gist-Eww-dens.dx' % (pattern % frame)).grid) * voxel_size**3)
                tso.append(np.sum(Grid('%s/gist-dTStrans-dens.dx' % (pattern % frame)).grid) * voxel_size**3)
                tst.append(np.sum(Grid('%s/gist-dTSorient-dens.dx' % (pattern % frame)).grid) * voxel_size**3)
                go.append(np.sum(Grid('%s/gist-gO.dx' % (pattern % frame)).grid))

            axarr[0].plot(md_frames, go, color='#1f77b4', label='MD sim. (triplicates) (%d ns)' % time_ns)
            axarr[0].scatter(md_frames, go, color='#1f77b4')
            axarr[1].plot(md_frames, esw, color='#1f77b4', label='MD sim. (triplicates) (%d ns)' % time_ns)
            axarr[1].scatter(md_frames, esw, color='#1f77b4')
            axarr[2].plot(md_frames, eww, color='#1f77b4', label='MD sim. (triplicates) (%d ns)' % time_ns)
            axarr[2].scatter(md_frames, eww, color='#1f77b4')
            axarr[3].plot(md_frames, tso, color='#1f77b4', label='MD sim. (triplicates) (%d ns)' % time_ns)
            axarr[3].scatter(md_frames, tso, color='#1f77b4')
            axarr[4].plot(md_frames, tst, color='#1f77b4', label='MD sim. (triplicates) (%d ns)' % time_ns)
            axarr[4].scatter(md_frames, tst, color='#1f77b4')

        axarr2 = [ax.twinx() for ax in axarr]

        for pattern in wk_patterns:
            tso = []
            tst = []
            go = []
            esw = []
            eww = []

            for frame in wk_frames:
                esw.append(np.sum(Grid('%s/gist_%s-Esw-dens.dx' % (pattern, frame)).grid) * voxel_size**3)
                eww.append(np.sum(Grid('%s/gist_%s-Eww-dens.dx' % (pattern, frame)).grid) * voxel_size**3)
                tso.append(np.sum(Grid('%s/gist_%s-dTStrans-dens.dx' % (pattern, frame)).grid) * voxel_size**3)
                tst.append(np.sum(Grid('%s/gist_%s-dTSorient-dens.dx' % (pattern, frame)).grid) * voxel_size**3)
                go.append(np.sum(Grid('%s/gist_%s-gO.dx' % (pattern, frame)).grid))

            axarr2[0].plot(wk_frames, go, color='#ff7f0e', label='Waterkit sim. (triplicates)')
            axarr2[0].scatter(wk_frames, go, color='#ff7f0e')
            axarr2[1].plot(wk_frames, esw, color='#ff7f0e', label='Waterkit sim. (triplicates)')
            axarr2[1].scatter(wk_frames, esw, color='#ff7f0e')
            axarr2[2].plot(wk_frames, eww, color='#ff7f0e', label='Waterkit sim. (triplicates)')
            axarr2[2].scatter(wk_frames, eww, color='#ff7f0e')
            axarr2[3].plot(wk_frames, tso, color='#ff7f0e', label='Waterkit sim. (triplicates)')
            axarr2[3].scatter(wk_frames, tso, color='#ff7f0e')
            axarr2[4].plot(wk_frames, tst, color='#ff7f0e', label='Waterkit sim. (triplicates)')
            axarr2[4].scatter(wk_frames, tst, color='#ff7f0e')
        
        lines_1, labels_1 = axarr[0].get_legend_handles_labels()
        lines_2, labels_2 = axarr2[0].get_legend_handles_labels()

        if len(wk_patterns) == 1:
            lines = [lines_1[0]] + lines_2
            labels = [labels_1[0]] + labels_2
        else:
            lines = [lines_1[0]] + [lines_2[0]]
            labels = [labels_1[0]] + [labels_2[0]]

        axarr[0].legend(lines, labels, bbox_to_anchor=(0., 1.15, 1., .102), loc='upper left',
                        ncol=2, mode="expand", borderaxespad=0., fontsize=15)

        axarr[0].set_ylabel(r"gO (bulk density)", fontsize=15, labelpad=15)
        axarr[1].set_ylabel(r"E$_{sw}$ (kcal/mol)", fontsize=15, labelpad=15)
        axarr[2].set_ylabel(r"E$_{ww}$ (kcal/mol)", fontsize=15, labelpad=15)
        axarr[3].set_ylabel(r"TS$_t$ (kcal/mol)", fontsize=15, labelpad=15)
        axarr[4].set_ylabel(r"TS$_o$ (kcal/mol)", fontsize=15, labelpad=15)        
        axarr[4].set_xlabel("#Frames analyzed with GIST", fontsize=15)

    plt.savefig(filename, dpi=300, bbox_inches="tight", transparent=True)

    plt.show()


def plot_water_energies(filename, df_exp, colors, df_pred, show=True,
                        accessible_only=False, in_pocket_only=False, close_to_edges=False):

    if in_pocket_only:
        fig, axarr = plt.subplots(6, figsize=(20, 6*6))
    else:
        fig, axarr = plt.subplots(6, figsize=(30, 6*6))

    criteria = {'accessible': accessible_only,
                'edges': close_to_edges}

    if in_pocket_only:
        criteria['pocket'] = in_pocket_only

    df_exp = select_water_molecules(df_exp, criteria)
    df_pred = select_water_molecules(df_pred, criteria)

    tmp = df_exp[df_exp["reference"] != "gist"].groupby("cluster")
    occ = df_exp["cluster"].value_counts(sort=False)
    
    n_axis = [0, 1, 2, 3, 4, 5]
    terms = ['gO', 'esw', 'eww', 'tst', 'tso', 'dG']
    
    for n, term in zip(n_axis, terms):
        s = tmp[term].apply(min)
        order = s.sort_values().index
        is_present = ['red'] * len(order.values)
        n_present = []
        for i, o in enumerate(order):
            tmp_2 = df_exp[df_exp["cluster"] == o]
            n_present.append(occ[o])
            for _, l in tmp_2.iterrows():
                axarr[n].scatter(i + 1, l[term], c=colors[l.reference])
                v = df_pred[df_pred["cluster"] == o][term].values

                if v.size > 0:
                    axarr[n].scatter([i + 1] * len(v), v, facecolors='none', s=100, edgecolors="#ff7f0e")
                    is_present[i] = 'green'

        axarr[n].xaxis.set_ticks(range(1, len(order.values) + 1))
        axarr[n].set_xticklabels(list(order), rotation='vertical')
        ax_t = axarr[n].secondary_xaxis('top')
        ax_t.xaxis.set_ticks(range(1, len(order.values) + 1))
        ax_t.set_xticklabels(n_present)
        for ticklabel, tickcolor in zip(ax_t.get_xticklabels(), is_present):
            ticklabel.set_color(tickcolor)

    axarr[0].set_ylabel(r"gO", fontsize=20, labelpad=20)
    axarr[1].set_ylabel(r"E$_{sw}$ (kcal/mol)", fontsize=20, labelpad=20)
    axarr[2].set_ylabel(r"E$_{ww}$ (kcal/mol)", fontsize=20, labelpad=20)
    axarr[3].set_ylabel(r"TS$_t$ (kcal/mol)", fontsize=20, labelpad=20)
    axarr[4].set_ylabel(r"TS$_o$ (kcal/mol)", fontsize=20, labelpad=20)
    axarr[5].set_ylabel(r"$\Delta$G (kcal/mol)", fontsize=20, labelpad=20)
    axarr[5].set_xlabel("Hydration Site ID", fontsize=20)
    
    plt.savefig(filename, bbox_inches="tight", dpi=300, transparent=True)
    if show:
        plt.show()
    else:
        plt.close()
    
    plt.cla()
    plt.clf()
    plt.close('all')
    plt.close(fig)

    del fig
    del axarr

    gc.collect()


def create_pymol_sessions(references, directories, df_exp, df_pred,
                          accessible_only=False, in_pocket_only=False, close_to_edges=False):
    colors = ["marine", "orange", "sand", "deepteal", "lightmagenta", "deeppurple", "deepolive"]

    cmd.delete("all")
    curdir = os.getcwd()

    criteria = {'accessible': accessible_only,
                'edges': close_to_edges}

    if in_pocket_only:
        criteria['pocket'] = in_pocket_only

    df_exp = select_water_molecules(df_exp, criteria)
    df_pred = select_water_molecules(df_pred, criteria)

    for reference in references:
        reference_name = reference.split("/")[-1]
        
        df_exp_selected = df_exp[df_exp["reference"] == reference_name]
        u = Universe("cluster_average_reference_%s.pdb" % reference_name)
        
        selected_water = u.atoms[df_exp_selected.water_id - 1]
        selected_water.residues.resids = df_exp_selected["cluster"]
        selected_water.write("cluster_average_reference_renum_%s.pdb" % reference_name)
        
        selected_water.tempfactors = df_exp_selected["gO"]
        selected_water.write("cluster_average_reference_renum_gO_%s.pdb" % reference_name)
        selected_water.tempfactors = df_exp_selected["esw"]
        selected_water.write("cluster_average_reference_renum_esw_%s.pdb" % reference_name)
        selected_water.tempfactors = df_exp_selected["eww"]
        selected_water.write("cluster_average_reference_renum_eww_%s.pdb" % reference_name)
        selected_water.tempfactors = df_exp_selected["tst"]
        selected_water.write("cluster_average_reference_renum_tst_%s.pdb" % reference_name)
        selected_water.tempfactors = df_exp_selected["tso"]
        selected_water.write("cluster_average_reference_renum_tso_%s.pdb" % reference_name)
        selected_water.tempfactors = df_exp_selected["dG"]
        selected_water.write("cluster_average_reference_renum_dG_%s.pdb" % reference_name)

    for directory in directories:
        directory_name = directory.split("/")[-1]

        df_pred_selected = df_pred[df_pred["reference"] == directory_name]
        
        df_pred_found = df_pred_selected[(df_pred_selected["cluster"] != '') & (~df_pred_selected["cluster"].isna())]
        df_pred_not_found = df_pred_selected[(df_pred_selected["cluster"] == '') | (df_pred_selected["cluster"].isna())]
        
        u = Universe('%s/cluster.pdb' % directory)
        # FOUND
        if not df_pred_found.empty:
            selected_water = u.atoms[df_pred_found.water_id - 1]
            selected_water.residues.resids = df_pred_found["cluster"]
            selected_water.write("%s/cluster_found.pdb" % directory)    
            selected_water.tempfactors = df_pred_found["gO"]
            selected_water.write("%s/cluster_found_gO.pdb" % directory)
            selected_water.tempfactors = df_pred_found["esw"]
            selected_water.write("%s/cluster_found_esw.pdb" % directory)
            selected_water.tempfactors = df_pred_found["eww"]
            selected_water.write("%s/cluster_found_eww.pdb" % directory)
            selected_water.tempfactors = df_pred_found["tst"]
            selected_water.write("%s/cluster_found_tst.pdb" % directory)
            selected_water.tempfactors = df_pred_found["tso"]
            selected_water.write("%s/cluster_found_tso.pdb" % directory)
            selected_water.tempfactors = df_pred_found["dG"]
            selected_water.write("%s/cluster_found_dG.pdb" % directory)

        # NOT FOUND
        if not df_pred_not_found.empty:
            selected_water = u.atoms[df_pred_not_found.water_id - 1]
            selected_water.write("%s/cluster_not_found.pdb" % directory)
            selected_water.tempfactors = df_pred_not_found["gO"]
            selected_water.write("%s/cluster_not_found_gO.pdb" % directory)
            selected_water.tempfactors = df_pred_not_found["esw"]
            selected_water.write("%s/cluster_not_found_esw.pdb" % directory)
            selected_water.tempfactors = df_pred_not_found["eww"]
            selected_water.write("%s/cluster_not_found_eww.pdb" % directory)
            selected_water.tempfactors = df_pred_not_found["tst"]
            selected_water.write("%s/cluster_not_found_tst.pdb" % directory)
            selected_water.tempfactors = df_pred_not_found["tso"]
            selected_water.write("%s/cluster_not_found_tso.pdb" % directory)
            selected_water.tempfactors = df_pred_not_found["dG"]
            selected_water.write("%s/cluster_not_found_dG.pdb" % directory)
        
        pml_str = ""
        pml_str += "load ../../protein_prepared.pdb\n"
        pml_str += 'util.color_deep("white", \'protein_prepared\', 0)\n'
        pml_str += 'util.cnc("protein_prepared",_self=cmd)\n'
        pml_str += 'set sphere_scale, 0.2\n'
        
        pml_str += "\n"
        for reference in references:
            reference_name = reference.split("/")[-1]
            pml_str += "load ../../cluster_average_reference_renum_%s.pdb\n" % (reference_name)
        pml_str += 'load cluster_found.pdb\n'
        pml_str += 'load cluster_not_found.pdb\n'
        
        pml_str += "\n"
        for reference in references:
            reference_name = reference.split("/")[-1]
            pml_str += 'load ../../cluster_average_reference_renum_gO_%s.pdb\n' % (reference_name)
        pml_str += 'load cluster_found_gO.pdb\n'
        pml_str += 'load cluster_not_found_gO.pdb\n'
        pml_str += 'group gO, cluster_*_gO*\n'

        pml_str += "\n"
        for reference in references:
            reference_name = reference.split("/")[-1]
            pml_str += 'load ../../cluster_average_reference_renum_esw_%s.pdb\n' % (reference_name)
        pml_str += 'load cluster_found_esw.pdb\n'
        pml_str += 'load cluster_not_found_esw.pdb\n'
        pml_str += 'group Esw, cluster_*_esw*\n'
        
        pml_str += "\n"
        for reference in references:
            reference_name = reference.split("/")[-1]
            pml_str += 'load ../../cluster_average_reference_renum_eww_%s.pdb\n' % (reference_name)
        pml_str += 'load cluster_found_eww.pdb\n'
        pml_str += 'load cluster_not_found_eww.pdb\n'
        pml_str += 'group Eww, cluster_*_eww*\n'
        
        pml_str += "\n"
        for reference in references:
            reference_name = reference.split("/")[-1]
            pml_str += 'load ../../cluster_average_reference_renum_tst_%s.pdb\n' % (reference_name)
        pml_str += 'load cluster_found_tst.pdb\n'
        pml_str += 'load cluster_not_found_tst.pdb\n'
        pml_str += 'group dTS trans, cluster_*_tst*\n'
        
        pml_str += "\n"
        for reference in references:
            reference_name = reference.split("/")[-1]
            pml_str += 'load ../../cluster_average_reference_renum_tso_%s.pdb\n' % (reference_name)
        pml_str += 'load cluster_found_tso.pdb\n'
        pml_str += 'load cluster_not_found_tso.pdb\n'
        pml_str += 'group dTS orient, cluster_*_tso*\n'

        pml_str += "\n"
        for reference in references:
            reference_name = reference.split("/")[-1]
            pml_str += 'load ../../cluster_average_reference_renum_dG_%s.pdb\n' % (reference_name)
        pml_str += 'load cluster_found_dG.pdb\n'
        pml_str += 'load cluster_not_found_dG.pdb\n'
        pml_str += 'group dG, cluster_*_dG*\n'
        
        pml_str += "\n"
        pml_str += 'label cluster_average_reference_renum_gO*, b\n'
        pml_str += 'label cluster_average_reference_renum_esw*, b\n'
        pml_str += 'label cluster_average_reference_renum_eww*, b\n'
        pml_str += 'label cluster_average_reference_renum_tst*, b\n'
        pml_str += 'label cluster_average_reference_renum_tso*, b\n'
        pml_str += 'label cluster_average_reference_renum_dG*, b\n'
        pml_str += 'label cluster_found_*, b\n'
        pml_str += 'label cluster_not_found_*, b\n'
        pml_str += 'set sphere_scale, 0.3, cluster_found*\n'
        pml_str += 'set sphere_scale, 0.3, cluster_not_found*\n'
        for reference, color in zip(references, colors):
            reference_name = reference.split("/")[-1]
            pml_str += 'color %s, cluster_average_reference_renum*%s\n' % (color, reference_name)
        pml_str += 'color forest, cluster_found*\n'
        pml_str += 'color red, cluster_not_found*\n'
        pml_str += 'show lines\n'
        pml_str += 'show sphere, cluster*\n'
        
        with open("%s/pymol.pml" % (directory), "w") as w:
            w.write(pml_str)

        # Create PyMOL session
        os.chdir(directory)

        cmd.load("pymol.pml", quiet=False)

        filename = "pymol"
        if in_pocket_only:
            filename += "_pocket"
        if accessible_only:
            filename += "_accessible"
        filename += ".pse"

        cmd.save(filename)

        # Cleaning
        files_to_delete = glob.glob("cluster_*.pdb")
        for file_to_delete in files_to_delete:
            os.remove(file_to_delete)
        os.remove("pymol.pml")

        os.chdir(curdir)

        cmd.delete("all")

    files_to_delete = glob.glob("cluster_average_reference_renum_*.pdb")
    for file_to_delete in files_to_delete:
        os.remove(file_to_delete)

