# %% [code]
# %% [code]

import numpy as np
import pandas as pd
import math

import matplotlib.pyplot as plt
import seaborn as sns
import os

import time

from script_cfg import *

################################################################## MAIN VARIABLES
    
class TRAIN:
    LINES_NB = 1643680
    SEQ_UNIQUE = LINES_NB / 2
    CHUNKSIZE = 82184 * 2
            
    DATASET = {
        '15k_2A3': 15000, '15k_DMS': 15000,
        'DasLabBigLib_OneMil_15K_REP_2A3': 15000, 'DasLabBigLib_OneMil_15K_REP_DMS': 65000, 'DasLabBigLib_OneMil_Coronavirus_genomes_SARS_related_betacoronaviruses_splitA_2A3': 70000, 'DasLabBigLib_OneMil_Coronavirus_genomes_SARS_related_betacoronaviruses_splitA_DMS': 20000,
        'DasLabBigLib_OneMil_Coronavirus_genomes_SARS_related_betacoronaviruses_splitB_2A3': 60000, 'DasLabBigLib_OneMil_Coronavirus_genomes_SARS_related_betacoronaviruses_splitB_DMS': 70000, 'DasLabBigLib_OneMil_Coronavirus_genomes_SARS_related_betacoronaviruses_splitC_2A3': 30000, 'DasLabBigLib_OneMil_Coronavirus_genomes_SARS_related_betacoronaviruses_splitC_DMS': 73298,
        'DasLabBigLib_OneMil_OpenKnot_Round_2_train_2A3': 61473, 'DasLabBigLib_OneMil_OpenKnot_Round_2_train_DMS': 133716, 'DasLabBigLib_OneMil_RFAM_windows_100mers_2A3': 29227, 'DasLabBigLib_OneMil_RFAM_windows_100mers_DMS': 13686, 'DasLabBigLib_OneMil_RFAM_REP_2A3': 15541, 'DasLabBigLib_OneMil_RFAM_REP_DMS': 15541,
        'DasLabBigLib_OneMil_RNAmake_designs_2A3': 55252, 'DasLabBigLib_OneMil_RNAmake_designs_DMS': 41566, 'DasLabBigLib_OneMil_RNAmake_designs_delete_longest_flex_helix_2A3': 87580, 'DasLabBigLib_OneMil_RNAmake_designs_delete_longest_flex_helix_DMS': 46014, 'DasLabBigLib_OneMil_RNAmake_designs_insert_bps_into_longest_flex_helix_2A3': 95890, 'DasLabBigLib_OneMil_RNAmake_designs_insert_bps_into_longest_flex_helix_DMS': 66866,
        'DasLabBigLib_OneMil_Single_nt_mutants_OpenKnot_Pilot_PK50_2A3': 16990, 'DasLabBigLib_OneMil_Single_nt_mutants_OpenKnot_Pilot_PK50_DMS': 51000, 'DasLabBigLib_OneMil_Replicates_from_previous_libraries_2A3': 16990, 'DasLabBigLib_OneMil_Replicates_from_previous_libraries_DMS': 16990,
        'OpenKnot1_Twist_2A3_EternaPlayers': 50192, 'OpenKnot1_Twist_DMS_EternaPlayers': 16182,
        'PK50_Wu_Twist_2A3': 2729, 'PK50_Wu_Twist_DMS': 2729, 'PK50_CustomArray_2A3': 2729, 'PK50_CustomArray_DMS': 2729, 'PK50_Twist_2A3': 2729, 'PK50_Twist_DMS': 2729, 'PK50_Twist_epPCR_2A3': 2729, 'PK50_Twist_epPCR_DMS': 2729, 'PK50_AltChemMap_NovaSeq_2A3': 18911, 'PK50_AltChemMap_NovaSeq_DMS': 2729, 'PK90_CustomArray_DMS': 2729, 'PK90_CustomArray_2A3': 2729, 'PK90_Twist_2A3': 2729, 'PK90_Twist_DMS': 2729, 'PK90_Twist_epPCR_2A3': 2729, 'PK90_Twist_epPCR_DMS': 2729,
        'SL5_M2seq_2A3': 2729, 'SL5_M2seq_DMS': 2729,
    }


################################################################## PLOTS ON TRAIN

def plot_dist(df, cols, title):
    ##### Figure
    fig, ax = plt.subplots(nrows=2, ncols=len(cols), figsize=(20, 4))
    fig.suptitle(title, color='#3f3f3f', y=1.1, fontweight='bold', fontsize='large')
    fig.subplots_adjust(left=0.05, bottom=0.05, right=0.95, top=0.95, wspace=0.2, hspace=0.2)

    ##### Data
    for j, col in zip(range(len(cols)), cols):      
        sns.histplot(data=df[cols], x=col, kde=True, bins=50, edgecolor=None, ax=ax[0, j], color='#008bf8ff')
        ax[0, j].set_title(col.capitalize(), loc='left', pad=10, fontdict={'fontsize':'large', 'color':'#3f3f3f'})
        sns.boxplot(data=df[cols], x=col, orient='h', saturation=0.9, showmeans=True,meanprops=dict(markerfacecolor='w', markeredgecolor='w', marker='o'), medianprops=dict(color='w', alpha=1, linewidth=1), 
                    boxprops=dict(color='#008bf8ff', linewidth=False), whiskerprops=dict(color='#008bf8ff', linewidth=1), capprops=dict(color='#151515', linewidth=1), showfliers=True, fliersize=3, flierprops=dict(markerfacecolor='#89fc00ff', markeredgecolor='#89fc00ff', marker='d'),
                    ax=ax[1, j])

    ##### Style
    for i in range(2):
        for j in range(len(cols)):
            ax[i, j].set_xlabel('', loc='center', color='#3f3f3f', fontsize='small')
            ax[i, j].set_ylabel('', loc='center', color='#3f3f3f', fontsize='small', rotation='vertical', labelpad=4)
            ax[i, j].grid(axis="both", lw=0.5, ls=':')
            ax[i, j].spines[['right', 'top']].set_visible(False)
            ax[i, j].spines['left'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
            ax[i, j].spines['bottom'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
            ax[i, j].tick_params(axis='both', color='#3f3f3f', colors='#3f3f3f', labelsize='small')


def plot_reactivity(df, err=False, title='Reactivity per base-position on several sequences'):
    ncols = 10
    if len(df) <= 10:
        ncols=len(df)
    
    if err == True:
        cols = CFG.COLS_REACT_ERR
    else:
        cols = CFG.COLS_REACT
    
    fig, ax = plt.subplots(nrows=1, ncols=ncols, figsize=(20, 2))
    fig.suptitle(title, color='#3f3f3f', y=1.1, fontweight='bold', fontsize='large')
    fig.subplots_adjust(left=0.05, bottom=0.05, right=0.95, top=0.95, wspace=0.2, hspace=0.2)

    for j in range(ncols):
        sns.lineplot(df[cols].values[j], ls='-', lw='0.5', marker='', color='#dc0073ff', ax=ax[j])
        ax[j].set_xlim(xmin=0, xmax=len(df.sequence.values[j]))
        ax[j].hlines(y=0, xmin=0, xmax=len(df.sequence.values[j]), ls='-', lw=0.5, colors='#04e762ff')
        if err == False:
            ax[j].hlines(y=1, xmin=0, xmax=len(df.sequence.values[j]), ls='-', lw=0.5, colors='#04e762ff')
        ax[j].set_xlabel('', loc='center', color='#3f3f3f', fontsize='small')
        ax[j].set_ylabel('', loc='center', color='#3f3f3f', fontsize='small', rotation='vertical', labelpad=4)
        ax[j].grid(axis="both", lw=0.5, ls=':')
        ax[j].spines[['right', 'top']].set_visible(False)
        ax[j].spines['left'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
        ax[j].spines['bottom'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
        ax[j].tick_params(axis='both', color='#3f3f3f', colors='#3f3f3f', labelsize='x-small')
        
        
def plot_hist_datasets_size(figsize = (10, 2)):
    fig = plt.figure(figsize=figsize)
    fig.suptitle("Distribution of the datasets' size",color='#3f3f3f', y=1.05, fontweight='bold', fontsize='medium')
    fig.subplots_adjust(left=0.05, bottom=0.05, right=0.95, top=0.95, wspace=0.2, hspace=0.2)

    mosaic = '''AC
                AB
                '''
    mosaiq_ = ['A', 'B', 'C']
    ax = fig.subplot_mosaic(mosaic)

    ##### Data
    sns.histplot(TRAIN.DATASET.values(), kde=True, bins=56, edgecolor=None, ax=ax['A'])
    sns.boxplot(list(TRAIN.DATASET.values()), orient='h', saturation=0.9, 
            showmeans=True, meanprops=dict(markerfacecolor='w', markeredgecolor='w', marker='o'), medianprops=dict(color='w', alpha=1, linewidth=2),
            boxprops=dict(linewidth=False), whiskerprops=dict(color='#151515', linewidth=1), capprops=dict(color='#151515', linewidth=1),
            showfliers=True, fliersize=3, flierprops=dict(markerfacecolor='#6a6a6a', markeredgecolor='#6a6a6a', marker='d'),
            ax=ax['B'])

    ##### Style
    for i in ['A', 'B']:
        ax[i].grid(axis="both", lw=0.5, ls=':')
        ax[i].spines[['right', 'top']].set_visible(False)
        ax[i].spines['left'].set(color='#3f3f3f', linewidth=0.5)
        ax[i].spines['bottom'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
        ax[i].tick_params(axis='both', color='#3f3f3f', colors='#3f3f3f', labelsize='x-small')
        ax[i].legend(loc='best', fontsize='small', labelcolor='w').get_frame().set(visible=True, facecolor='w', edgecolor='w', linewidth=1, alpha=0.7);
        
    ax['C'].spines[['right', 'top', 'left', 'bottom']].set_visible(False)
    ax['C'].tick_params(axis='both', color='w', colors='w')

def plot_hist_stats(react_stats, cols, title):
    ##### Figure
    fig = plt.figure(figsize=(20, 5))
    fig.suptitle(title, color='#3f3f3f', y=1.1, fontweight='bold', fontsize='large')
    fig.subplots_adjust(left=0.05, bottom=0.05, right=0.95, top=0.95, wspace=0.2, hspace=0.4)
    mosaic = '''ABCDEF
                GHIJKL
                '''
    mosaiq_hist = ['A', 'B', 'C', 'D', 'E', 'F']
    mosaiq_box = ['G', 'H', 'I', 'J', 'K', 'L']
    ax = fig.subplot_mosaic(mosaic)

    for h, b, index in zip(mosaiq_hist, mosaiq_box, ['not_null', 'min', 'max', 'mean', 'median_est', 'std_est']):  
        data = react_stats[cols].loc[index]
        
        sns.histplot(data=data, kde=True, bins=50, edgecolor=None, ax=ax[h]);
        ax[h].set_title(index.capitalize(), loc='left', pad=10, fontdict={'fontsize':'large', 'color':'#3f3f3f'})
        sns.boxplot(data=data, orient='h', saturation=0.9, showmeans=True,meanprops=dict(markerfacecolor='w', markeredgecolor='w', marker='o'),
                medianprops=dict(color='w', alpha=1, linewidth=2), boxprops=dict(linewidth=False), whiskerprops=dict(color='#151515', linewidth=1), capprops=dict(color='#151515', linewidth=1),
                showfliers=True, fliersize=3, flierprops=dict(markerfacecolor='#6a6a6a', markeredgecolor='#6a6a6a', marker='d'),
                ax=ax[b])

    for i in mosaiq_hist + mosaiq_box:
        ax[i].set_xlabel(ax[i].get_xlabel(), loc='center', color='#3f3f3f', fontsize='small')
        ax[i].set_ylabel(ax[i].get_ylabel(), loc='center', color='#3f3f3f', fontsize='small', rotation='vertical', labelpad=4)
        ax[i].set_xlim(xmin=None, xmax=None)
        ax[i].grid(axis="both", lw=0.5, ls=':')
        ax[i].spines[['right', 'top']].set_visible(False)
        ax[i].spines['left'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
        ax[i].spines['bottom'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
        ax[i].tick_params(axis='both', color='#3f3f3f', colors='#3f3f3f', labelsize='small')

def plot_reads_SN(train_partial):
    ''' Plot a scatter between reads and signal_to_noise with SF filter as hue
    '''
    fig = plt.figure(figsize=(8, 3))
    ax = plt.subplot()
    sns.scatterplot(data=train_partial, x='signal_to_noise', y='reads', hue='SN_filter', palette=['#008bf8ff', '#89fc00ff'], ax=ax);
    ax.set_xlabel(ax.get_xlabel(), loc='center', color='#3f3f3f', fontsize='small')
    ax.set_ylabel(ax.get_ylabel(), loc='center', color='#3f3f3f', fontsize='small', rotation='vertical', labelpad=4)
    ax.set_title('Reads vs Signal_to_Noise', loc='left', pad=10, fontdict={'fontsize':'medium', 'color':'#3f3f3f'})
    ax.grid(axis="both", lw=0.5, ls=':')
    ax.spines[['right', 'top']].set_visible(False)
    ax.spines['left'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
    ax.spines['bottom'].set(color='#3f3f3f', linewidth=0.5, position=('axes', 0))
    ax.tick_params(axis='both', color='#3f3f3f', colors='#3f3f3f', labelsize='x-small')
    ax.legend(loc='best', fontsize='small', labelcolor='#3f3f3f').get_frame().set(visible=True, facecolor='w', edgecolor='w', linewidth=1, alpha=0.7);
    
################################################################## TEST

class TEST:
    LINES_NB = 1343823
    PRED = 269796671
    CHUNKSIZE = 1000000

def get_sample_test_part():
    
    sample_chunks = pd.read_csv(CFG.SAMPLE, chunksize = TEST.CHUNKSIZE, dtype=CFG.DTYPES)
    sample_chunks_nb = 0
    sample_chunks_len = []

    for sample_part in sample_chunks:
        sample_chunks_nb += 1
        sample_chunks_len.append(len(sample_part))
        break

    test_chunks = pd.read_csv(CFG.TEST, chunksize = TEST.CHUNKSIZE)
    test_chunks_nb = 0
    test_chunks_len = []

    for test_part in test_chunks:
        test_chunks_nb += 1
        test_chunks_len.append(len(test_part))
        break

    sum(sample_chunks_len), sum(test_chunks_len)
    
    return sample_part, test_part